Download scripts/assemble.py from SupremeD/aether-phase1-build: direct link, hf CLI and curl.
- Browser
- Download file 5.67 kB
-
https://huggingface.co/SupremeD/aether-phase1-build/resolve/main/scripts/assemble.py
- Command line
-
hf download hf://SupremeD/aether-phase1-build/scripts/assemble.py
-
curl -L -o assemble.py https://huggingface.co/SupremeD/aether-phase1-build/resolve/main/scripts/assemble.py
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") | |