File size: 4,680 Bytes
a03335c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
#!/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")