modelbeast/docs/FABLE_STARTUP_PACKET.md

22 KiB
Raw Blame History

FABLE STARTUP PACKET — TRELLIS.2 + Hunyuan MLX build

Read this once, top to bottom, before touching anything. It consolidates five prep-only recon passes (all inspection, no GPU job launched, no installs, no production code touched). Every path here is on m3ultra (m3ultra@100.89.131.57) unless noted. Where a recon was uncertain, it says (UNVERIFIED) — do not treat those as facts.

Two baselines you will see quoted, do not conflate them:

  • 124.8 s = our current torch-MPS production path on the anatomy bench (the number to beat).
  • 1004.6 s / 3.09M tris / 17.82 GB = PR#175's reported MPS E2E on an M4 Max 36 GB (different hardware, different mesh size, and NOT the MLX path). Not comparable to 124.8 s.

1 · STATE ON DISK — what prep already put on the box

Everything below already exists. Do not redo it.

Staging dir: ~/Documents/trellis2-mlx-staging/

Path What it is Status
trellis2-apple/ Full clone of pedronaugusto/trellis2-apple, HEAD 6055b86 (2026-04-22), full 17-commit history deepened, MIT. 39 MB. Nothing installed. read-only clone
prep_local_draft.py 9638 B preprocessing draft (§4). py_compile OK under rmbg venv. Not wired in. draft, do-not-run-yet-in-prod
convert_weights.py Draft torch→MLX weight converter (§5). Loads ~14 GB — do not run now. draft
parity_dump_hook.py Draft monkeypatch fixture dumper (§5). Loads models+GPU — do not run now. draft
inspect_safetensors.py Stdlib safetensors header parser (used to produce the §5 layout table). utility, safe
__pycache__/ Compile artifacts from the py_compile checks. ignore

Venvs & operators already on the box (do not reinstall)

  • rmbg venv ~/Documents/MODELBEAST/venvs/rmbg/ (uv, CPython 3.11). Has torch 2.13.0, torchvision 0.28.0, transformers 5.13.1, timm, kornia, einops, safetensors, hf_hub, Pillow, numpy, scipy, scikit-image. No cv2 / rembg / birefnet-pkg / upscaler. Installed by ~/Documents/MODELBEAST/scripts/install_rmbg.sh.
  • bg_remove_local operator ~/Documents/MODELBEAST/server/operators/bg_remove_local/run.py — runs RMBG-2.0, pins venvs/rmbg/bin/python.

Weights already in HF cache

  • microsoft/TRELLIS.2-4B~/.cache/huggingface/hub/models--microsoft--TRELLIS.2-4B/snapshots/af44b45.../ckpts/ (7 of 8 pipeline modules).
  • sparse_structure_decoder (8th module) comes from microsoft/TRELLIS-image-large snapshot 25e0d31.../ckpts/ss_dec_conv3d_16l8_fp16.safetensors.
  • briaai/RMBG-2.0 — 844 MB, present.
  • Hunyuan MLX converted weights — ~/.cache/huggingface/hub/models--dgrauet--hunyuan3d-2.1-mlx/.
  • NOT present: ZhengPeng7/BiRefNet (Apache/MIT, ~900 MB) — must be fetched once for the commercial preprocessing path.

2 · TRELLIS.2 MLX — ADOPT-vs-BUILD DECISION

VERDICT: ADOPT the Jourloy PR#175 mlx_backend/ tree as the MLX vendor base. Do not build from scratch, and do not wholesale-swap the 124.8 s production path — gate any promotion on a measured A0 run.

Why PR#175 and not the already-cloned trellis2-apple: the two share the identical 16-file / ~3,114-LOC mlx_backend/ (same author — pedronaugusto co-authors PR#175). PR#175 (Jourloy/TRELLIS.2, opened 2026-07-17, +7,419/258 across 62 files) is a strict superset: same MLX modules plus the correctness oracle we lack (tests/test_mlx_parity.py, 4 unit parities at rtol 1e-63e-5), a backend probe/fallback resolver (trellis2/backends.py), a full CLI (scripts/generate_asset.py), macOS setup/probe/weight tooling, and the --backend mlx-experimental flag so optimized modules drop in behind a flag while MPS keeps E2E green. So PR#175 dominates trellis2-apple on every axis. Keep vendor/trellis-mac (shivampkumar, no MLX) only as the torch-MPS E2E oracle.

Why NOT a blind swap (both recons agree): the stock MLX path will very likely regress as-is, for two structural reasons:

  • Un-fused sparse conv. mlx_backend/sparse_conv.py is the naive gather→matmul→sum, 512 MB-chunked — it does not use the flex_gemm Metal fast path, which our recon measured at 1.72× faster on this shape. This is the ~10× crux, reimplemented slow in MLX.
  • CPU host bounce. mlx_backend/adapters.py does torch→numpy→mx→numpy→torch on every model call — per sampler step (default ~12 steps × 3 stages × 2 for CFG) and at every sparse-conv boundary. The upstream torch pipeline still owns sampling orchestration + mesh extraction; MLX models are injected as adapters.
  • Zero benchmarks in either repo. No evidence the MLX path beats — or matches — 124.8 s. Do not benchmark the stock --backend mlx-experimental path as representative; it's a correct scaffold, not a fast path.
  • Parity gap in trellis2-apple, closed by PR#175. trellis2-apple has NO MLX parity test; PR#175 adds test_mlx_parity.py. Adopt PR#175 specifically to get that oracle.

What is safe to reuse directly (low risk, no sparse-conv tax): dinov3.py, norm.py (two-pass LayerNorm32 PT-parity), rope.py (interleaved complex-mul parity), transformer_block.py, structure_decoder.py (dense mx.conv_general). Rewrite/replace first: sparse_conv.py (fused Metal gather-GEMM-scatter), the adapters.py boundary (move the Euler/CFG loop inside MLX to kill the bounce), sparse_ops.py (numpy np.where host sync in upsample).

Known smells to verify before trusting a module (UNVERIFIED in recon): (1) structure_decoder.py _remap_structure_decoder_weights is a no-op feeding dict-stored conv weights — untested whether nn.Module.load_weights populates plain-dict children; (2) hand-rolled softplus/quad_lerp in vae_decoders.py:1204 FlexiDualGrid decode; (3) full non-windowed attention (attention.py) may diverge from upstream windowed sparse attn; (4) docstrings claim "custom Metal kernel" but code is plain MLX built-in SDPA — misleading comments, no hand-written kernel anywhere in mlx_backend/. mx.compile is already applied to the DiT block loop and structure-decoder graph (with try/except fallback).

Upstreamability: LOW — plan to fork, not to depend on merge. PR#175 is unmerged, CLA-gated, 0 reviews, only the CLA bot has commented. Issue #74 (non-CUDA backends) has no maintainer response. Vendor it and pin the commit; basing our kernels on #175's layout keeps our work a clean, eventually-upstreamable diff.

First 3 concrete commands

# 1. Get the superset fork (has test_mlx_parity.py; the cloned trellis2-apple does not)
cd ~/Documents/trellis2-mlx-staging && \
  git clone https://github.com/Jourloy/TRELLIS.2 trellis2-jourloy && cd trellis2-jourloy

# 2. Run the correctness ORACLE first — cheap, no heavy GPU E2E, validates
#    LayerNorm32 / apply_rope / SDPA / MlxSparseConv3d vs torch before any gen.
#    (Needs an isolated venv with mlx + the deps; test no-ops off-Mac via importorskip.)
pytest tests/test_mlx_parity.py -v

# 3. The single GPU job — the measured A0 gate. Isolated venv, the 3 pedronaugusto
#    Metal pkgs --no-build-isolation, reuse the existing TRELLIS.2-4B weights.
#    ONE gen, anatomy bench, seed 42, 1024-cascade; compare wall-clock to 124.8 s
#    and quality to the fal renders in ~/Documents/trellis2-bench/.
pip install -r requirements_macos.txt --no-build-isolation   # torch>=2.11, mlx>=0.31, mtlgemm/mtldiffrast/cumesh
# then: scripts/generate_asset.py --backend mlx-experimental --seed 42 --pipeline-type 1024_cascade <anatomy_img>

Decision rule after A0: within ~1.5× of 124.8 s AND parity holds → patch-and-tune, adopt as production path behind the flag. If 35× slower (the likely outcome given un-fused conv + host bounce) → keep torch-MPS in nodes.json for production, lift mlx_backend/ as the module-by-module scaffold, and land the kernel rewrites (fused Metal sparse conv → in-MLX sampler loop → parity fixtures) in that order. Also: fork to Gitea monster/trellis2-mrp-mlx, keep origin upstream for rebases (house pattern).


3 · HUNYUAN TUNE — edit-target list (ready to apply)

All paths under ~/Documents/MODELBEAST/vendor/hunyuan3d-mlx/. Static inspection only (GPU was busy). Ship each edit with a PT-parity test (≤1e-4 rel) + a BENCHMARKS.md row.

Headline: head_dim pad-to-64 does NOT apply — every attention site is already a 64-multiple (64 or 128). The corridorkey padding win is out of scope; do not build pad wrappers. mx.compile count in the whole MLX port = 0 — that is the entire opportunity surface (corridorkey precedent: fixed-shape compile = 1.47× M3 / 1.11× M1, bit-identical, an 8-line change vendor/corridorkey-mlx@4c660df).

# Target file:line Change Expected Risk
1 Paint UNet forward (called 3×/step, ~15 steps) hy3dpaint/hunyuanpaintpbr_mlx/unet/unet_mlx.py:158 (__call__); loop inference.py:254, calls :314,328,339; steps default inference.py:72 Wrap forward in mx.compile (add compile=True ctor flag, self._fwd = mx.compile(self._forward_impl)). No MoE / no .item() → compile-safe. ~1.11.3× on ~112 s paint low; last chunk may have fewer views → 1 retrace (pad final chunk to fixed n_chunk)
2 DiT dense blocks 014 (100 forwards = 50 steps × CFG 2) hy3dshape/.../denoisers/hunyuandit_mlx.py:379 (HunYuanDiTBlock.__call__); driver pipeline_mlx.py:160, loop :143, steps :95 mx.compile the dense-block forward (fixed shape (B∈{1,2},4097,2048)). Leave 6 MoE blocks (1520) eager until #3. ~1.11.3× on ~149 s shape (dense = 15/21 depth) low
3 MoE .item() sync tax (PREREQUISITE for full-DiT compile) mlx_arsenal/moe/moe.py:110if not mx.any(w>0).item(): continue Remove the .item() early-skip; always run all 8 experts, let w*out (w=0) zero them. Deletes ~4800 syncs/gen AND makes the block static → unlocks compiling whole HunYuanDiTPlain.__call__ (hunyuandit_mlx.py:514). 4800 syncs/gen + unblocks #2 to 21/21 blocks mlx_arsenal is a SHARED dep — do NOT edit in place. Vendor a local MoELayer subclass override, or upstream to arsenal. MoEGate routing is already static/compile-clean.
4 Paint DINOv2-giant attn (manual→fused, once/gen ×40 layers) hy3dpaint/hunyuanpaintpbr_mlx/dino_mlx.py:5860 (manual q@kᵀ softmax, fp32) Replace with mx.fast.scaled_dot_product_attention(q,k,v,scale=self.scale) + existing transpose/reshape. head_dim=64 hits fast path; fp32 supported. one-shot; kills N² materialization (fused d64 = 1.51× on M3) low
5 Paint VAE decode (once/gen) hy3dpaint/hunyuanpaintpbr_mlx/vae_mlx.py decode; manual softmax :8184 (head_dim=512) mx.compile the decode (fuses the softmax). Prefer compile over mx.fast swap — head_dim=512 likely misses MLX fused sdpa (kernel ≤256), would fall back anyway. one-shot, minor low
6 (CONDITIONAL) Shape-VAE geo cross-attn hy3dshape/.../autoencoders/model_mlx.py:206 (CrossAttentionBlock); driver _query_sdf_volume :506, chunk :510 mx.compile per-chunk cross-attn (fixed 10000-query × 4096-latent). Only final ragged chunk retraces. smallmoderate Verify decode isn't dominated by the Metal SDF/marching-cubes step before investing

Rasterizer: NO opportunity. hy3dpaint/DifferentiableRenderer/mesh_render_mlx.py is a thin adapter over the mlx_arsenal.rasterize Metal kernel (already fused); Python side is glue + numpy round-trips. Skip for both levers.

DTYPE: keep fp16 as the global default. Port is fp16 everywhere (one deliberate exception: paint DINOv2-giant upcast to fp32 in load_model.py — accumulation-depth fix, not a format preference). bf16 buys zero speed on M3/M4/M5 (equal TFLOPS) and is ~20% slower on M1 — and M1 Ultra (johnking@100.91.239.7) is the second production lane and hunyuan's headline "runs on M1" selling point. Do NOT globally switch to bf16. (This is the opposite of the TRELLIS.2 port, where bf16 range matters.)

Execution order: 1 → 3 → 2(dense, then whole-forward once 3 lands) → 4/5/6 (one-shot cleanups last).


4 · PREPROCESSING — draft, plug-in point, license, TODOs

North star (meshgod/config.py:65-70, STYLE_3D): a single isolated object, centered, fully visible (no crop); plain flat neutral light-gray bg; even soft shadowless studio light; neutral symmetrical A-pose.

Draft: ~/Documents/trellis2-mlx-staging/prep_local_draft.py (9638 B, py_compile OK under rmbg venv, uses only PIL/numpy/torch/torchvision/transformers, not wired in). Signature prep_local(in_path, out_path, commercial=False):

  1. _maybe_upscale — if min(w,h) < 1024, Lanczos up to short-edge 1024, on RGB, BEFORE matte.
  2. _matte — BiRefNet if commercial else RMBG-2.0; real inference mirroring bg_remove_local/run.py exactly (Resize((1024,1024))→ToTensor→Normalize(imagenet); model(inp)[-1].sigmoid(); → L mask); model cached per-process. 37. _compose_center_square — composite onto NEUTRAL_GRAY=(200,200,200) RGB (not alpha) via mask, tight bbox crop from thresholded mask (MASK_THRESH=8), center, square-pad to SUBJECT_FRAC=0.88, resize 1024 Lanczos.

Run (GPU-free-time session): venvs/rmbg/bin/python prep_local_draft.py --in X --out Y [--commercial].

Order lesson (verified live 2026-07-13): upscale FIRST, matte LAST — SeedVR strips alpha (RGBA→RGB), so matting first gets flattened. Draft honors this.

Plug-in point: ~/Documents/meshgod/meshgod/server.py:350 (local lane, server.py:346-350):

paths[i] = stage3_reconstruct.mb_bg_remove(p, pp)   # <- replace with prep_local(p, pp, commercial=...)

prep_local supersedes mb_bg_remove (does the full north-star transform in-process on rmbg venv, not just a farm-round-trip cutout to RMBG-2.0). Secondary optional precede-point: fal lane server.py:359 (run prep_local locally, then skip fal clean_bg). Companion anchors to update: config.py:143 (BG_REMOVE_MODEL), web/index.html:229 label.

LICENSE RULE (hard):

Model HF id License Commercial ship? Local?
RMBG-2.0 briaai/RMBG-2.0 CC-BY-NC 4.0 NO — non-commercial / personal one-offs only Yes (844 MB)
BiRefNet ZhengPeng7/BiRefNet MIT YES — safe for game/product assets No — fetch once
Default for anything game-bound = BiRefNet. RMBG-2.0 stays default only for local/personal one-offs. (Note: RMBG-2.0 is itself a BiRefNet-architecture checkpoint — same code, differs only in weights + license.)

Open TODOs (in draft):

  1. BiRefNet weights fetch — first --commercial run must huggingface-cli download ZhengPeng7/BiRefNet (~900 MB, MIT, ungated). Until then only RMBG-2.0 works.
  2. Better upscaler — PIL Lanczos is the zero-dep default (never invents detail). Swap for local faithful SR: MODELBEAST seedvr2_upscale operator (weights present: models--numz--SeedVR2_comfyUI) or Real-ESRGAN. Must stay BEFORE _matte.
  3. Lighting/exposure normalization (D2) — optional auto-levels + gray-world WB on RGB before matte (fights TRELLIS dark-patch mode). Stubbed so v0 silhouette stays identical to the fal path for A/B.
  4. Tunables to A/B once GPU free: NEUTRAL_GRAY=(200,200,200), SUBJECT_FRAC=0.88, MASK_THRESH=8 (all chosen from north-star text).
  5. Wire-in / flag plumbing — add req.commercial, replace server.py:350, expose in web/index.html, decide local-lane upscale policy (server.py:511 currently forces upscale=False).

5 · WEIGHT CONVERSION + PARITY

Drafts in staging: convert_weights.py, parity_dump_hook.py, inspect_safetensors.py. Both loaders touch ~14 GB / GPU — run in a later session, not now.

Conversion pattern (from Hunyuan's convert_realesrgan.py:34-58)

torch.load(weights_only=True)keys pass through unchanged (module-tree names already match MLX list indexing, no remap dict) → the only tensor mutation is conv transpose (PyTorch (O,I,H,W)→MLX channels-last (O,H,W,I), detected by ndim==4; generalizes to 5D Conv3d as (0,2,3,4,1)) → mx.save_safetensors.

TRELLIS.2-4B layout — the transpose decision (resolved from source, not guessed)

There are two Conv3d conventions on disk; the correct action is opposite for each:

pipeline key file (ckpts/) dtype conv action
sparse_structure_flow_model ss_flow_img_dit_1_3B_64_bf16 (640 t, 2.58 GB) BF16 none (pure Linear DiT)
sparse_structure_decoder (from TRELLIS-image-large) ss_dec_conv3d_16l8_fp16 (74 t) F32/F16 20 dense Conv3d → TRANSPOSE (0,2,3,4,1)
shape_slat_flow_model_512 / _1024 slat_flow_img2shape_dit_1_3B_{512,1024}_bf16 BF16 none
shape_slat_decoder shape_dec_next_dc_f16c32_fp16 (292 t, 948 MB) F16 40 sparse Conv3d → NO transpose (already channels-last)
tex_slat_flow_model_512 / _1024 slat_flow_imgshape2tex_dit_1_3B_{512,1024}_bf16 BF16 none
tex_slat_decoder tex_dec_next_dc_f16c32_fp16 (284 t, 948 MB) F16 40 sparse Conv3d → NO transpose
  • ss_dec_conv3d = dense nn.Conv3d, on-disk (Co,Ci,Kd,Kh,Kw) → needs the 5D transpose. All 20 5-D tensors.
  • *_next_dc = o-voxel SparseConv3d, on-disk already (Co,Kd,Kh,Kw,Ci)load verbatim, NO transpose. Proven by source: conv_flex_gemm.py:33-34 and conv_none.py:40-41 both permute(0,2,3,4,1) at construction, and the checkpoint was saved from that path. MLX forward must reshape (Co,Kd,Kh,Kw,Ci)→(Kvol,Ci,Co) for the per-offset gather-GEMM (contract at conv_none.py:106).
  • All 5 DiTs: pure nn.Linear (ndim 2) + attention + RMSNorm — zero conv, zero transpose; only a BF16→target-dtype cast (fp16 on M1 per fleet matrix). Block arch (30 blocks, identical across all 5): hidden 1536, 12 heads × head_dim 128, MLP 8192, cross-ctx 1024, adaLN-zero modulation [9216]. head_dim 128 ≥ 64 → fused-SDPA fast path applies, NO padding. Keep fused to_qkv [4608,1536] / to_kv [3072,1024] packed and slice in-kernel.

convert_weights.py encodes this per-file policy (dense-transpose for ss_dec_conv3d, keep for *_next_dc, passthrough for DiTs) with a kernel-dims sanity-assert. Uses safetensors.torch.load_file + mx.save_safetensors. Publish converted safetensors to Gitea/HF so re-runs skip conversion.

Parity fixture plan

Stage boundaries all live in run() of trellis2/pipelines/trellis2_image_to_3d.py:488-596; each fixture = the return value of a sub-method → clean non-invasive monkeypatch, zero vendor edits. parity_dump_hook.py monkeypatches the 7 sub-methods, torch.saves each return to $TRELLIS2_DUMP_DIR/<name>.pt, then runpy-delegates to vendor generate.py. Unset TRELLIS2_DUMP_DIR → patches are no-ops.

# Fixture Return / hook (file:line in trellis2_image_to_3d.py) Shape/notes
A/A Image-cond embed 512 / 1024 get_cond return :177 (called :539/:540; disambiguate by resolution arg → cond_512.pt/cond_1024.pt) [1,N,1024] DINOv3 ViT-L/16 tokens; neg = zeros_like
B Sparse-structure coords sample_sparse_structure return :235 (call :542) int [M,4] (b,x,y,z); subs: z_s [1,C,64³] :219, occupancy :227
C Shape-SLat latent sample_shape_slat_cascade :364 (:567) / non-cascade :275 SparseTensor feats [K,32]+coords [K,4], un-normalized
D Tex-SLat latent sample_tex_slat :432 (:574) SparseTensor feats [K,32], un-normalized, sampled w/ concat_cond=shape_slat
E1 Decoded shape mesh + subs decode_shape_slat :389 (:470) List[Mesh] + List[SparseTensor]
E2 Decoded tex voxels decode_tex_slat :453 (:471) SparseTensor feats [K,6] PBR (base_color 0:3, metallic 3:4, roughness 4:5, alpha 5:6)
E3 Final mesh decode_latent :486 (:592) MeshWithVoxel: .vertices,.faces,.coords,.attrs,.voxel_size,.origin

Invocation (later, GPU): TRELLIS2_DUMP_DIR=~/Documents/trellis2-bench/parity_fixtures python parity_dump_hook.py <anatomy_img> --pipeline-type 1024_cascade --seed 42 --texture-size 2048. Fix seed=42 (torch.manual_seed :538). Parity target ≤1e-4 rel. Order cheap→expensive: A/A → B → C (the FlexGEMM sparse-conv crux) → D → E.

GOTCHA: C/D fixtures are captured after the std/mean de-normalization inside each sampler — the MLX port must apply the same shape_slat_normalization/tex_slat_normalization (32-dim mean/std, verbatim in pipeline.json) at the identical point or they diverge by a constant affine. sample_tex_slat also re-normalizes shape_slat back down (:407-409) before concatenating as concat_cond — replicate that round-trip.


  1. Read this packet + skim the 3 staging drafts (prep_local_draft.py, convert_weights.py, parity_dump_hook.py) so you know what already exists. Do NOT re-clone trellis2-apple or re-inventory the rmbg venv.
  2. Clone the Jourloy superset (git clone https://github.com/Jourloy/TRELLIS.2 ~/Documents/trellis2-mlx-staging/trellis2-jourloy) — §2 command 1. This is the vendor base; trellis2-apple lacks the parity oracle.
  3. Run the parity oracle pytest tests/test_mlx_parity.py -v in an isolated venv (no heavy GPU E2E). Confirms LayerNorm32/RoPE/SDPA/SparseConv parity before trusting any module. This is the cheapest signal available and needs no A0.
  4. Fork to Gitea monster/trellis2-mrp-mlx, keep origin upstream, pin the PR#175 head commit (it's unmerged/CLA-gated — do not depend on merge).
  5. Queue the single A0 GPU gate for when the box is free — the ONE gen (anatomy bench, seed 42, 1024-cascade), compare to 124.8 s and the ~/Documents/trellis2-bench/ fal renders. Do not benchmark the stock MLX path as representative before the sparse-conv + adapter rewrites; expect a regression. Apply the §2 decision rule to the result.
  6. In parallel while GPU is busy (all CPU/inspection, no contention): start Hunyuan Target 1 (paint UNet mx.compile, §3) — cleanest, no blockers — with its PT-parity test; and/or dry-run prep_local_draft.py on a sample image under the rmbg venv (§4). Both are safe without the A0 result.
  7. Before any commercial preprocessing run: huggingface-cli download ZhengPeng7/BiRefNet (§4 TODO 1) — otherwise only the non-commercial RMBG-2.0 path works.

Prime directives: parity-first (every kernel/edit ships a ≤1e-4-rel PT-parity test + a BENCHMARKS.md row); do NOT edit shared mlx_arsenal in place (§3 Target 3); default BiRefNet for game-bound assets (§4). The GPU A0 gen is the ONLY job that must wait for a free box.