From 5e8c8ad247504864151286cea86af00f7bcc2cc2 Mon Sep 17 00:00:00 2001 From: Aditya Sharma Date: Wed, 26 Aug 2026 22:12:58 +0530 Subject: [PATCH 1/2] Add challenge 114: Fused AdamW (Medium) --- .../medium/114_fused_adamw/challenge.html | 101 +++++++++ .../medium/114_fused_adamw/challenge.py | 191 ++++++++++++++++++ .../medium/114_fused_adamw/starter/starter.cu | 5 + .../114_fused_adamw/starter/starter.cute.py | 20 ++ .../114_fused_adamw/starter/starter.jax.py | 21 ++ .../114_fused_adamw/starter/starter.mojo | 22 ++ .../starter/starter.pytorch.py | 18 ++ .../114_fused_adamw/starter/starter.triton.py | 20 ++ 8 files changed, 398 insertions(+) create mode 100644 challenges/medium/114_fused_adamw/challenge.html create mode 100644 challenges/medium/114_fused_adamw/challenge.py create mode 100644 challenges/medium/114_fused_adamw/starter/starter.cu create mode 100644 challenges/medium/114_fused_adamw/starter/starter.cute.py create mode 100644 challenges/medium/114_fused_adamw/starter/starter.jax.py create mode 100644 challenges/medium/114_fused_adamw/starter/starter.mojo create mode 100644 challenges/medium/114_fused_adamw/starter/starter.pytorch.py create mode 100644 challenges/medium/114_fused_adamw/starter/starter.triton.py diff --git a/challenges/medium/114_fused_adamw/challenge.html b/challenges/medium/114_fused_adamw/challenge.html new file mode 100644 index 00000000..e54b2a14 --- /dev/null +++ b/challenges/medium/114_fused_adamw/challenge.html @@ -0,0 +1,101 @@ +

+ Implement a single fused AdamW optimizer step. Given a flat parameter vector + params, its gradient grad, and the optimizer's running + first and second moment estimates m and v, update + params, m, and v in place for one training + step. AdamW is the optimizer behind most modern LLM training runs (GPT, LLaMA, + and friends); fusing its handful of elementwise operations into a single kernel + avoids several redundant read–write passes over what can be a very large + parameter tensor, which is exactly what a fused=True optimizer does + under the hood. +

+ +

+ For step t (1-indexed), the update is: +

+

+ \[ + \begin{align} + \theta &\leftarrow \theta \cdot (1 - \eta \cdot \lambda) \\ + m &\leftarrow \beta_1 m + (1 - \beta_1) g \\ + v &\leftarrow \beta_2 v + (1 - \beta_2) g^2 \\ + \hat{m} &= \frac{m}{1 - \beta_1^{t}}, \quad + \hat{v} = \frac{v}{1 - \beta_2^{t}} \\ + \theta &\leftarrow \theta - \eta \cdot \frac{\hat{m}}{\sqrt{\hat{v}} + \epsilon} + \end{align} + \] +

+

+ where θ is params, g is + grad, η is the learning rate lr, + λ is weight_decay, and + β1, β2, + ε are beta1, beta2, + eps. The weight decay is decoupled: it is applied directly + to params rather than added into grad, which is the + detail that distinguishes AdamW from plain Adam with L2 regularization. +

+ +

Implementation Requirements

+ + +

Example 1

+
+N = 2, t = 1
+params = [1.0, -1.0]
+grad   = [0.1, -0.1]
+m      = [0.0,  0.0]
+v      = [0.0,  0.0]
+lr = 0.1, beta1 = 0.9, beta2 = 0.999, eps = 1e-8, weight_decay = 0.0
+
+m_new = 0.9*0 + 0.1*grad      = [0.01, -0.01]
+v_new = 0.999*0 + 0.001*grad^2 = [1e-5,  1e-5]
+m_hat = m_new / (1 - 0.9)      = [0.1,  -0.1]
+v_hat = v_new / (1 - 0.999)    = [0.01,  0.01]
+step  = lr * m_hat / (sqrt(v_hat) + eps) = [0.1, -0.1]
+
+params -> [0.9, -0.9]
+
+ +

Example 2

+
+N = 1, t = 1
+params = [1.0]
+grad   = [0.5]
+m      = [0.0]
+v      = [0.0]
+lr = 0.1, beta1 = 0.9, beta2 = 0.999, eps = 1e-8, weight_decay = 0.01
+
+params after decoupled decay: 1.0 * (1 - 0.1*0.01) = 0.999
+m_new = 0.05, v_new = 0.00025
+m_hat = 0.5,  v_hat = 0.25
+step  = 0.1 * 0.5 / (0.5 + 1e-8) = 0.1
+
+params -> 0.999 - 0.1 = 0.899
+
+ +

Constraints

+ diff --git a/challenges/medium/114_fused_adamw/challenge.py b/challenges/medium/114_fused_adamw/challenge.py new file mode 100644 index 00000000..61d53c58 --- /dev/null +++ b/challenges/medium/114_fused_adamw/challenge.py @@ -0,0 +1,191 @@ +import ctypes +from typing import Any, Dict, List + +import torch +from core.challenge_base import ChallengeBase + + +class Challenge(ChallengeBase): + name = "Fused AdamW" + atol = 1e-05 + rtol = 1e-05 + num_gpus = 1 + access_tier = "free" + + def reference_impl( + self, + params: torch.Tensor, + grad: torch.Tensor, + m: torch.Tensor, + v: torch.Tensor, + N: int, + lr: float, + beta1: float, + beta2: float, + eps: float, + weight_decay: float, + t: int, + ): + assert params.shape == (N,) + assert grad.shape == (N,) + assert m.shape == (N,) + assert v.shape == (N,) + assert params.dtype == grad.dtype == m.dtype == v.dtype == torch.float32 + assert t >= 1 + + # Decoupled weight decay: applied directly to the parameters, not folded + # into the gradient (this is what distinguishes AdamW from Adam + L2 reg). + params.mul_(1.0 - lr * weight_decay) + + # Update biased first and second raw moment estimates. + m.copy_(beta1 * m + (1.0 - beta1) * grad) + v.copy_(beta2 * v + (1.0 - beta2) * grad * grad) + + # Bias-correct the moment estimates. + bias_correction1 = 1.0 - beta1**t + bias_correction2 = 1.0 - beta2**t + m_hat = m / bias_correction1 + v_hat = v / bias_correction2 + + # Adaptive step. + params.sub_(lr * m_hat / (v_hat.sqrt() + eps)) + + def get_solve_signature(self) -> Dict[str, tuple]: + return { + "params": (ctypes.POINTER(ctypes.c_float), "inout"), + "grad": (ctypes.POINTER(ctypes.c_float), "in"), + "m": (ctypes.POINTER(ctypes.c_float), "inout"), + "v": (ctypes.POINTER(ctypes.c_float), "inout"), + "N": (ctypes.c_int, "in"), + "lr": (ctypes.c_float, "in"), + "beta1": (ctypes.c_float, "in"), + "beta2": (ctypes.c_float, "in"), + "eps": (ctypes.c_float, "in"), + "weight_decay": (ctypes.c_float, "in"), + "t": (ctypes.c_int, "in"), + } + + def _make_test_case( + self, + N, + lr=0.1, + beta1=0.9, + beta2=0.999, + eps=1e-8, + weight_decay=0.01, + t=1, + zero_state=True, + zero_grad=False, + param_range=(-2.0, 2.0), + grad_range=(-1.0, 1.0), + seed=0, + ): + device = self.device + dtype = torch.float32 + gen = torch.Generator(device=device).manual_seed(seed) + + params = torch.empty(N, device=device, dtype=dtype).uniform_(*param_range, generator=gen) + if zero_grad: + grad = torch.zeros(N, device=device, dtype=dtype) + else: + grad = torch.empty(N, device=device, dtype=dtype).uniform_(*grad_range, generator=gen) + if zero_state: + m = torch.zeros(N, device=device, dtype=dtype) + v = torch.zeros(N, device=device, dtype=dtype) + else: + # Simulate an optimizer that has already taken some steps: m carries + # the sign of typical gradients, v is a small positive magnitude. + m = torch.empty(N, device=device, dtype=dtype).uniform_(-0.5, 0.5, generator=gen) + v = torch.empty(N, device=device, dtype=dtype).uniform_(0.0, 0.5, generator=gen) + + return { + "params": params, + "grad": grad, + "m": m, + "v": v, + "N": N, + "lr": lr, + "beta1": beta1, + "beta2": beta2, + "eps": eps, + "weight_decay": weight_decay, + "t": t, + } + + def generate_example_test(self) -> Dict[str, Any]: + device = self.device + dtype = torch.float32 + params = torch.tensor([1.0, -1.0], device=device, dtype=dtype) + grad = torch.tensor([0.1, -0.1], device=device, dtype=dtype) + m = torch.zeros(2, device=device, dtype=dtype) + v = torch.zeros(2, device=device, dtype=dtype) + return { + "params": params, + "grad": grad, + "m": m, + "v": v, + "N": 2, + "lr": 0.1, + "beta1": 0.9, + "beta2": 0.999, + "eps": 1e-8, + "weight_decay": 0.0, + "t": 1, + } + + def generate_functional_test(self) -> List[Dict[str, Any]]: + tests = [] + + # Single element, first step, fresh (zero) optimizer state. + tests.append(self._make_test_case(N=1, t=1, zero_state=True, seed=1)) + + # Zero gradient, fresh (zero) optimizer state: m and v stay exactly zero, + # so the adaptive step is zero and only weight decay moves the parameters. + tests.append(self._make_test_case(N=8, t=5, zero_state=True, zero_grad=True, seed=2)) + + # Zero gradient with warm (nonzero) prior state: no new gradient signal + # arrives, but existing momentum still drives an adaptive step while m + # and v exponentially decay toward zero. + tests.append(self._make_test_case(N=8, t=5, zero_state=False, zero_grad=True, seed=11)) + + # Zero weight decay: reduces to plain Adam. + tests.append(self._make_test_case(N=16, t=1, weight_decay=0.0, zero_state=True, seed=3)) + + # Warm state, mid-training step count: bias correction is neither ~0 nor ~1. + tests.append(self._make_test_case(N=32, t=50, zero_state=False, seed=4)) + + # Very large step count: bias corrections both saturate close to 1. + tests.append(self._make_test_case(N=16, t=100000, zero_state=False, seed=5)) + + # Negative parameters and negative gradients throughout. + tests.append( + self._make_test_case( + N=16, + t=1, + zero_state=True, + param_range=(-2.0, -0.5), + grad_range=(-1.0, -0.2), + seed=6, + ) + ) + + # Non-default betas: catches implementations that hardcode 0.9 / 0.999. + tests.append( + self._make_test_case(N=64, t=10, beta1=0.8, beta2=0.99, zero_state=False, seed=7) + ) + + # Larger eps: shifts the denominator enough to matter numerically. + tests.append(self._make_test_case(N=64, t=1, eps=1e-2, zero_state=True, seed=8)) + + # Very small learning rate: update should be tiny but nonzero. + tests.append(self._make_test_case(N=64, t=1, lr=1e-6, zero_state=True, seed=9)) + + # Larger, realistic-ish size with default hyperparameters. + tests.append(self._make_test_case(N=1024, t=200, zero_state=False, seed=10)) + + return tests + + def generate_performance_test(self) -> Dict[str, Any]: + # 4096 x 4096 weight matrix flattened (LLaMA-7B-style hidden size), a + # realistic single-tensor slice of an optimizer step over a large model. + return self._make_test_case(N=4096 * 4096, t=1000, zero_state=False, seed=0) diff --git a/challenges/medium/114_fused_adamw/starter/starter.cu b/challenges/medium/114_fused_adamw/starter/starter.cu new file mode 100644 index 00000000..75e66e1d --- /dev/null +++ b/challenges/medium/114_fused_adamw/starter/starter.cu @@ -0,0 +1,5 @@ +#include + +// params, grad, m, v are device pointers +extern "C" void solve(float* params, const float* grad, float* m, float* v, int N, float lr, + float beta1, float beta2, float eps, float weight_decay, int t) {} diff --git a/challenges/medium/114_fused_adamw/starter/starter.cute.py b/challenges/medium/114_fused_adamw/starter/starter.cute.py new file mode 100644 index 00000000..c956a74d --- /dev/null +++ b/challenges/medium/114_fused_adamw/starter/starter.cute.py @@ -0,0 +1,20 @@ +import cutlass +import cutlass.cute as cute + + +# params, grad, m, v are tensors on the GPU +@cute.jit +def solve( + params: cute.Tensor, + grad: cute.Tensor, + m: cute.Tensor, + v: cute.Tensor, + N: cute.Int32, + lr: cute.Float32, + beta1: cute.Float32, + beta2: cute.Float32, + eps: cute.Float32, + weight_decay: cute.Float32, + t: cute.Int32, +): + pass diff --git a/challenges/medium/114_fused_adamw/starter/starter.jax.py b/challenges/medium/114_fused_adamw/starter/starter.jax.py new file mode 100644 index 00000000..a4f5b1e0 --- /dev/null +++ b/challenges/medium/114_fused_adamw/starter/starter.jax.py @@ -0,0 +1,21 @@ +import jax +import jax.numpy as jnp + + +# params, grad, m, v are tensors on device +@jax.jit +def solve( + params: jax.Array, + grad: jax.Array, + m: jax.Array, + v: jax.Array, + N: int, + lr: float, + beta1: float, + beta2: float, + eps: float, + weight_decay: float, + t: int, +) -> tuple[jax.Array, jax.Array, jax.Array]: + # return (params, m, v) tensors directly + pass diff --git a/challenges/medium/114_fused_adamw/starter/starter.mojo b/challenges/medium/114_fused_adamw/starter/starter.mojo new file mode 100644 index 00000000..b8438399 --- /dev/null +++ b/challenges/medium/114_fused_adamw/starter/starter.mojo @@ -0,0 +1,22 @@ +from std.gpu.host import DeviceContext +from std.gpu import block_dim, block_idx, thread_idx +from std.memory import UnsafePointer +from std.math import ceildiv + + +# params, grad, m, v are device pointers +@export +def solve( + params: UnsafePointer[Float32, MutExternalOrigin], + grad: UnsafePointer[Float32, MutExternalOrigin], + m: UnsafePointer[Float32, MutExternalOrigin], + v: UnsafePointer[Float32, MutExternalOrigin], + N: Int32, + lr: Float32, + beta1: Float32, + beta2: Float32, + eps: Float32, + weight_decay: Float32, + t: Int32, +) raises: + pass diff --git a/challenges/medium/114_fused_adamw/starter/starter.pytorch.py b/challenges/medium/114_fused_adamw/starter/starter.pytorch.py new file mode 100644 index 00000000..cad634c0 --- /dev/null +++ b/challenges/medium/114_fused_adamw/starter/starter.pytorch.py @@ -0,0 +1,18 @@ +import torch + + +# params, grad, m, v are tensors on the GPU +def solve( + params: torch.Tensor, + grad: torch.Tensor, + m: torch.Tensor, + v: torch.Tensor, + N: int, + lr: float, + beta1: float, + beta2: float, + eps: float, + weight_decay: float, + t: int, +): + pass diff --git a/challenges/medium/114_fused_adamw/starter/starter.triton.py b/challenges/medium/114_fused_adamw/starter/starter.triton.py new file mode 100644 index 00000000..9436fa01 --- /dev/null +++ b/challenges/medium/114_fused_adamw/starter/starter.triton.py @@ -0,0 +1,20 @@ +import torch +import triton +import triton.language as tl + + +# params, grad, m, v are tensors on the GPU +def solve( + params: torch.Tensor, + grad: torch.Tensor, + m: torch.Tensor, + v: torch.Tensor, + N: int, + lr: float, + beta1: float, + beta2: float, + eps: float, + weight_decay: float, + t: int, +): + pass From 0bbfe96898240d59b02b3059192f52b0a9dd1efc Mon Sep 17 00:00:00 2001 From: Aditya Sharma Date: Wed, 30 Sep 2026 20:22:38 +0530 Subject: [PATCH 2/2] Renumber fused AdamW challenge to 120 --- .../medium/{114_fused_adamw => 120_fused_adamw}/challenge.html | 0 .../medium/{114_fused_adamw => 120_fused_adamw}/challenge.py | 0 .../{114_fused_adamw => 120_fused_adamw}/starter/starter.cu | 0 .../{114_fused_adamw => 120_fused_adamw}/starter/starter.cute.py | 0 .../{114_fused_adamw => 120_fused_adamw}/starter/starter.jax.py | 0 .../{114_fused_adamw => 120_fused_adamw}/starter/starter.mojo | 0 .../starter/starter.pytorch.py | 0 .../starter/starter.triton.py | 0 8 files changed, 0 insertions(+), 0 deletions(-) rename challenges/medium/{114_fused_adamw => 120_fused_adamw}/challenge.html (100%) rename challenges/medium/{114_fused_adamw => 120_fused_adamw}/challenge.py (100%) rename challenges/medium/{114_fused_adamw => 120_fused_adamw}/starter/starter.cu (100%) rename challenges/medium/{114_fused_adamw => 120_fused_adamw}/starter/starter.cute.py (100%) rename challenges/medium/{114_fused_adamw => 120_fused_adamw}/starter/starter.jax.py (100%) rename challenges/medium/{114_fused_adamw => 120_fused_adamw}/starter/starter.mojo (100%) rename challenges/medium/{114_fused_adamw => 120_fused_adamw}/starter/starter.pytorch.py (100%) rename challenges/medium/{114_fused_adamw => 120_fused_adamw}/starter/starter.triton.py (100%) diff --git a/challenges/medium/114_fused_adamw/challenge.html b/challenges/medium/120_fused_adamw/challenge.html similarity index 100% rename from challenges/medium/114_fused_adamw/challenge.html rename to challenges/medium/120_fused_adamw/challenge.html diff --git a/challenges/medium/114_fused_adamw/challenge.py b/challenges/medium/120_fused_adamw/challenge.py similarity index 100% rename from challenges/medium/114_fused_adamw/challenge.py rename to challenges/medium/120_fused_adamw/challenge.py diff --git a/challenges/medium/114_fused_adamw/starter/starter.cu b/challenges/medium/120_fused_adamw/starter/starter.cu similarity index 100% rename from challenges/medium/114_fused_adamw/starter/starter.cu rename to challenges/medium/120_fused_adamw/starter/starter.cu diff --git a/challenges/medium/114_fused_adamw/starter/starter.cute.py b/challenges/medium/120_fused_adamw/starter/starter.cute.py similarity index 100% rename from challenges/medium/114_fused_adamw/starter/starter.cute.py rename to challenges/medium/120_fused_adamw/starter/starter.cute.py diff --git a/challenges/medium/114_fused_adamw/starter/starter.jax.py b/challenges/medium/120_fused_adamw/starter/starter.jax.py similarity index 100% rename from challenges/medium/114_fused_adamw/starter/starter.jax.py rename to challenges/medium/120_fused_adamw/starter/starter.jax.py diff --git a/challenges/medium/114_fused_adamw/starter/starter.mojo b/challenges/medium/120_fused_adamw/starter/starter.mojo similarity index 100% rename from challenges/medium/114_fused_adamw/starter/starter.mojo rename to challenges/medium/120_fused_adamw/starter/starter.mojo diff --git a/challenges/medium/114_fused_adamw/starter/starter.pytorch.py b/challenges/medium/120_fused_adamw/starter/starter.pytorch.py similarity index 100% rename from challenges/medium/114_fused_adamw/starter/starter.pytorch.py rename to challenges/medium/120_fused_adamw/starter/starter.pytorch.py diff --git a/challenges/medium/114_fused_adamw/starter/starter.triton.py b/challenges/medium/120_fused_adamw/starter/starter.triton.py similarity index 100% rename from challenges/medium/114_fused_adamw/starter/starter.triton.py rename to challenges/medium/120_fused_adamw/starter/starter.triton.py