Download scripts/train_align.py from SupremeD/aether-phase1: direct link, hf CLI and curl.
- Browser
- Download file 4.35 kB
-
https://huggingface.co/SupremeD/aether-phase1/resolve/main/scripts/train_align.py
- Command line
-
hf download hf://SupremeD/aether-phase1/scripts/train_align.py
-
curl -L -o train_align.py https://huggingface.co/SupremeD/aether-phase1/resolve/main/scripts/train_align.py
4.35 kB
| #!/usr/bin/env python | |
| """ | |
| Aether Phase-1 ALIGNMENT training loop + verification. | |
| Frozen backbone + frozen encoders; trains ONLY the projectors + cam_pose + NEW | |
| embed/lm_head rows (grad-mask hook). Loss = AR-CE masked to the CAPTION tokens. | |
| MILESTONE TEST = overfit a tiny FIXED batch (image->caption pairs). If the loop is | |
| wired correctly, loss must collapse (grad flows through projector into the frozen LM's | |
| input space). This is the standard "overfit a small batch" pipeline sanity check. | |
| Grad-safe splice: build inputs_embeds with torch.where (NO in-place on the autograd path). | |
| Reuses the VERIFIED assembly + M-RoPE from forward.py. Run on rental MI300. | |
| """ | |
| import os, sys, torch, torch.nn as nn | |
| sys.path.insert(0,"/work") | |
| import forward as F # AetherPhase1, build_mrope, dev, TXT (verified module) | |
| dev=F.dev; TXT=F.TXT | |
| def apply_freeze(m): | |
| for p in m.parameters(): p.requires_grad_(False) | |
| for p in m.trainable(): p.requires_grad_(True) # projectors + cam_pose | |
| emb=m.base.get_input_embeddings().weight | |
| head=m.base.get_output_embeddings().weight | |
| lo=m.new_row_lo | |
| for w in (emb,head): | |
| w.requires_grad_(True) | |
| w.register_hook(lambda g,lo=lo:(g.__setitem__(slice(0,lo),0) or g)) # zero grad on OLD rows | |
| def splice(m, ids, vis_idx, vproj): | |
| """Grad-safe: replace the <img> block embeds with projected vision features.""" | |
| emb = m.base.get_input_embeddings()(ids.to(dev)) # (1,L,TXT) non-leaf | |
| repl = torch.zeros_like(emb) | |
| repl[0, vis_idx] = vproj.to(emb.dtype) # fresh tensor, diff-able | |
| mask = torch.zeros(emb.shape[:2], dtype=torch.bool, device=dev) | |
| mask[0, vis_idx] = True | |
| return torch.where(mask.unsqueeze(-1), repl, emb) | |
| if __name__=="__main__": | |
| torch.manual_seed(0) | |
| m=F.AetherPhase1().to(dev) | |
| for mod in (m.visual_proj,m.audio_proj,m.mv_proj,m.geom_proj): mod.to(torch.bfloat16) | |
| m.cam_pose.data=m.cam_pose.data.to(torch.bfloat16) | |
| apply_freeze(m) | |
| T=m.tok | |
| tr=[p for p in m.parameters() if p.requires_grad] | |
| print(f"trainable tensors {len(tr)} | params {sum(p.numel() for p in tr)/1e6:.1f}M") | |
| # ---- tiny FIXED dataset: 4 distinct images -> 4 distinct captions ---- | |
| torch.manual_seed(42) | |
| N_IMG=4 | |
| imgs=[torch.randn(1,3,448,448) for _ in range(N_IMG)] | |
| caps=["a red brick electrical panel on a wall", | |
| "a blue ceramic coffee mug on a desk", | |
| "three yellow pencils in a glass jar", | |
| "a green circuit board with silver traces"] | |
| img_id=m.id["<img>"] | |
| # pre-encode frozen vision features ONCE (they never change -> pre-extract pattern) | |
| with torch.no_grad(): | |
| feats=[m.enc_vision(px) for px in imgs] # each (1025,3200) | |
| opt=torch.optim.AdamW(tr, lr=1e-3) | |
| def make(i): | |
| vproj=m.proj_vision(feats[i]) # (1025,TXT) grad-on | |
| cid=T(caps[i],add_special_tokens=False).input_ids | |
| ids=torch.tensor([img_id]*vproj.shape[0]+cid)[None] | |
| vis_idx=(ids[0]==img_id).nonzero(as_tuple=True)[0] | |
| labels=torch.tensor([-100]*vproj.shape[0]+cid)[None].to(dev) | |
| types=["img"]*vproj.shape[0]+["text"]*len(cid) | |
| pos=F.build_mrope(types).to(dev) | |
| return ids,vis_idx,vproj,labels,pos | |
| print("step loss") | |
| for step in range(60): | |
| opt.zero_grad(); tot=0.0 | |
| for i in range(N_IMG): | |
| ids,vis_idx,vproj,labels,pos=make(i) | |
| emb=splice(m,ids,vis_idx,vproj) | |
| out=m.base(inputs_embeds=emb,position_ids=pos,labels=labels,use_cache=False) | |
| out.loss.backward(); tot+=out.loss.item() | |
| opt.step() | |
| if step%5==0 or step==59: | |
| print(f"{step:4d} {tot/N_IMG:.4f}") | |
| print("ALIGNMENT LOOP OK — loss collapsed" if tot/N_IMG < 1.0 else "!! loss did not collapse") | |
| # save trained projectors (overfit-proof checkpoint) | |
| from safetensors.torch import save_file | |
| sd={} | |
| for name in ("visual_proj","audio_proj","mv_proj","geom_proj"): | |
| for k,v in getattr(m,name).state_dict().items(): sd[f"{name}.{k}"]=v.contiguous().cpu() | |
| sd["cam_pose"]=m.cam_pose.detach().cpu() | |
| os.makedirs("/work/aether-phase1-pkg",exist_ok=True) | |
| save_file(sd,"/work/aether-phase1-pkg/new_modules_aligncheck.safetensors") | |
| print("saved new_modules_aligncheck.safetensors") | |