summaryrefslogtreecommitdiff
path: root/experiments/analyze_kp_innovation_short.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/analyze_kp_innovation_short.py')
-rw-r--r--experiments/analyze_kp_innovation_short.py8
1 files changed, 8 insertions, 0 deletions
diff --git a/experiments/analyze_kp_innovation_short.py b/experiments/analyze_kp_innovation_short.py
index 4c70575..5f21dea 100644
--- a/experiments/analyze_kp_innovation_short.py
+++ b/experiments/analyze_kp_innovation_short.py
@@ -82,12 +82,20 @@ def main():
if len(tracking[rule]) != 20 or any(value is None for value in tracking[rule]):
raise ValueError(f"MT-1 {rule} tracking trajectory is incomplete")
for row, values in zip(records[rule]["epochs"], tracking[rule]):
+ mixed = row.get("mixed_apical")
+ if mixed is None:
+ raise ValueError(f"MT-1 {rule} mixed-apical trajectory is incomplete")
trajectory_values.extend([
float(row["train_loss"]),
float(values["mean_feedback_forward_cosine"]),
float(values["mean_feedback_forward_relative_error"]),
float(values["min_feedback_forward_cosine"]),
float(values["max_feedback_forward_relative_error"]),
+ float(mixed["teaching_rms"]),
+ float(mixed["instruction_rms"]),
+ float(mixed["raw_apical_rms"]),
+ float(mixed["innovation_rms"]),
+ float(mixed["traffic_rms"]),
])
all_finite = all(record["final"]["finite"] for record in records.values())