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 /sdil | |
| parent | 5cfce81898d1a8ca08beac2c68f73dfb9f6447d5 (diff) | |
audit: isolate BCI role learner from environment map
Diffstat (limited to 'sdil')
| -rw-r--r-- | sdil/bci.py | 19 |
1 files changed, 13 insertions, 6 deletions
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) |
