nuget server logo nuget api documents
↑

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

GeneEmbeddingLayer

Full name SMRUCC.genomics.Analysis.GEARS.Layers.GeneEmbeddingLayer Assembly SMRUCC.genomics.Analysis.GEARS Members 11

基因身份嵌入层(Gene identity embedding)

00 Remarks

对应 readme Step 2 中的可学习基因身份向量 e_i ∈ R^d(类似 word embedding, 维度通常取 16~256)。

本层同时承担 GEARS「扰动基因集合编码器」的第一段工作:

  1. 把输入的扰动 multi-hot 标记 p 作用到嵌入表上,得到被扰动基因的嵌入子集 p_i · e_i;
  2. 该结果随后交给 GlobalPoolingLayer 做 Deep Sets 均值池化,

得到与基因顺序无关的全局扰动向量 z_pert,再拼接到每一个节点的特征上。

嵌入表的梯度通过池化路径回传并在这里累积;由于优化器持有梯度张量引用, 反向传播只允许原地累加,不允许重新分配梯度张量。

01 Syntax

SMRUCC.genomics.Analysis.GEARS.Layers.GeneEmbeddingLayer

02 Methods

NameOverloadsSummary
.ctor 1 创建基因身份嵌入层
Forward 1 前向传播:输出被扰动基因掩码之后的嵌入矩阵
Backward 1 反向传播:按扰动掩码把梯度累积回嵌入表
GetParameters 1 获取本层可训练参数(基因身份嵌入表)
GetGradients 1 获取本层参数梯度(嵌入表梯度)

03 Properties

NameOverloadsSummary
NumGenes 1 基因数量
EmbeddingDim 1 嵌入向量维度
Embeddings 1 获取基因身份嵌入表(供模型拼接节点特征时读取当前值)

04 Fields

NameOverloadsSummary
embedding 1 基因身份嵌入表 [numGenes, embeddingDim]
embeddingGrad 1 嵌入表梯度 [numGenes, embeddingDim]
lastFlag 1 上一次前向传播使用的扰动标记

05 Members

method .ctor #
#ctor(Int32, Int32, Double, Nullable(Of Int32), String)

创建基因身份嵌入层

Parameters
NameTypeDescription
numGenesInt32

基因数量

embeddingDimInt32

嵌入向量维度

scaleDouble

初始化缩放系数

seedNullable(Of Int32)

随机初始化种子;给定种子可保证实验可复现

nameString

层名称

method Forward #
Forward(Tensor)

前向传播:输出被扰动基因掩码之后的嵌入矩阵

Parameters
NameTypeDescription
inputTensor

扰动 multi-hot 标记,可以是 [numGenes, 1] 的二维张量,也可以是 [numGenes] 的一维张量

Returns

掩码后的嵌入 [numGenes, embeddingDim],其中未被扰动的基因整行为 0

method Backward #
Backward(Tensor)

反向传播:按扰动掩码把梯度累积回嵌入表

Parameters
NameTypeDescription
gradientTensor

上游梯度 [numGenes, embeddingDim]

Returns

相对于输入扰动标记的梯度(形状与上一次 GeneEmbeddingLayer.Forward() 的输入一致)

method GetParameters #
GetParameters

获取本层可训练参数(基因身份嵌入表)

Returns

参数张量列表

method GetGradients #
GetGradients

获取本层参数梯度(嵌入表梯度)

Returns

梯度张量列表

property NumGenes #
NumGenes

基因数量

Returns

嵌入表的行数

property EmbeddingDim #
EmbeddingDim

嵌入向量维度

Returns

嵌入表的列数

property Embeddings #
Embeddings

获取基因身份嵌入表(供模型拼接节点特征时读取当前值)

Returns

嵌入张量 [numGenes, embeddingDim]

field embedding #
embedding

基因身份嵌入表 [numGenes, embeddingDim]

field embeddingGrad #
embeddingGrad

嵌入表梯度 [numGenes, embeddingDim]

field lastFlag #
lastFlag

上一次前向传播使用的扰动标记