#!/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 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[""] # 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")