From de233443dbf06f7a8e3b04e7ec7e990096e6abb0 Mon Sep 17 00:00:00 2001 From: dlyr3 <152425764+dlyr3@users.noreply.github.com> Date: Sun, 21 Jun 2026 22:57:38 +0200 Subject: [PATCH 1/3] benchmark fix --- benchmarks/bench_optimizer_step.py | 9 ++++++++- 1 file changed, 8 insertions(+), 1 deletion(-) diff --git a/benchmarks/bench_optimizer_step.py b/benchmarks/bench_optimizer_step.py index 62ced5e6..2fbfdf49 100644 --- a/benchmarks/bench_optimizer_step.py +++ b/benchmarks/bench_optimizer_step.py @@ -58,14 +58,20 @@ def main( gen = torch.Generator(device="cuda").manual_seed(seed) params = [] + grads = [] for dims in shapes: param = torch.nn.Parameter(torch.randn(dims, device="cuda", dtype=torch_dtype, generator=gen)) - param.grad = torch.randn(dims, device="cuda", dtype=torch_dtype, generator=gen) + grads.append(torch.randn(dims, device="cuda", dtype=torch_dtype, generator=gen)) params.append(param) + def set_grads(): + for param, grad in zip(params, grads): + param.grad = grad + module = heavyball if library is Library.heavyball else torch.optim step = getattr(module, optimizer)(params, **kwargs).step for _ in range(warmup): + set_grads() step() times = [] @@ -73,6 +79,7 @@ def main( torch.cuda.synchronize() start = perf_counter() for _ in range(steps): + set_grads() step() torch.cuda.synchronize() times.append((perf_counter() - start) / steps) From 80545ff7942c9b3a6bfad4253a50df8809a32b31 Mon Sep 17 00:00:00 2001 From: dlyr3 <152425764+dlyr3@users.noreply.github.com> Date: Sun, 21 Jun 2026 23:57:22 +0200 Subject: [PATCH 2/3] Revert "benchmark fix" This reverts commit de233443dbf06f7a8e3b04e7ec7e990096e6abb0. --- benchmarks/bench_optimizer_step.py | 9 +-------- 1 file changed, 1 insertion(+), 8 deletions(-) diff --git a/benchmarks/bench_optimizer_step.py b/benchmarks/bench_optimizer_step.py index 2fbfdf49..62ced5e6 100644 --- a/benchmarks/bench_optimizer_step.py +++ b/benchmarks/bench_optimizer_step.py @@ -58,20 +58,14 @@ def main( gen = torch.Generator(device="cuda").manual_seed(seed) params = [] - grads = [] for dims in shapes: param = torch.nn.Parameter(torch.randn(dims, device="cuda", dtype=torch_dtype, generator=gen)) - grads.append(torch.randn(dims, device="cuda", dtype=torch_dtype, generator=gen)) + param.grad = torch.randn(dims, device="cuda", dtype=torch_dtype, generator=gen) params.append(param) - def set_grads(): - for param, grad in zip(params, grads): - param.grad = grad - module = heavyball if library is Library.heavyball else torch.optim step = getattr(module, optimizer)(params, **kwargs).step for _ in range(warmup): - set_grads() step() times = [] @@ -79,7 +73,6 @@ def set_grads(): torch.cuda.synchronize() start = perf_counter() for _ in range(steps): - set_grads() step() torch.cuda.synchronize() times.append((perf_counter() - start) / steps) From ca58638150b64e52a4bbaaf915a60c284127684a Mon Sep 17 00:00:00 2001 From: dlyr3 <152425764+dlyr3@users.noreply.github.com> Date: Sun, 21 Jun 2026 23:58:58 +0200 Subject: [PATCH 3/3] codex fix --- benchmarks/bench_optimizer_step.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/benchmarks/bench_optimizer_step.py b/benchmarks/bench_optimizer_step.py index 62ced5e6..90b6225b 100644 --- a/benchmarks/bench_optimizer_step.py +++ b/benchmarks/bench_optimizer_step.py @@ -50,7 +50,7 @@ def main( ): shapes = DEFAULT_SHAPES if shape is None else tuple(map(parse_shape, shape)) torch_dtype = getattr(torch, dtype) - kwargs = {"compile_step": compile_step} if library is Library.heavyball else {} + kwargs = {"compile_step": compile_step, "consume_grad": False} if library is Library.heavyball else {} if fused is not None and library is Library.torch: kwargs["fused"] = fused if update_precond is not None and library is Library.heavyball: