aether-phase1 / scripts /save_package.py
SupremeD's picture
Phase-1 assembled+wired: 4-modality forward + factorized 3D-spatial M-RoPE verified
311660a verified
Raw History Blame Contribute Delete
4.96 kB
#!/usr/bin/env python
"""
Package the Aether Phase-1 assembled+wired model for HF (resumable, portable).
Does NOT re-upload the 16GB base or 12GB InternViT — pins them in a manifest so the
model re-assembles deterministically on ANY box (6900XT / Kaggle / Colab / rental).
Writes /work/aether-phase1-pkg/ : tokenizer + new_modules.safetensors + MANIFEST.json.
Light: builds only the tokenizer + the small new modules (no 32GB encoder load).
"""
import os, json, torch, torch.nn as nn
from transformers import AutoTokenizer
from safetensors.torch import save_file
OUT="/work/aether-phase1-pkg"; os.makedirs(OUT, exist_ok=True)
TXT=4096
torch.manual_seed(0) # deterministic projector init for reproducible resume
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 Projector(nn.Module):
def __init__(self,i,o):
super().__init__(); self.net=nn.Sequential(nn.Linear(i,o),nn.GELU(),nn.Linear(o,o))
# --- tokenizer (base + expanded vocab) ---
tok=AutoTokenizer.from_pretrained("/work/base",trust_remote_code=True)
old=len(tok)
tok.add_special_tokens({"additional_special_tokens":build_special()})
new=len(tok)
tok.save_pretrained(OUT+"/tokenizer")
print(f"tokenizer: {old} -> {new} (+{new-old})")
# --- new trainable modules (fresh, seeded) ---
mods={"visual_proj":Projector(3200,TXT),"audio_proj":Projector(1280,TXT),
"mv_proj":Projector(3200,TXT),"geom_proj":Projector(8,TXT)}
sd={}
for name,m in mods.items():
for k,v in m.state_dict().items(): sd[f"{name}.{k}"]=v.contiguous()
sd["cam_pose"]=torch.zeros(1,1,TXT)
save_file(sd, OUT+"/new_modules.safetensors")
print(f"new_modules: {len(sd)} tensors, {sum(v.numel() for v in sd.values())/1e6:.1f}M params")
# --- manifest: everything needed to re-assemble the exact model anywhere ---
manifest={
"name":"aether-phase1",
"desc":"Qwen3-VL-8B base + InternViT-6B vision + MiMo-Audio + TRELLIS-SLAT 3D, "
"5 projectors, expanded vocab, factorized/3D-spatial M-RoPE. Phase-1 assembled+wired.",
"total_params_B":15.16, "trainable_params_B":1.38, "vram_bf16_GB":32,
"base":{"repo":"SupremeD/leeworld-aether-base-pure","revision":"main",
"arch":"Qwen3-VL-8B","text_hidden":TXT},
"encoders":{
"vision":{"repo":"OpenGVLab/InternViT-6B-448px-V2_5","revision":"main","hidden":3200,
"tokens_per_448img":1025,"license":"MIT"},
"audio":{"repo":"XiaomiMiMo/MiMo-Audio-Tokenizer","code":"XiaomiMiMo/MiMo-Audio-7B-Base",
"revision":"main","hidden":1280,"n_mels":128,"input":"PACKED (total_frames,128)",
"call":"encoder.encode(mel,lens,use_quantizer=False)","frame_downsample":4,"license":"MIT"},
"geom":{"repo":"JeffreyXiang/TRELLIS-image-large","fork":"CalebisGross/TRELLIS-AMD",
"ckpt":"ckpts/slat_enc_swin8_B_64l8_fp16.safetensors","latent":8,"resolution":64,
"in_channels":1024,"attn":"sdpa","sparse_backend":"torchsparse","spconv":"NOT required",
"note":"SLatEncoder attention-only; build torchsparse from source on ROCm","license":"MIT"}},
"projectors":{"visual":[3200,TXT],"audio":[1280,TXT],"mv":[3200,TXT],"geom":[8,TXT],
"type":"Linear-GELU-Linear","cam_pose":[1,1,TXT]},
"vocab":{"old":old,"new":new,"new_row_lo":old,
"special":"21 structural + 256 timestamp + 4352 MiMo-RVQ audio"},
"freeze":"backbone+encoders frozen; trainable = 5 projectors + cam_pose + NEW embed/lm_head rows "
"[new_row_lo:] via grad-mask hook",
"mrope":{"impl":"Qwen3-VL native 3-channel M-RoPE position_ids (temporal,H,W) — no kernel surgery",
"text":"isotropic sequential (t=h=w)","img":"2D grid (t const, h,w)",
"audio":"scaled-1D time (t=i, h=w=start)","3d_app":"per-view 2D grid (t=view)",
"3d_geom":"3D-SPATIAL voxel (X,Y,Z) binned"},
"shims":["transformers.PreTrainedModel.all_tied_weights_keys={} (settable)",
"flash_attn varlen SDPA shim","ATTN_BACKEND=sdpa SPARSE_BACKEND=torchsparse"],
"verified":{"assembly":"15.16B params, 32GB VRAM, ASSEMBLY OK",
"forward":"seq 5542, img1025+audio100+mv4100+geom300 spliced, "
"M-RoPE max [230,230,230], logits (1,5542,156296) finite, FORWARD OK",
"date":"2026-09-27","hardware":"MI300 gfx942 ROCm (rental aefinal)"},
"resume":"clone this repo; pull base+encoders per repos above; run assemble.py then forward.py; "
"load new_modules.safetensors into the projectors+cam_pose; begin sliver alignment.",
}
json.dump(manifest, open(OUT+"/MANIFEST.json","w"), indent=2)
print("MANIFEST.json written")
print("PACKAGE READY:", OUT)