diff options
Diffstat (limited to 'experiments/analyze_kp_innovation_short.py')
| -rw-r--r-- | experiments/analyze_kp_innovation_short.py | 30 |
1 files changed, 26 insertions, 4 deletions
diff --git a/experiments/analyze_kp_innovation_short.py b/experiments/analyze_kp_innovation_short.py index 2e7b9f4..89d21ff 100644 --- a/experiments/analyze_kp_innovation_short.py +++ b/experiments/analyze_kp_innovation_short.py @@ -17,6 +17,20 @@ def mean_early(values): return sum(float(value) for value in values[:count]) / count +def numeric_leaves(value): + """Yield every numeric audit value while excluding boolean flags.""" + if isinstance(value, bool) or value is None: + return + if isinstance(value, (int, float)): + yield float(value) + elif isinstance(value, dict): + for child in value.values(): + yield from numeric_leaves(child) + elif isinstance(value, (list, tuple)): + for child in value: + yield from numeric_leaves(child) + + def main(): parser = argparse.ArgumentParser() parser.add_argument("--input_dir", default="results/kp_innovation_short") @@ -102,10 +116,6 @@ def main(): float(mixed["traffic_rms"]), ]) - all_finite = all(record["final"]["finite"] for record in records.values()) - all_finite = all_finite and all(math.isfinite(value) for value in ( - trajectory_values + list(accuracies.values()) + [innovation_early, - innovation_raw_early])) initial_ratio_errors = [] predictor_residual_ratios = [] total_macs = {} @@ -125,6 +135,18 @@ def main(): diagnostics["innovation"]["mean_feedback_forward_cosine"]) late_feedback_cosine = sum(float(value["mean_feedback_forward_cosine"]) for value in tracking["innovation"][10:]) / 10 + audit_values = trajectory_values + [innovation_early, innovation_raw_early, + final_feedback_cosine, + late_feedback_cosine, + matched_norm_error] + for record in records.values(): + audit_values.extend(numeric_leaves(record["final"])) + audit_values.extend(numeric_leaves(record["diagnostics"])) + audit_values.extend(numeric_leaves(record["traffic_calibration"])) + audit_values.extend(numeric_leaves(record["predictor_warmup"])) + all_finite = all(record["final"]["finite"] for record in records.values()) + all_finite = all_finite and all( + math.isfinite(value) for value in audit_values) checks = { "records_trajectories_and_diagnostics_finite": all_finite, |
