SupremeD's picture
Upload folder using huggingface_hub
bed4b7b verified
Raw History Blame Contribute Delete
7.78 kB
#!/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": ("<img>", "</img>"), # 2D vision (InternViT) tokens
"audio": ("<audio>", "</audio>"), # audio encoder tokens (understanding-in)
"3d_app": ("<3d_app>", "</3d_app>"), # 3D appearance (multi-view->InternViT)
"3d_geom": ("<3d_geom>", "</3d_geom>"), # 3D geometry (TRELLIS-SLAT) — ISOLATED from 3d_app
# #1 temporal video + 3D structural enclosures
"frame": ("<frame_start>", "<frame_end>"),
"view": ("<view_start>", "<view_end>"),
"pose": ("<pose_start>", "<pose_end>"),
"geom_enc":("<geometry_start>", "<geometry_end>"),
# #2 reasoning-block infrastructure (never masked in generation)
"think": ("<think>", "</think>"),
}
TIMESTAMP_TOKENS = [f"<timestamp_{i}>" for i in range(256)] # video timeline bins
REJECT_TOKENS = ["<corrupt_audio>", "<malformed_3d>", "<corrupt_image>"] # #3 sensory-reject
# NEVER-MASK set (must survive generation): <think></think> + 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 <img>/<audio>/<3d_app>/<3d_geom> block's
# placeholder ids with the corresponding projector outputs, spliced into inputs_embeds as
# DISTINCT contiguous blocks (3d_app and 3d_geom stay separate, per review). Loss = autoregressive
# CE masked to TEXT-ANSWER tokens only -> gradients flow through projectors + new embed rows.
if __name__ == "__main__":
print("surgery scaffold — set encoder dims/N_AUDIO_RVQ from rental configs, then instantiate + save.")