JAM-20K-single / training_code /verify_execution.py
Recharge23's picture
Release separate task 4 and task 5 Franka-JAM 20K adapters
5e5ffd1 verified
Raw History Blame Contribute Delete
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()