aether-phase1 / scripts /train_cached.py
SupremeD's picture
Upload scripts/train_cached.py with huggingface_hub
a03335c verified
Raw History Blame Contribute Delete
4.68 kB
#!/usr/bin/env python
"""
Aether Phase-1 CACHED alignment trainer — trains from pre-extracted features with
NO ENCODERS LOADED. This is the portable path: base LM (or 4-bit QLoRA of it) + the
tiny projectors fit a 16GB card, so alignment/adapters run on 6900XT / Kaggle-T4 / Colab.
Consumes /work/slivers_cache/{index.jsonl, <modality>/<id>.pt} from preextract.py.
Proves loss drops on REAL cached features across ALL THREE modalities at once.
"""
import os, sys, json, torch, torch.nn as nn
sys.path.insert(0,"/work")
import forward as F # build_mrope, Projector, build_special, TXT, dev (no encoder load on import)
from transformers import AutoModelForImageTextToText, AutoTokenizer
dev=F.dev; TXT=F.TXT; CACHE="/work/slivers_cache"
class CachedAether(nn.Module):
"""Base LM + projectors + expanded vocab. NO encoders (features come pre-extracted)."""
def __init__(self):
super().__init__()
self.base=AutoModelForImageTextToText.from_pretrained("/work/base",torch_dtype=torch.bfloat16,trust_remote_code=True)
self.tok=AutoTokenizer.from_pretrained("/work/base",trust_remote_code=True)
old=len(self.tok); self.tok.add_special_tokens({"additional_special_tokens":F.build_special()})
self.new_vocab=len(self.tok); self.base.resize_token_embeddings(self.new_vocab); self.new_row_lo=old
self.visual_proj=F.Projector(3200,TXT); self.audio_proj=F.Projector(1280,TXT)
self.mv_proj=F.Projector(3200,TXT); self.geom_proj=F.Projector(8,TXT)
self.cam_pose=nn.Parameter(torch.zeros(1,1,TXT))
self.id={k:self.tok.convert_tokens_to_ids(k) for k in ["<img>","<audio>","<3d_app>","<3d_geom>"]}
def trainable(self):
for m in (self.visual_proj,self.audio_proj,self.mv_proj,self.geom_proj): yield from m.parameters()
yield self.cam_pose
def apply_freeze(m):
for p in m.parameters(): p.requires_grad_(False)
for p in m.trainable(): p.requires_grad_(True)
lo=m.new_row_lo
for w in (m.base.get_input_embeddings().weight, m.base.get_output_embeddings().weight):
w.requires_grad_(True); w.register_hook(lambda g,lo=lo:(g.__setitem__(slice(0,lo),0) or g))
PROJ={"vision":("visual_proj","<img>"),"audio":("audio_proj","<audio>"),"geom":("geom_proj","<3d_geom>")}
def splice(m, ids, idxpos, vproj):
emb=m.base.get_input_embeddings()(ids.to(dev))
repl=torch.zeros_like(emb); repl[0,idxpos]=vproj.to(emb.dtype)
mask=torch.zeros(emb.shape[:2],dtype=torch.bool,device=dev); mask[0,idxpos]=True
return torch.where(mask.unsqueeze(-1),repl,emb)
if __name__=="__main__":
torch.manual_seed(0)
m=CachedAether().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)
print("[cached] base+projectors loaded, NO ENCODERS. VRAM GB:",round(torch.cuda.memory_allocated()/1e9,1))
items=json.load(open(f"{CACHE}/index.jsonl"))
print(f"[cached] {len(items)} pre-extracted items:",{r['modality'] for r in items})
T=m.tok
def make(r):
projname,ptok=PROJ[r["modality"]]; pid=m.id[ptok]; proj=getattr(m,projname)
blob=torch.load(r["feat"],map_location=dev)
if r["modality"]=="geom":
feat=blob["feat"].to(dev); xyz=blob["xyz"].to(dev); v=proj(feat.to(torch.bfloat16))+m.cam_pose.squeeze(0)
else:
feat=blob.to(dev); xyz=None; v=proj(feat.to(torch.bfloat16))
cid=T(r["caption"],add_special_tokens=False).input_ids
ids=torch.tensor([pid]*v.shape[0]+cid)[None]
pos_ids=(ids[0]==pid).nonzero(as_tuple=True)[0]
labels=torch.tensor([-100]*v.shape[0]+cid)[None].to(dev)
ttype=("3d_geom" if r["modality"]=="geom" else ("audio" if r["modality"]=="audio" else "img"))
types=[ttype]*v.shape[0]+["text"]*len(cid)
pos=F.build_mrope(types,geom_xyz=xyz).to(dev)
return ids,pos_ids,v,labels,pos
opt=torch.optim.AdamW([p for p in m.parameters() if p.requires_grad],lr=1e-3)
STEPS=int(os.environ.get("STEPS","200"))
print(f"step loss (STEPS={STEPS})")
for step in range(STEPS):
opt.zero_grad(); tot=0.0
for r in items:
ids,pos_ids,v,labels,pos=make(r)
out=m.base(inputs_embeds=splice(m,ids,pos_ids,v),position_ids=pos,labels=labels,use_cache=False)
out.loss.backward(); tot+=out.loss.item()
opt.step()
if step%20==0 or step==STEPS-1: print(f"{step:4d} {tot/len(items):.4f}")
ok = tot/len(items) < 1.0
print("CACHED MULTIMODAL ALIGNMENT OK (no encoders) — loss collapsed" if ok else "!! loss did not collapse")