File size: 13,511 Bytes
311660a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 | #!/usr/bin/env python
"""
Aether Phase-1 FULL MULTIMODAL FORWARD WIRING + verification.
Builds on the VERIFIED assembly (assemble.py). Extends the vision-only forward
(proven: logits (1,1031,151671)) to ALL FOUR modality paths, and adds #4 the
factorized / 3D-spatial RoPE via Qwen3-VL's native 3-channel M-RoPE position_ids
(temporal, height, width) -> no attention-kernel surgery, fully portable.
Modality routing (each spliced as a DISTINCT block into inputs_embeds):
<img> : InternViT -> visual_proj (3200->4096) pos = 2D grid (t const, h, w)
<audio> : MiMo enc -> audio_proj (1280->4096) pos = scaled-1D time (t=i, h=w=0)
<3d_app> : InternViT -> mv_proj (3200->4096) pos = per-view 2D grid (t=view, h, w)
<3d_geom> : TRELLIS-SLAT-> geom_proj (8->4096)+cam pos = 3D-SPATIAL voxel (X,Y,Z)
text : embed_tokens pos = sequential (t=h=w=idx)
Loss target (later): AR-CE masked to TEXT-ANSWER tokens. This file only proves a
full forward produces finite logits with every path active + M-RoPE positions set.
Run on rental MI300 (aefinal). Encoder 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
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
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
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)]
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
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)
self.geom_proj = Projector(8, TXT)
self.cam_pose = nn.Parameter(torch.zeros(1,1,TXT))
# resolve placeholder token ids once
self.id = {k:self.tok.convert_tokens_to_ids(k) for k in
["<img>","<audio>","<3d_app>","<3d_geom>"]}
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
# ---------------- encoders -> text-space token blocks ----------------
@torch.no_grad()
def enc_vision(self, pixel_values):
"""(N,3,448,448) -> (N*1025, 3200) InternViT last_hidden_state."""
out = self.vision(pixel_values=pixel_values.to(dev,torch.bfloat16))
h = out.last_hidden_state if hasattr(out,"last_hidden_state") else out[0]
return h.reshape(-1, h.shape[-1]) # (N*T, 3200)
@torch.no_grad()
def enc_audio(self, mel_packed, lens):
"""PACKED mel (total_frames, n_mels=128) + lens -> (sum_frames/4, 1280)."""
z = self.audio.encode(mel_packed.to(dev,torch.bfloat16), lens.to(dev), use_quantizer=False)
if isinstance(z,(tuple,list)): z=z[0]
return z.reshape(-1, z.shape[-1]) # (Ta, 1280)
@torch.no_grad()
def enc_geom(self, slat_coords, slat_feats):
"""Sparse voxel coords (M,4 [b,z,y,x]) + feats (M,1024) -> SLAT latent (M,8)."""
from trellis.modules.sparse import SparseTensor
gdt = next(self.geom.parameters()).dtype
st = SparseTensor(feats=slat_feats.to(dev,gdt), coords=slat_coords.to(dev).int())
z = self.geom(st)
feats = z.feats if hasattr(z,"feats") else z
coords = z.coords if hasattr(z,"coords") else slat_coords.to(dev)
return feats.reshape(-1, feats.shape[-1]), coords # (M,8),(M,4)
def _cast(self,h,mod): return h.to(next(mod.parameters()).dtype)
def proj_vision(self,h): return self.visual_proj(self._cast(h,self.visual_proj))
def proj_mv(self,h): return self.mv_proj(self._cast(h,self.mv_proj))
def proj_audio(self,h): return self.audio_proj(self._cast(h,self.audio_proj))
def proj_geom(self,h): return self.geom_proj(self._cast(h,self.geom_proj)) + self.cam_pose.squeeze(0)
# ---------------- #4 factorized / 3D-spatial M-RoPE positions ----------------
def build_mrope(token_types, geom_xyz=None, view_of=None):
"""
Build Qwen3-VL M-RoPE position_ids of shape (3, 1, L): channels (temporal,H,W).
token_types: list len L, each in {"text","img","audio","3d_app","3d_geom"}.
Contiguous same-modality runs form ONE block; text advances all 3 channels by 1.
text : t=h=w = running max +1 (sequential, isotropic)
img : t=const(block start), (h,w)=row/col over ceil(sqrt(n)) grid
audio : t=start+i (scaled-1D time), h=w=start -> pure temporal axis
3d_app : t=view index, (h,w)=grid within the view (needs view_of per token)
3d_geom : (t,h,w) = quantized voxel (Z,Y,X) from geom_xyz -> 3D-SPATIAL RoPE
Returns LongTensor (3,1,L). Guarantees causal monotonicity of the block start.
"""
import math
L=len(token_types); pos=torch.zeros(3,1,L,dtype=torch.long)
cur=0; i=0
while i<L:
tt=token_types[i]; j=i
while j<L and token_types[j]==tt: j+=1
n=j-i; start=cur
if tt=="text":
for k in range(n):
pos[:,0,i+k]=start+k
cur=start+n
elif tt in("img","3d_app"):
g=max(1,math.ceil(math.sqrt(n)))
for k in range(n):
r,c=divmod(k,g)
t = start + (view_of[i+k] if (tt=="3d_app" and view_of) else 0)
pos[0,0,i+k]=t; pos[1,0,i+k]=start+r; pos[2,0,i+k]=start+c
cur=start+max(g, (view_of[j-1]+1 if (tt=="3d_app" and view_of) else 1))
elif tt=="audio":
for k in range(n):
pos[0,0,i+k]=start+k; pos[1,0,i+k]=start; pos[2,0,i+k]=start
cur=start+n
elif tt=="3d_geom":
xyz=geom_xyz # (n,3) already quantized ints, small range
mx=int(xyz.max().item()) if xyz.numel() else 0
for k in range(n):
pos[0,0,i+k]=start+int(xyz[k,0]); pos[1,0,i+k]=start+int(xyz[k,1]); pos[2,0,i+k]=start+int(xyz[k,2])
cur=start+mx+1
if os.environ.get("MROPE_DEBUG"):
print(f" [mrope] block tt={tt:8s} n={n:5d} start={start:5d} -> cur={cur:5d}")
i=j
return pos
# ---------------- full multimodal forward ----------------
def forward_multimodal(m, ids, blocks):
"""
ids: (1,L) LongTensor with placeholder ids at each modality slot.
blocks: dict modality-> list of (mask_indices_tensor, embed_tensor(n,TXT)).
Returns logits, position_ids.
"""
emb = m.base.get_input_embeddings()(ids.to(dev)) # (1,L,TXT)
# id -> SHORT modality name expected by build_mrope ("<img>"->"img", ...)
short={"<img>":"img","<audio>":"audio","<3d_app>":"3d_app","<3d_geom>":"3d_geom"}
inv={m.id[k]:short[k] for k in short}
types=[inv.get(t,"text") for t in ids[0].tolist()]
for modality, items in blocks.items():
if modality.startswith("_"): continue # helper payloads (e.g. _geom_xyz)
for idx, e in items:
emb[0, idx] = e.to(emb.dtype)
# positions
geom_xyz = blocks.get("_geom_xyz")
pos = build_mrope(types, geom_xyz=geom_xyz).to(dev)
out = m.base(inputs_embeds=emb, position_ids=pos, use_cache=False)
return out.logits, pos
if __name__=="__main__":
torch.set_grad_enabled(False)
m=AetherPhase1().to(dev); m.eval()
# new modules default to fp32; match the bf16 base/encoders for the forward
for mod in (m.visual_proj, m.audio_proj, m.mv_proj, m.geom_proj):
mod.to(torch.bfloat16)
m.cam_pose.data = m.cam_pose.data.to(torch.bfloat16)
print("assembled. building a full 4-modality sequence...")
T=m.tok
# --- synthetic per-modality encoder inputs (shapes = real; values random) ---
px = torch.randn(1,3,448,448) # 1 image
mvx = torch.randn(4,3,448,448) # 4-view 3D appearance
mel = torch.randn(400,128); alen=torch.tensor([400]) # 0.4k packed mel frames
# sparse geom: 300 active voxels in a 64^3 grid
M=300; vz=torch.randint(0,64,(M,3)); b=torch.zeros(M,1)
slat_coords=torch.cat([b,vz.flip(-1).float()],1) # [b,z,y,x]
slat_feats =torch.randn(M,1024)
with torch.no_grad():
ve = m.proj_vision(m.enc_vision(px)) # (1025,TXT)
mv = m.proj_mv(m.enc_vision(mvx)) # (4*1025,TXT)
au = m.proj_audio(m.enc_audio(mel,alen)) # (~100,TXT)
gf, gc = m.enc_geom(slat_coords, slat_feats)
ge = m.proj_geom(gf) # (M',TXT)
nV,nMV,nA,nG = ve.shape[0], mv.shape[0], au.shape[0], ge.shape[0]
print(f"encoded tokens -> img {nV} | mv(3d_app) {nMV} | audio {nA} | geom {nG}")
# --- lay out one sequence: text <img>.. text <audio>.. text <3d_app>.. text <3d_geom>.. text ---
txt = lambda s: T(s, add_special_tokens=False).input_ids
seq = txt("Describe: ") + [m.id["<img>"]]*nV + txt(" and the sound ") + [m.id["<audio>"]]*nA
seq += txt(" plus the asset ") + [m.id["<3d_app>"]]*nMV + txt(" whose shape ") + [m.id["<3d_geom>"]]*nG
seq += txt(" — answer:")
ids = torch.tensor(seq)[None]
# index masks for each modality (in order they appear)
def where(tok): return (ids[0]==m.id[tok]).nonzero(as_tuple=True)[0]
blocks = {
"img": [(where("<img>"), ve)],
"audio": [(where("<audio>"), au)],
"3d_app": [(where("<3d_app>"),mv)],
"3d_geom": [(where("<3d_geom>"),ge)],
}
# geom xyz for RoPE = quantized voxel coords of the surviving latents, small-binned
gcz = gc[:, [1,2,3]].float() # z,y,x
gcz = (gcz / gcz.max().clamp(min=1) * 15).long() # bin to 0..15
blocks["_geom_xyz"] = gcz
logits, pos = forward_multimodal(m, ids, blocks)
print(f"seq len {ids.shape[1]} | spliced img{nV}+audio{nA}+mv{nMV}+geom{nG}")
print(f"position_ids (3,1,L) max-per-channel: {pos[:,0].max(1).values.tolist()}")
print(f"logits {tuple(logits.shape)} | finite={bool(torch.isfinite(logits).all())} "
f"| mean {logits.float().mean().item():.3f}")
print("FULL MULTIMODAL FORWARD OK" if torch.isfinite(logits).all() else "!! NON-FINITE LOGITS")
|