diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 19:56:12 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 19:56:12 -0500 |
| commit | 3c19f11acf3ebe290df439a3cb0558bdab0b3522 (patch) | |
| tree | 66bd0a671cd5b7c023e4f20e2e9d9c19b962306d | |
| parent | 5cfce81898d1a8ca08beac2c68f73dfb9f6447d5 (diff) | |
audit: isolate BCI role learner from environment map
| -rw-r--r-- | experiments/bci_smoke.py | 9 | ||||
| -rw-r--r-- | sdil/bci.py | 19 |
2 files changed, 21 insertions, 7 deletions
diff --git a/experiments/bci_smoke.py b/experiments/bci_smoke.py index 4724d26..1fda89e 100644 --- a/experiments/bci_smoke.py +++ b/experiments/bci_smoke.py @@ -104,7 +104,14 @@ def check_temporal_difference_role_vectorizer(): cfg.episodes_per_day, cfg.n_neurons, generator=generator) xi = torch.empty_like(soma).bernoulli_( 0.5, generator=generator).mul_(2).sub_(1) - target = model.causal_role_targets(soma, xi) + cursor_plus = (soma + cfg.perturb_sigma * xi) @ model.role + cursor_minus = (soma - cfg.perturb_sigma * xi) @ model.role + target = model.causal_role_targets(cursor_plus, cursor_minus, xi) + saved_role = model.role.clone() + model.role.zero_() + assert torch.equal( + target, model.causal_role_targets(cursor_plus, cursor_minus, xi)) + model.role.copy_(saved_role) mean_target = target.mean(0) cosine = torch.nn.functional.cosine_similarity( mean_target, model.role, dim=0).item() diff --git a/sdil/bci.py b/sdil/bci.py index 2978e5a..e20bbb6 100644 --- a/sdil/bci.py +++ b/sdil/bci.py @@ -165,18 +165,18 @@ class BCISDIL: directional_derivative = (loss_plus - loss_minus) / (2.0 * sigma) return -directional_derivative.unsqueeze(1) * perturbations - def causal_role_targets(self, soma, perturbations): + def causal_role_targets(self, cursor_plus, cursor_minus, perturbations): """Estimate each cell's signed causal effect on the BCI cursor. Unlike ``causal_targets``, this target contains no instantaneous error magnitude. A scalar antithetic cursor difference is tagged by the locally available perturbation at each cell; its expectation is the - experimenter-unknown causal role vector. + experimenter-unknown causal role vector. The learner receives the two + scalar cursor observations and cannot read the environment's role map. """ sigma = self.cfg.perturb_sigma - plus = (soma + sigma * perturbations) @ self.role - minus = (soma - sigma * perturbations) @ self.role - directional_derivative = (plus - minus) / (2.0 * sigma) + directional_derivative = ( + cursor_plus - cursor_minus) / (2.0 * sigma) return directional_derivative.unsqueeze(1) * perturbations def exact_causal_direction(self, soma, target): @@ -255,8 +255,15 @@ class BCISDIL: causal_target = None if did_perturb: if self.cfg.feedback == "performance_velocity": + sigma = self.cfg.perturb_sigma + # These two scalar values are environment observations. + # ``causal_role_targets`` has no access to ``self.role``. + cursor_plus = ( + base_soma + sigma * perturbations) @ self.role + cursor_minus = ( + base_soma - sigma * perturbations) @ self.role causal_target = self.causal_role_targets( - base_soma, perturbations) + cursor_plus, cursor_minus, perturbations) else: causal_target = self.causal_targets( base_soma, self.cfg.target, perturbations) |
