#!/usr/bin/env python """ Aether Phase-1 surgery: graft 4 modality pathways onto the Qwen3-VL-8B base. Runs on the rental MI300 (aefinal). Operates on a DUPLICATE of the pristine base. Implements the reviewed architecture: 5 isolated projectors, interleaved token blocks, vocab expansion for MiMo RVQ audio tokens, freeze scheme (backbone + encoders frozen; projectors + NEW audio embed rows trainable). STATUS: architecture scaffold — fill encoder-load specifics (InternViT / MiMo-audio / TRELLIS-SLAT) against their actual configs on the rental, then run. """ import torch, torch.nn as nn from transformers import AutoModelForCausalLM, AutoTokenizer # --- REQUIRED shim (verified 2026-09-27): transformers 5.17 expects all_tied_weights_keys; # InternViT/older remote code lacks it. Add before loading any trust_remote_code encoder. --- from transformers.modeling_utils import PreTrainedModel as _PTM if not hasattr(_PTM, "all_tied_weights_keys"): _PTM.all_tied_weights_keys = property(lambda self: (getattr(self,"_tied_weights_keys",{}) or {}) if not isinstance(getattr(self,"_tied_weights_keys",{}),(list,tuple)) else {k:k for k in getattr(self,"_tied_weights_keys",[])}) BASE = "/mnt/thefleet/models/aether-base-pure-BACKUP" # DUPLICATE of pristine; never mutate HF/original INTERNVIT = "OpenGVLab/InternViT-6B-448px-V2_5" # MIT — verify LICENSE file before graft MIMO_AUDIO = "XiaomiMiMo/MiMo-Audio-7B-Base" # MIT — audio encoder + RVQ tokenizer TRELLIS = "JeffreyXiang/TRELLIS-image-large" # MIT — SLAT geometry encoder (verify LICENSE) # ---- Special-token contract (shared with the data schema, artifact B) ---- # Full structural set (Derek 2026-09-27): modality blocks + frame/view/pose/geometry markers # + reasoning + reject tokens. All added in the vocab expansion (cheap; embeddings resized anyway). SPECIAL = { "img": ("", ""), # 2D vision (InternViT) tokens "audio": (""), # audio encoder tokens (understanding-in) "3d_app": ("<3d_app>", ""), # 3D appearance (multi-view->InternViT) "3d_geom": ("<3d_geom>", ""), # 3D geometry (TRELLIS-SLAT) — ISOLATED from 3d_app # #1 temporal video + 3D structural enclosures "frame": ("", ""), "view": ("", ""), "pose": ("", ""), "geom_enc":("", ""), # #2 reasoning-block infrastructure (never masked in generation) "think": ("", ""), } TIMESTAMP_TOKENS = [f"" for i in range(256)] # video timeline bins REJECT_TOKENS = ["", "", ""] # #3 sensory-reject # NEVER-MASK set (must survive generation): + all block/reject markers. # POSITIONS (#4): modality-routed — text+2D-vision keep Qwen M-RoPE; audio=scaled-1D; # 3D-geometry=3D-Spatial-RoPE seeded from camera-pose (X,Y,Z), NOT linear text index. # MiMo RVQ audio *discrete* tokens (for audio OUTPUT): expand vocab by the codebook size. N_AUDIO_RVQ = 4352 # VERIFIED: MiMo-Audio-Tokenizer = 20 quantizers, codebooks [1024,1024,128*18] -> sum=4352 # (per-level distinct tokens; confirm flattened-vs-delay scheme from MiMo audio-token format on rental) class Projector(nn.Module): """Fresh 2-layer MLP mapping an encoder's hidden dim -> Qwen text hidden dim. GELU.""" def __init__(self, in_dim, out_dim, hidden=None): super().__init__() hidden = hidden or out_dim self.net = nn.Sequential(nn.Linear(in_dim, hidden), nn.GELU(), nn.Linear(hidden, out_dim)) def forward(self, x): return self.net(x) class AetherPhase1(nn.Module): """ Backbone = Qwen3 LM (frozen). Four encoders (frozen) -> five NEW trainable pieces: visual_proj, audio_proj (+ new audio embed rows), 3d_fusion_proj (+ cam-pose), geom_proj. Modality features are injected as DISTINCT interleaved token blocks (never pre-merged). """ def __init__(self, base, txt_hidden): super().__init__() self.lm = base # Qwen3-VL-8B text backbone (frozen) # --- 2D vision: swap ViT -> InternViT-6B, fresh projector (old proj discarded) --- self.vision_encoder = None # load InternViT-6B; freeze self.visual_proj = Projector(3200, txt_hidden) # InternViT-6B hidden ~3200 -> txt (verify dim) # --- audio: MiMo encoder + fresh projector; new RVQ output embeds handled on lm_head/embed --- self.audio_encoder = None # MiMo-Audio encoder (frozen). Input = PACKED (total_frames, n_mels=128); # call encoder.encode(packed, lens, use_quantizer=False) -> (B, frames/4, 1280) self.audio_proj = Projector(1280, txt_hidden) # VERIFIED on MI300: encoder d_model=1280 -> txt(4096) # --- 3D appearance: reuse vision_encoder on multi-view + camera-pose embed + fusion proj --- self.cam_pose_embed = nn.Parameter(torch.zeros(1, 1, txt_hidden)) # learned camera-pose enc self.threed_fusion_proj = Projector(3200, txt_hidden) # multi-view InternViT tokens -> txt # --- 3D geometry: TRELLIS-SLAT encoder + ISOLATED geometry projector --- self.geom_encoder = None # load TRELLIS SLAT encoder; freeze self.geom_proj = Projector(8, txt_hidden) # VERIFIED on ROCm: TRELLIS-SLAT latent=8 -> txt(4096) def trainable_parameters(self): """Phase-1 trainable set: 5 projectors + cam-pose + NEW audio embed rows (below).""" mods = [self.visual_proj, self.audio_proj, self.threed_fusion_proj, self.geom_proj] for m in mods: for p in m.parameters(): yield p yield self.cam_pose_embed def expand_audio_vocab(model, tokenizer, n_rvq): """Add MiMo RVQ audio tokens; resize embed_tokens + lm_head. Return the NEW row index range so those rows can be kept UNFROZEN while the rest of embed/lm_head stays frozen.""" old = len(tokenizer) audio_toks = [f"<|audio_{i}|>" for i in range(n_rvq)] + [t for pair in SPECIAL.values() for t in pair] tokenizer.add_special_tokens({"additional_special_tokens": audio_toks}) model.resize_token_embeddings(len(tokenizer)) return old, len(tokenizer) # [old, new) = newly-initialized rows -> keep trainable def apply_freeze(model, new_row_lo, new_row_hi): """Freeze EVERYTHING, then re-enable: the 5 projectors/cam-pose + only the NEW embed/lm_head rows.""" for p in model.parameters(): p.requires_grad_(False) # projectors + cam-pose for p in model.trainable_parameters(): p.requires_grad_(True) # NEW audio-token rows in embed_tokens + lm_head must learn their base coordinates (review pt 2) emb = model.lm.get_input_embeddings().weight head = model.lm.get_output_embeddings().weight # register grad hooks that zero-out gradients for all-but-new rows (partial-unfreeze) def mask_grad(g): g[:new_row_lo] = 0; return g emb.requires_grad_(True); emb.register_hook(mask_grad) head.requires_grad_(True); head.register_hook(mask_grad) # NOTE on routing: at forward time, replace each /