Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 3 additions & 11 deletions examples/simpl_demo.ipynb
Original file line number Diff line number Diff line change
Expand Up @@ -237,17 +237,9 @@
" element must be 0. The Kalman filter runs independently within each trial to\n",
" prevent smoothing across trial boundaries (e.g. between separate recording\n",
" sessions). If None, all data is treated as a single trial. By default None.\n",
"align_to_behavior : bool or str, optional\n",
" How to linearly align (via CCA) the decoded latent positions to the behavioral\n",
" coordinate system after each E-step. Default: ``\"trajectory\"``. Options:\n",
"\n",
" - ``\"trajectory\"`` (default) — align the decoded trajectory ``mu_s`` directly\n",
" to ``Xb``.\n",
" - ``\"fields\"`` — align based on peak positions of receptive fields. CAn be useful\n",
" in 1D where the latent position distribution can be bimodal. Unstable / not\n",
" recommended if fields are likely to have multiple peaks\n",
" - ``True`` — alias for ``\"trajectory\"``.\n",
" - ``False`` — no alignment.\n",
"align_to_behavior : bool, optional\n",
" If True, linearly align (via CCA) decoded latent positions to the behavioral\n",
" coordinate system after each E-step. By default True.\n",
"resume : bool, optional\n",
" If True, continue training from the current state without re-initialising. The\n",
" ``Y``, ``Xb``, and ``time`` arguments are ignored when resuming — training\n",
Expand Down
68 changes: 22 additions & 46 deletions src/simpl/simpl.py
Original file line number Diff line number Diff line change
Expand Up @@ -190,7 +190,7 @@ def fit(
time: np.ndarray | None,
n_iterations: int = 5,
trial_boundaries: np.ndarray | None = None,
align_to_behavior: bool | str = "trajectory",
align_to_behavior: bool = True,
resume: bool = False,
save_full_history: bool = False,
early_stopping: bool = True,
Expand Down Expand Up @@ -243,17 +243,9 @@ def fit(
element must be 0. The Kalman filter runs independently within each trial to
prevent smoothing across trial boundaries (e.g. between separate recording
sessions). If None, all data is treated as a single trial. By default None.
align_to_behavior : bool or str, optional
How to linearly align (via CCA) the decoded latent positions to the behavioral
coordinate system after each E-step. Default: ``"trajectory"``. Options:

- ``"trajectory"`` (default) — align the decoded trajectory ``mu_s`` directly
to ``Xb``.
- ``"fields"`` — align based on peak positions of receptive fields. CAn be useful
in 1D where the latent position distribution can be bimodal. Unstable / not
recommended if fields are likely to have multiple peaks
- ``True`` — alias for ``"trajectory"``.
- ``False`` — no alignment.
align_to_behavior : bool, optional
If True, linearly align (affine transformation) decoded latent positions to the behavioral
coordinate system after each iterations E-step. By default True.
resume : bool, optional
If True, continue training from the current state without re-initialising. The
``Y``, ``Xb``, and ``time`` arguments are ignored when resuming — training
Expand Down Expand Up @@ -897,7 +889,7 @@ def plot_prediction(
# ──────────────────────────────────────────────────────────────────────────

def _E_step(self, Y: jax.Array, F: jax.Array) -> dict:
"""E-step: decode latent positions and optionally align to behavior.
"""E-step: decode latent positions and align to behavior.

Parameters
----------
Expand All @@ -921,26 +913,17 @@ def _E_step(self, Y: jax.Array, F: jax.Array) -> dict:

# Manifold alignment (fit-time only)
self._substatus("decode··· aligning")
align_dict = {}
if self.align_mode_ == "fields":
current_peaks = utils.get_field_peaks(F, self.xF_)
source, target = current_peaks, self.Falign_peaks_
elif self.align_mode_ == "trajectory":
source, target = E["mu_s"], self.Xalign_
else:
source, target = None, None

if source is not None:
if self.is_1D_angular:
angle, _ = utils.cca_angular(source, target)
E["X"] = utils._wrap_minuspi_pi(E["mu_s"] + angle)
align_dict = {"intercept": jnp.atleast_1d(angle)}
else:
coef, intercept = utils.cca(source, target)
E["X"] = E["mu_s"] @ coef.T + intercept
align_dict = {"coef": coef, "intercept": intercept}
else:
if not self.align_to_behavior_:
E["X"] = E["mu_s"]
align_dict = {}
elif self.is_1D_angular:
angle, _ = utils.cca_angular(E["mu_s"], self.Xb_)
E["X"] = utils._wrap_minuspi_pi(E["mu_s"] + angle)
align_dict = {"intercept": jnp.atleast_1d(angle)}
else:
coef, intercept = utils.cca(E["mu_s"], self.Xb_)
E["X"] = E["mu_s"] @ coef.T + intercept
align_dict = {"coef": coef, "intercept": intercept}

E.update(align_dict)
return E
Expand Down Expand Up @@ -1239,8 +1222,6 @@ def _run_iteration_zero(self, verbose: bool) -> None:

self._fit_iteration()
self.FX_first_iteration_ = self.M_["FX"]
if self.align_mode_ == "fields":
self.Falign_peaks_ = utils.get_field_peaks(self.M_["F"], self.xF_)

if verbose:
print() # newline after header
Expand Down Expand Up @@ -1415,7 +1396,7 @@ def _init_from_data(self, Y, Xb, time, trial_boundaries, align_to_behavior) -> N
def _init_infrastructure(
self,
trial_boundaries,
align_to_behavior=None,
align_to_behavior=True,
spike_mask=None,
) -> None:
"""Set up Kalman filter, masks, alignment, and coordinate registry.
Expand All @@ -1428,8 +1409,8 @@ def _init_infrastructure(
----------
trial_boundaries : array-like or None
Trial boundary indices passed through to ``_validate_trial_boundaries``.
align_to_behavior : bool or str or None
Alignment mode (``True``/``"trajectory"``/``"fields"``/``None``).
align_to_behavior : bool
Whether to align decoded positions to the behavioral trajectory.
spike_mask : array-like or None
If provided, use this mask instead of generating a fresh speckled mask.
Used when loading from saved results.
Expand Down Expand Up @@ -1482,14 +1463,9 @@ def _init_infrastructure(
"Adjust val_frac or speckle_block_size_seconds."
)

if align_to_behavior is True:
align_to_behavior = "trajectory"
if align_to_behavior and align_to_behavior not in ("trajectory", "fields"):
raise ValueError(
f"align_to_behavior must be True, False, 'trajectory', or 'fields', got {align_to_behavior!r}"
)
self.align_mode_ = align_to_behavior if align_to_behavior else None
self.Xalign_ = self.Xb_ if self.align_mode_ else None
if not isinstance(align_to_behavior, (bool, np.bool_)):
raise TypeError(f"align_to_behavior must be a bool, got {align_to_behavior!r}")
self.align_to_behavior_ = bool(align_to_behavior)
self._kde = kde.kde_angular if self.is_1D_angular else kde.kde

self.lastF_, self.lastX_ = None, None
Expand Down Expand Up @@ -1940,7 +1916,7 @@ def _build_dataset_attrs(self, trial_boundaries) -> dict:
"speed_prior": np.nan if self.speed_prior_ is None else self.speed_prior_,
"behavior_prior": np.nan if self.behavior_prior is None else self.behavior_prior,
"is_1D_angular": int(self.is_1D_angular),
"align_mode": self.align_mode_ or "none",
"align_to_behavior": int(self.align_to_behavior_),
"val_frac": self.val_frac,
"speckle_block_size_seconds": self.speckle_block_size_seconds,
"save_full_history": int(getattr(self, "save_full_history_", False)),
Expand Down
61 changes: 15 additions & 46 deletions tests/test_simpl.py
Original file line number Diff line number Diff line change
Expand Up @@ -609,56 +609,26 @@ def test_cca_runs(self, small_simpl_model):
assert model.iteration_ >= 1
assert "X" in model.E_

def test_align_to_behavior(self, demo_data):
model = self._make_model(demo_data, align_to_behavior=True)
assert model.Xalign_ is not None
assert "X" in model.E_

def test_no_alignment(self, demo_data):
model = self._make_model(demo_data, align_to_behavior=False)
assert model.Xalign_ is None
assert "X" in model.E_

def test_align_trajectory_mode(self, demo_data):
model = self._make_model(demo_data, align_to_behavior="trajectory")
assert model.align_mode_ == "trajectory"
assert model.Xalign_ is not None
def test_aligns_to_behavior(self, demo_data):
model = self._make_model(demo_data)
assert model.align_to_behavior_ is True
assert "coef" in model.E_
assert "intercept" in model.E_

def test_align_fields_mode(self, demo_data):
model = self._make_model(demo_data, align_to_behavior="fields")
assert model.align_mode_ == "fields"
assert hasattr(model, "Falign_peaks_")
assert model.Falign_peaks_.shape == (model.N_neurons_, model.D_)
assert "coef" in model.E_
assert "intercept" in model.E_
def test_alignment_can_be_disabled(self, demo_data):
model = self._make_model(demo_data, align_to_behavior=False)
assert model.align_to_behavior_ is False
assert "coef" not in model.E_
assert "intercept" not in model.E_
np.testing.assert_array_equal(model.E_["X"], model.E_["mu_s"])

def test_align_invalid_mode_raises(self, demo_data):
with pytest.raises(ValueError, match="align_to_behavior"):
self._make_model(demo_data, align_to_behavior="invalid")
@pytest.mark.parametrize("value", ["trajectory", "fields", 1, None])
def test_alignment_rejects_non_boolean_values(self, demo_data, value):
with pytest.raises(TypeError, match="align_to_behavior must be a bool"):
self._make_model(demo_data, align_to_behavior=value)

def test_align_angular(self):
"""Field-based angular alignment uses rotation, not CCA."""
rng = np.random.default_rng(42)
T, N_neurons = 2000, 15
time = np.arange(T) * 0.02
Xb = np.linspace(-np.pi, np.pi, T, endpoint=False)[:, None]
# Simulate spikes with angular tuning
preferred = np.linspace(-np.pi, np.pi, N_neurons, endpoint=False)
rates = np.exp(3 * np.cos(Xb - preferred[None, :]))
Y = rng.poisson(rates * 0.02)

model = SIMPL(is_1D_angular=True, bin_size=np.pi / 32, speed_prior=0.1, kernel_bandwidth=0.3)
model.fit(Y, Xb, time, n_iterations=1, align_to_behavior="fields")
assert model.align_mode_ == "fields"
assert "intercept" in model.E_
# X should be wrapped to [-pi, pi)
assert np.all(model.X_ >= -np.pi)
assert np.all(model.X_ < np.pi)

def test_align_angular_trajectory_mode(self):
"""Trajectory-based angular alignment also uses rotation."""
"""Angular alignment uses rotation."""
rng = np.random.default_rng(42)
T, N_neurons = 2000, 15
time = np.arange(T) * 0.02
Expand All @@ -668,8 +638,7 @@ def test_align_angular_trajectory_mode(self):
Y = rng.poisson(rates * 0.02)

model = SIMPL(is_1D_angular=True, bin_size=np.pi / 32, speed_prior=0.1, kernel_bandwidth=0.3)
model.fit(Y, Xb, time, n_iterations=1, align_to_behavior="trajectory")
assert model.align_mode_ == "trajectory"
model.fit(Y, Xb, time, n_iterations=1)
assert "intercept" in model.E_
assert "coef" not in model.E_

Expand Down
Loading