diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-06-24 05:44:25 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-06-24 05:44:25 -0500 |
| commit | 8102b84160d83b0b221505d68a915c746a686ef4 (patch) | |
| tree | 3b118031611d1e43b68d935c6c4ba3895747edf7 | |
| parent | b73408289e875460af32c069546a588c4e5f6354 (diff) | |
Add stream ACT OGB runner
| -rw-r--r-- | README.md | 39 | ||||
| -rw-r--r-- | rrog/cli.py | 5 | ||||
| -rw-r--r-- | rrog/collect_results.py | 26 | ||||
| -rw-r--r-- | rrog/runspecs.py | 12 | ||||
| -rw-r--r-- | rrog/train_ogb_graphprop.py | 29 | ||||
| -rwxr-xr-x | scripts/run_ogb_act_task.sh | 119 | ||||
| -rwxr-xr-x | scripts/run_ogb_act_two_gpu.sh | 87 |
7 files changed, 309 insertions, 8 deletions
@@ -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 |
