summaryrefslogtreecommitdiff
path: root/experiments/conv_run.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/conv_run.py')
-rw-r--r--experiments/conv_run.py16
1 files changed, 14 insertions, 2 deletions
diff --git a/experiments/conv_run.py b/experiments/conv_run.py
index 9b41dd2..98c1f19 100644
--- a/experiments/conv_run.py
+++ b/experiments/conv_run.py
@@ -235,8 +235,14 @@ def run(args):
"epochs": [],
}
+ sync(args.device)
+ total_start = time.time()
+ predictor_warmup_wall = 0.0
+ apical_warmup_wall = 0.0
loader_state = train.g.get_state().clone()
if config is not None and config.learn_P and args.predictor_warmup_steps:
+ sync(args.device)
+ warmup_start = time.time()
iterator = iter(train)
for _ in range(args.predictor_warmup_steps):
try:
@@ -249,8 +255,12 @@ def run(args):
forward["hiddens"], config.eta_P, config.nuisance_scale)
counters["predictor_warmup_examples"] += x.shape[0]
train.g.set_state(loader_state)
+ sync(args.device)
+ predictor_warmup_wall = time.time() - warmup_start
if config is not None and args.a_warmup_steps:
+ sync(args.device)
+ warmup_start = time.time()
iterator = iter(train)
warmup_metrics = []
for _ in range(args.a_warmup_steps):
@@ -277,14 +287,14 @@ def run(args):
"steps": args.a_warmup_steps,
"last": warmup_metrics[-1],
}
+ sync(args.device)
+ apical_warmup_wall = time.time() - warmup_start
step = 0
train_wall = 0.0
eval_wall = 0.0
validation_evaluations = 0
test_evaluations = 0
- sync(args.device)
- total_start = time.time()
for epoch in range(args.epochs):
lr = scheduled_lr(args.lr, epoch, args)
output_lr = (scheduled_lr(args.output_lr, epoch, args)
@@ -393,6 +403,8 @@ def run(args):
"work": work_report(net, args.mode, counters),
"hardware": training_hardware,
"timing": {
+ "predictor_warmup_wall_s": predictor_warmup_wall,
+ "apical_warmup_wall_s": apical_warmup_wall,
"train_wall_s": train_wall,
"evaluation_wall_s": eval_wall,
"total_timed_wall_s": total_wall,