nuget server logo nuget api documents
↑

API Docs / SMRUCC.genomics.Analysis.GEARS / GEARSTrainer

GEARSTrainer

Full name SMRUCC.genomics.Analysis.GEARS.Training.GEARSTrainer Assembly SMRUCC.genomics.Analysis.GEARS Members 13

GEARS 模型训练器

00 Remarks

训练流程严格对应 readme §7.2 的伪代码:

  1. 构建扰动标记 p(组合扰动为 multi-hot);
  2. 构建初始节点特征 h0 = [x̄ ‖ p ‖ e ‖ z_pert];
  3. 多层消息传递;
  4. 解码得到 Δ 预测,损失取 MSE(Δ̂, Δ);
  5. 反向传播并用 Adam 更新参数。

归一化约定:输入表达按基因做 Z-score(减 controlMean 除 controlSD), Δ 标签同样除以 controlSD。预测时把 Δ̂ 乘回 controlSD 即可还原到原始表达尺度。

01 Syntax

SMRUCC.genomics.Analysis.GEARS.Training.GEARSTrainer

02 Methods

NameOverloadsSummary
.ctor 1 创建训练器
Train 1 执行训练
ApplyRegularization 1 把 L2 正则项的梯度累加到参数梯度上(权重衰减)
Evaluate 1 在给定样本集上评估模型的平均 MSE(不改变模型参数)

03 Properties

NameOverloadsSummary
LossCurve 1 训练过程中每个 epoch 的平均损失
Parameters 1 模型可训练参数(交给优化器原地更新)
Gradients 1 模型参数梯度

04 Fields

NameOverloadsSummary
model 1 待训练的模型
graphData 1 基因调控图
optimizer 1 Adam 优化器
controlMean 1 control 表达均值
controlSD 1 control 表达标准差(归一化尺度)
l2Lambda 1 L2 正则化系数

05 Members

method .ctor #
#ctor(GEARSModel, GeneRegulatoryGraph, Double(), Double(), Single, Double)

创建训练器

Parameters
NameTypeDescription
modelGEARSModel

GEARS 模型

graphDataGeneRegulatoryGraph

基因调控图

controlMeanDouble()

control 表达均值 [numGenes]

controlSDDouble()

control 表达标准差 [numGenes]

learningRateSingle

Adam 学习率

l2LambdaDouble

L2 正则化系数,0 表示不启用正则

method Train #
Train(List(Of PerturbSeqSample), Int32, Int32)

执行训练

Parameters
NameTypeDescription
samplesList(Of PerturbSeqSample)

训练样本集合

epochsInt32

训练轮数

printEveryInt32

每隔多少个 epoch 打印一次损失;0 表示不打印

Returns

损失曲线(每个 epoch 的平均 MSE)

method ApplyRegularization #
ApplyRegularization

把 L2 正则项的梯度累加到参数梯度上(权重衰减)

method Evaluate #
Evaluate(List(Of PerturbSeqSample))

在给定样本集上评估模型的平均 MSE(不改变模型参数)

Parameters
NameTypeDescription
samplesList(Of PerturbSeqSample)

评估样本集合

Returns

平均均方误差(归一化尺度)

property LossCurve #
LossCurve

训练过程中每个 epoch 的平均损失

Returns

损失曲线

property Parameters #
Parameters

模型可训练参数(交给优化器原地更新)

Returns

参数张量列表

property Gradients #
Gradients

模型参数梯度

Returns

梯度张量列表

field model #
model

待训练的模型

field graphData #
graphData

基因调控图

field optimizer #
optimizer

Adam 优化器

field controlMean #
controlMean

control 表达均值

field controlSD #
controlSD

control 表达标准差(归一化尺度)

field l2Lambda #
l2Lambda

L2 正则化系数