NuCLR#

class pretrain.models.NuCLR(ctx_duration, latent_step, bin_size, dim, self_heads, dim_head, lin_dropout, t_layers, st_layers, rot_ratio=0.5, rotate_value=True, attn_biases=True)[source]#

Bases: torch.nn.modules.module.Module

Contrastive unit encoder over population context [Arora et al., 2025].

Reference implementation: nerdslab/nuclr, with the deviations listed in src/pretrain/models/nuclr/README.md.

forward(bins, spax_seqlen)[source]#
Parameters:
  • bins (Tensor) – Chained tensor of binned spikes (*num_neurons, num_bins)

  • spax_seqlen (Tensor) – seqlen to constrain spatial attention

Return type:

Tensor