diff --git a/challenges/medium/120_fused_adamw/challenge.html b/challenges/medium/120_fused_adamw/challenge.html
new file mode 100644
index 00000000..e54b2a14
--- /dev/null
+++ b/challenges/medium/120_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
+
+ - Do not change the function signature.
+ - Do not use external libraries beyond what is available in the starter.
+ params, m, and v are updated in place;
+ there is no separate output buffer.
+ - Apply the decoupled weight decay to
params using its
+ pre-update value, before computing the moment updates and the adaptive
+ step, matching the order shown above.
+ t is always ≥ 1, so the bias-correction denominators are
+ never zero.
+
+
+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
+
+ - 1 ≤
N ≤ 16,777,216
+ - 1 ≤
t ≤ 100,000
+ params, grad, m, and v
+ are float32
+ - Values in
params and grad are in the range
+ [−10, 10]
+ - 0 <
lr ≤ 1
+ - 0 ≤
beta1, beta2 < 1
+ - 1e-12 ≤
eps ≤ 1e-2
+ - 0 ≤
weight_decay ≤ 1
+ v is always non-negative; m may be negative
+ - Performance is measured with
N = 4,096×4,096 = 16,777,216
+
diff --git a/challenges/medium/120_fused_adamw/challenge.py b/challenges/medium/120_fused_adamw/challenge.py
new file mode 100644
index 00000000..61d53c58
--- /dev/null
+++ b/challenges/medium/120_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/120_fused_adamw/starter/starter.cu b/challenges/medium/120_fused_adamw/starter/starter.cu
new file mode 100644
index 00000000..75e66e1d
--- /dev/null
+++ b/challenges/medium/120_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/120_fused_adamw/starter/starter.cute.py b/challenges/medium/120_fused_adamw/starter/starter.cute.py
new file mode 100644
index 00000000..c956a74d
--- /dev/null
+++ b/challenges/medium/120_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/120_fused_adamw/starter/starter.jax.py b/challenges/medium/120_fused_adamw/starter/starter.jax.py
new file mode 100644
index 00000000..a4f5b1e0
--- /dev/null
+++ b/challenges/medium/120_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/120_fused_adamw/starter/starter.mojo b/challenges/medium/120_fused_adamw/starter/starter.mojo
new file mode 100644
index 00000000..b8438399
--- /dev/null
+++ b/challenges/medium/120_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/120_fused_adamw/starter/starter.pytorch.py b/challenges/medium/120_fused_adamw/starter/starter.pytorch.py
new file mode 100644
index 00000000..cad634c0
--- /dev/null
+++ b/challenges/medium/120_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/120_fused_adamw/starter/starter.triton.py b/challenges/medium/120_fused_adamw/starter/starter.triton.py
new file mode 100644
index 00000000..9436fa01
--- /dev/null
+++ b/challenges/medium/120_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