#!/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=["","","","<3d_app>","","<3d_geom>","", "","","","","","", "","","","","","",""] S += [f"" 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")