summaryrefslogtreecommitdiff
path: root/rrog
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-06-24 05:44:25 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-06-24 05:44:25 -0500
commit8102b84160d83b0b221505d68a915c746a686ef4 (patch)
tree3b118031611d1e43b68d935c6c4ba3895747edf7 /rrog
parentb73408289e875460af32c069546a588c4e5f6354 (diff)
Add stream ACT OGB runner
Diffstat (limited to 'rrog')
-rw-r--r--rrog/cli.py5
-rw-r--r--rrog/collect_results.py26
-rw-r--r--rrog/runspecs.py12
-rw-r--r--rrog/train_ogb_graphprop.py29
4 files changed, 64 insertions, 8 deletions
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"))