import os,json,copy from pathlib import Path import torch,torch.distributed as dist from safetensors.torch import load_file from jam_real.load_exact import load_exact from jam_real.adapt_model import make_batch from jam_real.multitask_finetune import copy_adapters from jam_real.fast_execution import sum_gradients_fast rank=int(os.environ['RANK']);local=int(os.environ['LOCAL_RANK']) torch.cuda.set_device(local);device=torch.device('cuda',local);torch.set_num_threads(4) dist.init_process_group('nccl');torch.backends.cuda.matmul.allow_tf32=False D=Path('/root/autodl-tmp/realrobot_task45_20260922/task_4') model=load_exact('/dev/shm/franka_20k_20260922/Franka-JAM.safetensors',D/'repo/code/JAM/configs/franka',device,torch.bfloat16) named={n:p for n,p in model.named_parameters() if p.requires_grad};ps=list(named.values()) copy_adapters(named,load_file(str(D/'run/final.safetensors')));model.train() rows=[r for r in json.loads((D/'cache/manifest.json').read_text())['rows'] if r['split']=='train'] raw=[load_file(str(D/'cache'/r['file'])) for r in rows[rank*4:rank*4+4]] def grad(ckpt): model.zero_grad(set_to_none=True);torch.manual_seed(301+rank) model._mot_driver.mot_checkpoint_mixed_attn=ckpt b=make_batch(raw,device,torch.bfloat16);b['use_gradient_checkpointing']=ckpt with torch.autocast('cuda',dtype=torch.bfloat16): loss=model.compute_loss(**b,lambda_video=1.,lambda_action=1.)['loss'] loss.backward();return float(loss.detach()),torch.cat([p.grad.detach().flatten() for p in ps]) l0,g0=grad(True);l1,g1=grad(False) rel=float((g0-g1).norm()/g0.norm());cos=float(torch.nn.functional.cosine_similarity(g0,g1,dim=0)) # BF16 backward accumulation order differs with recomputation; accept <1% and cosine >.9999. assert abs(l0-l1)<1e-6 and rel<.01 and cos>.9999,(l0,l1,rel,cos) del g0,g1 # The two-rank flattened reduction must equal separate SUM collectives exactly. small=[torch.nn.Parameter(torch.zeros(7+i,device=device)) for i in range(3)] for p in small:p.grad=torch.randn_like(p) expected=[p.grad.clone() for p in small] for x in expected:dist.all_reduce(x,op=dist.ReduceOp.SUM) sum_gradients_fast(small,2)() assert all(torch.equal(p.grad,x) for p,x in zip(small,expected)) # Check fused/non-fused AdamW arithmetic with matching moments and step counts. p=torch.nn.Parameter(torch.randn(10000,device=device));q=torch.nn.Parameter(p.detach().clone()) o=torch.optim.AdamW([p],lr=1e-5,betas=(.9,.95),weight_decay=.01,foreach=False) f=torch.optim.AdamW([q],lr=1e-5,betas=(.9,.95),weight_decay=.01,fused=True) for _ in range(5): g=torch.randn_like(p);p.grad=g;q.grad=g.clone();o.step();f.step() assert torch.allclose(p,q,atol=5e-7,rtol=1e-6) out={'rank':rank,'status':'PASS','checkpointed_loss':l0,'uncheckpointed_loss':l1,'gradient_relative_error':rel,'gradient_cosine':cos,'flat_sum_exact':True,'fused_adam_max_abs_diff':float((p-q).abs().max())} print(json.dumps(out),flush=True) dist.destroy_process_group()