aether-phase1 / scripts /train_align.py
SupremeD's picture
Upload scripts/train_align.py with huggingface_hub
d8acb64 verified
Raw History Blame Contribute Delete
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")