SupremeD's picture
Upload folder using huggingface_hub
bed4b7b verified
Raw History Blame Contribute Delete
5.67 kB
#!/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")