CryoFM2 是一款基于流的生成式基础模型,专为冷冻电镜密度图设计。该模型在精选的 EMDB 半图上进行预训练,以学习高质量冷冻电镜密度的通用先验知识,并可通过微调适配下游任务。
该模型学习从简单高斯分布到复杂冷冻电镜密度分布之间的连续映射,从而实现稳定的生成与灵活的适配。CryoFM2 还可作为贝叶斯先验,与任务特定的似然函数自然结合,支持各向异性感知精修、非均匀重建以及受控密度修饰等应用。
CryoFM2 在精选的 EMDB 半图上进行预训练,以学习高质量冷冻电镜密度的通用先验知识。该模型可针对多种下游任务进行微调,如密度图增强和后处理。
预训练架构:
微调架构(适用于 EMhancer/EMReady 风格的后处理):
本目录已按 PyTorch + torch_npu 框架将 CryoFM2 迁移至华为昇腾 NPU(Ascend910),并完成服务化推理交付。
| 项目 | 配置 |
|---|---|
| 硬件 | Ascend910 (Atlas 800),本机 2 卡,每卡 61.3 GB |
| OS | openEuler aarch64 / 麒麟 / Ubuntu |
| CANN | 8.5.1(source /usr/local/Ascend/ascend-toolkit/set_env.sh) |
| Python | 3.11.14 |
| PyTorch | 2.9.0 + torch_npu 2.9.0.post1+gitee7ba04 |
| safetensors / numpy / Pillow | 0.7.0 / 1.26.4 / 12.2.0 |
| HTTP 框架 | fastapi 0.123.10 + uvicorn 0.46.0 + pydantic 2.13.3 |
cryofm-v2/
├── inference.py # 服务化推理脚本(HTTP 服务 + 自测模式)
├── make_screenshots.py # 截图/水印生成脚本(maggie_Ha 水印)
├── requirements.txt # 运行环境依赖(基于本机实测版本)
├── workflow.txt # 适配流程文本(用于生成适配过程截图)
├── result.txt # 完整中文测试日志
├── cryofm2_npu/ # 自包含推理包(昇腾适配)
│ ├── __init__.py # 对外导出 API
│ ├── unet3d.py # 3D UNet 模型(复刻官方架构,权重严格加载)
│ ├── scheduler.py # 流匹配调度器 FMScheduler
│ ├── sampler.py # Euler / midpoint 采样
│ ├── model.py # 无条件/条件模型前向封装
│ └── mrc_io.py # MRC 文件读写(自包含,MRC2014)
├── cryofm2-pretrain/ # 无条件预训练权重(config.yaml + model.safetensors)
├── cryofm2-emhancer/ # EMhancer 风格微调权重(in_channels=3, class_embedding)
├── cryofm2-emready/ # EMReady 风格微调权重(in_channels=3, class_embedding)
├── outputs/ # 推理输出
│ ├── cryofm2-pretrain_selftest.mrc # 自测生成的 64³ 密度图
│ ├── selftest.log # 完整中文自测日志
│ └── server.log # HTTP 服务日志
└── assets/ # 部署截图(全部带 maggie_Ha 水印)
├── npu_device_call.png # NPU 设备调用截图
├── model_result.png # 测试结果截图
├── agent_workflow.png # 适配过程截图
├── cryofm2_overview.jpg
├── cryofm2_arch-pretrain.jpg
└── cryofm2_arch-finetune.jpgsource /usr/local/Ascend/ascend-toolkit/set_env.sh
pip install -r requirements.txt# 默认变体 cryofm2-pretrain,端口 8090
python inference.py --variant cryofm2-pretrain --port 8090
# 条件模型(EMhancer / EMReady)
python inference.py --variant cryofm2-emhancer --output-tag 1 --port 8091
python inference.py --variant cryofm2-emready --output-tag 0 --port 8092# 健康检查
curl http://127.0.0.1:8090/health
# 生成 1 个 64³ 密度图(50 步 Euler)
curl -X POST http://127.0.0.1:8090/generate \
-H "Content-Type: application/json" \
-d '{"num_samples":1,"num_steps":50,"method":"euler","side_shape":64}'python inference.py --variant cryofm2-pretrain --selftest --num-steps 50| 指标 | 值 |
|---|---|
| 模型 | cryofm2-pretrain(3D UNet,约 168 M 参数,64³ 输入,in_channels=2) |
| NPU 设备 | Ascend910_9362,61.3 GB × 2(详见 assets/npu_device_call.png) |
| 推理方法 | 50 步 Euler 流匹配采样生成 1 个 64³ 密度图 |
| 单次完整推理耗时 | 6423 ms(50 步,含全部 UNet 前向 + 数值积分) |
| 平均单步耗时 | 128.5 ms/步 |
| HTTP 10 步调用耗时 | 约 1304 ms |
| 权重加载耗时 | 约 1184 ms(首次 safetensors 反序列化) |
| 输出密度图 | outputs/cryofm2-pretrain_selftest.mrc(64³,voxel=1.5 Å) |
| 数值范围 | mean=-0.10, std=2.18, min=-9.09, max=11.12 |
| 精度 | 与 diffusers 参考逐位对比 max_abs_diff < 1e-5 |
| 测试结果截图 | assets/model_result.png |
| 适配过程截图 | assets/agent_workflow.png |
推理一次耗时(核心指标):50 步 Euler 完整生成 ≈ 6423 ms。即一次
POST /generate接口调用完整返回耗时。 降低--num-steps可显著提速(10 步约 1.3 s);提高--num-steps(如 200 步)可提升生成质量。
_from_deprecated_attn_block=True + upcast_softmax;自包含实现用手工 bmm 复刻(fp32 softmax),规避 NPU 上 SDPA 行为差异。einops 'b c d h w -> b (d h w) c' 展平等价于 x.permute(0,2,3,4,1).reshape(b,d*h*w,c),不能直接 x.reshape(b,d*h*w,c)(channel 未转置,diffusers 内部不会出错但自包含实现容易踩坑,曾导致注意力输出 diff≈5)。inference.py 在初始化阶段用 64³ 真实输入做一次预热,保证后续计时不被首次编译干扰。config.yaml 是 mmengine 风格,含 !!python/tuple 标签,需在 yaml.SafeLoader 注册 tuple 构造函数。load_state_dict(strict=True) 直接成功。torch.npu.set_device(0) + torch.device("npu:0"),全程中文日志输出耗时。三张交付截图(位于 assets/)均带 maggie_Ha 水印(半透明斜向重复 + 右下角签名):
| 文件 | 内容 | 水印 |
|---|---|---|
npu_device_call.png | NPU 设备调用信息(torch.npu 设备、版本、初始化日志) | ✅ maggie_Ha |
model_result.png | 测试结果(耗时、数值统计、推理指标) | ✅ maggie_Ha |
agent_workflow.png | 适配过程(模型分析 → 算子评估 → 自包含实现 → NPU 迁移 → 服务化) | ✅ maggie_Ha |
如需重新生成截图:
python make_screenshots.py \
--npu-info npu_info.txt \
--result result.txt \
--workflow workflow.txt在使用 CryoFM2 之前,您需要配置环境并安装相应软件包。请按照以下步骤开始操作:
# Clone the repository
git clone https://github.com/ByteDance-Seed/cryofm.git
cd cryofm
# Create a new conda environment for CryoFM (recommended)
conda create -n cryofm python=3.10 -y
conda activate cryofm
# Install CryoFM
pip install .从预训练模型中生成样本,以探索学习到的数据分布:
预训练模型:
import torch
from mmengine import Config
from cryofm.core.utils.mrc_io import save_mrc
from cryofm.core.utils.sampling_fm import sample_from_fm
from cryofm.projects.cryofm2.lit_modules import CryoFM2Uncond
# Update the path to your model directory
model_dir = "path/to/cryofm-v2/cryofm2-pretrain"
cfg = Config.fromfile(f"{model_dir}/config.yaml")
lit_model = CryoFM2Uncond.load_from_safetensors(f"{model_dir}/model.safetensors", cfg=cfg)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
lit_model = lit_model.to(device)
lit_model.eval()
def v_xt_t(_xt, _t):
return lit_model(_xt, _t)
# Enable bfloat16 for faster inference if your GPU supports it
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
out = sample_from_fm(
v_xt_t,
lit_model.noise_scheduler,
method="euler",
num_steps=200,
num_samples=3,
device=lit_model.device,
side_shape=64
)
# Apply normalization if configured
if hasattr(lit_model.cfg, "z_scale") and lit_model.cfg.z_scale.mean is not None:
out = out * lit_model.cfg.z_scale.std + lit_model.cfg.z_scale.mean
# Save generated samples
for i in range(3):
save_mrc(out[i].float().cpu().numpy(), f"sample-{i}.mrc", voxel_size=1.5)微调模型(EMhancer/EMReady):
import torch
from mmengine import Config
from cryofm.core.utils.mrc_io import save_mrc
from cryofm.core.utils.sampling_fm import sample_from_fm
from cryofm.projects.cryofm2.lit_modules import CryoFM2Cond
# Choose style: "emhancer" or "emready"
style = "emhancer"
model_dir = f"path/to/cryofm-v2/cryofm2-{style}"
cfg = Config.fromfile(f"{model_dir}/config.yaml")
lit_model = CryoFM2Cond.load_from_safetensors(f"{model_dir}/model.safetensors", cfg=cfg)
output_tag = 1 if style == "emhancer" else 0
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
lit_model = lit_model.to(device)
lit_model.eval()
def v_xt_t(_xt, _t):
bs = _xt.shape[0]
unconditional_generation_conds = {
"input_cond": None,
"output_cond": torch.tensor([output_tag] * bs).to(device),
"vol_cond": None, # dimension should be [bs, d, h, w]
}
return lit_model(_xt, _t, generation_conds=unconditional_generation_conds)
# Enable bfloat16 for faster inference if your GPU supports it
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
out = sample_from_fm(
v_xt_t,
lit_model.noise_scheduler,
method="euler",
num_steps=200,
num_samples=3,
device=lit_model.device,
side_shape=64
)
# Apply normalization if configured
if hasattr(lit_model.cfg, "z_scale") and lit_model.cfg.z_scale.mean is not None:
out = out * lit_model.cfg.z_scale.std + lit_model.cfg.z_scale.mean
# Save generated samples
for i in range(3):
save_mrc(out[i].float().cpu().numpy(), f"{style}-sample-{i}.mrc", voxel_size=1.5)CryoFM2 支持利用预训练模型作为贝叶斯先验,对密度图执行多种修正操作。支持的算子包括:
基本用法:
python -m cryofm.projects.cryofm2.uncond_sampling \
-i1 half_map_1.mrc \
-i2 half_map_2.mrc \
-o ./output \
--model-dir path/to/cryofm-v2/cryofm2-pretrain \
--op denoise \
--norm-grad \
--use-lamb-w对于修复任务,您需要提供一个 RELION starfile 路径:
python -m cryofm.projects.cryofm2.uncond_sampling \
-i1 half_map_1.mrc \
-i2 half_map_2.mrc \
-o ./output \
--model-dir path/to/cryofm-v2/cryofm2-pretrain \
--op inpaint \
--data-starfile-path path/to/relion_data.star \
--norm-grad \
--use-lamb-wCryoFM2 提供针对不同风格的密度图增强微调模型,类似于 EMhancer 和 EMReady。
python -m cryofm.projects.cryofm2.cond_sampling \
-i input_map.mrc \
-o ./output_emhancer \
--model-dir path/to/cryofm-v2/cryofm2-emhancer \
--output-tag 1python -m cryofm.projects.cryofm2.cond_sampling \
-i input_map.mrc \
-o ./output_emready \
--model-dir path/to/cryofm-v2/cryofm2-emready \
--output-tag 0 \
--cfg-weight 0.5参数:
-i:输入密度图文件(MRC 格式)-o:输出目录--model-dir:包含 config.yaml 和 model.safetensors 的模型目录路径--output-tag:风格标签(1 为 EMhancer,0 为 EMReady)--cfg-weight:无分类器引导权重(可选,默认值随模型不同而异)accelerate launch 在多个 GPU 上加速推理:
NCCL_DEBUG=ERROR accelerate launch --num_processes=${NUM_GPUS} --main_process_port=8881 \
python -m cryofm.projects.cryofm2.cond_sampling ...--bf16 标志以降低内存占用并加速推理。本模型面向科学研究和结构生物学应用。使用者应:
如果您觉得 CryoFM2 有用,请引用:
@article{
Li2025.12.29.696802,
author={Li, Yilai and Yuan, Jing and Zhou, Yi and Wang, Zhenghua and Chen, Suyi and Yang, Fengyu and Ling, Haibin and Kovalsky, Shahar Z and Zheng, Xiaoqing and Gu, Quanquan},
title={A Generative Foundation Model for Cryo-EM Densities},
elocation-id={2025.12.29.696802},
year={2025},
doi={10.64898/2025.12.29.696802},
publisher={Cold Spring Harbor Laboratory},
URL={https://www.biorxiv.org/content/early/2025/12/29/2025.12.29.696802},
eprint={https://www.biorxiv.org/content/early/2025/12/29/2025.12.29.696802.full.pdf},
journal={bioRxiv}
}本模型采用 Apache 2.0 许可证发布。详细条款请参阅 LICENSE 文件。
本工作由字节跳动 Seed 团队开发。如需了解更多信息,请访问: