GRU#

class ts1.models.single_session.GRU(bin_size=0.02, depth=2, hidden_dim=64, dropout=0.2, bidirectional=False)[source]#

Bases: core.model.BaseModel

Gated recurrent unit network mapping binned spike counts to task outputs.

Notation: \(B\) = batch size, \(T_{in}\) = input time bins, \(N\) = units, \(D_{out}\) = task output dim, \(T_{out}\) = output time steps, \(D\) = hidden dim (\(2D\) when bidirectional).

configure_readout() must be called before inference; it fixes \(D_{out}\) and the output shape.

  1. input_fn(): bin raw spikes into \((T_{in}, N)\).

  2. forward(): project \((B, T_{in}, N)\) to \((B, T_{in}, D)\), run through depth GRU layers, adaptively pool to \((B, T_{out}, D)\), and project to \((B, 1, D_{out})\) or \((B, T_{out}, D_{out})\) depending on the target resolution.

Parameters:
  • bin_size (float) – Width of each time bin in seconds.

  • depth (int) – Number of stacked GRU layers.

  • hidden_dim (int) – Hidden state size \(D\) per GRU layer.

  • dropout (float) – Dropout probability between GRU layers (disabled for single-layer models).

  • bidirectional (bool) – If True, use a bidirectional GRU; the readout input dim becomes \(2D\).

configure_readout(readout_spec)[source]#

Fix \(D_{out}\) and build the temporal pool and readout head.

The output shape depends on the target resolution:

  • Sequence-level (\(T_{out}=1\)): pool to a single time step, project \(D \to D_{out}\), reshape to \((B, 1, D_{out})\).

  • Timestep-level: pool to \(T_{out}\) behaviour frames, project \(D \to D_{out}\), reshape to \((B, T_{out}, D_{out})\).

Parameters:

readout_spec (ReadoutSpec) – Task specification carrying \(D_{out}\) and the target resolution.

input_fn(data)[source]#

Bin spikes.

Parameters:

data (Data) – Trial data containing raw spike times and unit metadata.

Return type:

dict[str, Tensor]

Returns:

Dict with model_inputs.spikes of shape \((T_{in}, N)\).

forward(spikes)[source]#

Map binned spikes to task predictions.

Parameters:

spikes (Tensor) – \((B, T_{in}, N)\) binned spike counts.

Return type:

Tensor

Returns:

\((B, 1, D_{out})\) for sequence-level tasks or \((B, T_{out}, D_{out})\) for timestep-level tasks.

classmethod create_search_space(trial, cfg)[source]#

Map out the model’s Optuna search space.

Call trial.suggest_*; the names suggested become the keys process_tunable_params() receives.

Parameters:
  • trial (Trial) – Optuna trial to register suggestions on.

  • cfg (DictConfig) – The run config, for values the space depends on.

classmethod process_tunable_params(tune_params)[source]#

Turn suggested hyperparameters into config overrides.

Runs before the config is filled, so this is where a suggestion is mapped onto the config path it sets (batch_size_log2 -> batch_size), a value is derived from another, or a default is supplied for something not being tuned.

Parameters:

tune_params (dict) – The names create_search_space() suggested, with their values.

Return type:

dict

Returns:

The overrides to apply to the config. The default returns them unchanged.