File size: 4,348 Bytes
d8acb64
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
86
87
88
89
90
91
92
93
94
#!/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")