fMRI → 3D:用 TRELLIS.2 替换 MinD-3D 的 3D 生成端(实验记录)
在 fMRI-Shape 的 sub-01 上,把 MinD-3D 的 3D 生成端(Argus3D 自回归 GPT)换成 TRELLIS.2-4B。思路是:fMRI 只需预测 TRELLIS.2 的图像条件(DINOv3 ViT-L/16 token,1029×1024),3D 先验全部来自冻结的 TRELLIS.2。
本仓库只包含代码、配置、指标和日志,不含模型权重、网格和数据。
结果(sub-01,测试集)
测试集 104 个物体,其中 101 个在 ShapeNetCore.v2.PC15k 中有真值点云,下表只统计这 101 个。
| 方法 | CD ×1e-3 ↓ | CD 中位数 | EMD ↓ | F-score@0.02 ↑ | 逐物体优于 MinD-3D |
|---|---|---|---|---|---|
| MinD-3D 官方权重(sub-01) | 14.92 | 10.90 | 0.115 | 0.186 | — |
| fMRI → 预测 token → TRELLIS.2 | 12.20 | 8.90 | 0.101 | 0.275 | 54 / 101 |
| fMRI → 预测 token → 检索最近邻训练物体的真实 token → TRELLIS.2 | 9.61 | 6.36 | 0.095 | 0.326 | 67 / 101 |
| 上限:真实刺激帧(第 24 帧)→ TRELLIS.2 | 1.20 | 0.81 | 0.041 | 0.678 | — |
分类别的 CD(×1e-3):
| 类别 | MinD-3D | 预测 token | 检索 | 上限 |
|---|---|---|---|---|
| 02691156 飞机 | 11.76 | 10.32 | 8.35 | 0.41 |
| 02828884 长椅 | 17.85 | 8.18 | 6.55 | 0.75 |
| 02933112 柜子 | 11.22 | 18.13 | 11.95 | 1.93 |
| 02958343 汽车 | 4.02 | 5.95 | 6.07 | 1.36 |
| 03001627 椅子 | 16.32 | 9.53 | 10.50 | 2.21 |
| 03211117 显示器 | 15.53 | 15.43 | 7.86 | 1.19 |
| 03636649 灯 | 21.36 | 28.43 | 21.77 | 2.21 |
| 03691459 音箱 | 9.43 | 11.59 | 9.57 | 2.26 |
| 04090263 步枪 | 33.71 | 1.77 | 1.33 | 0.36 |
| 04256520 沙发 | 9.94 | 14.46 | 7.21 | 0.83 |
| 04379243 桌子 | 9.16 | 16.17 | 13.12 | 0.82 |
| 04401088 电话 | 12.22 | 9.47 | 3.75 | 0.46 |
| 04530566 船 | 20.97 | 8.13 | 15.58 | 0.75 |
token 层面的指标(映射器 v0,测试集)
| MSE | 检索 top-1 / top-5(在 104 个测试物体中) | 最近邻类别准确率 | |
|---|---|---|---|
| 直接输出训练集平均 token | 0.116 | 1% / 5% | — |
| 作弊参照:真实类别的平均 token | 0.074 | 12.5% / 56% | — |
| 映射器 v0,对比损失权重 0.1 | 0.116 | 13.5% / 51% | 49% |
| 映射器 v0,对比损失权重 1.0 | 0.368 | 6.7% / 37% | 44% |
结论
- 3D 生成端不是瓶颈:真实图像条件下 CD 为 1.20e-3,约为 MinD-3D 的 1/12。
- fMRI 可以解码出物体和类别信息:检索 top-1 13.5%(随机 1%),最近邻类别准确率 49%(13 类,随机约 7.7%)。
- 直接回归 1029×1024 的 token 会严重过拟合:1200 个训练样本,最佳验证点出现在第 500–1000 步。回归结果会退化成类似平均的 token。
- 检索比直接回归更好:喂给 TRELLIS.2 的 token 一定来自真实分布,形状更干净。代价是只能输出训练集中出现过的形状。
- 当前主要误差来源是类别选错,约一半测试物体的最近邻属于错误类别。下一步重点是提升 fMRI 的语义解码精度。
方法细节
评测(code/eval_shape.py)
- 沿用 MinD-3D 的
tools/get_cd.py:每个形状采 2048 点,用pc_norm归一化(质心居中、最大半径缩放到 0.5);CD 为双向最近邻平方距离均值之和。 - 原协议中的 PCA 对齐换成统一的朝向搜索:2 种向上轴映射(y-up / z-up)× 8 个方位角(每 45°),取 CD 最小者。所有方法使用相同搜索。
- EMD 为 2048 点上的精确匹配(Hungarian),F-score 阈值 0.02。
- 由于对齐方式不同,数值不能直接与论文表格比较。
真值:ShapeNetCore.v2.PC15k(15000 点,y-up),来自 ModelScope。
TRELLIS.2 设置
- 512³ 管线:sparse structure flow(32³)→ shape SLAT flow 512 → shape decoder →
fill_holes,seed 42,只生成几何。 - 稀疏结构为空时,依次换种子重试,最多 4 次。v0 的 104 个物体全部成功。
阶段 2 的条件 token(code/extract_dinov3_tokens.py)
- 1622 个刺激视频,每个取第 0、24、…、168 帧。
- 预处理与 TRELLIS.2 的
preprocess_image相同(RMBG-2.0 去背景、裁剪、alpha 合成到黑底)。 - 再用 DINOv3 ViT-L/16 在 512 分辨率下编码,并做 layer norm,得到 (8, 1029, 1024),fp16 保存。
映射器 v0(code/mapper/train_mapper.py)
- 输入:MinD-3D sub-01 官方权重中的 fMRI ViT-L 编码器(冻结,
forward_encoder_wo_pred)。每个物体 10 帧,每帧 257×1024;特征预先计算(code/mind3d/precompute_fmri_feats.py)。 - 训练时:随机取 6 帧 fMRI,加帧嵌入和位置嵌入;特征加噪 0.1,随机丢弃 20% 的 memory token。
- 测试时:固定使用第 2–7 帧,与 MinD-3D 相同。
- 模型:1029 个查询向量的 Transformer 解码器,6 层,宽度 768,12 个头,59.3M 参数。输出为训练集平均 token 加残差,再做 layer norm。
- 目标:第 24 帧(与上限实验一致)的 DINOv3 token。
- 损失:token MSE + λ × 池化 token 的对称 InfoNCE(τ=0.05),λ 取 0.1 / 1.0。
- 优化:AdamW,lr 3e-4,wd 0.05,batch 32,6000 步,cosine 调度。
- 数据划分:从 1300 个训练物体中留出 100 个做验证,按验证集余弦相似度选 checkpoint;测试集只在最后评估一次。
- 检索基线:用预测 token 的池化向量,在 1200 个训练物体的真实 token 中找最近邻,再把该物体第 24 帧的真实 token 送进 TRELLIS.2。
MinD-3D 基线(code/mind3d/gen_testset_mind3d.py)
- 官方 sub-01 权重,fMRI 第 2–7 帧,每个物体
set_random_seed(100),top-k 250。 - 官方 demo 的两个物体与官方给出的网格一致(Chamfer-L1 0.0083 / 0.0119)。
仓库结构
code/
eval_shape.py 统一评测(CD / EMD / F-score + 朝向搜索)
upper_bound_trellis2.py 上限实验:真实刺激帧 → TRELLIS.2
extract_dinov3_tokens.py 阶段 2 目标:全部刺激视频的 DINOv3 条件 token
gen_from_tokens.py 由预测或检索得到的 token 生成 TRELLIS.2 网格
mapper/train_mapper.py fMRI → token 映射器
mind3d/gen_testset_mind3d.py MinD-3D 官方权重跑全测试集
mind3d/precompute_fmri_feats.py 冻结 fMRI 编码器的特征预计算
mind3d/setup_py312.patch MinD-3D 在 Python 3.12 下的编译修复
build_trellis2*.sh TRELLIS.2 扩展编译脚本(H100/H200,sm_90)
configs/ TRELLIS.2 本地管线配置(512 / 完整)
results/ 各方法的逐物体指标、映射器结果(含检索到的最近邻)、上限实验统计
logs/ 映射器训练日志与评测日志
环境
- Python 3.12,NGC PyTorch(CUDA 12),GPU 为 H100/H200。
- TRELLIS.2 使用 flex_gemm 和 flash_attn。
- 权重来源:TRELLIS.2-4B(HF / ModelScope);DINOv3 ViT-L/16 和 RMBG-2.0 来自 ModelScope;sparse structure decoder 来自 TRELLIS-image-large。
下一步
- 改进检索:top-k 加权平均或投票;比较 SigLIP 2、DINO CLS 等检索空间。
- 解冻 fMRI 编码器微调,加入 ROI 加噪增强和 sigmoid 对比损失,提升类别准确率。
- 用以 fMRI 为条件的扩散先验生成 DINO token,替代直接回归。
- 检查 MinD-3D 基线在训练集与测试集上的差距,判断是否过拟合。
- Downloads last month
- -
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support