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