File size: 5,669 Bytes
bed4b7b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
#!/usr/bin/env python
"""
Aether Phase-1 ASSEMBLY: load base + 3 verified encoders into one model, expand vocab,
wire 5 projectors, apply freeze scheme, report. Runs on rental MI300 (aefinal).
Milestone = does the whole thing instantiate + freeze correctly (not a full forward yet).
Encoder loads use the recipes verified 2026-09-27.
"""
import os, sys, types, torch, torch.nn as nn, torch.nn.functional as F
os.environ["ATTN_BACKEND"]="sdpa"; os.environ["SPARSE_ATTN_BACKEND"]="sdpa"
os.environ["XFORMERS_DISABLED"]="1"; os.environ["SPARSE_BACKEND"]="torchsparse"

# --- shims (verified required) ---
from transformers.modeling_utils import PreTrainedModel as _PTM
# settable default: InternViT (older code) lacks it -> reads {}; Qwen3-VL (5.17) SETS it in post_init -> must be settable
if not hasattr(_PTM,"all_tied_weights_keys"):
    _PTM.all_tied_weights_keys = {}
fa=types.ModuleType("flash_attn")
def _favarlen(q,k,v,cu_q,cu_k,mq,mk,dropout_p=0.0,softmax_scale=None,causal=False,**kw):
    cq=cu_q.tolist(); outs=[]
    for i in range(len(cq)-1):
        s,e=cq[i],cq[i+1]
        o=F.scaled_dot_product_attention(q[s:e].transpose(0,1)[None],k[s:e].transpose(0,1)[None],v[s:e].transpose(0,1)[None],is_causal=causal,scale=softmax_scale)
        outs.append(o[0].transpose(0,1))
    return torch.cat(outs,0)
fa.flash_attn_varlen_func=_favarlen; fa.flash_attn_func=lambda *a,**k:None; sys.modules["flash_attn"]=fa

from transformers import AutoModel, AutoModelForImageTextToText, AutoTokenizer
from safetensors.torch import load_file
dev="cuda"; TXT=4096

class Projector(nn.Module):
    def __init__(self,i,o):
        super().__init__(); self.net=nn.Sequential(nn.Linear(i,o),nn.GELU(),nn.Linear(o,o))
    def forward(self,x): return self.net(x)

def load_base():
    print("[base] loading Qwen3-VL-8B...")
    m=AutoModelForImageTextToText.from_pretrained("/work/base",torch_dtype=torch.bfloat16,trust_remote_code=True)
    tok=AutoTokenizer.from_pretrained("/work/base",trust_remote_code=True)
    return m,tok

def load_internvit():
    print("[vision] loading InternViT-6B...")
    return AutoModel.from_pretrained("/work/encoders/InternViT-6B-448px-V2_5",trust_remote_code=True,torch_dtype=torch.bfloat16)

def load_mimo_audio():
    print("[audio] loading MiMo-Audio encoder...")
    sys.path.insert(0,"/work/MiMo-Audio-src/src")
    from mimo_audio_tokenizer import MiMoAudioTokenizer, MiMoAudioTokenizerConfig
    cfgp="/work/encoders/MiMo-Audio-Tokenizer"; cfg=MiMoAudioTokenizerConfig.from_pretrained(cfgp)
    m=MiMoAudioTokenizer(cfg); m.load_state_dict(load_file(cfgp+"/model.safetensors"),strict=False)
    return m.encoder            # frozen encoder only (packed-mel -> 1280)

def load_trellis_slat():
    print("[3d] loading TRELLIS-SLAT encoder...")
    sys.path.insert(0,"/work/TRELLIS-AMD")
    from trellis.models.structured_latent_vae.encoder import SLatEncoder
    enc=SLatEncoder(resolution=64,in_channels=1024,model_channels=768,latent_channels=8,num_blocks=12,num_heads=12,mlp_ratio=4,attn_mode="swin",window_size=8,use_fp16=True)
    enc.load_state_dict(load_file("/work/encoders/TRELLIS-image-large/ckpts/slat_enc_swin8_B_64l8_fp16.safetensors"),strict=False)
    return enc

# --- special tokens (full structural set) ---
def build_special():
    S=["<img>","</img>","<audio>","</audio>","<3d_app>","</3d_app>","<3d_geom>","</3d_geom>",
       "<frame_start>","<frame_end>","<view_start>","<view_end>","<pose_start>","<pose_end>",
       "<geometry_start>","<geometry_end>","<think>","</think>","<corrupt_audio>","<malformed_3d>","<corrupt_image>"]
    S += [f"<timestamp_{i}>" for i in range(256)]
    S += [f"<|audio_{i}|>" for i in range(4352)]   # MiMo RVQ tokens
    return S

class AetherPhase1(nn.Module):
    def __init__(self):
        super().__init__()
        self.base, self.tok = load_base()
        old_vocab = len(self.tok)
        self.tok.add_special_tokens({"additional_special_tokens": build_special()})
        self.new_vocab = len(self.tok)
        self.base.resize_token_embeddings(self.new_vocab)
        self.new_row_lo = old_vocab                       # [lo, new_vocab) = trainable rows
        self.vision = load_internvit()
        self.audio  = load_mimo_audio()
        self.geom   = load_trellis_slat()
        self.visual_proj = Projector(3200, TXT)
        self.audio_proj  = Projector(1280, TXT)
        self.mv_proj     = Projector(3200, TXT)           # 3D appearance (multi-view via vision)
        self.geom_proj   = Projector(8, TXT)              # 3D geometry (SLAT latent=8)
        self.cam_pose    = nn.Parameter(torch.zeros(1,1,TXT))
    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(model):
    for p in model.parameters(): p.requires_grad_(False)
    for p in model.trainable(): p.requires_grad_(True)
    emb=model.base.get_input_embeddings().weight; head=model.base.get_output_embeddings().weight
    lo=model.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))

if __name__=="__main__":
    torch.set_grad_enabled(True)
    m=AetherPhase1().to(dev)
    apply_freeze(m)
    tot=sum(p.numel() for p in m.parameters())
    tr =sum(p.numel() for p in m.parameters() if p.requires_grad)
    print(f"ASSEMBLED. total params {tot/1e9:.2f}B | trainable {tr/1e6:.1f}M ({100*tr/tot:.3f}%)")
    print(f"vocab {m.new_vocab} (new rows {m.new_row_lo}..{m.new_vocab})")
    print("VRAM alloc GB:", round(torch.cuda.memory_allocated()/1e9,1))
    print("ASSEMBLY OK")