diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-23 07:53:33 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-23 07:53:33 -0500 |
| commit | ce10335dcc3aa12f4f1e6dd53352d7e71d0cccf7 (patch) | |
| tree | 1389b61c3e6a1dbdabe7df557128c16e5a3f2122 /sdil | |
| parent | c34675664b949ac7c3aa7a61e714d69183d0ea6f (diff) | |
bci: audit v2 outcome-surprise signatures
Diffstat (limited to 'sdil')
| -rw-r--r-- | sdil/bci_v2.py | 15 | ||||
| -rw-r--r-- | sdil/bci_v2_metrics.py | 303 |
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 |
