import torch import torch.nn as nn class VirtualCell(nn.Module): def __init__(self, expr_dim, emb_dim, hidden_dim): super(VirtualCell, self).__init__() self.delta_net = nn.Sequential( nn.Linear(emb_dim, hidden_dim), nn.LayerNorm(hidden_dim), nn.GELU(), nn.Linear(hidden_dim, expr_dim), ) nn.init.zeros_(self.delta_net[-1].weight) nn.init.zeros_(self.delta_net[-1].bias) def forward(self, ctrl_expr, pert_emb): pred_delta = self.delta_net(pert_emb) return ctrl_expr + pred_delta