summaryrefslogtreecommitdiff
path: root/rrog/collect_results.py
diff options
context:
space:
mode:
Diffstat (limited to 'rrog/collect_results.py')
-rw-r--r--rrog/collect_results.py26
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))