#!/usr/bin/env python3 """ Reference for the full stage-1 flow-Euler sampling loop, to validate the C++ trellis2_ss_flow_sample against the real FlowEulerGuidanceIntervalSampler. Builds SparseStructureFlowModel in float32, loads the DINOv3 cond from a .dinodata file, draws fixed noise, runs the sampler, and writes a self-describing binary `ss_sample_ref.bin`: magic : 8 bytes "SSSAMP01" int32 : resolution, in_channels, cond_tokens, cond_channels, steps float32 : guidance_strength, guidance_rescale, gi_min, gi_max, rescale_t, sigma_min float32 : noise [in_channels * R^3] channel-major float32 : cond [cond_tokens * cond_channels] token-major float32 : latent[in_channels * R^3] channel-major (reference z_s) Usage: python ref_ss_sample.py [--device mps|cpu] [--dinodata .../MushroomBoy.dinodata] """ import argparse import json import os import struct import sys os.environ.setdefault("ATTN_BACKEND", "sdpa") os.environ.setdefault("SPARSE_ATTN_BACKEND", "sdpa") os.environ.setdefault("SPARSE_CONV_BACKEND", "none") os.environ.setdefault("PYTORCH_ENABLE_MPS_FALLBACK", "1") SHIV = os.environ.get("TRELLIS2_PY", "/trellis2") sys.path.insert(0, SHIV) import numpy as np import torch DEFAULT_CKPT = os.environ.get("TRELLIS2_CKPT", os.path.join( os.path.dirname(__file__), "..", "models", "TRELLIS.2-4B/ckpts/ss_flow_img_dit_1_3B_64_bf16")) def load_dinodata(path): with open(path, "rb") as f: assert f.read(8) == b"DINOCOND", "bad magic" _v, _d, ndim = struct.unpack("0: {(z_np>0).mean()*100:.2f}%)") noise_cm = noise_np.reshape(Cin, -1).reshape(-1) z_cm = z_np.reshape(Cin, -1).reshape(-1) cond_tm = cond_np.reshape(-1) with open(args.out, "wb") as f: f.write(b"SSSAMP01") f.write(struct.pack("<5i", R, Cin, Lkv, Cctx, params["steps"])) f.write(struct.pack("<6f", params["guidance_strength"], params["guidance_rescale"], params["guidance_interval"][0], params["guidance_interval"][1], params["rescale_t"], 1e-5)) f.write(noise_cm.astype("