Download training_code/verify_execution.py from Recharge23/JAM-20K-single: direct link, hf CLI and curl.
- Browser
- Download file 2.93 kB
-
https://huggingface.co/Recharge23/JAM-20K-single/resolve/main/training_code/verify_execution.py
- Command line
-
hf download hf://Recharge23/JAM-20K-single/training_code/verify_execution.py
-
curl -L -o verify_execution.py https://huggingface.co/Recharge23/JAM-20K-single/resolve/main/training_code/verify_execution.py
2.93 kB
| 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() | |