"""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)
