From 8102b84160d83b0b221505d68a915c746a686ef4 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Wed, 24 Jun 2026 05:44:25 -0500 Subject: Add stream ACT OGB runner --- rrog/cli.py | 5 ++++- rrog/collect_results.py | 26 +++++++++++++++++++++++++- rrog/runspecs.py | 12 +++++++++++- rrog/train_ogb_graphprop.py | 29 ++++++++++++++++++++++++----- 4 files changed, 64 insertions(+), 8 deletions(-) (limited to 'rrog') diff --git a/rrog/cli.py b/rrog/cli.py index 9c826e3..5055b2e 100644 --- a/rrog/cli.py +++ b/rrog/cli.py @@ -52,7 +52,8 @@ def build_command(args) -> list[str]: run_args = dict(spec.default_args) for key in [ "epochs", "hidden", "bs", "seed", "T", "n_sup", "halt_max_steps", "halt_target", - "halt_min_steps", "halt_loss_threshold", "q_warmup_epochs", + "halt_min_steps", "halt_loss_threshold", "halt_exploration_prob", "q_warmup_epochs", + "act_train_mode", "eval_every", "max_train_batches", "max_eval_batches", "num_workers", "ema", "lr", "lam_q", "device", ]: @@ -117,7 +118,9 @@ def main(): rp.add_argument("--halt_target", choices=["soft", "binary", "exact", "loss"]) rp.add_argument("--halt_min_steps", type=int) rp.add_argument("--halt_loss_threshold", type=float) + rp.add_argument("--halt_exploration_prob", type=float) rp.add_argument("--q_warmup_epochs", type=int) + rp.add_argument("--act_train_mode", choices=["stream", "trace"]) rp.add_argument("--eval_every", type=int) rp.add_argument("--max_train_batches", type=int) rp.add_argument("--max_eval_batches", type=int) 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)) diff --git a/rrog/runspecs.py b/rrog/runspecs.py index 285348f..17c5ef4 100644 --- a/rrog/runspecs.py +++ b/rrog/runspecs.py @@ -99,6 +99,7 @@ def _ogb_graphprop_gin_command(args: dict[str, object]) -> list[str]: "--halt_target", str(args.get("halt_target", "loss")), "--halt_loss_threshold", str(args.get("halt_loss_threshold", 0.2)), "--halt_exploration_prob", str(args.get("halt_exploration_prob", 0.1)), + "--act_train_mode", str(args.get("act_train_mode", "stream")), ]) if args.get("q_warmup_epochs") is not None: cmd.extend(["--q_warmup_epochs", str(args["q_warmup_epochs"])]) @@ -175,7 +176,16 @@ for _task in OGB_MOL_TASKS: task=_task, view=_view, compute="rrog-act", - default_args={"task": _task, "view": _view, "compute": "rrog-act", "T": 1, "n_sup": 3, "epochs": 100}, + default_args={ + "task": _task, + "view": _view, + "compute": "rrog-act", + "T": 1, + "n_sup": 3, + "epochs": 100, + "act_train_mode": "stream", + "lam_q": 0.1, + }, command_builder=_ogb_graphprop_gin_command, ), ]) diff --git a/rrog/train_ogb_graphprop.py b/rrog/train_ogb_graphprop.py index 387ef3c..0961359 100644 --- a/rrog/train_ogb_graphprop.py +++ b/rrog/train_ogb_graphprop.py @@ -367,7 +367,7 @@ def _split_nodes(t, ptr): return [t[ptr[i].item():ptr[i + 1].item()].detach() for i in range(ptr.numel() - 1)] -def act_train_step(model, state, replacement_batch, opt, dev, args, metric): +def act_train_step(model, state, replacement_batch, opt, dev, args, epoch, metric): replacement = replacement_batch.to_data_list() batch_size = len(replacement) if state is None: @@ -410,8 +410,12 @@ def act_train_step(model, state, replacement_batch, opt, dev, args, metric): target = ((per_graph_loss <= args.halt_loss_threshold) & has_label).to(logits.dtype) else: raise ValueError(args.halt_target) - q_loss = nn.functional.binary_cross_entropy_with_logits(q, target) - loss = pred_loss + 0.5 * args.lam_q * q_loss + if epoch <= args.q_warmup_epochs: + q_loss = pred_loss.detach() * 0.0 + loss = pred_loss + else: + q_loss = nn.functional.binary_cross_entropy_with_logits(q, target) + loss = pred_loss + 0.5 * args.lam_q * q_loss y_det, z_det = y.detach(), z.detach() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) @@ -499,7 +503,10 @@ def train_epoch(model, loader, opt, dev, args, act_state, ema_state, epoch, metr if args.max_train_batches and i >= args.max_train_batches: break if args.compute == "rrog-act": - m = act_trace_train_step(model, batch, opt, dev, args, epoch, metric) + if args.act_train_mode == "stream": + act_state, m = act_train_step(model, act_state, batch, opt, dev, args, epoch, metric) + else: + m = act_trace_train_step(model, batch, opt, dev, args, epoch, metric) update_ema_state(ema_state, model, args.ema) metrics.append(m) continue @@ -540,6 +547,7 @@ def main(): ap.add_argument("--halt_loss_threshold", type=float, default=0.2) ap.add_argument("--halt_exploration_prob", type=float, default=0.1) ap.add_argument("--q_warmup_epochs", type=int, default=0) + ap.add_argument("--act_train_mode", choices=["stream", "trace"], default="stream") ap.add_argument("--ema", type=float, default=0.0) ap.add_argument("--seed", type=int, default=0) ap.add_argument("--num_workers", type=int, default=0) @@ -643,8 +651,17 @@ def main(): print(msg, flush=True) ema_tag = f"_ema{args.ema:g}" if args.ema > 0 else "" + act_tag = "" + if args.compute == "rrog-act": + target_tag = args.halt_target + if args.halt_target == "loss": + target_tag += f"{args.halt_loss_threshold:g}" + act_tag = ( + f"_{args.act_train_mode}_hm{args.halt_max_steps}_hmin{args.halt_min_steps}_" + f"{target_tag}_lq{args.lam_q:g}_hex{args.halt_exploration_prob:g}_qw{args.q_warmup_epochs}" + ) tag = ( - f"{args.dataset}_{args.view}_{args.compute}_T{T}_ns{args.n_sup}_" + f"{args.dataset}_{args.view}_{args.compute}_T{T}_ns{args.n_sup}{act_tag}_" f"h{args.hidden}_e{args.epochs}{ema_tag}_s{args.seed}" ) rep = { @@ -673,6 +690,8 @@ def main(): "halt_min_steps": args.halt_min_steps, "halt_target": args.halt_target, "halt_loss_threshold": args.halt_loss_threshold, + "halt_exploration_prob": args.halt_exploration_prob, + "act_train_mode": args.act_train_mode, "view": args.view, }, }, os.path.join(OUT, f"ckpt_{tag}.pt")) -- cgit v1.2.3