Download src/ml/VirtualCell.py from Quazim0t0/SpikeWhale-SNN-216M: direct link, hf CLI and curl.
- Browser
- Download file 612 Bytes
-
https://huggingface.co/Quazim0t0/SpikeWhale-SNN-216M/resolve/main/src/ml/VirtualCell.py
- Command line
-
hf download hf://Quazim0t0/SpikeWhale-SNN-216M/src/ml/VirtualCell.py
-
curl -L -o VirtualCell.py https://huggingface.co/Quazim0t0/SpikeWhale-SNN-216M/resolve/main/src/ml/VirtualCell.py
612 Bytes
| 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 |