diff options
Diffstat (limited to 'rrog/train_ogb_graphprop.py')
| -rw-r--r-- | rrog/train_ogb_graphprop.py | 29 |
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")) |
