YAML Metadata Warning:empty or missing yaml metadata in repo card
Check out the documentation for more information.
SymFold:基于 RNA 语言模型的 RNA 二级结构预测
SymFold 是一个研究型 RNA 二级结构预测项目。它将 RNA 序列的二级结构表示为对称二值 contact map:给定长度为 (L) 的序列,模型预测 (L\times L) 矩阵中每一对核苷酸是否配对。
项目同时维护两条训练路线:
- 直接判别式预测:一次前向直接输出 contact logits,适合快速实验和当前主要消融;
- 离散 Flow Matching:从带噪二值 contact map 逐步去噪,通过 τ-leap CTMC 采样生成结构,用于研究生成式结构预测。
当前研究重点是:GB.RNA 表示、pair representation、结构重复数据、序列变体泛化、长程配对及非 canonical 配对的影响。
1. 项目要解决什么问题
输入是一条 RNA 序列,例如:
GGCUCACCAAGGCG...
输出是其二级结构的配对集合或等价的 contact map。数据标签原始形式为 dot-bracket;项目会解析普通 stem 与多层 pseudoknot bracket,再转为对称 contact map。支持的 bracket tier 包括 () [] {} <> 及大小写字母配对。
[ C_{ij}=C_{ji}=\begin{cases} 1,& \text{nucleotide }i\text{ 与 }j\text{ 配对}\ 0,& \text{otherwise} \end{cases} ]
实现依据: dot-bracket 解析及伪结 tier 定义见 symfold/data/dotbracket.py:1-57;Parquet 样本读取、contact-map 构建和 batch padding 见 symfold/data/datasets.py:24-97。
2. 总体架构
RNA sequence
│
├── RNA encoder(GB.RNA / RNA-FM / RiNALMo)
│ ├── per-nucleotide hidden states: [B, L, H]
│ └── selected attention maps: [B, A, L, L]
│
├── 直接判别式路线
│ hidden/attention fusion → [B, L, L] contact logits
│
└── 离散 Flow Matching 路线
pair condition + noisy x_t + time t
→ DiT-style pair backbone → [B, L, L] contact logits
→ τ-leap CTMC sampling → contact probabilities
2.1 共用 RNA encoder 适配层
RNAEncoderFeatureExtractor 统一封装本地 RNA-FM、RiNALMo 与 GB.RNA:
- 输出逐核苷酸 1D 表示
[B,L,H]; - 输出最后若干 encoder layer 的 attention map
[B,A,L,L]; - 通过
mask排除 padding; - 支持冻结 encoder,或在解冻时启用 gradient checkpointing 降低显存。
GB.RNA 走仓库内置的 RNABert/tokenizer 路径;其输入按单碱基 token 化,且显式校验 token 数为 (L+2)。
实现依据: encoder 类型识别与加载见 symfold/models/rnafm_encoder.py:28-91;token 对齐与特征/attention 输出见 symfold/models/rnafm_encoder.py:117-194;配置键兼容逻辑见 symfold/models/rnafm_encoder.py:197-210。
2.2 直接判别式模型(当前主实验模型)
discriminative_contact_map 的前向输入只有 List[str],一次输出:
logits: [B, L, L]
mask: [B, L]
模型流程:
- encoder hidden 经
LayerNorm + Linear投影到pair_dim; - 对任意位置对 ((i,j)),使用对称平均 ((h_i+h_j)/2) 构造 sequence pair feature;
- encoder attention 用
1×1 Conv投影为 pair feature; - 通过双向 gate 与 FiLM 调制融合两条路径;
- 用 interaction MLP 和可选的单层 depthwise
3×3pair smoother 建模局部一致性; - 输出对称 logits。
这条路线不使用 flow noising 或 CTMC,因此训练和评估显著快于生成式路线。
实现依据: gated fusion 见 symfold/models/discriminative_contact_map_model.py:19-50;局部 smoother 见 symfold/models/discriminative_contact_map_model.py:53-76;模型组装与对称 logits 输出见 symfold/models/discriminative_contact_map_model.py:79-149;训练损失与循环见 symfold/train_supervised_pair.py:54-86,186-268。
2.3 离散 Flow Matching 模型
Flow Matching 路线使用:
[ x_t\sim\mathrm{Bernoulli}((1-t)\rho_0+t x_1) ]
其中 (x_1) 为真实 contact map,(x_t) 是时间 (t) 的对称二值带噪状态。网络预测 (p(x_1=1\mid x_t,t,\mathrm{RNA}))。
其 pair condition 显式融合:
- (h_i,h_j,|h_i-h_j|,h_i\odot h_j);
- encoder attention;
- log-distance embedding;
- 可选 AU/GC/GU/other/unknown base-pair type embedding。
随后在 pair space 上运行 DiT-style backbone;当 patch_size>1 时,主干在 patch space 计算,再回到全分辨率进行可选 refinement。
实现依据: pair representation 见 symfold/models/flow_matching_model.py:28-119;模型组装、patch/unpatch 与 logits 输出见 symfold/models/flow_matching_model.py:180-307;Bernoulli bridge、BCE/Dice/degree loss 与 CTMC 采样见 symfold/models/flow_matching.py:72-188。
3. 训练、解码与指标
3.1 损失
直接判别式训练使用 masked BCE-with-logits,并可附加:
- 正样本权重
pos_weight; - focal reweighting;
- Dice loss;
- soft degree penalty(抑制一个碱基配给多个 partner)。
Flow Matching 使用对应的 masked flow loss。
实现依据: 直接模型损失见 symfold/train_supervised_pair.py:54-86;Flow Matching 损失见 symfold/models/flow_matching.py:78-131。
3.2 Greedy 结构投影
默认评估可开启 greedy projection:
- 排除 padding 和小于
min_sequence_separation的 pair; - 取概率不低于 threshold 的候选边;
- 按概率降序;
- 每个碱基最多保留一个 partner。
这保证了 at-most-one-partner,但不保证 non-crossing,也不会强制 canonical pairing。
实现依据: 解码见 symfold/metrics.py:21-51;Precision、Recall、F1、MCC、Accuracy 计算见 symfold/metrics.py:54-86。
3.3 Flow Matching 评估
Flow 路线评估时先执行 CTMC 多步采样,再根据一个或多个阈值计算指标;如果提供 threshold grid,会选验证 F1 最优阈值。
实现依据: 采样评估、阈值扫描与 greedy decode 见 symfold/evaluate_flow_matching.py:21-150。
4. 数据集与实验设计
完整的数据说明、样本数、派生关系和使用边界见 data/README.md。核心数据如下:
| 数据 | 当前规模 | 作用 |
|---|---|---|
data/bprna-spot0/ |
train 10,814 / val 1,300 / test 1,305 | 默认可比基线 split |
data/bprna-spot0-trainfilter/ |
train 8,860 | spot0 train 内,以 CD-HIT + 同长度结构 Jaccard 过滤近重复;val/test 不变 |
data/bprna-spot0-structdedup098/ |
train 7,240 | spot0 train 内,以全对全 bpRNA-align norm_score>=0.98 去重;val/test 字节一致 |
data/bprna-genfilter/ |
6,521 / 970 / 971 | 从全 bpRNA 池重建的 cluster-disjoint split;不可与 spot0 直接当作只改 train 的对照 |
data/bprna-full/ |
train 81,991 | 近乎完整 bpRNA-1m 训练池;保留 spot0 val/test |
data/bprna-new/test.parquet |
5,401 | 跨 RNA family 泛化测试 |
data/archiveii.512/test.parquet |
3,865 | 外部 benchmark 测试 |
data/rnastralign.512/ |
train/val/test | RNAStrAlign 训练与同分布测试基准 |
4.1 结构去重实验的含义
structdedup098 使用 spot0 train 的全对全 bpRNA-align 相似度。所有 norm_score>=0.98 connected component 仅保留最长的代表序列:10,814 条训练样本变为 7,240 条,移除 3,574 条;validation/test 与原 spot0 完全一致。
实现与证据: 去重 manifest 为 data/bprna-spot0-structdedup098/deduplication_manifest.json:1-13;构建逻辑见 scripts/build_alignment_struct_dedup.py:50-125;全对全对齐分片格式和归一化得分计算见 scripts/run_bprna_align_full.py:132-242。
5. 当前完成的工作与研究状态
以下状态记录于 2026-08-06。运行状态会变化,具体以
outputs/<run>/logs/train.log为准。
5.1 已完成的工程与分析
- 完成了 GB.RNA、RNA-FM、RiNALMo 的统一 encoder 接入;
- 实现离散 Flow Matching 和直接判别式两条训练路线,并由统一入口按
trainer.type分发; - 完成
pair_dim=512的 GB.RNA 直接模型和 Flow Matching 实验配置; - 完成 spot0 train 的全对全 bpRNA-align 分析,并构建
structdedup098结构去重训练集; - 完成判别式 GB.RNA
pair_dim=512 + covariation模型的 bad-case、训练模板迁移、错误 pair、碱基类型、跨度和来源家族分析; - 生成 bad-case contact-map 图和报告,便于检查模板漂移、过配、漏配、stem register shift、长程 partner 错误及 non-canonical pair 偏置。
主要分析入口:
docs/DISCRIMINATIVE_BADCASE_TEMPLATE_TRANSFER_REPORT.md:高结构相似 test/train 对的模板转移与序列差异;docs/DISCRIMINATIVE_BADCASE_ERROR_DIAGNOSIS.md:密度、跨度、阈值和 decoder 诊断;docs/DISCRIMINATIVE_BADCASE_MISPAIR_CONTEXT_ANALYSIS.md:逐 pair 碱基、局部 context 与错误 partner 分析;outputs/badcase_discriminative_gbrna_cov_pair512_template_transfer_analysis/:对齐感知 contact-map 图与 JSON 明细。
5.2 当前运行中的主要消融
| GPU | 配置 | 目的 | 核心区别 |
|---|---|---|---|
cuda:0 |
configs/discriminative_contact_map_gbrna_spot0_structdedup098_cov_pair512_cuda0_400.yaml |
结构去重下的冻结 GB.RNA 对照 | GB.RNA 冻结,BF16,pair_dim=512 |
cuda:1 |
configs/discriminative_contact_map_gbrna_spot0_structdedup098_cov_pair512_cuda1_unfrozen_400.yaml |
结构去重下的序列变体泛化测试 | GB.RNA 全量解冻,gradient checkpointing,FP32,encoder LR 5e-6 |
两组都使用 bprna-spot0-structdedup098/train.parquet,但 validation/test 保持原 bprna-spot0,因此可以隔离“结构去重”和“是否解冻 GB.RNA”的影响。
配置依据: 冻结实验见 configs/discriminative_contact_map_gbrna_spot0_structdedup098_cov_pair512_cuda0_400.yaml:1-101;解冻实验见 configs/discriminative_contact_map_gbrna_spot0_structdedup098_cov_pair512_cuda1_unfrozen_400.yaml:1-106。
5.3 当前已知问题
bad-case 分析显示,模型错误不是单一阈值问题,主要混合了:
- 训练集中没有足够相似结构模板的 coverage/OOD 样本;
- 高结构相似但序列变化后,pair prediction 与正确训练模板脱钩;
- stem register 的局部端点偏移;
- 长程 partner 漏检或错误重连;
- 稀疏结构的过配/少配;
- non-canonical
otherGT pair 在错误样本中富集,而预测更偏 canonical AU/GC/GU。
这些是当前研究假设与消融方向,不应被视作已经解决的功能。
6. 目录结构
symfold/
├── configs/ # 每个实验的 YAML 配置
├── data/ # 训练、评测、去重与相似度分析数据
│ └── README.md # 数据目录详细说明
├── models/ # 本地 RNA encoder 权重,例如 gbrna1.6B/
├── symfold/
│ ├── data/ # Parquet loader、dot-bracket/contact-map、sampler、增强
│ ├── models/ # encoder adapter、direct model、flow model、DiT backbone
│ ├── train.py # 单阶段/多阶段统一入口
│ ├── train_staged.py # 多 config 顺序训练编排器
│ ├── train_supervised_pair.py # 判别式单阶段训练器
│ ├── train_flow_matching.py # Flow Matching 单阶段训练器
│ ├── evaluate_flow_matching.py
│ ├── metrics.py
│ └── visualize.py
├── scripts/ # 数据构建、bpRNA-align 与诊断脚本
│ └── README.md # 脚本用途、保留/归档建议与删除检查清单
├── docs/ # 架构、数据、训练和 bad-case 分析报告
├── outputs/ # 每次运行的日志、checkpoint、曲线和可视化
└── requirements.txt # 当前环境的精确 pip 依赖
统一入口会根据 YAML 中的 trainer.type 调用:
flow_matching→symfold.train_flow_matching;direct_contact_map→symfold.train_supervised_pair。
实现依据: symfold/train.py:1-58。
7. 从零复现
7.1 前置条件
- Linux x86_64;
- Python
3.10; - NVIDIA GPU;
- 对 GPU 训练,使用兼容 PyTorch CUDA
13.0wheel 的驱动; - 本地预训练权重目录,例如
models/gbrna1.6B/; - 本地 Parquet 数据目录
data/。
当前 requirements.txt 固定了实际运行环境版本,包括 torch==2.12.1+cu130、transformers==5.13.0、multimolecule==0.2.0、numpy==2.2.6、pandas==2.3.3 和 pyarrow==24.0.0。
7.2 安装
cd /path/to/symfold
python3.10 -m venv .venv
source .venv/bin/activate
python -m pip install --upgrade pip
python -m pip install -r requirements.txt
快速验证:
python - <<'PY'
import torch, transformers, multimolecule, pandas, pyarrow
print('torch:', torch.__version__)
print('cuda build:', torch.version.cuda)
print('cuda available:', torch.cuda.is_available())
print('transformers:', transformers.__version__)
print('pandas:', pandas.__version__, 'pyarrow:', pyarrow.__version__)
PY
requirements.txt使用 PyTorch CUDA 13.0 的官方 wheel index。没有 GPU 或驱动不兼容时,请按目标平台的 PyTorch 官方安装方式替换 PyTorch 相关行,再安装其余依赖。
7.3 准备本地资源
至少确认:
models/gbrna1.6B/config.json
data/bprna-spot0/train.parquet
data/bprna-spot0/validation.parquet
data/bprna-spot0/test.parquet
配置中的 model.rna_encoder_path 必须指向真实权重目录。当前 GB.RNA 实验配置使用绝对路径 /efs/dannyyan/symfold/models/gbrna1.6B;迁移到新机器时请将其改为本地实际路径。
7.4 运行直接判别式训练
冻结 GB.RNA 的结构去重对照:
python -m symfold.train \
--config configs/discriminative_contact_map_gbrna_spot0_structdedup098_cov_pair512_cuda0_400.yaml
解冻 GB.RNA 的结构去重实验:
python -m symfold.train \
--config configs/discriminative_contact_map_gbrna_spot0_structdedup098_cov_pair512_cuda1_unfrozen_400.yaml
开始前请将 YAML 中的 experiment.device 改为本机可用 GPU。解冻 GB.RNA 显存与计算开销明显更高;当前配置使用 FP32 和 gradient checkpointing,因为长序列 GB.RNA 全量反向曾出现 BF16 non-finite gradient。
配置依据: 解冻设置、独立 encoder 学习率及 FP32 选择见 configs/discriminative_contact_map_gbrna_spot0_structdedup098_cov_pair512_cuda1_unfrozen_400.yaml:35-45,77-94。
7.5 运行 Flow Matching 训练
例如运行 GB.RNA、共变增强、pair_dim=512 的 Flow Matching 配置:
python -m symfold.train_flow_matching \
--config configs/flow_matching_gbrna_spot0_covariation_pair512_unfrozen_400.yaml
也可以走统一入口:
python -m symfold.train \
--config configs/flow_matching_gbrna_spot0_covariation_pair512_unfrozen_400.yaml
后者在没有明确 trainer.type 时,会因配置含 flow_matching 字段而选择 Flow Matching 路线。实现依据: symfold/train.py:25-52。
7.6 多阶段训练
多个 config 可以作为同一个训练实验顺序执行。所有 config 必须使用相同的 trainer.type,例如都使用 flow_matching,或者都使用 direct_contact_map。
python -m symfold.train \
--config /path/to/phase1.yaml /path/to/phase2.yaml
也可以直接调用编排器:
python -m symfold.train_staged \
--config /path/to/phase1.yaml /path/to/phase2.yaml
每个 config 的 train.num_epochs 表示该阶段新增的 epoch 数。比如两个 config 都是 400,实际执行为:
Phase 1: epoch 0–399
Phase 2: epoch 400–799
第二阶段自动从共享 run 目录中的 checkpoints/last.pt 恢复模型、optimizer、scheduler、epoch、global step 和 best metric。多个阶段共用:
outputs/<run>/
├── checkpoints/
├── logs/train.log
├── logs/history.json
├── dashboards/stage_01_*.png
└── dashboards/stage_02_*.png
每个阶段单独生成一个 dashboard:
- loss 只绘制当前阶段,不跨阶段连接;
- Val/Test F1、Precision、Recall 等评估指标按累计 epoch 接续;
- 所有阶段的训练历史写入同一个
history.json,并记录stage字段; - 不在
outputs/顶层额外生成train_*.log; - 不再自动生成
logs/curves/明细目录。
Flow Matching 示例:
python -m symfold.train_flow_matching \
--config configs/flow_phase1.yaml configs/flow_phase2.yaml
判别式示例:
python -m symfold.train_supervised_pair \
--config configs/direct_phase1.yaml configs/direct_phase2.yaml
train.py 是推荐的统一入口;train_staged.py 是多阶段编排实现;两个具体 trainer 负责各自单阶段的模型训练。它们不是四套独立训练逻辑,使用其中一个入口即可。
实现依据: 统一入口见 symfold/train.py:18-66;多阶段编排见 symfold/train_staged.py:57-151;判别式阶段参数见 symfold/train_supervised_pair.py:32-51,170-240;Flow Matching 阶段参数见 symfold/train_flow_matching.py:43-64,81-145。
7.7 续训
单阶段或多阶段都可以通过 --resume-run-dir 继续写入已有 run。多阶段续训时,第二阶段仍会自动从该目录的 checkpoints/last.pt 接续。
python -m symfold.train \
--config /path/to/continue.yaml \
--resume-run-dir /path/to/existing_run
实现依据: checkpoint 保存与恢复字段见 symfold/utils.py:155-182。
8. 配置指南
8.1 训练范式
trainer:
type: direct_contact_map # 或 flow_matching
直接路线还须指定:
model:
type: discriminative_contact_map # 或 legacy supervised_pair
8.2 常用 encoder 配置
model:
rna_encoder_path: /absolute/path/to/gbrna1.6B
rna_encoder_freeze: false
rna_encoder_num_attn_layers: 4
rna_encoder_gradient_checkpointing: true
当设置 rna_encoder_lr 时,optimizer 会把 encoder 与下游模块拆成不同学习率 param group。
optim:
lr: 2.0e-4
rna_encoder_lr: 5.0e-6
实现依据: optimizer 分组见 symfold/utils.py:100-123;warmup/cosine scheduler 见 symfold/utils.py:126-152。
8.3 输出目录
每次训练默认写入:
outputs/<experiment.name>_<CST timestamp>/
├── logs/
│ ├── train.log
│ ├── history.json
│ └── events.out.tfevents.*
├── checkpoints/
│ ├── best.pt
│ └── last.pt
├── visualizations/
├── training_dashboard.png # 单阶段兼容输出
└── dashboards/ # 多阶段时每个 config 一个 dashboard
├── stage_01_*.png
└── stage_02_*.png
用户当前约定是不创建运行目录之外的额外 train_*.log;请以 <run_dir>/logs/train.log 为唯一训练日志。
9. 结果解读与限制
- F1 是严格 exact-pair 指标:stem 两端只偏移 1–2 nt 仍会同时计为 FP 与 FN。
- 当前 greedy decoder 每个碱基最多一个 partner,但没有 non-crossing 约束;因此它可能放大 logits 中错误 partner 的排序。
- 直接判别式主模型的 sequence pair feature 是对称平均,不是完整的四路 pair interaction;这正是当前 sequence-variant 泛化研究的重点。
structdedup098只去除 train 内高相似结构,并不使 test 自动成为与 train 完全不相似的集合;应结合 bad-case、cluster 和外部集结果解释。data/README.md、启动时使用的 YAML 副本,以及运行目录中的logs/train.log才是复现实验条件的最终事实来源;当前训练代码不会自动复制 YAML 到运行目录。
10. 相关文档
data/README.md:各数据目录、规模、构建方式与实验边界;scripts/README.md:20 个脚本的用途、输入输出、保留/归档建议与删除前检查清单;docs/DISCRIMINATIVE_BADCASE_TEMPLATE_TRANSFER_REPORT.md:结构近邻 train/test 模板转移、序列差异与 bad case;docs/DISCRIMINATIVE_BADCASE_ERROR_DIAGNOSIS.md:判别式模型错误模式;docs/DISCRIMINATIVE_BADCASE_MISPAIR_CONTEXT_ANALYSIS.md:逐 pair 碱基与局部上下文分析;configs/:每个可复现实验的参数、数据路径和设备设置。