summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/larkum_public_pilot.py16
1 files changed, 12 insertions, 4 deletions
diff --git a/experiments/larkum_public_pilot.py b/experiments/larkum_public_pilot.py
index 67773c8..7e1be93 100644
--- a/experiments/larkum_public_pilot.py
+++ b/experiments/larkum_public_pilot.py
@@ -82,17 +82,18 @@ def _predict_trials(
def _template_predict(
branch_trials: list[np.ndarray],
- instructions: np.ndarray,
+ train_instructions: np.ndarray,
train_indices: np.ndarray,
+ test_instructions: np.ndarray,
test_indices: Iterable[int],
) -> dict[int, np.ndarray]:
templates: dict[int, np.ndarray] = {}
for side in (0, 1):
- selected = [branch_trials[i] for i in train_indices if instructions[i] == side]
+ selected = [branch_trials[i] for i in train_indices if train_instructions[i] == side]
if not selected:
raise ValueError(f"training fold has no instruction-{side} trials")
templates[side] = np.mean(selected, axis=0)
- return {int(i): templates[int(instructions[i])] for i in test_indices}
+ return {int(i): templates[int(test_instructions[i])] for i in test_indices}
def _summarize(trace: np.ndarray) -> np.ndarray:
@@ -158,7 +159,13 @@ def _session_trials(path: Path) -> list[Trial]:
model, _ = _fit_predictor(sal_spines, sal_branches, sal_instructions, train_indices)
crossfit_prediction.update(_predict_trials(model, sal_spines, sal_instructions, test_indices))
crossfit_template.update(
- _template_predict(sal_branches, sal_instructions, train_indices, test_indices)
+ _template_predict(
+ sal_branches,
+ sal_instructions,
+ train_indices,
+ sal_instructions,
+ test_indices,
+ )
)
all_sal_indices = np.arange(n_sal)
@@ -171,6 +178,7 @@ def _session_trials(path: Path) -> list[Trial]:
sal_branches,
sal_instructions,
all_sal_indices,
+ dcz["instructions"],
range(len(dcz["spines"])),
)