diff options
Diffstat (limited to 'rrog/collect_results.py')
| -rw-r--r-- | rrog/collect_results.py | 26 |
1 files changed, 25 insertions, 1 deletions
diff --git a/rrog/collect_results.py b/rrog/collect_results.py index 4c1390c..85925ba 100644 --- a/rrog/collect_results.py +++ b/rrog/collect_results.py @@ -87,6 +87,17 @@ def _compute_label(rep: dict) -> str: compute = str(rep["compute"]) if _is_classic_baseline(rep): label = "classic" + elif compute == "rrog-act": + target = str(rep.get("halt_target", "")) + if target == "loss": + target += f"{float(rep.get('halt_loss_threshold', 0.0) or 0.0):g}" + label = ( + f"{compute}-{rep.get('act_train_mode', 'stream')}-T{rep.get('T')}-ns{rep.get('n_sup')}" + f"-hm{rep.get('halt_max_steps')}-min{rep.get('halt_min_steps')}" + f"-{target}-lq{float(rep.get('lam_q', 0.0) or 0.0):g}" + f"-hex{float(rep.get('halt_exploration_prob', 0.0) or 0.0):g}" + f"-qw{rep.get('q_warmup_epochs', 0)}" + ) else: label = f"{compute}-T{rep.get('T')}-ns{rep.get('n_sup')}" ema = float(rep.get("ema", 0.0) or 0.0) @@ -186,6 +197,7 @@ def print_tables(args) -> None: base_val_mu = None val_scores, test_scores, val_deltas, test_deltas = [], [], [], [] + adaptive_test_scores, adaptive_test_deltas = [], [] adaptive_steps = [] for rep, base_rep in paired: val = _score(rep, "val") @@ -204,11 +216,21 @@ def print_tables(args) -> None: test_scores.append(test) val_deltas.append(direction * (val - base_val)) test_deltas.append(direction * (test - base_test)) + adaptive_test = _score(rep, "test_adaptive") + if adaptive_test is not None: + adaptive_test_scores.append(adaptive_test) + adaptive_test_deltas.append(direction * (adaptive_test - base_test)) if rep.get("adaptive_steps") is not None: adaptive_steps.append(float(rep["adaptive_steps"])) if not test_scores: continue + adaptive_test_cell = "" + if adaptive_test_scores: + adaptive_test_cell = ( + f"{_fmt(_mean(adaptive_test_scores), args.digits)} " + f"({_fmt(_mean(adaptive_test_deltas), args.digits)})" + ) delta_rows.append([ task, view, @@ -217,6 +239,7 @@ def print_tables(args) -> None: str(len(test_scores)), f"{_fmt(_mean(val_scores), args.digits)} ({_fmt(_mean(val_deltas), args.digits)})", f"{_fmt(_mean(test_scores), args.digits)} ({_fmt(_mean(test_deltas), args.digits)})", + adaptive_test_cell, _fmt(_mean(adaptive_steps), 2) if adaptive_steps else "", ]) @@ -224,7 +247,8 @@ def print_tables(args) -> None: print(_markdown_table(["task", "backbone", "metric", "n", "val", "test"], baseline_rows)) print("\nDelta vs matching classic") print(_markdown_table([ - "task", "backbone", "compute", "metric", "n", "val score (delta)", "test score (delta)", "steps" + "task", "backbone", "compute", "metric", "n", "val score (delta)", "test score (delta)", + "adaptive test (delta)", "steps" ], delta_rows)) |
