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) 矩阵中每一对核苷酸是否配对。

项目同时维护两条训练路线:

  1. 直接判别式预测:一次前向直接输出 contact logits,适合快速实验和当前主要消融;
  2. 离散 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]

模型流程:

  1. encoder hidden 经 LayerNorm + Linear 投影到 pair_dim
  2. 对任意位置对 ((i,j)),使用对称平均 ((h_i+h_j)/2) 构造 sequence pair feature;
  3. encoder attention 用 1×1 Conv 投影为 pair feature;
  4. 通过双向 gate 与 FiLM 调制融合两条路径;
  5. 用 interaction MLP 和可选的单层 depthwise 3×3 pair smoother 建模局部一致性;
  6. 输出对称 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:

  1. 排除 padding 和小于 min_sequence_separation 的 pair;
  2. 取概率不低于 threshold 的候选边;
  3. 按概率降序;
  4. 每个碱基最多保留一个 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 other GT 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_matchingsymfold.train_flow_matching
  • direct_contact_mapsymfold.train_supervised_pair

实现依据: symfold/train.py:1-58


7. 从零复现

7.1 前置条件

  • Linux x86_64;
  • Python 3.10
  • NVIDIA GPU;
  • 对 GPU 训练,使用兼容 PyTorch CUDA 13.0 wheel 的驱动;
  • 本地预训练权重目录,例如 models/gbrna1.6B/
  • 本地 Parquet 数据目录 data/

当前 requirements.txt 固定了实际运行环境版本,包括 torch==2.12.1+cu130transformers==5.13.0multimolecule==0.2.0numpy==2.2.6pandas==2.3.3pyarrow==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. 结果解读与限制

  1. F1 是严格 exact-pair 指标:stem 两端只偏移 1–2 nt 仍会同时计为 FP 与 FN。
  2. 当前 greedy decoder 每个碱基最多一个 partner,但没有 non-crossing 约束;因此它可能放大 logits 中错误 partner 的排序。
  3. 直接判别式主模型的 sequence pair feature 是对称平均,不是完整的四路 pair interaction;这正是当前 sequence-variant 泛化研究的重点。
  4. structdedup098 只去除 train 内高相似结构,并不使 test 自动成为与 train 完全不相似的集合;应结合 bad-case、cluster 和外部集结果解释。
  5. data/README.md、启动时使用的 YAML 副本,以及运行目录中的 logs/train.log 才是复现实验条件的最终事实来源;当前训练代码不会自动复制 YAML 到运行目录。

10. 相关文档

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support