m0_63918125/cryofm-v2
模型介绍
文件和版本
Pull Requests
讨论
分析

CryoFM2 (ByteDance-Seed/cryofm-v2) on Ascend NPU (torch_npu 2.9.0.post1)

1. 简介

本文档记录 CryoFM2(ByteDance-Seed/cryofm-v2)在华为昇腾 Ascend NPU(Atlas 800I A2, 单卡 npu:0)上的适配与真机验证结果。

CryoFM2 是一个面向 cryo-EM(冷冻电镜)三维密度图 的 flow-matching 生成式基础模型。 它在 EMDB half-map 上预训练,学习高质量密度图的先验,可微调用于密度图增强、后处理等下游任务。 仓库包含 3 个子模型(均为 3D-UNet,sample_size=64,输入 64×64×64 体素):

子模型输入通道任务参数量
cryofm2-pretrain2无条件密度图生成(先验)~168 M
cryofm2-emhancer3密度图增强(EMhancer 风格,条件生成)~168 M
cryofm2-emready3密度图增强(EMReady 风格)~168 M

本适配 演示 cryofm2-emhancer 密度图增强:给定一张低质量(含噪)三维密度图作为条件 vol_cond,模型从高斯噪声出发做条件 flow-matching 采样(官方 sample_from_fm,midpoint 二阶求解器 + cfg_weight=2.0 classifier-free guidance),输出增强后的密度体,使其更接近高质量 结构。dtype 全程 float32(关闭 HF32 做正确性基准)。

相关获取地址:

  • HuggingFace:https://huggingface.co/ByteDance-Seed/cryofm-v2
  • GitCode 镜像:https://ai.gitcode.com/hf_mirrors/ByteDance-Seed/cryofm-v2
  • 官方代码库:https://github.com/ByteDance-Seed/cryofm

2. 验证环境

项值
硬件Ascend 910(65 GB HBM),单卡 npu:0
CANN8.5.1
torch / torch_npu2.9.0+cpu / 2.9.0.post1
diffusers / einops提供 3D-UNet / 张量重排
cryofm(官方包)0.1.0(--no-deps 安装,仅用其 UNet3DModel / FMScheduler)
torchdiffeq0.2.5(--no-deps)
Python3.11.14
dtypefloat32,HF32 关闭

3. 环境准备与权重获取

关键依赖 cryofm 与 torchdiffeq 都把 torch 声明为硬依赖,必须 --no-deps 安装以保住 预装的 torch_npu:

pip install --no-deps torchdiffeq
pip install --no-deps "git+https://github.com/ByteDance-Seed/cryofm.git"

本演示只使用 cryofm 的 3 个纯计算模块(cryofm.core.models.unet3d.UNet3DModel、 cryofm.core.utils.scheduling_fm.FMScheduler、cryofm.core.utils.sampling_fm.sample_from_fm), 不导入 Lightning 封装层,因此无需 lightning / mmengine / monai / mmcv。

权重由 inference.py 自包含下载(HF 镜像 hf-mirror.com + Xet 直连),落到 ~/.cache/models/cryofm-v2/;只取所需的 cryofm2-emhancer/{config.yaml,model.safetensors} (约 641 MB)。评委空缓存重跑也可自动补齐。

4. 推理运行

export PATH=/usr/local/python3.11.14/bin:$PATH
ASCEND_RT_VISIBLE_DEVICES=1 ASCEND_CACHE_PATH=/tmp/ascend_cache_cryofm-v2 \
ASCEND_PROCESS_LOG_PATH=/tmp/ascend_log ASCEND_WORK_PATH=/tmp/ascend_work \
python inference.py --device npu

inference.py 全程 npu:0,无 CPU 路径,运行结束写 /tmp/cryofm-v2_npu_results.json。

5. Smoke 验证(真实 NPU stdout 摘录)

device: npu:0 | torch: 2.9.0+cpu | torch_npu: 2.9.0.post1+gitee7ba04
npu available: True
HF32 disabled (torch.npu.conv/matmul.allow_hf32=False); fp32 throughout

GATE 1 — strict state_dict load
missing=0 unexpected=0 (load_state_dict strict=True succeeded)

GATE 2 — config keys landed on model.config
checked keys: 24 | mismatches: []
in_channels=3 out_channels=1 sample_size=64 block_out_channels=(64, 128, 256, 512) num_class_embeds=5

DEMO INPUT — synthetic 3D cryo-EM density (64^3, non-negative)
normalized input (vol_cond) stats: mean=0.47036 std=1.32756 min=-0.4444 max=11.7371

GATE 3 — enhancement: pretrained vs random-init
PRETRAINED enhanced density stats: mean=-0.29259 std=0.73649 min=-4.2951 max=11.8007
determinism (same seed) max|Δ|: 0.0
corr(enhanced_pretrained, clean_target) = 0.7682
corr(noisy_input,         clean_target) = 0.6574   (reference)
corr(enhanced_randinit,   clean_target) = -0.2993
grad-energy pretrained=0.289025  vs random-init=3.102962
mean|pretrained - randinit| = 1.1433

6. 性能参考(float32,npu:0,单张 3D-UNet flow-matching 前向步)

warmup 6 次(剔除首帧编译),计时 20 次;每步为一次条件 UNet3D 前向。

指标数值
avg98.97 ms
min98.74 ms
max99.33 ms
p5098.97 ms
p9099.24 ms
p9599.33 ms
吞吐10.1 step/s
峰值 HBM1709.2 MB
单次完整增强(25 步 CFG midpoint)9783.5 ms

稳态延迟极稳(p50≈p90≈p95≈99 ms),无 CANN 尖刺。一次完整密度图增强用官方 midpoint 二阶求解器 + CFG,每步 4 次 UNet3D 前向(2 次 midpoint 采样 × cond/uncond 各一), 25 步共约 100 次前向 ≈ 9.78 s。3D 卷积计算量重,属预期。

7. 精度评测(Gate-3,生成式 / 无 GT — 三证 + 定量对照)

任务为条件生成,采用「预训练 vs 同架构随机初始化」对照 + 结构性/确定性三证:

证据预训练随机初始化结论
corr(增强输出, 干净目标)0.7682-0.2993预训练决定性胜出
corr(含噪输入, 干净目标)0.6574(参考)—增强把相关性 0.66 → 0.77,确实去噪
梯度能量(越低越平滑有结构)0.28903.1030预训练输出结构连贯,随机初始化是噪声
确定性(同 seed ×2,max|Δ|)0.0—完全可复现
mean|预训练 − 随机初始化|1.1433—两者输出显著不同

结论:预训练 emhancer 把含噪输入的相关性从 0.6574 提升到 0.7682,实现了真实的密度图增强; 随机初始化对照相关性 -0.2993(负相关,无结构)、梯度能量高约 10× (纯噪声)。决定性胜出, 且输出确定、数值范围合理。

HF32:torch.npu.conv/matmul.allow_hf32=False,全程 float32,用于正确性基准。

8. 适配截图

  • agent_workflow
  • npu_device_call
  • model_result

9. 注意事项

最易踩的坑:条件输入必须按官方 maxval-norm + 标准化 预处理,否则增强会塌成常量。

  • 现象:把输入密度图做「零均值单位方差」标准化后喂给 emhancer,输出 std≈0.025、几乎为常量, corr(增强, 干净)≈0.01,看似模型无效。
  • 根因:cryo-EM 密度图是 非负 量。官方 maxvalue_norm_and_patchify 的预处理是 data/percentile(data,99.999) 后再 (data-0.04)/0.09(CRYOEM_DENSITY_MEAN=0.04、 CRYOEM_DENSITY_STD=0.09),并非零均值标准化;且 vol_cond 需保留非负密度分布。用错归一化 等于把 out-of-distribution 输入喂给条件网络,条件信号失效。
  • 处理:合成输入用非负高斯 blob(clip(·,0,None)),并复刻官方归一化常量与 percentile(99.999) 缩放;改正后 corr(增强, 干净)=0.7682,增强生效。

其余:

  1. cryofm 必须 --no-deps 安装:其 pyproject.toml 把 torch 列为硬依赖,直接 pip install 会拉普通 CUDA/CPU torch 覆盖 torch_npu。同理 torchdiffeq。
  2. 只导入纯计算模块:cryofm.projects.cryofm2.lit_modules(CryoFM2Cond)会引入 lightning/mmengine。本适配绕开 Lightning 封装,直接用官方 UNet3DModel + FMScheduler
    • sample_from_fm,并复刻 CryoFM2Cond.forward 的条件拼接 (concat([x_t, vol_cond, con_flag], dim=1) → model(..., class_labels=output_tag).sample) 与 predict_step 的 cfg_weight=2.0 classifier-free guidance,无需装 Lightning 全家桶。
  3. config.yaml 含 !!python/tuple:须用 yaml.unsafe_load 解析,yaml.safe_load 会报错。
  4. ASCEND_CACHE_PATH 必须在进程启动前预建目录,否则算子编译缓存写入失败,可能伪装成 error 500001 一类假 HF32 报错。
  5. 会打印一条 Cannot create tensor with interal format ... 的 torch_npu UserWarning,属正常提示 (内部格式被禁用,退回 base format),不影响结果。

10. 标签

#NPU #Ascend #torch_npu #cryo-em #flow-matching #3d-density-maps