细胞扰动预测框架 · 模块化架构设计

Perturbation Prediction Framework · Modular Component Architecture
将 GEARS、线性基线、scGPT 等算法拆解为可独立组合的数据层、知识层、预测器层与评测层
🧩 组件驱动 🔁 旧代码并行兼容 📐 Protocol 接口 ✅ 194 tests passed
项目: DeepCell Brain · iterMix 更新: 2026-08-13 状态: Steps 1–7 全部完成
1

背景与目标

Motivation & Goals · 为什么做模块化

1.1 从单点算法到可组合框架

此前的代码以 GEARS 为中心组织:训练脚本 train_gears.py、 评测脚本 evaluate_gears_run.py 都对 GEARS 的 PertData / PyG 图格式有硬依赖。 要接入第二个算法(如 scGPT),需要复制整条流程并手工改写,数据格式、评测指标难以共享。

本次重构的目标是:任意算法都能通过声明所需知识组件来组装运行, 评测数字可以横向比较,agent 循环可以自主组合新方法,而无需人工介入每一条链路。

1.2 三条核心设计原则

框架无关的数据容器 可插拔知识组件 Protocol 驱动的预测器接口 旧代码并行保留 不重写已验证的统计逻辑 不提前耦合尚未使用的功能
验收标准:score_predictor(LinearPredictor, dataset, "discovery") 产出的 deg_pearson_macro 与直接公式计算逐位一致(误差 < 1e-6), 已在 tests/test_scorer.py::test_deg_pearson_matches_direct_computation 验证。
2

架构总览

Architecture Overview · 四层结构

2.1 四层横向流水线

整个框架由四层组成(如图 1 所示)。数据层提供统一的 numpy 容器, 知识组件层封装外部知识(GO 图、预训练权重等), 预测器层通过 Protocol 声明需要哪些知识并实现推断, 评测与训练层对任意预测器打出格式兼容的 metrics dict。

扰动预测框架 · 四层模块化架构 Perturbation Prediction Framework · Four-Layer Modular Architecture 1 数据层 Data Layer PerturbationDataset ctrl_expression carved_heldout unseen_cells ctrl_mean split_manifest 2 知识组件层 Knowledge Layer KnowledgeComponent GOGraphComponent PretrainedEmbedding (可插拔) required_knowledge 按需声明 3 预测器层 Predictor Layer PerturbationPredictor MeanPredictor LinearPredictor PerturbedMeanPredictor NearestPerturbation GEARSPredictor 4 评测 & 训练 Evaluation & Training score_predictor() train_with_harness() metrics dict (与旧格式兼容) train_done.json 哨兵机制 通用数据容器 可插拔知识组件 Protocol 驱动 数字兼容验证
图 1 · 四层模块化架构 — 数据层 → 知识组件层 → 预测器层 → 评测&训练层

2.2 新增文件布局

目录文件职责
data/dataset.pyPerturbationDataset dataclass,save/load
loaders.pyfrom_carved_pkl() / from_h5ad()
knowledge/base.pyKnowledgeComponent Protocol
go_graph.pyGOGraphComponent,包装 GEARS PertData
predictors/base.pyPerturbationPredictor Protocol
wrappers.py4 个无训练基线包装器
gears_predictor.pyGEARSPredictor 完整链路
evaluation/scorer.pyscore_predictor() 通用打分
training/harness.pytrain_with_harness() 通用训练外壳
兼容性: 旧文件 train_gears.pyevaluate_gears_run.pysrc/itermix_brain/ 下所有文件保持原样,新旧两套并行运行,等新链路在服务器数字验证后再归档旧脚本。
3

数据层

Data Layer · PerturbationDataset

3.1 PerturbationDataset 字段设计

所有字段均为 numpy ndarray 或 Python 原生类型,与 PyTorch / PyG / GEARS 完全解耦。 如图 2 所示,ctrl_mean__post_init__ 自动计算, 所有预测器共用同一份均值基线,消除对齐误差。

PerturbationDataset · 内部字段设计 Universal data container — all fields are plain numpy + Python native types from_carved_pkl() GEARS 已有 pkl + PertData from_h5ad() AnnData 原始 h5ad 文件 D PerturbationDataset dataclass · framework-agnostic · save / load via .npy + .npz + JSON ctrl_expression (n_ctrl, n_genes) float32 对照组细胞表达矩阵 perturbed_cells dict[str, ndarray] 训练可见扰动细胞 carved_heldout dict[str, ndarray] 雕刻留出评测细胞 unseen_cells dict[str, ndarray] 条件级完全留出细胞 ctrl_mean (n_genes,) float32 __post_init__ 自动计算 split_manifest SplitManifest Tier 成员关系 conditions_for_tier(tier) heldout_cells_for_condition(cond) fit_X_conditions() save(path) / load(path)
图 2 · PerturbationDataset 内部字段与两个加载入口
字段形状 / 类型含义来源
ctrl_expression(n_ctrl, n_genes) float32对照组细胞表达矩阵GEARS ctrl_expression / h5ad
perturbed_cellsdict[str, ndarray]训练可见扰动细胞(雕刻后)dataset_processed 训练条件
carved_heldoutdict[str, ndarray]单/组合条件留出评测细胞carved_heldout_graphs.pkl
unseen_cellsdict[str, ndarray]条件级完全留出(validation unseen)dataset_processed val 条件
ctrl_mean(n_genes,) float32对照均值,自动计算__post_init__
split_manifestSplitManifestTier 成员关系norman2019-split-v1.json

3.2 两个加载入口

注意: from_h5ad() 依赖 manifest 里的 cell ID 列表做雕刻。 若 manifest 用 build_split() 不带 cell_ids_by_condition 构建, 则 carved_heldout 为空,需先用实际 cell ID 重新生成 v2 manifest。
4

知识组件层

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"🔲 规划中图神经网络变体
扩展方式: 新增知识类型只需新建文件、继承 Protocol, 在 knowledge/__init__.py 注册,无需修改任何现有代码。
5

预测器层

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 转换), 不重复实现任何统计逻辑

类名namerequired_knowledge特点
MeanPredictormean[]预测全零 delta(最弱基线)
PerturbedMeanPredictorperturbed_mean[]所有条件预测同一均值偏移,高分低区分
LinearPredictorlinear[]Ridge 回归,加法可组合基线
NearestPerturbationPredictornearest_perturbation[]迁移最近共表达扰动的 delta

5.3 GEARSPredictor — 完整训练链路

封装原 train_gears.py 的完整流程(apply_heldout_filtercarve_inner_valget_dataloaderswap_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 阶段一并解决。
组件拼装示例 · GEARS vs scGPT 路径对比 Component Assembly · How algorithms declare and consume knowledge PerturbationDataset 统一输入 GOGraphComponent component_type = "go_graph" G GEARSPredictor required_knowledge = ["go_graph"] fit(dataset, {"go_graph": go}) → carve inner-val → GEARS train() → save model.pt + config.pkl predict_delta(cond) → (n_genes,) PretrainedEmbeddingComponent component_type = "pretrained_embeddings" S ScGPTPredictor required_knowledge = ["pretrained_embeddings"] fit(dataset, {"pretrained_embeddings": emb}) → fine-tune scGPT on perturbed_cells → save weights + config predict_delta(cond) → (n_genes,) 规划中 score_predictor(predictor, dataset, tier) 对任意 PerturbationPredictor 输出相同格式的 metrics dict vs
图 3 · 组件拼装对比 — GEARSPredictor(需要 GO 图)vs ScGPTPredictor(需要预训练 embedding),两者向下汇合至同一个 score_predictor
6

评测与训练

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_ciBootstrap 2000 次,95% CI单跑可信度
direction_accuracy_top_degTop-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.py30ctrl_mean 精度、save/load 圆回 +条件名含 + 的 npz 转义
tests/test_predictors.py26Protocol 合规、predict_delta 形状/dtype、save/load 等值
tests/test_scorer.py14deg_pearson 与直接公式逐位一致(误差 <1e-6)
在服务器上的验收计划:from_carved_pkl() 加载 baseline-gears/seed0 的产物, 用 LinearPredictor + score_predictor 计算 discovery tier, 与 metrics_discovery_linear.json 的已知数字进行逐行比对, 差值应 <1e-4(浮点累积误差上界)。
7

扩展路径

Extension Roadmap · 接入新算法

7.1 接入 scGPT 只需三步

7.2 下一步工作清单

技术债: GEARSPredictor 的 predict_delta() 需要条件有 inference graphs(training 或 carved),不支持完全未见基因的 zero-shot 推断。 这个限制在接入 scGPT 时会自然消除(scGPT 用 embedding 而非图做推断), 届时再回头升级 GEARS 路径。