46 KiB
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_localoperator~/Documents/MODELBEAST/server/operators/bg_remove_local/run.py— runs RMBG-2.0, pinsvenvs/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 frommicrosoft/TRELLIS-image-largesnapshot25e0d31.../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-6–3e-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.pyis the naive gather→matmul→sum, 512 MB-chunked — it does not use theflex_gemmMetal 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.pydoestorch→numpy→mx→numpy→torchon 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-experimentalpath as representative; it's a correct scaffold, not a fast path. - Parity gap in
trellis2-apple, closed by PR#175.trellis2-applehas NO MLX parity test; PR#175 addstest_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 3–5× 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.1–1.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 0–14 (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 (15–20) eager until #3. |
~1.1–1.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:110 — if 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:58–60 (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 :81–84 (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. |
small–moderate | 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):
_maybe_upscale— ifmin(w,h) < 1024, Lanczos up to short-edge 1024, on RGB, BEFORE matte._matte— BiRefNet ifcommercialelse RMBG-2.0; real inference mirroringbg_remove_local/run.pyexactly (Resize((1024,1024))→ToTensor→Normalize(imagenet);model(inp)[-1].sigmoid(); →Lmask); model cached per-process. 3–7._compose_center_square— composite ontoNEUTRAL_GRAY=(200,200,200)RGB (not alpha) via mask, tight bbox crop from thresholded mask (MASK_THRESH=8), center, square-pad toSUBJECT_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):
- BiRefNet weights fetch — first
--commercialrun musthuggingface-cli download ZhengPeng7/BiRefNet(~900 MB, MIT, ungated). Until then only RMBG-2.0 works. - Better upscaler — PIL Lanczos is the zero-dep default (never invents detail). Swap for local faithful SR: MODELBEAST
seedvr2_upscaleoperator (weights present:models--numz--SeedVR2_comfyUI) or Real-ESRGAN. Must stay BEFORE_matte. - 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.
- 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). - Wire-in / flag plumbing — add
req.commercial, replaceserver.py:350, expose inweb/index.html, decide local-lane upscale policy (server.py:511currently forcesupscale=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= densenn.Conv3d, on-disk(Co,Ci,Kd,Kh,Kw)→ needs the 5D transpose. All 20 5-D tensors.*_next_dc= o-voxelSparseConv3d, on-disk already(Co,Kd,Kh,Kw,Ci)→ load verbatim, NO transpose. Proven by source:conv_flex_gemm.py:33-34andconv_none.py:40-41bothpermute(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 atconv_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-zeromodulation [9216]. head_dim 128 ≥ 64 → fused-SDPA fast path applies, NO padding. Keep fusedto_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.
6 · RECOMMENDED FIRST HOUR
- 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-clonetrellis2-appleor re-inventory the rmbg venv. - 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-applelacks the parity oracle. - Run the parity oracle
pytest tests/test_mlx_parity.py -vin 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. - Fork to Gitea
monster/trellis2-mrp-mlx, keeporiginupstream, pin the PR#175 head commit (it's unmerged/CLA-gated — do not depend on merge). - 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. - 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-runprep_local_draft.pyon a sample image under the rmbg venv (§4). Both are safe without the A0 result. - 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.
7 · PREP ROUND 2 — VERIFICATION RESULTS (2026-07-19)
Round-2 deep-verify of §1–§6 against live source + on-disk safetensors headers. Still inspection-only: no GPU gen, no weights loaded, no prod code touched (
py_compile/ header-parse /pytest --collect-only). Where §7 says CONFIRMED / CORRECTION / CONTRADICTION, trust §7 over the original section.
7.1 · ADOPT confidence — §2 stance → CONFIRMED (HIGH)
-
Clone DONE & pinnable.
~/Documents/trellis2-mlx-staging/trellis2-jourloy/, HEAD754d403(2026-07-17, Jourloy — matches "PR#175"), branchmainis the PR head, tree clean. Pin754d403when forking. -
Superset CONFIRMED — stronger than §2 claimed. All 4 named artifacts present (
tests/test_mlx_parity.py,trellis2/backends.py,scripts/generate_asset.py,requirements_macos.txt) plus 4 extra test files (test_backend_fallbacks,test_background_preprocess,test_generate_asset,test_model_revisions) and a newtrellis2/model_revisions.py(HF revision pins).mlx_backend= 3172 LOC / 16 files, allpy_compileclean. -
mlx_backendnumerics =trellis2-applebyte-for-byte.diff -rq: 14/16 files identical; the only 2 that differ (dinov3.py,pipeline.py) add HF revision-pinning / offline kwargs only — zero algorithm change. §2's "identical mlx_backend" holds for numerics; jourloy is a strict superset (adds reproducibility plumbing on top of the same math). -
--backend mlx-experimentalcleanly wired (generate_asset.py:333→create_mlx_pipeline,:343) behind a real Metal-completeness probe gate (:170-187→backends.probe_metal_backends()) that refuses a half-built stack instead of silently degrading. MPS stays the E2E path. -
§2's 4 UNVERIFIED smells — now resolved:
# smell round-2 disposition 1 structure_decoder no-op remap + dict-child conv load BENIGN by construction (leaves pre-alloc'd mx.zeros; MLX param tree recurses dict/list;load_weights(strict=True)replaces leaves). Residual = exact key-name match only, not exercised by parity test → run a 2-sec non-GPU load smoke on ss_dec.2 vae_decoders softplus/quad_lerp BENIGN — it's the numerically-stable softplus identity `max(x,0)+log1p(exp(- 3 non-windowed global attention NEEDS-RUNTIME-CHECK / POTENTIALLY REAL — highest-value item. _sdpadoes full global SDPA, no window mask. If upstream sparse flow is windowed → numeric divergence and O(N²) blowup at 1024-cascade. Tiny SDPA parity test validates the primitive, not the windowing semantics. Verify vs upstream sparse-attn config before trusting shape/tex output.4 "custom Metal kernel" docstrings CONFIRMED BENIGN (cosmetic) — grep `metal_kernel -
Bonus (confirms §2): sparse-conv perf crux is REAL —
sparse_conv.pyis naive gather→matmul→sum, 512 MB-chunked with a per-chunkmx.evalhost sync, O(batch·D³) dense LUT, noflex_gemm. Matches the ~10× claim. -
Confidence HIGH, not lowered: nothing contradicted the ADOPT case; several findings strengthened it. Two items still need a runtime signal (smell 3 windowed-attn, smell 1 dict-child load) — both cheap, non-GPU; run before the A0 gate (see 7.6).
7.2 · WEIGHT-CONVERT — transpose policy VERIFIED vs disk; converter has 2 bugs + 1 double-transpose contradiction
All 8 modules parsed header-only (no tensor bytes). Tensor counts match §5 verbatim (640/640/640/640/640/292/284/74). Classifier found ZERO ambiguous 5D tensors — conventions separate cleanly by shape.
| module(s) | dtype on disk | 5D convs | §5 policy | matches disk |
|---|---|---|---|---|
| 5 DiTs | BF16 ×640 ea | 0 (215 Linear + 120 RMSNorm γ + 305 1D) | passthrough + dtype cast | ✅ |
| shape_dec / tex_dec | F16 | 40 sparse ea, (Co,K,K,K,Ci) |
NO transpose | ✅ |
| ss_dec | F32×38 + F16×36 (2 F32 convs = input_layer + out_layer.2) |
20 dense, (Co,Ci,K,K,K) |
transpose (0,2,3,4,1) |
✅ |
Transpose/convention classification in §5 holds on disk with zero contradictions. The
converter's else: raise RuntimeError on any unclassified 5D fails loud, so no silent
transpose-convention corruption is possible. But three items must be fixed before any convert run:
- 🔴 BUG 1 (BLOCKING)
convert_weights.py:59v.numpy()ontorch.bfloat16→TypeError: unsupported ScalarType BFloat16. Crashes on all 5 BF16 DiTs (the bulk of the pipeline); only survives on F16/F32 decoders. Loud, not silent. Fix:v.float().numpy(), or bit-reinterpret viauint16, or drop torch and usemx.load/mx.save_safetensors. - 🟠 BUG 2 (SILENT precision loss) uniform
target_dtypedefault"bf16"(:46,81-83,94) downcasts the decoders under default args: shape_dec/tex_dec F16→bf16 (lose 3 mantissa bits, strictly worse + ~20% slower on M1); ss_dec's 2 F32 convs → bf16 (lose 16 bits on the most sensitive first/last SDF layers). Contradicts §5's own dtype column (decoders = F16). Fix: per-file dtype policy — DiTs→bf16 (fp16 on M1), decoders→preserve on-disk dtype. - ⚠️ CONTRADICTION — ss_dec DOUBLE-TRANSPOSE (silent-corruption risk). §5/converter transpose
ss_dec
(0,2,3,4,1)at load. But the live MLXstructure_decoder.pystores PT-format(Co,Ci,kD,kH,kW)verbatim and transposes at forward-call time (conv3d,structure_decoder.py:14). If the converter also transposes ss_dec, the decoder double-transposes → garbage, silently. Resolve before converting: convert ss_dec to PT-format on disk (NO load transpose) to match this file, OR move the transpose — pick exactly one, not both. - Port dependency (not the converter): the MLX sparse-conv module MUST do the
(Co,K,K,K,Ci)→(Kvol,Ci,Co)reshape at load (conv_none.py:106) or the*_next_dcdecoders emit garbage. Converter correctly leaves this to the module.
7.3 · PARITY HOOKS — all 8 targets match live source, ZERO corrections to the hook
parity_dump_hook.py needs no edits. Source trellis2_image_to_3d.py is frozen at upstream
75fbf01 (2026-06-05) + 1 benign local Mac cuda-guard (:590-591, downstream of every hook →
shifts no fixture line). No post-packet drift. Verified def/return/call (corrections to §5's
fixture table where it drifted):
| fixture | def | return | call-site(s) |
|---|---|---|---|
| get_cond (A/A′) | 164 | 181 (no-neg) / 183-186 (dict) | 539, 540 |
| sample_sparse_structure (B) | 188 | 235 | 542 |
| sample_shape_slat (C, non-cascade) | 237 | 275 | 547, 557 |
| sample_shape_slat_cascade (C′) | 277 | 364 (tuple slat,hr_res; transform takes r[0]) |
567, 579 |
| sample_tex_slat (D) | 391 | 432 | 551, 561, 574, 586 |
| decode_shape_slat (E1) | 366 | 389 | 470 |
| decode_tex_slat (E2) | 434 | 453 | 471 |
| decode_latent (E3) | 456 (@torch.no_grad :455) |
486 | 592 |
- §5 CORRECTION: the fixture table cites get_cond "return :177" — that line is the model
call
cond = self.image_cond_model(image), not the return. Real returns are :181 / :183-186. Cosmetic (hook wraps by method name), but corrected. - Normalization gotcha CONFIRMED:
shape_slat.pt/tex_slat.ptcaptured POST-denorm (:271-273,:360-362,:428-430); tex sampler re-normalizes shape_slat back down at:407-409beforeconcat_cond(:419). Port must replicate both the on-return de-norm and the denorm→renorm round-trip or fixtures diverge by a constant affine (silent ≤1e-4 failure). - UNCERTAIN — flag for exec session: the hook's
__main__runpy-delegates to vendorgenerate.py; that file's existence + CLI contract (--pipeline-type/--seed/--texture-size) was not checked. Verify before the dump run — a missing/renamed delegate means the hook installs but never fires.
7.4 · PREPROCESSING — BiRefNet PRESENT; D2 IMPLEMENTED
- BiRefNet fetched (§1 "NOT present" TODO → DONE). Cache:
/Users/m3ultra/.cache/huggingface/hub/models--ZhengPeng7--BiRefNet/snapshots/e2bf8e4460fc8fa32bba5ea4d94b3233d367b0e4/. Size 424 MB (model.safetensors= 444,473,596 B) — §1/§4's "~900 MB" was high. License MIT (READMElicense: mit) — §1/§4's "Apache-2.0" label was wrong. Ungated, no login. - BiRefNet branch VERIFIED CORRECT (not stubbed): RMBG-2.0 is a BiRefNet checkpoint (identical
auto_map/architectures); forwardmodel(inp)[-1].sigmoid()holds (birefnet.py:2087/2090); deps satisfied by the rmbg venv — no cv2 needed (birefnet.pyhas zero opencv refs). Only metadata was hardened; no code change. - D2 (light normalization) IMPLEMENTED behind
normalize_light=False(default OFF → v0 silhouette stays A/B-identical to fal). New_normalize_light(prep_local_draft.py:84, PIL/numpy only): 1/99-pct auto-levels (flat-channel guarded) + gray-world WB; wired at step 1b (after upscale, before matte — honors "matte LAST"). CLI--normalize-light. Functional smoke passed on a synthetic cast image. - Draft status:
~/Documents/trellis2-mlx-staging/prep_local_draft.pynow 236 lines,py_compileOK + imports clean under rmbg venv. Signatureprep_local(in_path, out_path, commercial=False, normalize_light=False). Staging only — still NOT wired into prod (§4.5 wire-in + tunable A/B still open; need a GPU-free box for real A/B).
7.5 · MLX TEST VENV — oracle is one command from running
- Venv READY:
/Users/m3ultra/Documents/trellis2-mlx-staging/mlx-test-venv(CPython 3.12.13, uv).mlx==0.32.0(+mlx-metal0.32.0),torch==2.13.0(prebuilt wheel, no compile),numpy2.5.1,pytest9.1.1. (safetensors/transformers/einops also installed but the oracle uses none of them — droppable for a leaner venv.) - Collects clean:
--collect-only→ 4 tests collected, 0 import errors, 0 skips (0 skips proves bothimportorskip("mlx.core")andimportorskip("torch")succeeded on the box). Not yet executed — a real pass is still pending (low-risk caveat:np.asarray(mlx_array)under numpy 2.5.1 surfaces only on the real run, not collection). - GPU/Metal? 4 tests run microsecond ops on tiny tensors. MLX's default device is Metal (GPU), so
they touch the GPU — but this is trivial library compute, NOT a generation/inference job, uses
no Metal build packages (mtlgemm/mtldiffrast/cumesh) and no compile. Won't meaningfully contend
for the box; force CPU with
mx.set_default_device(mx.cpu)if ever desired. - EXACT run command (use the venv python — a bare
pytestuses the wrong interpreter, correcting §6/§2):
Expected: 4 passed. Baked tolerances: LayerNormssh -n m3ultra@100.89.131.57 'cd ~/Documents/trellis2-mlx-staging/trellis2-jourloy && \ ../mlx-test-venv/bin/python -m pytest tests/test_mlx_parity.py -v'2e-5, RoPE1e-6, SDPA3e-5, SparseConv2e-5(all tighter than the ≤1e-4 target).
7.6 · UPDATED FIRST MOVES — deltas to §6
- §6 step 2 (clone Jourloy): DONE —
trellis2-jourloy/@754d403, tree clean. Skip the clone; justgit checkout 754d403/ pin it when forking. - §6 step 3 (run parity oracle): venv READY, oracle collects clean. Now a one-command move — run
the exact command in 7.5. Do NOT use the bare
pytest tests/...from §6/§2 (wrong interpreter). - §6 step 7 / §4 TODO 1 (BiRefNet fetch): DONE — skip. Present at the 7.4 path (424 MB, MIT).
- NEW — before any weight-convert run (§5): fix converter Bug 1 (blocking bf16 crash) + Bug 2 (silent decoder downcast), AND resolve the ss_dec double-transpose contradiction (7.2). Else DiT convert crashes, decoders lose precision, or ss_dec silently corrupts.
- NEW — before the A0 GPU gate: run the two cheap non-GPU runtime checks — (a) ss_dec dict-child weight-load smoke, and especially (c) global-vs-windowed attention vs the upstream sparse-attn config (the one smell that can be a true correctness bug, not just slow). Both fit CPU/inspection while the box is busy.
- Unchanged: §6 step 5 (single A0 GPU gate — seed 42, 1024-cascade, compare 124.8 s) is still the only job that must wait for a free box; §6 step 6 (Hunyuan Target 1) still safe in parallel.
8 · CORRECTION — the adopt path needs NO weight conversion (2026-07-19, verified in source)
Round-2 flagged three bugs in convert_weights.py (bf16 crash, decoder downcast, ss_dec double-transpose). All three are moot: the converter is not needed for the adopted Jourloy base, and using it would corrupt. Verified by reading Jourloy mlx_backend/:
- Weights load directly from the original TRELLIS.2-4B safetensors.
pipeline.py:116/141/162→load_safetensors("<path>.safetensors")→model.load_weights(...).load_safetensors(__init__.py:26) is justmx.load(path), keeping native bf16/fp16; an optionaldtype=casts float weights (this is the M1 fp16 path — not a separate converter). - Remap = key-rename ONLY, zero tensor transform.
remap_flow_model_weights/remap_vae_decoder_weightsonly rewrite key strings (nn.Sequential.N.→.layers.N.; MlxSparseLinear.linear.wrap). No transpose, no dtype change to any tensor. - All conv transposing is IN-MODULE at forward.
structure_decoder.py:14transposes the dense ss_dec (Co,Ci,K,K,K)→(Co,K,K,K,Ci) every forward;sparse_conv.py:130w.transpose(1,2,0)for the sparse GEMM. So the modules EXPECT PyTorch-format weights. - Why the converter is a trap: it transposes ss_dec at load; the module transposes again at forward → double-transpose. For conv1/conv2 (Ci==Co==channels) the SHAPE still matches, so
load_weightssucceeds silently with wrong values; for input_layer (Ci=latent≠Co) it would shape-mismatch and error. Mixed silent/loud corruption.
ACTION for Fable: do not run convert_weights.py for the Jourloy path (it's now banner-warned in staging). For M1 fp16, pass dtype=mx.float16 to load_safetensors. §5's transpose table remains correct as a description of the on-disk layout, but the CONVERSION step it implied is unnecessary — Jourloy consumes originals directly. A converter is only relevant if you ever build a from-scratch decoder that does not transpose at forward.
9 · EXECUTION LOG — Fable session 1 (2026-07-19)
Done this session (all committed; BENCHMARKS.md has the rows):
- Parity oracle: 4/4 PASS (2.55s, CPU-only, mlx-test-venv).
- Jourloy venv built via scripts/setup_macos.sh (PYTHON_BIN=uv 3.11; Metal pkgs compiled).
- A0 GATE PASSED — ADOPT CONFIRMED: mlx-experimental beat MPS (216.6s vs 248.8s wall) at 14× bake density. It's a HYBRID (Metal flex_gemm sparse ops + MLX dense); naive MLX sparse-conv never runs. 75.4GB peak.
- Hunyuan mlx-tune branch @ 4b368ee (pushed to partly): T1 UNet mx.compile (env-gated HY3D_MLX_COMPILE, parity 3e-5, 1.04× micro), T3 StaticMoELayer (bit-identical, kills ~4800 syncs, compile-unlocked), T4 fused DINO SDPA (1.41× micro, parity 1e-7). E2E flag-on run in flight vs the 260s defaults baseline.
- Dark-patch culprit narrowed: reproduces on BOTH independent pipelines (trellis-mac MPS + Jourloy MLX) → shared pedronaugusto Metal texel-sampling (mtldiffrast/BVH) or decoded attrs. fal is clean on the same weights → attrs likely fine → Metal sampling = prime suspect. KDTree pure-Python bake control still grinding (2h+, alive, 300% CPU) — its render decides.
NEXT (in order): read KDTree verdict when it lands → Hunyuan e2e A/B number → wire prep_local into MESHGOD behind a flag → Target 2 (whole-DiT compile now that T3 unblocked it) → promote mlx-experimental into MODELBEAST as a trellis2_mlx operator behind a flag once the bake-quality question is settled.
10 · DARK-PATCH VERDICT: Metal texel-sampling convicted (2026-07-19)
Decisive experiment: dumped the decoded tex voxels (E2 fixture,
1,398,769 voxels × 6 PBR attrs) and rendered base_color as a direct
front-projection — trellis2-bench/attrs_front.png. The attrs are
near-fal quality (clean anatomy palette; dark voxels only 5.6%, all
interior). Both local pipelines (trellis-mac MPS and Jourloy MLX) produce
the same mottled dark bake from these clean attrs, and both share the
pedronaugusto mtldiffrast/mtlbvh texel→voxel sampling. ⇒ the bake's
attr-sampling is the sole quality gap. The 3h+ KDTree run is now merely
confirmatory (left running).
Designed fix (next session's opening move)
The attrs live on a REGULAR voxel grid — no BVH/KDTree needed at all:
texel → (xatlas UV → triangle barycentric → world pos) →
grid[(pos-origin)/voxel_size] direct index (+8-neighbor occupancy
search). One vectorized gather over all texels; replaces the suspect
ray-cast sampling entirely; implementable in torch-MPS or MLX in ~a page.
Validate against attrs_front.png colors + the fal render. If it works,
BOTH local paths get fal-class texture quality at zero bake-speed risk.
Parity fixture set: COMPLETE
All 8 fixtures captured at trellis2-bench/parity_fixtures/ (cond_512/
1024, ss_coords, shape_slat, shape_dec, tex_slat, tex_voxels, mesh_final;
seed 42, 1024-cascade). Hook env needed: PYTORCH_ENABLE_MPS_FALLBACK=1
ATTN_BACKEND=sdpa SPARSE_ATTN_BACKEND=sdpa SPARSE_CONV_BACKEND=flex_gemm.
Fork home (added 2026-07-19, post-§10)
The vendored Jourloy/PR#175 tree now lives on our Gitea:
ssh://git@100.71.119.27:222/monster/trellis-2-mrp-mlx.git (note the
hyphens), branch main, pin tag pr175-head-754d403. m3 staging clone has
it as remote partly; origin stays github/Jourloy for rebases. All our
kernel/bake work lands as commits on this fork.
11 · BAKE FIX PROVEN (2026-07-19 evening) — dark patches eliminated
trellis2-bench/render_vbake5_front.png: full anatomy figure, correct
colors, ZERO dark patches — from the same decoder attrs the Metal baker
mangles. Method: cKDTree IDW sampling of the voxel grid at mesh vertices
(k=4, distance² weights, 1.5-voxel radius, single-nearest fallback) →
vertex-colored GLB. Sampling cost: 0.3s for 270K verts / 1.4M voxels.
Two integration gotchas (cost hours; do not rediscover):
- Axis spaces. Voxel grid + raw pipeline mesh = Z-up-long-axis in
[-0.5,0.5]³; GLB-loaded meshes = glTF Y-up. Transform before sampling:
(x,y,z)_gltf → (x,−z,y)_voxel. The "mangled blob" renders were ONLY this (figure viewed down its own axis) — no data corruption anywhere. - glTF vertex colors are LINEAR — do not pre-gamma (double-gamma = washed out). Store linear, let the renderer do the transfer.
Also learned: raw MeshWithVoxel.vertices/faces need to_glb's weld/init
before use as a plain triangle soup; output_3d.obj inherits the raw
convention. And xatlas hangs (2h+) on écorché-topology meshes — the UV
path needs chart budgeting or per-island parallel unwrap.
Production path (next session)
- Option A (fast to ship): full-density vertex-colored GLB as the raw output; MESHGOD's Blender finish farm already decimates + can bake verts→texture (standard Blender bake) — sidesteps xatlas entirely.
- Option B (fal-parity): fast bucket rasterizer (written, in fast_bake_test.py) + xatlas-alternative UV → 2048² texture bake with this sampler. Needs the écorché-safe unwrap.
- Wire as
--baker python-fastin the Jourloy fork + trellis-mac, make it the default over the convicted Metal texel path; port sampler to MLX (trivial — gather + weighted sum) for the pure-MLX lane.
§11 addendum — vertex baker PRODUCTIONIZED (2026-07-19 night)
Wired and validated in BOTH pipelines:
vendor/trellis-mac@ 55fdd7b:TRELLIS2_BAKER=vertexenv mode in generate.py (Metal + kdtree paths untouched; sentinel-gated).- Jourloy fork @ 746e727 (pushed to monster/trellis-2-mrp-mlx): CLI
--baker vertex.
E2E validation (anatomy, seed 42, 1024-cascade, uncapped density):
130.0s total, bake 1s, 232s wall, 20.8GB peak — full 2.79M-face
vertex-colored GLB, ZERO dark patches, best local quality yet
(render_vprod_front.png). vs Metal-bake baseline: same speed envelope,
quality gap closed for albedo. Not yet carried: metallic/roughness maps
(vertex colors are albedo-only) — MR needs the UV path (Option B) or a
COLOR_1 convention. MESHGOD finish-farm verts→texture bake = the
remaining Option-A step for game-budget assets.
12 · RUNG-3 KERNEL DAY — RESULTS (2026-07-20)
Headline: 216.6s → 71.96s e2e (3.0×), peak 75.4GB → 26.5GB (2.8×), parity held (raw mesh Δ0.03%), full gate EXIT=0. All on the fork (monster/trellis-2-mrp-mlx): fa972aa + 7860148 + 0db816b.
What actually moved the needle (in order of impact):
- Vertex baker properly wired into generate_asset (scheduler
fallthrough fixed —
--baker vertexwas silently running the 5h pure-python kdtree path): "other" bucket 69.4s → 8.3s, and the 75GB peak turned out to be the METAL BAKE's, not the conv's. - Pure-MLX sampler loop (mlx_samplers.py): torch CPU loop + per- forward host bounces → one conversion per stage, batched dense CFG, once-per-stage concat_cond. Sampling 56.8→54.4s (~4% — forwards are compute-bound) but it is the correctness-proven, compile-ready architecture, and removes ~72 host round-trips/gen.
- Validator taught the vertex-colored asset class (COLOR_0 without UVs = legit; normals materialized at export).
Honest negatives (documented so nobody re-chases them):
- mx.compile on the step forward: ~0 on M3 Ultra (GEMM-bound; matches Hunyuan T1/T2 findings).
- Fused Metal sparse-conv kernel: built, parity 8e-4, but stock chunking already bounds memory at decoder scale and beats the scalar kernel on speed. Kept as TRELLIS2_METAL_SPCONV=1 opt-in reference (simdgroup tiling = the future version if decoders ever matter for time).
- Remaining gap to H100 (~17s at 1024): inside the DiT forwards (54s sampling ≈ 2.3× ideal FLOPs) — attention/MLP internals, a future deep-dive, not orchestration.
Env knobs (all default-sane): TRELLIS2_MLX_SAMPLER=0 (revert sampler), TRELLIS2_MLX_COMPILE=1, TRELLIS2_METAL_SPCONV=1. M1 Ultra validation: in flight (26.5GB peak fits with 100GB headroom).
Rung-3 addendum — M1 Ultra validated (2026-07-20)
| box | e2e wall | peak mem |
|---|---|---|
| m3ultra | 71.96 s | 26.5 GB |
| m1ultra | 181.5 s | 17.2 GB |
The M1 Ultra — which could not run TRELLIS.2 at all before this fork (torch-MPS path was m3-only in practice) — now generates full-density vertex-baked assets in ~3 min at 17GB. Install: clone monster/trellis-2-mrp-mlx to ~/trellis2-mlx + scripts/setup_macos.sh (PYTHON_BIN=uv 3.11) + gated weights rsynced from m3's HF cache (no HF login needed on workers). Both Ultras now serve the fastest local TRELLIS.2 in existence. TRELLIS.2 fleet: m3 primary (72s), m1 second lane (182s), zero dollars per asset.