GEARS 模型训练器
GEARSTrainer
00 Remarks
训练流程严格对应 readme §7.2 的伪代码:
- 构建扰动标记
p(组合扰动为 multi-hot); - 构建初始节点特征
h0 = [x̄ ‖ p ‖ e ‖ z_pert]; - 多层消息传递;
- 解码得到 Δ 预测,损失取 MSE(Δ̂, Δ);
- 反向传播并用 Adam 更新参数。
归一化约定:输入表达按基因做 Z-score(减 controlMean 除 controlSD), Δ 标签同样除以 controlSD。预测时把 Δ̂ 乘回 controlSD 即可还原到原始表达尺度。
01 Syntax
SMRUCC.genomics.Analysis.GEARS.Training.GEARSTrainer
02 Methods
| Name | Overloads | Summary |
|---|---|---|
| .ctor | 1 | 创建训练器 |
| Train | 1 | 执行训练 |
| ApplyRegularization | 1 | 把 L2 正则项的梯度累加到参数梯度上(权重衰减) |
| Evaluate | 1 | 在给定样本集上评估模型的平均 MSE(不改变模型参数) |
03 Properties
| Name | Overloads | Summary |
|---|---|---|
| LossCurve | 1 | 训练过程中每个 epoch 的平均损失 |
| Parameters | 1 | 模型可训练参数(交给优化器原地更新) |
| Gradients | 1 | 模型参数梯度 |
04 Fields
05 Members
创建训练器
Parameters
| Name | Type | Description |
|---|---|---|
model | GEARSModel | GEARS 模型 |
graphData | GeneRegulatoryGraph | 基因调控图 |
controlMean | Double() | control 表达均值 [numGenes] |
controlSD | Double() | control 表达标准差 [numGenes] |
learningRate | Single | Adam 学习率 |
l2Lambda | Double | L2 正则化系数,0 表示不启用正则 |
Train(
List(Of PerturbSeqSample), Int32, Int32)执行训练
Parameters
| Name | Type | Description |
|---|---|---|
samples | List(Of PerturbSeqSample) | 训练样本集合 |
epochs | Int32 | 训练轮数 |
printEvery | Int32 | 每隔多少个 epoch 打印一次损失;0 表示不打印 |
Returns
损失曲线(每个 epoch 的平均 MSE)
ApplyRegularization
把 L2 正则项的梯度累加到参数梯度上(权重衰减)
Evaluate(
List(Of PerturbSeqSample))在给定样本集上评估模型的平均 MSE(不改变模型参数)
Parameters
| Name | Type | Description |
|---|---|---|
samples | List(Of PerturbSeqSample) | 评估样本集合 |
Returns
平均均方误差(归一化尺度)
LossCurve
训练过程中每个 epoch 的平均损失
Returns
损失曲线
Parameters
模型可训练参数(交给优化器原地更新)
Returns
参数张量列表
Gradients
模型参数梯度
Returns
梯度张量列表
model
待训练的模型
graphData
基因调控图
optimizer
Adam 优化器
controlMean
control 表达均值
controlSD
control 表达标准差(归一化尺度)
l2Lambda
L2 正则化系数