aether-phase1 / scripts /forward.py
SupremeD's picture
Phase-1 assembled+wired: 4-modality forward + factorized 3D-spatial M-RoPE verified
311660a verified
Raw History Blame Contribute Delete
13.5 kB
#!/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")