diff --git a/examples/simpl_demo.ipynb b/examples/simpl_demo.ipynb index 9f10a40..5496eb8 100755 --- a/examples/simpl_demo.ipynb +++ b/examples/simpl_demo.ipynb @@ -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", diff --git a/src/simpl/simpl.py b/src/simpl/simpl.py index 43f9010..ecb7cb6 100755 --- a/src/simpl/simpl.py +++ b/src/simpl/simpl.py @@ -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, @@ -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 @@ -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 ---------- @@ -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 @@ -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 @@ -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. @@ -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. @@ -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 @@ -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)), diff --git a/tests/test_simpl.py b/tests/test_simpl.py index 121a347..68f90c8 100644 --- a/tests/test_simpl.py +++ b/tests/test_simpl.py @@ -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 @@ -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_