summaryrefslogtreecommitdiff
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
parentb73408289e875460af32c069546a588c4e5f6354 (diff)
Add stream ACT OGB runner
-rw-r--r--README.md39
-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
-rwxr-xr-xscripts/run_ogb_act_task.sh119
-rwxr-xr-xscripts/run_ogb_act_two_gpu.sh87
7 files changed, 309 insertions, 8 deletions
diff --git a/README.md b/README.md
index 2b55983..a3c7e7d 100644
--- a/README.md
+++ b/README.md
@@ -80,6 +80,34 @@ Run all selected OGB molecular tasks serially on one GPU:
DEVICE=cuda:1 ./scripts/run_ogb_mol_all_tasks.sh
```
+Run the corrected stream-ACT sweep on two GPUs:
+
+```bash
+EPOCHS=100 SEEDS=0 ./scripts/run_ogb_act_two_gpu.sh
+```
+
+Defaults:
+
+- GPU0: `ogbg-molhiv ogbg-molbbbp ogbg-molsider ogbg-molbace`
+- GPU1: `ogbg-molesol ogbg-mollipo ogbg-moltox21 ogbg-molclintox`
+- Every task runs all 17 backbones.
+- ACT config: `T=1`, `n_sup=3`, `halt_max=8`, `halt_min=2`, `halt_target=loss`, `loss_threshold=0.2`, `halt_exploration=0.1`, `lam_q=0.1`, `q_warmup=0`, `act_train_mode=stream`.
+
+Optional ACT variants:
+
+```bash
+# Add FreeSolv as a separate regression stress test.
+TASKS_GPU0="ogbg-molfreesolv" TASKS_GPU1="" ./scripts/run_ogb_act_two_gpu.sh
+
+# Classification-only exact-halt target.
+TASKS_GPU0="ogbg-molhiv ogbg-molbbbp ogbg-molsider" \
+TASKS_GPU1="ogbg-molbace ogbg-moltox21 ogbg-molclintox" \
+HALT_TARGET=exact ./scripts/run_ogb_act_two_gpu.sh
+
+# More robust but longer seed sweep.
+SEEDS="0 1 2" ./scripts/run_ogb_act_two_gpu.sh
+```
+
Collect summaries:
```bash
@@ -104,3 +132,14 @@ For OGB molecular tasks, GINE and edge-aware backbones use OGB bond encodings.
- ZINC cycle-count cache is generated under `data/cycle_cache`.
- OGB datasets are downloaded under `data/ogb`.
- Override data/runs locations with `RROG_DATA_DIR` and `RROG_RUNS_DIR`.
+
+## Upload Results
+
+After a remote machine finishes:
+
+```bash
+git pull
+git add -f runs/*.json logs/*.log summaries/*.md
+git commit -m "Add stream ACT OGB results"
+git push
+```
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"))
diff --git a/scripts/run_ogb_act_task.sh b/scripts/run_ogb_act_task.sh
new file mode 100755
index 0000000..37eb1a6
--- /dev/null
+++ b/scripts/run_ogb_act_task.sh
@@ -0,0 +1,119 @@
+#!/usr/bin/env bash
+set -euo pipefail
+
+ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
+cd "${ROOT_DIR}"
+export PYTHONPATH="${ROOT_DIR}:${PYTHONPATH:-}"
+
+TASK="${TASK:-ogbg-molhiv}"
+DEVICE="${DEVICE:-cuda:0}"
+EPOCHS="${EPOCHS:-100}"
+SEEDS="${SEEDS:-${SEED:-0}}"
+HIDDEN="${HIDDEN:-128}"
+BS="${BS:-128}"
+LR="${LR:-}"
+EVAL_EVERY="${EVAL_EVERY:-10}"
+NUM_WORKERS="${NUM_WORKERS:-0}"
+T="${T:-1}"
+N_SUP="${N_SUP:-3}"
+HALT_MAX="${HALT_MAX:-8}"
+HALT_MIN="${HALT_MIN:-2}"
+HALT_TARGET="${HALT_TARGET:-loss}"
+HALT_LOSS_THRESHOLD="${HALT_LOSS_THRESHOLD:-0.2}"
+HALT_EXPLORATION_PROB="${HALT_EXPLORATION_PROB:-0.1}"
+LAM_Q="${LAM_Q:-0.1}"
+Q_WARMUP="${Q_WARMUP:-0}"
+ACT_TRAIN_MODE="${ACT_TRAIN_MODE:-stream}"
+EMA="${EMA:-0}"
+MAX_TRAIN_BATCHES="${MAX_TRAIN_BATCHES:-}"
+MAX_EVAL_BATCHES="${MAX_EVAL_BATCHES:-}"
+COLLECT="${COLLECT:-1}"
+VIEWS="${VIEWS:-gin gine gcn graphsage gatv2 graphconv transformer pna gen film resgated tag sgc cheb arma mf appnp}"
+
+mkdir -p runs logs summaries
+
+fmt_float() {
+ python3 - "$1" <<'PY'
+import sys
+print(f"{float(sys.argv[1]):g}")
+PY
+}
+
+result_path() {
+ local view="$1"
+ local seed="$2"
+ local target_tag="${HALT_TARGET}"
+ local loss_tag
+ local lam_tag
+ local hex_tag
+ local ema_tag=""
+ loss_tag="$(fmt_float "${HALT_LOSS_THRESHOLD}")"
+ lam_tag="$(fmt_float "${LAM_Q}")"
+ hex_tag="$(fmt_float "${HALT_EXPLORATION_PROB}")"
+ if [[ "${HALT_TARGET}" == "loss" ]]; then
+ target_tag="loss${loss_tag}"
+ fi
+ if [[ "$(fmt_float "${EMA}")" != "0" ]]; then
+ ema_tag="_ema$(fmt_float "${EMA}")"
+ fi
+ echo "runs/${TASK}_${view}_rrog-act_T${T}_ns${N_SUP}_${ACT_TRAIN_MODE}_hm${HALT_MAX}_hmin${HALT_MIN}_${target_tag}_lq${lam_tag}_hex${hex_tag}_qw${Q_WARMUP}_h${HIDDEN}_e${EPOCHS}${ema_tag}_s${seed}.json"
+}
+
+run_cell() {
+ local view="$1"
+ local seed="$2"
+ local out
+ out="$(result_path "${view}" "${seed}")"
+ if [[ -f "${out}" ]]; then
+ echo "[skip] ${out}"
+ return
+ fi
+
+ echo "[run] ${TASK} view=${view} compute=rrog-act mode=${ACT_TRAIN_MODE} T=${T} ns=${N_SUP} seed=${seed} device=${DEVICE}"
+ cmd=(
+ python3 -m rrog.cli run
+ --task "${TASK}"
+ --view "${view}"
+ --compute rrog-act
+ --epochs "${EPOCHS}"
+ --hidden "${HIDDEN}"
+ --bs "${BS}"
+ --T "${T}"
+ --n_sup "${N_SUP}"
+ --halt_max_steps "${HALT_MAX}"
+ --halt_min_steps "${HALT_MIN}"
+ --halt_target "${HALT_TARGET}"
+ --halt_loss_threshold "${HALT_LOSS_THRESHOLD}"
+ --halt_exploration_prob "${HALT_EXPLORATION_PROB}"
+ --lam_q "${LAM_Q}"
+ --q_warmup_epochs "${Q_WARMUP}"
+ --act_train_mode "${ACT_TRAIN_MODE}"
+ --eval_every "${EVAL_EVERY}"
+ --num_workers "${NUM_WORKERS}"
+ --seed "${seed}"
+ --device "${DEVICE}"
+ )
+ if [[ -n "${LR}" ]]; then
+ cmd+=(--lr "${LR}")
+ fi
+ if [[ "$(fmt_float "${EMA}")" != "0" ]]; then
+ cmd+=(--ema "${EMA}")
+ fi
+ if [[ -n "${MAX_TRAIN_BATCHES}" ]]; then
+ cmd+=(--max_train_batches "${MAX_TRAIN_BATCHES}")
+ fi
+ if [[ -n "${MAX_EVAL_BATCHES}" ]]; then
+ cmd+=(--max_eval_batches "${MAX_EVAL_BATCHES}")
+ fi
+ "${cmd[@]}"
+}
+
+for seed in ${SEEDS}; do
+ for view in ${VIEWS}; do
+ run_cell "${view}" "${seed}"
+ done
+done
+
+if [[ "${COLLECT}" == "1" ]]; then
+ python3 -m rrog.cli results --epochs "${EPOCHS}" | tee "summaries/ogb_graphprop_act_${TASK}_e${EPOCHS}.md"
+fi
diff --git a/scripts/run_ogb_act_two_gpu.sh b/scripts/run_ogb_act_two_gpu.sh
new file mode 100755
index 0000000..3c31176
--- /dev/null
+++ b/scripts/run_ogb_act_two_gpu.sh
@@ -0,0 +1,87 @@
+#!/usr/bin/env bash
+set -euo pipefail
+
+ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
+cd "${ROOT_DIR}"
+export PYTHONPATH="${ROOT_DIR}:${PYTHONPATH:-}"
+
+GPU0="${GPU0:-cuda:0}"
+GPU1="${GPU1:-cuda:1}"
+TASKS_GPU0="${TASKS_GPU0:-ogbg-molhiv ogbg-molbbbp ogbg-molsider ogbg-molbace}"
+TASKS_GPU1="${TASKS_GPU1:-ogbg-molesol ogbg-mollipo ogbg-moltox21 ogbg-molclintox}"
+EPOCHS="${EPOCHS:-100}"
+SEEDS="${SEEDS:-${SEED:-0}}"
+HALT_MAX="${HALT_MAX:-8}"
+HALT_MIN="${HALT_MIN:-2}"
+HALT_TARGET="${HALT_TARGET:-loss}"
+HALT_LOSS_THRESHOLD="${HALT_LOSS_THRESHOLD:-0.2}"
+HALT_EXPLORATION_PROB="${HALT_EXPLORATION_PROB:-0.1}"
+LAM_Q="${LAM_Q:-0.1}"
+Q_WARMUP="${Q_WARMUP:-0}"
+ACT_TRAIN_MODE="${ACT_TRAIN_MODE:-stream}"
+
+mkdir -p logs summaries
+
+fmt_float() {
+ python3 - "$1" <<'PY'
+import sys
+print(f"{float(sys.argv[1]):g}")
+PY
+}
+
+target_log_tag() {
+ local target_tag="${HALT_TARGET}"
+ if [[ "${HALT_TARGET}" == "loss" ]]; then
+ target_tag="loss$(fmt_float "${HALT_LOSS_THRESHOLD}")"
+ fi
+ echo "${ACT_TRAIN_MODE}_hm${HALT_MAX}_hmin${HALT_MIN}_${target_tag}_lq$(fmt_float "${LAM_Q}")_hex$(fmt_float "${HALT_EXPLORATION_PROB}")_qw${Q_WARMUP}_e${EPOCHS}_s${SEEDS// /-}"
+}
+
+run_queue() {
+ local device="$1"
+ shift
+ local tasks=("$@")
+ local task
+ local tag
+ tag="$(target_log_tag)"
+ for task in "${tasks[@]}"; do
+ if [[ -z "${task}" ]]; then
+ continue
+ fi
+ echo "[task] ${task} on ${device}"
+ TASK="${task}" DEVICE="${device}" EPOCHS="${EPOCHS}" SEEDS="${SEEDS}" \
+ HALT_MAX="${HALT_MAX}" HALT_MIN="${HALT_MIN}" HALT_TARGET="${HALT_TARGET}" \
+ HALT_LOSS_THRESHOLD="${HALT_LOSS_THRESHOLD}" HALT_EXPLORATION_PROB="${HALT_EXPLORATION_PROB}" \
+ LAM_Q="${LAM_Q}" Q_WARMUP="${Q_WARMUP}" \
+ ACT_TRAIN_MODE="${ACT_TRAIN_MODE}" COLLECT=0 \
+ ./scripts/run_ogb_act_task.sh 2>&1 | tee "logs/${task}_act_${tag}.log"
+ done
+}
+
+tasks0=()
+tasks1=()
+if [[ -n "${TASKS_GPU0}" ]]; then
+ read -r -a tasks0 <<< "${TASKS_GPU0}"
+fi
+if [[ -n "${TASKS_GPU1}" ]]; then
+ read -r -a tasks1 <<< "${TASKS_GPU1}"
+fi
+
+pids=()
+if (( ${#tasks0[@]} > 0 )); then
+ echo "[launch] ${GPU0}: ${tasks0[*]}"
+ run_queue "${GPU0}" "${tasks0[@]}" &
+ pids+=("$!")
+fi
+if (( ${#tasks1[@]} > 0 )); then
+ echo "[launch] ${GPU1}: ${tasks1[*]}"
+ run_queue "${GPU1}" "${tasks1[@]}" &
+ pids+=("$!")
+fi
+
+for pid in "${pids[@]}"; do
+ wait "${pid}"
+done
+
+echo "[done] collecting summaries"
+OGB_EPOCHS="${EPOCHS}" ./scripts/collect_results.sh