summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 19:56:12 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 19:56:12 -0500
commit3c19f11acf3ebe290df439a3cb0558bdab0b3522 (patch)
tree66bd0a671cd5b7c023e4f20e2e9d9c19b962306d
parent5cfce81898d1a8ca08beac2c68f73dfb9f6447d5 (diff)
audit: isolate BCI role learner from environment map
-rw-r--r--experiments/bci_smoke.py9
-rw-r--r--sdil/bci.py19
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)