summaryrefslogtreecommitdiff
path: root/sdil
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-23 07:53:33 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-23 07:53:33 -0500
commitce10335dcc3aa12f4f1e6dd53352d7e71d0cccf7 (patch)
tree1389b61c3e6a1dbdabe7df557128c16e5a3f2122 /sdil
parentc34675664b949ac7c3aa7a61e714d69183d0ea6f (diff)
bci: audit v2 outcome-surprise signatures
Diffstat (limited to 'sdil')
-rw-r--r--sdil/bci_v2.py15
-rw-r--r--sdil/bci_v2_metrics.py303
2 files changed, 316 insertions, 2 deletions
diff --git a/sdil/bci_v2.py b/sdil/bci_v2.py
index f2fdf6c..303a272 100644
--- a/sdil/bci_v2.py
+++ b/sdil/bci_v2.py
@@ -194,6 +194,7 @@ class BCIV2:
final_step,
plasticity_gain=1.0,
learn_role=True,
+ probe_role=True,
learn_predictor=True,
learn_critic=True,
critic_enabled=True,
@@ -265,7 +266,7 @@ class BCIV2:
eligibility = eligibility * active[:, None, None]
did_perturb = (
- learn_role and global_step % self.cfg.perturb_every == 0
+ probe_role and global_step % self.cfg.perturb_every == 0
)
causal_target = None
if did_perturb:
@@ -281,7 +282,7 @@ class BCIV2:
innovation[:, :, None] * eligibility
)[active].mean(0)
self.W += plasticity_gain * self.cfg.forward_eta * update
- if did_perturb:
+ if did_perturb and learn_role:
self.update_role(causal_target, active)
if learn_critic and critic_enabled:
self.update_critic(
@@ -356,6 +357,7 @@ def run_day_v2(
*,
plasticity_gain=1.0,
learn_role=True,
+ probe_role=True,
learn_predictor=True,
learn_critic=True,
critic_enabled=True,
@@ -372,6 +374,8 @@ def run_day_v2(
raise ValueError("v2 horizon exceeds the available trajectory")
state = model.initial_episode_state(episodes)
events = []
+ active_transitions = 0
+ role_probe_examples = 0
for step_index in range(horizon):
previous_soma = state["soma"]
state, record = model.step(
@@ -383,6 +387,7 @@ def run_day_v2(
final_step=step_index == horizon - 1,
plasticity_gain=plasticity_gain,
learn_role=learn_role,
+ probe_role=probe_role,
learn_predictor=learn_predictor,
learn_critic=learn_critic,
critic_enabled=critic_enabled,
@@ -399,6 +404,10 @@ def run_day_v2(
)
event["step"] = step_index
events.append(event)
+ active_count = int(record["active"].sum().item())
+ active_transitions += active_count
+ if record["did_perturb"]:
+ role_probe_examples += active_count
global_step += 1
if bool(state["active"].any()):
raise RuntimeError("every v2 episode must terminate")
@@ -407,4 +416,6 @@ def run_day_v2(
"success_rate": state["success"].float().mean().item(),
"events": events,
"global_step": global_step,
+ "active_transitions": active_transitions,
+ "role_probe_examples": role_probe_examples,
}
diff --git a/sdil/bci_v2_metrics.py b/sdil/bci_v2_metrics.py
new file mode 100644
index 0000000..5af6301
--- /dev/null
+++ b/sdil/bci_v2_metrics.py
@@ -0,0 +1,303 @@
+"""Frozen metrics for the independent actor--critic oral-B-v2 branch."""
+import math
+
+import torch
+
+from sdil.bci_metrics import (
+ grouped_classification_accuracy,
+ grouped_regression_correlation,
+ pearson,
+)
+
+
+def annotate_events(events, day, final_success, episode_offset, horizon):
+ """Attach stable groups, outcomes, and assay horizon to event records."""
+ final_success = torch.as_tensor(final_success, dtype=torch.bool)
+ episode_ids = (
+ torch.arange(final_success.numel(), dtype=torch.long)
+ + episode_offset
+ )
+ for event in events:
+ event["day"] = day
+ event["episode_ids"] = episode_ids.clone()
+ event["final_success"] = final_success.clone()
+ event["horizon"] = horizon
+ return events
+
+
+def flatten_active(events):
+ names = (
+ "soma",
+ "previous_soma",
+ "raw_apical",
+ "innovation",
+ "abs_error",
+ "previous_abs_error",
+ "td_innovation",
+ "value_prediction",
+ "terminal",
+ "outcome",
+ )
+ values = {name: [] for name in names}
+ values.update({"groups": [], "day": [], "horizon": []})
+ for event in events:
+ active = event["active"].bool()
+ if not bool(active.any()):
+ continue
+ count = int(active.sum())
+ for name in names:
+ values[name].append(event[name][active].double())
+ values["groups"].append(event["episode_ids"][active].long())
+ values["day"].append(
+ torch.full((count,), int(event["day"]), dtype=torch.long)
+ )
+ values["horizon"].append(
+ torch.full((count,), int(event["horizon"]), dtype=torch.long)
+ )
+ if not values["soma"]:
+ raise ValueError("no active oral-B-v2 events")
+ return {key: torch.cat(chunks, dim=0) for key, chunks in values.items()}
+
+
+def population_correlations(flat, left, right, mask=None):
+ if mask is None:
+ mask = torch.ones(flat[left].shape[0], dtype=torch.bool)
+ return [
+ pearson(flat[left][mask, cell], flat[right][mask, cell])
+ for cell in range(flat[left].shape[1])
+ ]
+
+
+def terminal_population(events):
+ """Return exactly one terminal record per annotated episode."""
+ records = {}
+ for event in events:
+ for row, group in enumerate(event["episode_ids"].tolist()):
+ if not bool(event["terminal"][row]):
+ continue
+ if group in records:
+ raise ValueError("multiple terminal events for one episode")
+ records[group] = {
+ "innovation": event["innovation"][row].double(),
+ "previous_soma": event["previous_soma"][row].double(),
+ "soma": event["soma"][row].double(),
+ "td_innovation": event["td_innovation"][row].double(),
+ "value_prediction": event["value_prediction"][row].double(),
+ "outcome": bool(event["outcome"][row]),
+ "horizon": int(event["horizon"]),
+ }
+ expected = {
+ int(group)
+ for event in events
+ for group in event["episode_ids"].tolist()
+ }
+ if set(records) != expected:
+ raise ValueError("every annotated episode needs one terminal event")
+ groups = sorted(records)
+ return {
+ "innovation": torch.stack([
+ records[group]["innovation"] for group in groups
+ ]),
+ "previous_soma": torch.stack([
+ records[group]["previous_soma"] for group in groups
+ ]),
+ "soma": torch.stack([
+ records[group]["soma"] for group in groups
+ ]),
+ "td_innovation": torch.stack([
+ records[group]["td_innovation"] for group in groups
+ ]),
+ "value_prediction": torch.stack([
+ records[group]["value_prediction"] for group in groups
+ ]),
+ "outcome": torch.tensor([
+ records[group]["outcome"] for group in groups
+ ], dtype=torch.bool),
+ "horizon": torch.tensor([
+ records[group]["horizon"] for group in groups
+ ], dtype=torch.long),
+ "groups": torch.tensor(groups, dtype=torch.long),
+ }
+
+
+def _paired_terminal_populations(mode_events):
+ populations = {
+ name: terminal_population(events)
+ for name, events in mode_events.items()
+ }
+ reference = populations["intact"]
+ for name, population in populations.items():
+ for key in ("outcome", "horizon", "groups"):
+ if not torch.equal(reference[key], population[key]):
+ raise ValueError(
+ f"challenge mode {name} is not paired on {key}"
+ )
+ return populations
+
+
+def _mean_by_label(values, labels, label):
+ selected = values[labels == label]
+ if not selected.numel():
+ return 0.0
+ return selected.double().mean().item()
+
+
+def signature_metrics_v2(train_events, challenge_mode_events, cfg, role):
+ """Measure nonterminal innovation and terminal outcome surprise."""
+ train = flatten_active(train_events)
+ role = torch.as_tensor(role, dtype=torch.float64)
+ nonterminal = ~train["terminal"].bool()
+ if not bool(nonterminal.any()):
+ raise ValueError("v2 training needs nonterminal events")
+
+ soma_corr = population_correlations(
+ train, "innovation", "soma", nonterminal
+ )
+ raw_corr = population_correlations(
+ train, "raw_apical", "soma", nonterminal
+ )
+
+ event_decoder = []
+ confidence_corr = []
+ for cell in range(cfg.n_plus + cfg.n_minus):
+ keep = [
+ index for index in range(cfg.n_neurons)
+ if index != cell
+ ]
+ labels = train["innovation"][nonterminal, cell] > 0
+ accuracy, distance = grouped_classification_accuracy(
+ train["previous_soma"][nonterminal][:, keep],
+ labels,
+ train["groups"][nonterminal],
+ )
+ event_decoder.append(accuracy)
+ confidence_corr.append(
+ pearson(
+ distance,
+ train["innovation"][nonterminal, cell],
+ )
+ )
+
+ difference = (
+ train["innovation"][nonterminal, :cfg.n_plus].mean(1)
+ - train["innovation"][
+ nonterminal,
+ cfg.n_plus:cfg.n_plus + cfg.n_minus,
+ ].mean(1)
+ )
+ improvement = (
+ train["previous_abs_error"][nonterminal]
+ - train["abs_error"][nonterminal]
+ )
+ improving = improvement > 0
+ worsening = improvement < 0
+ sign_index = 0.0
+ if bool(improving.any()) and bool(worsening.any()):
+ sign_index = 0.5 * (
+ difference[improving].mean()
+ - difference[worsening].mean()
+ ).item()
+ aligned = train["innovation"][nonterminal] @ role
+ velocity_corr = grouped_regression_correlation(
+ aligned, improvement, train["groups"][nonterminal]
+ )
+ error_corr = grouped_regression_correlation(
+ aligned,
+ train["abs_error"][nonterminal],
+ train["groups"][nonterminal],
+ )
+
+ populations = _paired_terminal_populations(challenge_mode_events)
+ intact = populations["intact"]
+ critic_lesion = populations["acute_critic_lesion"]
+ outcome_lesion = populations["acute_outcome_lesion"]
+ labels = intact["outcome"]
+ groups = intact["groups"]
+ residual_accuracy, _ = grouped_classification_accuracy(
+ intact["innovation"], labels, groups
+ )
+ previous_soma_accuracy, _ = grouped_classification_accuracy(
+ intact["previous_soma"], labels, groups
+ )
+ outcome_lesion_accuracy, _ = grouped_classification_accuracy(
+ outcome_lesion["innovation"], labels, groups
+ )
+
+ scores = {
+ name: population["innovation"] @ role
+ for name, population in populations.items()
+ }
+ intact_separation = (
+ _mean_by_label(scores["intact"], labels, True)
+ - _mean_by_label(scores["intact"], labels, False)
+ )
+ outcome_lesion_separation = (
+ _mean_by_label(scores["acute_outcome_lesion"], labels, True)
+ - _mean_by_label(
+ scores["acute_outcome_lesion"], labels, False
+ )
+ )
+ critic_contribution = (
+ scores["acute_critic_lesion"] - scores["intact"]
+ )
+ critic_prediction_corr = pearson(
+ critic_contribution, intact["value_prediction"]
+ )
+
+ success_fraction = labels.double().mean().item()
+ horizon_success = {}
+ for horizon in sorted(set(intact["horizon"].tolist())):
+ selected = intact["horizon"] == horizon
+ horizon_success[str(horizon)] = (
+ labels[selected].double().mean().item()
+ )
+
+ values = {
+ "mean_abs_residual_soma_corr":
+ sum(abs(value) for value in soma_corr) / len(soma_corr),
+ "mean_abs_raw_soma_corr":
+ sum(abs(value) for value in raw_corr) / len(raw_corr),
+ "raw_minus_residual_abs_soma_corr": (
+ sum(abs(value) for value in raw_corr) / len(raw_corr)
+ - sum(abs(value) for value in soma_corr) / len(soma_corr)
+ ),
+ "surrounding_event_decoder_balanced_acc":
+ sum(event_decoder) / len(event_decoder),
+ "decoder_distance_residual_corr":
+ sum(confidence_corr) / len(confidence_corr),
+ "causal_role_sign_inversion_index": sign_index,
+ "role_aligned_velocity_cv_corr": velocity_corr,
+ "role_aligned_error_cv_corr": error_corr,
+ "velocity_minus_error_abs_cv_corr":
+ abs(velocity_corr) - abs(error_corr),
+ "terminal_residual_outcome_balanced_acc": residual_accuracy,
+ "terminal_previous_soma_outcome_balanced_acc":
+ previous_soma_accuracy,
+ "terminal_residual_minus_previous_soma_acc":
+ residual_accuracy - previous_soma_accuracy,
+ "acute_outcome_lesion_outcome_balanced_acc":
+ outcome_lesion_accuracy,
+ "terminal_role_aligned_outcome_separation": intact_separation,
+ "acute_outcome_lesion_role_aligned_separation":
+ outcome_lesion_separation,
+ "terminal_outcome_separation_drop_under_acute_lesion":
+ intact_separation - outcome_lesion_separation,
+ "mean_critic_expectedness_contribution":
+ critic_contribution.mean().item(),
+ "critic_contribution_value_prediction_corr":
+ critic_prediction_corr,
+ "challenge_success_fraction": success_fraction,
+ "challenge_success_count": int(labels.sum()),
+ "challenge_failure_count": int((~labels).sum()),
+ "challenge_episodes": int(labels.numel()),
+ "nonterminal_training_events": int(nonterminal.sum()),
+ "terminal_training_events": int(train["terminal"].sum()),
+ "horizon_success_fraction": horizon_success,
+ }
+ flat_values = {
+ key: value for key, value in values.items()
+ if isinstance(value, (int, float))
+ }
+ if not all(math.isfinite(value) for value in flat_values.values()):
+ raise ValueError("non-finite oral-B-v2 signature metric")
+ return values