细胞扰动预测框架 · 模块化架构设计
将 GEARS、线性基线、scGPT 等算法拆解为可独立组合的数据层、知识层、预测器层与评测层
背景与目标
Motivation & Goals · 为什么做模块化1.1 从单点算法到可组合框架
此前的代码以 GEARS 为中心组织:训练脚本 train_gears.py、
评测脚本 evaluate_gears_run.py 都对 GEARS 的 PertData / PyG 图格式有硬依赖。
要接入第二个算法(如 scGPT),需要复制整条流程并手工改写,数据格式、评测指标难以共享。
本次重构的目标是:任意算法都能通过声明所需知识组件来组装运行, 评测数字可以横向比较,agent 循环可以自主组合新方法,而无需人工介入每一条链路。
1.2 三条核心设计原则
score_predictor(LinearPredictor, dataset, "discovery")
产出的 deg_pearson_macro 与直接公式计算逐位一致(误差 < 1e-6),
已在 tests/test_scorer.py::test_deg_pearson_matches_direct_computation 验证。
架构总览
Architecture Overview · 四层结构2.1 四层横向流水线
整个框架由四层组成(如图 1 所示)。数据层提供统一的 numpy 容器, 知识组件层封装外部知识(GO 图、预训练权重等), 预测器层通过 Protocol 声明需要哪些知识并实现推断, 评测与训练层对任意预测器打出格式兼容的 metrics dict。
2.2 新增文件布局
| 目录 | 文件 | 职责 |
|---|---|---|
data/ | dataset.py | PerturbationDataset dataclass,save/load |
loaders.py | from_carved_pkl() / from_h5ad() | |
knowledge/ | base.py | KnowledgeComponent Protocol |
go_graph.py | GOGraphComponent,包装 GEARS PertData | |
predictors/ | base.py | PerturbationPredictor Protocol |
wrappers.py | 4 个无训练基线包装器 | |
gears_predictor.py | GEARSPredictor 完整链路 | |
evaluation/ | scorer.py | score_predictor() 通用打分 |
training/ | harness.py | train_with_harness() 通用训练外壳 |
train_gears.py、evaluate_gears_run.py
及 src/itermix_brain/ 下所有文件保持原样,新旧两套并行运行,等新链路在服务器数字验证后再归档旧脚本。
数据层
Data Layer · PerturbationDataset3.1 PerturbationDataset 字段设计
所有字段均为 numpy ndarray 或 Python 原生类型,与 PyTorch / PyG / GEARS 完全解耦。
如图 2 所示,ctrl_mean 在 __post_init__ 自动计算,
所有预测器共用同一份均值基线,消除对齐误差。
| 字段 | 形状 / 类型 | 含义 | 来源 |
|---|---|---|---|
ctrl_expression | (n_ctrl, n_genes) float32 | 对照组细胞表达矩阵 | GEARS ctrl_expression / h5ad |
perturbed_cells | dict[str, ndarray] | 训练可见扰动细胞(雕刻后) | dataset_processed 训练条件 |
carved_heldout | dict[str, ndarray] | 单/组合条件留出评测细胞 | carved_heldout_graphs.pkl |
unseen_cells | dict[str, ndarray] | 条件级完全留出(validation unseen) | dataset_processed val 条件 |
ctrl_mean | (n_genes,) float32 | 对照均值,自动计算 | __post_init__ |
split_manifest | SplitManifest | Tier 成员关系 | norman2019-split-v1.json |
3.2 两个加载入口
from_carved_pkl() — 已有 GEARS 产物
从 carved_heldout_graphs.pkl 提取 .y 字段(真实表达量),
丢弃 .x(GEARS 专属扰动编码)和 .pert,
配合调用方传入的 ctrl_expression / perturbed_cells 构建 dataset。
适用于所有已完成的 GEARS 训练跑。
from_h5ad() — 原始 AnnData
直接读取 .h5ad 文件,依据 manifest 的 cell ID 列表对细胞进行雕刻分配:
discovery/validation 留出细胞 → carved_heldout,unseen 条件全部 → unseen_cells,
其余 → perturbed_cells。
适用于 Replogle 2022 等新数据集(无需先跑 GEARS)。
from_h5ad() 依赖 manifest 里的 cell ID 列表做雕刻。
若 manifest 用 build_split() 不带 cell_ids_by_condition 构建,
则 carved_heldout 为空,需先用实际 cell ID 重新生成 v2 manifest。
知识组件层
Knowledge Layer · 可插拔外部知识4.1 KnowledgeComponent Protocol
KnowledgeComponent 是结构子类型(runtime_checkable Protocol),
只要求实现 component_type: ClassVar[str] 这一类变量。
预测器通过 required_knowledge = ["go_graph", ...] 声明需要哪些组件,
训练外壳在调用 fit() 之前完成校验和传入。
不需要知识的预测器声明空列表,永远不会接触到此模块。
4.2 GOGraphComponent — 当前唯一实现
包装已加载的 GEARS PertData 对象(包含 GO / 共表达图张量),
通过 from_gears_pertdata(pert_data) 一行创建,
GEARSPredictor.fit() 从中取出 pert_data 运行原有训练循环。
| 组件类型 | component_type | 状态 | 适配算法 |
|---|---|---|---|
GOGraphComponent | "go_graph" | ✅ 已实现 | GEARSPredictor |
PretrainedEmbeddingComponent | "pretrained_embeddings" | 🔲 规划中 | ScGPTPredictor |
CoexpressionGraphComponent | "coexp_graph" | 🔲 规划中 | 图神经网络变体 |
knowledge/__init__.py 注册,无需修改任何现有代码。
预测器层
Predictor Layer · Protocol + 实现5.1 PerturbationPredictor Protocol
四个方法组成最小接口:fit(dataset, knowledge) 训练,
predict_delta(condition) → (n_genes,) float32 推断,
save(path) 持久化,load(path) 恢复。
scorer 和 harness 只调用这四个方法,与具体模型完全解耦。
5.2 四个无训练基线包装器
均委托给 src/itermix_brain/micro/baselines/simple_baselines.py 的已测实现,
包装层只做接口适配(fit_X_conditions() 拆包 + float32 转换),
不重复实现任何统计逻辑。
| 类名 | name | required_knowledge | 特点 |
|---|---|---|---|
MeanPredictor | mean | [] | 预测全零 delta(最弱基线) |
PerturbedMeanPredictor | perturbed_mean | [] | 所有条件预测同一均值偏移,高分低区分 |
LinearPredictor | linear | [] | Ridge 回归,加法可组合基线 |
NearestPerturbationPredictor | nearest_perturbation | [] | 迁移最近共表达扰动的 delta |
5.3 GEARSPredictor — 完整训练链路
封装原 train_gears.py 的完整流程(apply_heldout_filter →
carve_inner_val → get_dataloader → swap_val_loader →
epoch hook → gears.train() → save_model),
产出物格式与现有 evaluate_gears_run.py 完全兼容,
from_run_dir(path, go) 可从已有跑目录恢复。
predict_delta() 目前要求该条件至少有一组
inference graphs(来自训练或 carved),完全未见过的条件(certification tier)暂不支持
zero-shot 推断,留待 scGPT 阶段一并解决。
评测与训练
Evaluation & Training · 通用打分与训练外壳6.1 score_predictor() — 通用打分器
接受任意满足 Protocol 的预测器,从 dataset 的 carved_heldout / unseen_cells 取真实表达,
计算完整诊断指标套件,输出与旧 evaluate_gears_run.py metrics JSON 格式逐字段兼容,
现有报告模板、分析脚本无需修改即可复用。
| 指标 | 计算方式 | 用途 |
|---|---|---|
deg_pearson_macro | 预测 delta vs 观测 delta Pearson,条件级宏平均 | 主要指标 |
deg_pearson_ci | Bootstrap 2000 次,95% CI | 单跑可信度 |
direction_accuracy_top_deg | Top-20 DEG 符号一致率 | 方向准确性 |
perturbation_discrimination_macro | 归一化排名,0.5 = 随机 | 模式坍塌检测 |
split_half_ceiling_macro | 观测细胞 split-half Pearson | 数据上限参考 |
expression_mse | 细胞级 MSE | 补充指标 |
6.2 train_with_harness() — 通用训练外壳
负责目录创建、train_config.json(训练前写入,崩溃可审计)、
knowledge 完整性校验、predictor.fit() 调用、
以及成功后写入 train_done.json 哨兵。
inner-val / checkpoint 逻辑由各预测器内部处理,harness 只管外层生命周期。
6.3 验收数字(本地 CI)
全套 194 个测试通过,其中三个新测试文件共 70 个 case 覆盖三层核心模块:
| 测试文件 | case 数 | 关键验收点 |
|---|---|---|
tests/test_dataset.py | 30 | ctrl_mean 精度、save/load 圆回 +条件名含 + 的 npz 转义 |
tests/test_predictors.py | 26 | Protocol 合规、predict_delta 形状/dtype、save/load 等值 |
tests/test_scorer.py | 14 | deg_pearson 与直接公式逐位一致(误差 <1e-6) |
from_carved_pkl() 加载 baseline-gears/seed0 的产物,
用 LinearPredictor + score_predictor 计算 discovery tier,
与 metrics_discovery_linear.json 的已知数字进行逐行比对,
差值应 <1e-4(浮点累积误差上界)。
扩展路径
Extension Roadmap · 接入新算法7.1 接入 scGPT 只需三步
新建 PretrainedEmbeddingComponent
在 knowledge/pretrained_emb.py 实现:
component_type = "pretrained_embeddings",
存储 scGPT 模型权重路径和基因 embedding 矩阵 (n_genes, d_embed)。
新建 ScGPTPredictor
在 predictors/scgpt_predictor.py 实现,声明
required_knowledge = ["pretrained_embeddings"],
fit() 在 dataset.perturbed_cells 上微调 scGPT,
predict_delta() 调用 scGPT forward 返回 (n_genes,)。
用 train_with_harness 启动 + score_predictor 评测
接口不变,一行启动:
train_with_harness(ScGPTPredictor(), dataset, {"pretrained_embeddings": emb}, out_dir=...)
score_predictor(predictor, dataset, "discovery")
输出即可与 GEARS/linear baseline 数字直接比较。
7.2 下一步工作清单
predict_delta() 需要条件有
inference graphs(training 或 carved),不支持完全未见基因的 zero-shot 推断。
这个限制在接入 scGPT 时会自然消除(scGPT 用 embedding 而非图做推断),
届时再回头升级 GEARS 路径。