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 序列,例如:
输出是其二级结构的配对集合或等价的 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. 总体架构
1RNA sequence
2 │
3 ├── RNA encoder(GB.RNA / RNA-FM / RiNALMo)
4 │ ├── per-nucleotide hidden states: [B, L, H]
5 │ └── selected attention maps: [B, A, L, L]
6 │
7 ├── 直接判别式路线
8 │ hidden/attention fusion → [B, L, L] contact logits
9 │
10 └── 离散 Flow Matching 路线
11 pair condition + noisy x_t + time t
12 → DiT-style pair backbone → [B, L, L] contact logits
13 → τ-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],一次输出:
1logits: [B, L, L]
2mask: [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×3 pair 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:61-99;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
other GT pair 在错误样本中富集,而预测更偏 canonical AU/GC/GU。
这些是当前研究假设与消融方向,不应被视作已经解决的功能。
6. 目录结构
1symfold/
2├── configs/ # 每个实验的 YAML 配置
3├── data/ # 训练、评测、去重与相似度分析数据
4│ └── README.md # 数据目录详细说明
5├── models/ # 本地 RNA encoder 权重,例如 gbrna1.6B/
6├── symfold/
7│ ├── data/ # Parquet loader、dot-bracket/contact-map、sampler、增强
8│ ├── models/ # encoder adapter、direct model、flow model、DiT backbone
9│ ├── train.py # 单阶段/多阶段统一入口
10│ ├── train_staged.py # 多 config 顺序训练编排器
11│ ├── train_supervised_pair.py # 判别式单阶段训练器
12│ ├── train_flow_matching.py # Flow Matching 单阶段训练器
13│ ├── evaluate_flow_matching.py
14│ ├── metrics.py
15│ └── visualize.py
16├── scripts/ # 数据构建、bpRNA-align 与诊断脚本
17│ └── README.md # 脚本用途、保留/归档建议与删除检查清单
18├── docs/ # 架构、数据、训练和 bad-case 分析报告
19├── outputs/ # 每次运行的日志、checkpoint、曲线和可视化
20└── 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.0 wheel 的驱动;
- 本地预训练权重目录,例如
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 安装
1cd /path/to/symfold
2python3.10 -m venv .venv
3source .venv/bin/activate
4python -m pip install --upgrade pip
5python -m pip install -r requirements.txt
快速验证:
1python - <<'PY'
2import torch, transformers, multimolecule, pandas, pyarrow
3print('torch:', torch.__version__)
4print('cuda build:', torch.version.cuda)
5print('cuda available:', torch.cuda.is_available())
6print('transformers:', transformers.__version__)
7print('pandas:', pandas.__version__, 'pyarrow:', pyarrow.__version__)
8PY
requirements.txt 使用 PyTorch CUDA 13.0 的官方 wheel index。没有 GPU 或驱动不兼容时,请按目标平台的 PyTorch 官方安装方式替换 PyTorch 相关行,再安装其余依赖。
7.3 准备本地资源
至少确认:
1models/gbrna1.6B/config.json
2data/bprna-spot0/train.parquet
3data/bprna-spot0/validation.parquet
4data/bprna-spot0/test.parquet
配置中的 model.rna_encoder_path 必须指向真实权重目录。当前 GB.RNA 实验配置使用绝对路径 /efs/dannyyan/symfold/models/gbrna1.6B;迁移到新机器时请将其改为本地实际路径。
7.4 运行直接判别式训练
冻结 GB.RNA 的结构去重对照:
1python -m symfold.train \
2 --config configs/discriminative_contact_map_gbrna_spot0_structdedup098_cov_pair512_cuda0_400.yaml
解冻 GB.RNA 的结构去重实验:
1python -m symfold.train \
2 --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 配置:
1python -m symfold.train_flow_matching \
2 --config configs/flow_matching_gbrna_spot0_covariation_pair512_unfrozen_400.yaml
也可以走统一入口:
1python -m symfold.train \
2 --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。
1python -m symfold.train \
2 --config /path/to/phase1.yaml /path/to/phase2.yaml
也可以直接调用编排器:
1python -m symfold.train_staged \
2 --config /path/to/phase1.yaml /path/to/phase2.yaml
每个 config 的 train.num_epochs 表示该阶段新增的 epoch 数。比如两个 config 都是 400,实际执行为:
1Phase 1: epoch 0–399
2Phase 2: epoch 400–799
第二阶段自动从共享 run 目录中的 checkpoints/last.pt 恢复模型、optimizer、scheduler、epoch、global step 和 best metric。多个阶段共用:
1outputs/<run>/
2├── checkpoints/
3├── logs/train.log
4├── logs/history.json
5├── dashboards/stage_01_*.png
6└── dashboards/stage_02_*.png
每个阶段单独生成一个 dashboard:
- loss 只绘制当前阶段,不跨阶段连接;
- Val/Test F1、Precision、Recall 等评估指标按累计 epoch 接续;
- 所有阶段的训练历史写入同一个
history.json,并记录 stage 字段;
- 不在
outputs/ 顶层额外生成 train_*.log;
- 不再自动生成
logs/curves/ 明细目录。
Flow Matching 示例:
1python -m symfold.train_flow_matching \
2 --config configs/flow_phase1.yaml configs/flow_phase2.yaml
判别式示例:
1python -m symfold.train_supervised_pair \
2 --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 接续。
1python -m symfold.train \
2 --config /path/to/continue.yaml \
3 --resume-run-dir /path/to/existing_run
实现依据: checkpoint 保存与恢复字段见 symfold/utils.py:155-182。
8. 配置指南
8.1 训练范式
1trainer:
2 type: direct_contact_map # 或 flow_matching
直接路线还须指定:
1model:
2 type: discriminative_contact_map # 或 legacy supervised_pair
8.2 常用 encoder 配置
1model:
2 rna_encoder_path: /absolute/path/to/gbrna1.6B
3 rna_encoder_freeze: false
4 rna_encoder_num_attn_layers: 4
5 rna_encoder_gradient_checkpointing: true
当设置 rna_encoder_lr 时,optimizer 会把 encoder 与下游模块拆成不同学习率 param group。
1optim:
2 lr: 2.0e-4
3 rna_encoder_lr: 5.0e-6
实现依据: optimizer 分组见 symfold/utils.py:100-123;warmup/cosine scheduler 见 symfold/utils.py:126-152。
8.3 输出目录
每次训练默认写入:
1outputs/<experiment.name>_<CST timestamp>/
2├── logs/
3│ ├── train.log
4│ ├── history.json
5│ └── events.out.tfevents.*
6├── checkpoints/
7│ ├── best.pt
8│ └── last.pt
9├── visualizations/
10├── training_dashboard.png # 单阶段兼容输出
11└── dashboards/ # 多阶段时每个 config 一个 dashboard
12 ├── stage_01_*.png
13 └── 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. AI/研究助手实验操作手册
本节是给 AI 编程助手和研究协作者使用的实验 SOP。除非用户明确要求,助手只创建新配置、运行目录和分析结果,不覆盖正在使用的 YAML、checkpoint 或日志。
10.1 实验启动前的安全规则
- 先查看
git status --short、目标 YAML 和 outputs/*/logs/train.log,确认当前是否有未提交改动或正在运行的实验。
- 不要直接修改已有实验配置来做新实验;复制配置到
configs/dis/ablations/、configs/flowmatching/ 或其他明确的实验目录,并修改唯一的 experiment.name。
- 不要复用正在写入的
experiment.name 或 --resume-run-dir,除非用户明确要求续训。
- 新实验必须明确记录:数据 split、encoder checkpoint、是否冻结、设备、随机种子、训练 epoch、阈值策略和 checkpoint 路径。
- 训练日志以
<run_dir>/logs/train.log 为准;不要在 outputs/ 顶层额外创建平行 train_*.log。
10.2 实验类型总览
| 实验类型 | YAML 关键字段 | 推荐入口 | 适用目的 |
|---|
| 判别式 contact-map | trainer.type: direct_contact_map、model.type: discriminative_contact_map | python -m symfold.train | 当前主路线,一次 forward 输出 pair logits |
| 判别式 legacy baseline | trainer.type: direct_contact_map、model.type: supervised_pair | python -m symfold.train | 与旧版 supervised_pair 保持可比 |
| Fusion ablation | model.fusion_mode | python -m symfold.train | 比较 1d_only、2d_only、simple_fusion、cross_gated |
| Multi-task | model.multitask.enabled: true | python -m symfold.train | 联合学习 contact、pairedness 和 pair type |
| 离散 Flow Matching | trainer.type: flow_matching | python -m symfold.train 或 train_flow_matching | 研究生成式 contact-map 预测和 CTMC 采样 |
| 多阶段训练 | 多个同类型 YAML | python -m symfold.train --config phase1.yaml phase2.yaml | 冻结→解冻、不同学习率或续训阶段 |
统一入口按照 trainer.type 分发到两个 trainer,代码见 symfold/train.py:18-66;判别式模型类型分发见 symfold/train_supervised_pair.py:50-58。
10.3 推荐配置骨架
一个新的判别式实验至少应包含以下部分:
1trainer:
2 type: direct_contact_map
3
4experiment:
5 name: unique_experiment_name
6 seed: 42
7 output_dir: outputs
8 device: cuda:0
9
10data:
11 max_length: 512
12 train_datasets:
13 - name: train
14 path: data/bprna-spot0/train.parquet
15 val_datasets:
16 - name: val
17 path: data/bprna-spot0/validation.parquet
18 test_datasets:
19 - name: test
20 path: data/bprna-spot0/test.parquet
21
22model:
23 type: discriminative_contact_map
24 rna_encoder_path: /absolute/path/to/encoder
25 rna_encoder_freeze: true
26 rna_encoder_num_attn_layers: 4
27 pair_dim: 512
28
29loss:
30 pos_weight: 200.0
31 focal_gamma: 1.0
32 dice_weight: 0.05
33 degree_weight: 0.02
34
35optim:
36 lr: 2.0e-5
37 rna_encoder_lr: 5.0e-6
38 weight_decay: 0.01
39 scheduler: none
40
41train:
42 num_epochs: 400
43 eval_every_n_epochs: 20
44 test_every_n_epochs: 40
45 amp: true
46 amp_dtype: bfloat16
47
48eval:
49 threshold: 0.5
50 min_sequence_separation: 4
51 greedy_projection: true
路径和参数的实际解析位置如下:
- 数据读取、padding、length bucket:
symfold/data/datasets.py:24-108,163-208;
- encoder 参数:
symfold/models/rnafm_encoder.py:28-118,206-215;
- optimizer 和 scheduler:
symfold/utils.py:100-152;
- checkpoint 和续训:
symfold/utils.py:155-183;
- 判别式训练循环:
symfold/train_supervised_pair.py:354-534;
- Flow Matching 训练循环:
symfold/train_flow_matching.py:43-64,81-145,282-326。
10.4 现有配置模板如何选择
| 研究问题 | 配置模板 |
|---|
| GB.RNA 冻结 baseline | configs/discriminative_contact_map_gbrna_spot0_covariation_pair512_cuda0_bf16_frozen_400.yaml |
| GB.RNA 解冻 baseline | configs/discriminative_contact_map_gbrna_spot0_covariation_pair512_cuda0_bf16_unfrozen_800.yaml |
| 仅使用 1D hidden | configs/dis/ablations/1d_only_frozen_400.yaml 或 1d_only_unfrozen_800.yaml |
| 仅使用 2D attention | configs/dis/ablations/2d_only_frozen_400.yaml 或 2d_only_unfrozen_800.yaml |
| 简单拼接 fusion | configs/dis/ablations/simple_fusion_cuda1_moderate2x_unfrozen_400.yaml |
| Cross-Gated fusion | configs/dis/ablations/cross_gated_frozen_400.yaml 或 cross_gated_unfrozen_800.yaml |
| Multi-task 完整实验 | configs/dis/ablations/multitask_pair_type_cuda1_moderate2x_unfrozen_400.yaml |
| Flow Matching | configs/flowmatching/flow_matching_gbrna_spot0_covariation_pair512_cuda1_bf16_frozen_400.yaml |
discriminative_contact_map 当前支持的 fusion_mode 为 1d_only、2d_only、simple_fusion 和 cross_gated,模型组装见 symfold/models/discriminative_contact_map_model.py:108-229。
10.5 如何开启 Multi-task Learning
Multi-task 默认关闭。只需在 model 下加入:
1model:
2 multitask:
3 enabled: true
4 pairedness:
5 enabled: true
6 pair_type:
7 enabled: true
8 num_classes: 4
9 class_names: [AU, GC, GU, other]
并在 loss 下加入:
1loss:
2 multitask:
3 pairedness_weight: 0.1
4 pairedness_pos_weight: 1.0
5 pairedness_focal_gamma: 0.0
6 pair_type_weight: 0.2
7 pair_type_class_weights: [1.0, 1.0, 1.0, 2.0]
当前实现的任务定义为:
contact:位置对是否配对,仍是主任务和 checkpoint 选择指标;
pairedness:每个 nucleotide 是否参与任意配对;
pair_type:真实配对属于 AU、GC、GU 或 other。
辅助标签由 dot-bracket 和序列自动生成,见 symfold/data/dotbracket.py:61-81、symfold/data/datasets.py:57-78;联合 loss 和辅助评估见 symfold/train_supervised_pair.py:102-204。Multi-task 模型只在显式调用 return_auxiliary=True 时返回辅助 logits,默认 forward 接口保持兼容,见 symfold/models/discriminative_contact_map_model.py:189-229。
启动完整 Multi-task 实验:
1python -m symfold.train \
2 --config configs/dis/ablations/multitask_pair_type_cuda1_moderate2x_unfrozen_400.yaml
建议的消融顺序:
contact-only:不设置 model.multitask;
contact + pairedness:只打开 pairedness;
contact + pair_type:只打开 pair_type;
contact + pairedness + pair_type:使用完整配置。
10.6 如何运行单阶段、多阶段和续训
进入项目和环境:
1cd /efs/dannyyan/symfold
2source /efs/miniconda3/envs/symfold/bin/activate
单阶段:
python -m symfold.train --config configs/<experiment>.yaml
多阶段时,所有 YAML 必须使用相同的 trainer.type。每个阶段的 train.num_epochs 表示该阶段新增的 epoch 数:
1python -m symfold.train \
2 --config configs/phase1_frozen.yaml configs/phase2_unfrozen.yaml
如果阶段改变可训练参数或 optimizer 参数,必须明确决定是否重置 optimizer:
1python -m symfold.train_supervised_pair \
2 --config configs/phase2.yaml \
3 --stage-reset-optimizer \
4 --stage-resume-from outputs/<run>/checkpoints/last.pt
普通续训使用:
1python -m symfold.train \
2 --config configs/continue.yaml \
3 --resume-run-dir outputs/<existing_run>
多阶段编排、阶段 dashboard 和 checkpoint 接续逻辑见 symfold/train_staged.py:1-151;命令行参数定义见 symfold/train_supervised_pair.py:32-47 和 symfold/train_flow_matching.py:43-60。
10.7 结果检查 SOP
每次实验完成后,AI 应按以下顺序检查:
outputs/<run>/logs/train.log:确认配置意图、设备、数据规模、loss、val/test 指标和是否早停;
outputs/<run>/logs/history.json:检查完整训练曲线和阶段字段;
outputs/<run>/checkpoints/best.pt:正式评估优先使用 validation F1 最优 checkpoint;
outputs/<run>/checkpoints/last.pt:续训使用最后状态;
outputs/<run>/dashboards/ 或 training_dashboard.png:检查 loss、F1、precision、recall 是否异常;
- 使用 validation 选择 threshold,再在 test 上只评估一次;不要用 test 直接调 threshold。
判别式评估和 best checkpoint 保存见 symfold/train_supervised_pair.py:207-266,470-534;threshold scan 见 scripts/scan_discriminative_threshold.py:1-100。
Multi-task 实验除 contact F1/MCC 外,还应查看日志中的 pairedness_f1 和 pair_type_accuracy,但不能只凭辅助任务指标选择 checkpoint。
10.8 AI 修改或新增实验时的最小交付内容
每个新实验至少应包含:
- 一个独立 YAML;
- 唯一
experiment.name;
- 明确的 train/val/test 路径;
- 训练目标和主要改变点;
- 启动命令;
- 预期输出目录;
- 与 baseline 的对照指标;
- 若修改代码,文档中注明文件路径和行号。
不要把“修改了配置但没有启动训练”描述成实验结果;不要把 test threshold 选择结果当作正式泛化结果;不要在没有数据或日志证据时声称模型性能提升。
11. 相关文档