基因身份嵌入层(Gene identity embedding)
GeneEmbeddingLayer
00 Remarks
对应 readme Step 2 中的可学习基因身份向量 e_i ∈ R^d(类似 word embedding, 维度通常取 16~256)。
本层同时承担 GEARS「扰动基因集合编码器」的第一段工作:
- 把输入的扰动 multi-hot 标记
p作用到嵌入表上,得到被扰动基因的嵌入子集p_i · e_i; - 该结果随后交给
GlobalPoolingLayer做 Deep Sets 均值池化,
得到与基因顺序无关的全局扰动向量 z_pert,再拼接到每一个节点的特征上。
嵌入表的梯度通过池化路径回传并在这里累积;由于优化器持有梯度张量引用, 反向传播只允许原地累加,不允许重新分配梯度张量。
01 Syntax
SMRUCC.genomics.Analysis.GEARS.Layers.GeneEmbeddingLayer
02 Methods
| Name | Overloads | Summary |
|---|---|---|
| .ctor | 1 | 创建基因身份嵌入层 |
| Forward | 1 | 前向传播:输出被扰动基因掩码之后的嵌入矩阵 |
| Backward | 1 | 反向传播:按扰动掩码把梯度累积回嵌入表 |
| GetParameters | 1 | 获取本层可训练参数(基因身份嵌入表) |
| GetGradients | 1 | 获取本层参数梯度(嵌入表梯度) |
03 Properties
| Name | Overloads | Summary |
|---|---|---|
| NumGenes | 1 | 基因数量 |
| EmbeddingDim | 1 | 嵌入向量维度 |
| Embeddings | 1 | 获取基因身份嵌入表(供模型拼接节点特征时读取当前值) |
04 Fields
| Name | Overloads | Summary |
|---|---|---|
| embedding | 1 | 基因身份嵌入表 [numGenes, embeddingDim] |
| embeddingGrad | 1 | 嵌入表梯度 [numGenes, embeddingDim] |
| lastFlag | 1 | 上一次前向传播使用的扰动标记 |
05 Members
#ctor(
Int32, Int32, Double, Nullable(Of Int32), String)创建基因身份嵌入层
Parameters
| Name | Type | Description |
|---|---|---|
numGenes | Int32 | 基因数量 |
embeddingDim | Int32 | 嵌入向量维度 |
scale | Double | 初始化缩放系数 |
seed | Nullable(Of Int32) | 随机初始化种子;给定种子可保证实验可复现 |
name | String | 层名称 |
Forward(
Tensor)前向传播:输出被扰动基因掩码之后的嵌入矩阵
Parameters
| Name | Type | Description |
|---|---|---|
input | Tensor | 扰动 multi-hot 标记,可以是 [numGenes, 1] 的二维张量,也可以是 [numGenes] 的一维张量 |
Returns
掩码后的嵌入 [numGenes, embeddingDim],其中未被扰动的基因整行为 0
Backward(
Tensor)反向传播:按扰动掩码把梯度累积回嵌入表
Parameters
| Name | Type | Description |
|---|---|---|
gradient | Tensor | 上游梯度 [numGenes, embeddingDim] |
Returns
相对于输入扰动标记的梯度(形状与上一次 GeneEmbeddingLayer.Forward() 的输入一致)
GetParameters
获取本层可训练参数(基因身份嵌入表)
Returns
参数张量列表
GetGradients
获取本层参数梯度(嵌入表梯度)
Returns
梯度张量列表
NumGenes
基因数量
Returns
嵌入表的行数
EmbeddingDim
嵌入向量维度
Returns
嵌入表的列数
Embeddings
获取基因身份嵌入表(供模型拼接节点特征时读取当前值)
Returns
嵌入张量 [numGenes, embeddingDim]
embedding
基因身份嵌入表 [numGenes, embeddingDim]
embeddingGrad
嵌入表梯度 [numGenes, embeddingDim]
lastFlag
上一次前向传播使用的扰动标记