summaryrefslogtreecommitdiff
path: root/rrog/train_ogb_graphprop.py
diff options
context:
space:
mode:
Diffstat (limited to 'rrog/train_ogb_graphprop.py')
-rw-r--r--rrog/train_ogb_graphprop.py29
1 files changed, 24 insertions, 5 deletions
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"))