SNN / srcn_model.py
sensen0025's picture
Upload srcn_model.py with huggingface_hub
f3e30cc verified
Raw
History Blame Contribute Delete
6.66 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
def nmda_gate(v, mg=1.0):
return 1.0 / (1.0 + mg * torch.exp(-0.062 * v) / 3.57)
class FastSigmoidSurrogate(torch.autograd.Function):
@staticmethod
def forward(ctx, v, v_th):
ctx.save_for_backward(v, v_th)
return (v >= v_th).to(v.dtype)
@staticmethod
def backward(ctx, grad_output):
v, v_th = ctx.saved_tensors
diff = torch.abs(v - v_th)
grad_v = grad_output / ((1.0 + diff) ** 2)
return grad_v, None
class SRCNLayer(nn.Module):
def __init__(self, num_columns=256, neurons_per_column=512, num_partners=16):
super().__init__()
self.C = num_columns
self.M = neurons_per_column
self.K = num_partners
self.W_raw = nn.Parameter(
torch.randn(self.C, self.M, self.K, self.M) * (0.02 / (self.K * self.M) ** 0.5)
)
num_excitatory = int(self.C * 200 / 256)
num_inhibitory = self.C - num_excitatory
partner_indices = torch.zeros(self.C, self.K, dtype=torch.long)
partner_signs = torch.zeros(self.C, self.K)
half_k = self.K // 2
for c in range(self.C):
for k_idx in range(self.K):
partner_col = (c - half_k + k_idx) % self.C
partner_indices[c, k_idx] = partner_col
partner_signs[c, k_idx] = 1.0 if partner_col < num_excitatory else -1.0
self.register_buffer("partner_indices", partner_indices)
self.register_buffer("_partner_signs", partner_signs.view(self.C, 1, self.K, 1))
@property
def device(self):
return self.W_raw.device
def precompute_W(self):
return torch.abs(self.W_raw) * self._partner_signs + (1e-6 * self._partner_signs)
def forward(self, S_prev, V, V_th, I_ampa, I_nmda, I_inj,
tau_mem=0.9, epsilon=1e-4, a_target=0.015,
alpha_ampa=0.667, alpha_nmda=0.98, W_fp16=None):
if W_fp16 is None:
W_fp16 = self.precompute_W()
batch_size = S_prev.shape[0]
flat_indices = self.partner_indices.reshape(-1)
S_gathered = S_prev.index_select(1, flat_indices)
S_partners = S_gathered.view(batch_size, self.C, self.K, self.M)
S_partners_reshaped = S_partners.permute(1, 0, 2, 3).reshape(self.C, batch_size, self.K * self.M)
W_reshaped = W_fp16.reshape(self.C, self.M, self.K * self.M).transpose(1, 2)
I_syn = torch.bmm(S_partners_reshaped, W_reshaped).transpose(0, 1)
I_syn = torch.clamp(I_syn, min=-500.0, max=500.0)
I_syn_exc = torch.clamp(I_syn, min=0.0)
I_syn_inh = torch.clamp(I_syn, max=0.0)
I_ampa_next = alpha_ampa * I_ampa + I_syn_exc
nmda_gate_val = nmda_gate(V)
I_nmda_next = alpha_nmda * I_nmda + I_syn_exc * nmda_gate_val
I_nmda_next = torch.clamp(I_nmda_next, max=1000.0)
I_total = I_ampa_next + I_nmda_next + I_syn_inh
V_leaked = tau_mem * V + (1.0 - tau_mem) * (I_total + I_inj)
V_next = V_leaked * (1.0 - S_prev)
V_next = torch.clamp(V_next, min=-100.0, max=100.0)
S_next = FastSigmoidSurrogate.apply(V_next, V_th)
V_th_next = V_th + epsilon * (S_next - a_target)
V_th_next = torch.clamp(V_th_next, min=0.1, max=5.0)
return S_next, V_next, V_th_next, I_ampa_next, I_nmda_next
class TemporalPhaseEncoder(nn.Module):
def __init__(self, vocab_size, embed_dim=512, num_columns=256, neurons_per_column=512, gain=13.0):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim)
self.proj = nn.Linear(embed_dim, num_columns * neurons_per_column, bias=False)
self.C = num_columns
self.M = neurons_per_column
self.gain = gain
def forward(self, token_ids, timestep, freq=80.0, tau_mem=0.9):
embed = self.embedding(token_ids)
phase_offsets = self.proj(embed)
t = timestep * 0.001
I_inj = torch.sin(2.0 * 3.141592653589793 * freq * t + phase_offsets)
return (I_inj * self.gain).view(-1, self.C, self.M)
class SRCNv3_1B(nn.Module):
def __init__(self, vocab_size, num_columns=256, neurons_per_column=512,
num_partners=16, embed_dim=512, num_motor_pool_steps=8, encoder_gain=13.0,
motor_ratio=0.54):
super().__init__()
self.C = num_columns
self.M = neurons_per_column
self.num_motor_pool_steps = num_motor_pool_steps
motor_start = int(self.C * (1.0 - motor_ratio))
self.motor_start_col = motor_start
self.num_motor_neurons = (self.C - motor_start) * self.M
self.encoder = TemporalPhaseEncoder(vocab_size, embed_dim, self.C, self.M, gain=encoder_gain)
self.layer = SRCNLayer(self.C, self.M, num_partners)
self.vocab_head = nn.Sequential(
nn.LayerNorm(self.num_motor_neurons),
nn.Linear(self.num_motor_neurons, 4096, bias=True),
nn.ReLU(),
nn.Linear(4096, vocab_size, bias=True),
)
@property
def device(self):
return self.layer.W_raw.device
def get_projected_weights(self):
return self.layer.get_projected_weights()
def precompute_W(self):
return self.layer.precompute_W()
def forward_step(self, S_prev, V, V_th, I_ampa, I_nmda, input_token_id, timestep, W_fp16=None):
I_inj = self.encoder(input_token_id, timestep)
if W_fp16 is None:
W_fp16 = self.layer.precompute_W()
S_next, V_next, V_th_next, I_ampa_next, I_nmda_next = self.layer(
S_prev, V, V_th, I_ampa, I_nmda, I_inj,
0.9, 5e-5, 0.10, 0.667, 0.98, W_fp16
)
return S_next, V_next, V_th_next, I_ampa_next, I_nmda_next
def forward_token_with_psc(self, S, V, V_th, I_ampa, I_nmda, I_psc, token, t_start_tensor, W_fp16):
psc_hist = []
spikes_sum = 0.0
t_start = int(t_start_tensor.item())
for st in range(self.num_motor_pool_steps):
ts = t_start + st
I_inj = self.encoder(token, ts)
S, V, V_th, I_ampa, I_nmda = self.layer(
S, V, V_th, I_ampa, I_nmda, I_inj,
0.9, 5e-5, 0.10, 0.667, 0.98, W_fp16
)
spikes_sum = spikes_sum + S.sum()
Sm = S[:, self.motor_start_col:, :].reshape(S.shape[0], -1)
I_psc = (1.0 - 1.0 / 3.0) * I_psc + Sm
psc_hist.append(I_psc)
pooled_token = torch.stack(psc_hist, dim=0).mean(dim=0)
return S, V, V_th, I_ampa, I_nmda, I_psc, pooled_token, spikes_sum