"""Bounded H3 pretrained block-0 ROCm backward probe; no video-quality claim.""" import json, time, hashlib from pathlib import Path import torch from safetensors.torch import load_file, save_file from diffusers.models.transformers.transformer_minimax_h3 import MiniMaxH3TransformerBlock, MiniMaxH3RotaryPosEmbed torch.manual_seed(42) torch.cuda.reset_peak_memory_stats() raw = load_file('h3-block0.safetensors') state = {} for key, value in raw.items(): key = key.removeprefix('blocks.0.') if key == 'attn.qkv_proj.weight': for target, part in zip(('to_q', 'to_k', 'to_v'), value.chunk(3, 0)): state[f'attn.{target}.weight'] = part elif key == 'mlp.fc1.weight': gate, val = value.chunk(2, 0) state['ff.net.0.proj.weight'] = torch.cat((val, gate), 0) else: key = key.replace('attn.out_proj.', 'attn.to_out.0.').replace('attn.q_norm.', 'attn.norm_q.').replace('attn.k_norm.', 'attn.norm_k.').replace('mlp.fc2.', 'ff.net.2.') state[key] = value with torch.device('meta'): block = MiniMaxH3TransformerBlock(5376, 56, 128, 14336, 2688, 1e-5, 1e-5) block.load_state_dict(state, strict=True, assign=True) block = block.to(device='cuda', dtype=torch.bfloat16).requires_grad_(False) class LoRA(torch.nn.Module): def __init__(self, base): super().__init__() self.base = base self.a = torch.nn.Parameter(torch.randn(4, base.in_features, device='cuda') * .01) self.b = torch.nn.Parameter(torch.zeros(base.out_features, 4, device='cuda')) def forward(self, x): return self.base(x) + ((x.float() @ self.a.T) @ self.b.T).to(x.dtype) * 2 block.attn.to_q = LoRA(block.attn.to_q) block.attn.to_v = LoRA(block.attn.to_v) params = [p for p in block.parameters() if p.requires_grad] optimizer = torch.optim.Adam(params, lr=1e-4) x = (torch.randn(1, 32, 5376, device='cuda', dtype=torch.bfloat16) * .1).requires_grad_() temb = torch.randn(1, 2688, device='cuda') * .1 tags = torch.arange(32, device='cuda') % 3 positions = torch.zeros(32, 3, device='cuda'); positions[:, 2] = torch.arange(32, device='cuda') rope = MiniMaxH3RotaryPosEmbed().to('cuda')(positions) steps = [] for step in range(2): start = time.perf_counter() optimizer.zero_grad(); x.grad = None y = block(x, temb, tags, rope) loss = (y.float() - x.detach().float()).square().mean() loss.backward() assert torch.isfinite(loss) and torch.isfinite(x.grad).all() assert all(p.grad is not None and torch.isfinite(p.grad).all() for p in params) grad_norm = sum(p.grad.float().square().sum() for p in params).sqrt().item() assert grad_norm > 0 optimizer.step(); torch.cuda.synchronize() row = dict(step=step, loss=loss.item(), adapter_grad_norm=grad_norm, input_grad_norm=x.grad.float().norm().item(), seconds=time.perf_counter()-start) steps.append(row); print(json.dumps(row), flush=True) adapters = {k: v.detach().cpu() for k,v in block.named_parameters() if v.requires_grad} save_file(adapters, 'h3-block0-adapter.safetensors') result = dict(scope='pretrained H3 block 0 only, synthetic hidden states, no full video pipeline or distributed H3', revision='42ed227ee7df40d41602854ae760620d6eb651fe', block_parameters=sum(p.numel() for p in block.parameters())-sum(p.numel() for p in params), trainable_parameters=sum(p.numel() for p in params), sequence_length=32, steps=steps, gpu_peak_allocated_bytes=torch.cuda.max_memory_allocated(), gpu=torch.cuda.get_device_name(), torch_version=torch.__version__, checkpoint_mapping='strict keys; fused QKV split, SwiGLU gate/value swap; upstream numerical parity not yet qualified') Path('h3-block-result.json').write_text(json.dumps(result, indent=2)) print(json.dumps(result), flush=True)