docs: FABLE_STARTUP_PACKET — prep recon consolidated (adopt PR#175 MLX base, Hunyuan mx.compile targets, prep_local draft, weight-convert+parity plan)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
parent
ca4c802e40
commit
abfa344565
193
docs/FABLE_STARTUP_PACKET.md
Normal file
193
docs/FABLE_STARTUP_PACKET.md
Normal file
@ -0,0 +1,193 @@
|
|||||||
|
# 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-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.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
|
||||||
|
```bash
|
||||||
|
# 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)`:
|
||||||
|
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.
|
||||||
|
3–7. `_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`):
|
||||||
|
```python
|
||||||
|
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.save`s 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
|
||||||
|
|
||||||
|
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.
|
||||||
Loading…
Reference in New Issue
Block a user