diff --git a/challenges/medium/121_fused_logit_penalties/challenge.html b/challenges/medium/121_fused_logit_penalties/challenge.html new file mode 100644 index 00000000..6fd7417c --- /dev/null +++ b/challenges/medium/121_fused_logit_penalties/challenge.html @@ -0,0 +1,217 @@ +

+ Implement the fused logit penalty stage that LLM serving stacks such as vLLM, SGLang and + Hugging Face TGI run on every decoding step, right before sampling. Given a batch of + B in-flight requests, a logit row of V vocabulary entries per request, + the tokens each request has already seen, and per-request penalty coefficients, produce the + penalized logits in output. Every tensor is float32 except the token id + tensors, which are int32. +

+ +

+ For request \(b\), let \(c_{b,v}\) be how many times token \(v\) appears in + output_tokens[b] (the tokens generated so far), and let \(s_{b,v}\) be true when + token \(v\) appears anywhere in prompt_tokens[b] or output_tokens[b]. + Entries equal to -1 are padding and are ignored everywhere. Starting from + \(z = \texttt{logits}[b][v]\), apply the three penalties in this order: +

+
    +
  1. + Repetition (multiplicative, only for tokens that were seen): + if \(s_{b,v}\), then \(z \leftarrow z / r_b\) when \(z > 0\) and \(z \leftarrow z \cdot r_b\) + otherwise, where \(r_b = \texttt{repetition_penalty}[b]\). +
  2. +
  3. + Frequency (scales with the generated count): + \(z \leftarrow z - \texttt{frequency_penalty}[b] \cdot c_{b,v}\). +
  4. +
  5. + Presence (flat, applied once per generated token): + \(z \leftarrow z - \texttt{presence_penalty}[b]\) when \(c_{b,v} > 0\). +
  6. +
+

+ Note the asymmetry that makes this kernel interesting: the repetition penalty keys off + prompt and generated tokens, while the frequency and presence penalties key off + generated tokens only. Since \(V\) is far larger than the histories, scanning the token + lists once per vocabulary entry is hopelessly slow — build the per-request occurrence counts + first, then stream the logits. +

+ + + + + + + + + + one request (row b), V = 6 + + + prompt_tokens[b] + + + + 1 + 3 + -1 + + output_tokens[b] + + + + + 2 + 2 + 5 + -1 + -1 = padding, skipped + + + + scatter + count + + + per-request table over the vocabulary + v + count c + seen s + + 0 + 1 + 2 + 3 + 4 + 5 + + + + + 0 + 0 + 2 + 0 + 0 + 1 + 0 + 1 + 1 + 1 + 0 + 1 + + + + + + + + + + + + + + + + + + logits + + + 1.0 + -2.0 + 3.0 + 0.5 + 0.0 + -1.0 + + + + + + + + + + + r = 2.0, freq = 0.25, pres = 0.5 + + + output + + + 1.0 + -4.0 + 0.5 + 0.25 + 0.0 + -2.75 + + + + + + + + + v = 4 was never seen, so its logit passes through untouched + + +

Implementation Requirements

+ + +

Example

+

+ With B = 2, V = 6, P = 3, G = 4: +

+

+ Input logits (2×6): + \[ + \begin{bmatrix} 1.0 & -2.0 & 3.0 & 0.5 & 0.0 & -1.0 \\ 2.0 & 1.0 & -1.0 & 0.0 & 4.0 & -3.0 \end{bmatrix} + \] + prompt_tokens (2×3), output_tokens (2×4): + \[ + \begin{bmatrix} 1 & 3 & -1 \\ 0 & 0 & 2 \end{bmatrix} + \qquad + \begin{bmatrix} 2 & 2 & 5 & -1 \\ 4 & -1 & -1 & -1 \end{bmatrix} + \] + presence_penalty = \([0.5,\ 1.0]\), frequency_penalty = \([0.25,\ 0.5]\), + repetition_penalty = \([2.0,\ 1.0]\). +

+

+ Output output (2×6): + \[ + \begin{bmatrix} 1.0 & -4.0 & 0.5 & 0.25 & 0.0 & -2.75 \\ 2.0 & 1.0 & -1.0 & 0.0 & 2.5 & -3.0 \end{bmatrix} + \] +

+

+ Row 0 has \(s = \{1, 2, 3, 5\}\) and counts \(c_2 = 2\), \(c_5 = 1\). Token 1 is seen with a + negative logit, so it is multiplied: \(-2.0 \times 2.0 = -4.0\). Token 2 is seen with a positive + logit, so it is divided and then charged both additive penalties: + \(3.0 / 2.0 - 0.25 \cdot 2 - 0.5 = 0.5\). Token 4 is untouched. Row 1 has + \(\text{repetition_penalty} = 1.0\), so only token 4 changes: + \(4.0 - 0.5 \cdot 1 - 1.0 = 2.5\). +

+ +

Constraints

+ diff --git a/challenges/medium/121_fused_logit_penalties/challenge.py b/challenges/medium/121_fused_logit_penalties/challenge.py new file mode 100644 index 00000000..5641f2ed --- /dev/null +++ b/challenges/medium/121_fused_logit_penalties/challenge.py @@ -0,0 +1,321 @@ +import ctypes +from typing import Any, Dict, List + +import torch +from core.challenge_base import ChallengeBase + + +class Challenge(ChallengeBase): + name = "Fused Logit Penalties" + atol = 1e-05 + rtol = 1e-05 + num_gpus = 1 + access_tier = "free" + + def reference_impl( + self, + logits: torch.Tensor, + prompt_tokens: torch.Tensor, + output_tokens: torch.Tensor, + presence_penalty: torch.Tensor, + frequency_penalty: torch.Tensor, + repetition_penalty: torch.Tensor, + output: torch.Tensor, + B: int, + V: int, + P: int, + G: int, + ): + assert logits.shape == (B, V) + assert prompt_tokens.shape == (B, P) + assert output_tokens.shape == (B, G) + assert presence_penalty.shape == (B,) + assert frequency_penalty.shape == (B,) + assert repetition_penalty.shape == (B,) + assert output.shape == (B, V) + assert prompt_tokens.dtype == torch.int32 + assert output_tokens.dtype == torch.int32 + assert ( + logits.dtype + == presence_penalty.dtype + == frequency_penalty.dtype + == repetition_penalty.dtype + == output.dtype + ) + + dtype = logits.dtype + + prompt_idx = prompt_tokens.long() + prompt_valid = prompt_idx >= 0 + prompt_idx = torch.where(prompt_idx >= 0, prompt_idx, torch.zeros_like(prompt_idx)) + prompt_counts = torch.zeros((B, V), dtype=dtype, device=logits.device) + prompt_counts.scatter_add_(1, prompt_idx, prompt_valid.to(dtype)) + + output_idx = output_tokens.long() + output_valid = output_idx >= 0 + output_idx = torch.where(output_idx >= 0, output_idx, torch.zeros_like(output_idx)) + output_counts = torch.zeros((B, V), dtype=dtype, device=logits.device) + output_counts.scatter_add_(1, output_idx, output_valid.to(dtype)) + + seen = (prompt_counts + output_counts) > 0 + generated = output_counts > 0 + + rep = repetition_penalty.reshape(B, 1) + penalized = torch.where(logits > 0, logits / rep, logits * rep) + result = torch.where(seen, penalized, logits) + + result = result - frequency_penalty.reshape(B, 1) * output_counts + result = result - presence_penalty.reshape(B, 1) * generated.to(dtype) + + output.copy_(result) + + def get_solve_signature(self) -> Dict[str, tuple]: + return { + "logits": (ctypes.POINTER(ctypes.c_float), "in"), + "prompt_tokens": (ctypes.POINTER(ctypes.c_int), "in"), + "output_tokens": (ctypes.POINTER(ctypes.c_int), "in"), + "presence_penalty": (ctypes.POINTER(ctypes.c_float), "in"), + "frequency_penalty": (ctypes.POINTER(ctypes.c_float), "in"), + "repetition_penalty": (ctypes.POINTER(ctypes.c_float), "in"), + "output": (ctypes.POINTER(ctypes.c_float), "out"), + "B": (ctypes.c_int, "in"), + "V": (ctypes.c_int, "in"), + "P": (ctypes.c_int, "in"), + "G": (ctypes.c_int, "in"), + } + + def _build_case( + self, + logits: torch.Tensor, + prompt_tokens: torch.Tensor, + output_tokens: torch.Tensor, + presence: List[float], + frequency: List[float], + repetition: List[float], + ) -> Dict[str, Any]: + B, V = logits.shape + return { + "logits": logits, + "prompt_tokens": prompt_tokens, + "output_tokens": output_tokens, + "presence_penalty": torch.tensor(presence, device=self.device, dtype=torch.float32), + "frequency_penalty": torch.tensor(frequency, device=self.device, dtype=torch.float32), + "repetition_penalty": torch.tensor(repetition, device=self.device, dtype=torch.float32), + "output": torch.empty(B, V, device=self.device, dtype=torch.float32), + "B": B, + "V": V, + "P": prompt_tokens.shape[1], + "G": output_tokens.shape[1], + } + + def _pad_tail(self, tokens: torch.Tensor, lengths: torch.Tensor) -> torch.Tensor: + positions = torch.arange(tokens.shape[1], device=self.device).reshape(1, -1) + keep = positions < lengths.reshape(-1, 1) + return torch.where(keep, tokens, torch.full_like(tokens, -1)) + + def generate_example_test(self) -> Dict[str, Any]: + logits = torch.tensor( + [ + [1.0, -2.0, 3.0, 0.5, 0.0, -1.0], + [2.0, 1.0, -1.0, 0.0, 4.0, -3.0], + ], + device=self.device, + dtype=torch.float32, + ) + prompt_tokens = torch.tensor([[1, 3, -1], [0, 0, 2]], device=self.device, dtype=torch.int32) + output_tokens = torch.tensor( + [[2, 2, 5, -1], [4, -1, -1, -1]], device=self.device, dtype=torch.int32 + ) + return self._build_case( + logits, + prompt_tokens, + output_tokens, + presence=[0.5, 1.0], + frequency=[0.25, 0.5], + repetition=[2.0, 1.0], + ) + + def generate_functional_test(self) -> List[Dict[str, Any]]: + tests = [] + + # single sequence, tiny vocab, the one generated token repeats the prompt + tests.append( + self._build_case( + torch.tensor([[1.0, -1.0, 2.0, 0.0]], device=self.device, dtype=torch.float32), + torch.tensor([[2]], device=self.device, dtype=torch.int32), + torch.tensor([[2]], device=self.device, dtype=torch.int32), + presence=[0.5], + frequency=[0.75], + repetition=[2.0], + ) + ) + + # all penalties disabled (repetition_penalty = 1) => output must equal logits + tests.append( + self._build_case( + torch.tensor( + [[0.5, -0.5, 1.5], [-2.0, 3.0, -4.0]], device=self.device, dtype=torch.float32 + ), + torch.tensor([[0, 1], [2, 2]], device=self.device, dtype=torch.int32), + torch.tensor([[1, 1], [0, 2]], device=self.device, dtype=torch.int32), + presence=[0.0, 0.0], + frequency=[0.0, 0.0], + repetition=[1.0, 1.0], + ) + ) + + # fully padded prompt and output: no token has been seen + tests.append( + self._build_case( + torch.tensor([[-3.0, 2.0, 0.0, 7.0]], device=self.device, dtype=torch.float32), + torch.full((1, 3), -1, device=self.device, dtype=torch.int32), + torch.full((1, 2), -1, device=self.device, dtype=torch.int32), + presence=[1.0], + frequency=[1.0], + repetition=[2.0], + ) + ) + + # heavy repeats of a single token, mixed signs, repetition penalty below 1 + tests.append( + self._build_case( + torch.tensor( + [ + [-1.0, 4.0, -4.0, 0.0, 2.5, -2.5, 1.0, -1.5], + [3.0, 3.0, -3.0, -3.0, 0.5, 0.5, -0.5, -0.5], + [0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0], + ], + device=self.device, + dtype=torch.float32, + ), + torch.tensor( + [[0, 0, 0, 7], [1, -1, -1, -1], [3, 4, 5, 6]], + device=self.device, + dtype=torch.int32, + ), + torch.tensor( + [[1, 1, 1, 1, 1], [5, 5, 6, -1, -1], [-1, -1, -1, -1, 0]], + device=self.device, + dtype=torch.int32, + ), + presence=[0.25, 1.0, 0.5], + frequency=[0.5, 0.125, 2.0], + repetition=[0.5, 1.5, 2.0], + ) + ) + + # power-of-two vocab, all-zero logits + torch.manual_seed(0) + B, V, P, G = 2, 1024, 64, 32 + tests.append( + self._build_case( + torch.zeros(B, V, device=self.device, dtype=torch.float32), + torch.randint(0, V, (B, P), device=self.device, dtype=torch.int32), + torch.randint(0, 8, (B, G), device=self.device, dtype=torch.int32), + presence=[0.5, 1.5], + frequency=[0.25, 0.75], + repetition=[1.2, 1.8], + ) + ) + + # non-power-of-two vocab with ragged (trailing padded) prompts + torch.manual_seed(1) + B, V, P, G = 8, 255, 64, 17 + lengths = torch.randint(0, P + 1, (B,), device=self.device) + tests.append( + self._build_case( + torch.randn(B, V, device=self.device, dtype=torch.float32) * 3.0, + self._pad_tail( + torch.randint(0, V, (B, P), device=self.device, dtype=torch.int32), lengths + ), + torch.randint(0, 32, (B, G), device=self.device, dtype=torch.int32), + presence=[0.0, 0.5, 1.0, 1.5, 2.0, 0.25, 0.75, 1.25], + frequency=[1.0, 0.0, 0.5, 0.25, 2.0, 0.125, 0.0, 1.5], + repetition=[1.0, 1.1, 1.5, 2.0, 0.5, 1.25, 1.75, 1.05], + ) + ) + + # small non-power-of-two vocab, every row generates the same token repeatedly + torch.manual_seed(2) + B, V, P, G = 4, 100, 30, 30 + tests.append( + self._build_case( + torch.randn(B, V, device=self.device, dtype=torch.float32) * 2.0, + torch.randint(0, 10, (B, P), device=self.device, dtype=torch.int32), + torch.randint(0, 5, (B, G), device=self.device, dtype=torch.int32), + presence=[1.0, 0.5, 0.0, 2.0], + frequency=[0.5, 1.0, 0.25, 0.0], + repetition=[1.5, 2.0, 1.0, 0.75], + ) + ) + + # wide batch, very short histories + torch.manual_seed(3) + B, V, P, G = 64, 1024, 4, 2 + presence = torch.rand(B).tolist() + frequency = torch.rand(B).tolist() + repetition = (1.0 + torch.rand(B)).tolist() + tests.append( + self._build_case( + torch.randn(B, V, device=self.device, dtype=torch.float32) * 4.0, + torch.randint(0, V, (B, P), device=self.device, dtype=torch.int32), + torch.randint(0, V, (B, G), device=self.device, dtype=torch.int32), + presence=presence, + frequency=frequency, + repetition=repetition, + ) + ) + + # realistic decode step + torch.manual_seed(4) + B, V, P, G = 16, 4096, 256, 128 + presence = torch.rand(B).tolist() + frequency = torch.rand(B).tolist() + repetition = (0.5 + 1.5 * torch.rand(B)).tolist() + tests.append( + self._build_case( + torch.randn(B, V, device=self.device, dtype=torch.float32) * 2.5, + torch.randint(0, V, (B, P), device=self.device, dtype=torch.int32), + torch.randint(0, 512, (B, G), device=self.device, dtype=torch.int32), + presence=presence, + frequency=frequency, + repetition=repetition, + ) + ) + + # realistic LLM vocab + torch.manual_seed(5) + B, V, P, G = 32, 32000, 512, 256 + lengths = torch.randint(1, P + 1, (B,), device=self.device) + presence = torch.rand(B).tolist() + frequency = torch.rand(B).tolist() + repetition = (1.0 + torch.rand(B)).tolist() + tests.append( + self._build_case( + torch.randn(B, V, device=self.device, dtype=torch.float32) * 2.0, + self._pad_tail( + torch.randint(0, V, (B, P), device=self.device, dtype=torch.int32), lengths + ), + torch.randint(0, 2048, (B, G), device=self.device, dtype=torch.int32), + presence=presence, + frequency=frequency, + repetition=repetition, + ) + ) + + return tests + + def generate_performance_test(self) -> Dict[str, Any]: + torch.manual_seed(42) + B, V, P, G = 256, 128256, 1024, 512 + lengths = torch.randint(1, P + 1, (B,), device=self.device) + return self._build_case( + torch.randn(B, V, device=self.device, dtype=torch.float32) * 2.0, + self._pad_tail( + torch.randint(0, V, (B, P), device=self.device, dtype=torch.int32), lengths + ), + torch.randint(0, 8192, (B, G), device=self.device, dtype=torch.int32), + presence=torch.rand(B).tolist(), + frequency=torch.rand(B).tolist(), + repetition=(1.0 + torch.rand(B)).tolist(), + ) diff --git a/challenges/medium/121_fused_logit_penalties/starter/starter.cu b/challenges/medium/121_fused_logit_penalties/starter/starter.cu new file mode 100644 index 00000000..6a423604 --- /dev/null +++ b/challenges/medium/121_fused_logit_penalties/starter/starter.cu @@ -0,0 +1,7 @@ +#include + +// logits, prompt_tokens, output_tokens, presence_penalty, frequency_penalty, repetition_penalty, +// output are device pointers +extern "C" void solve(const float* logits, const int* prompt_tokens, const int* output_tokens, + const float* presence_penalty, const float* frequency_penalty, + const float* repetition_penalty, float* output, int B, int V, int P, int G) {} diff --git a/challenges/medium/121_fused_logit_penalties/starter/starter.cute.py b/challenges/medium/121_fused_logit_penalties/starter/starter.cute.py new file mode 100644 index 00000000..623786f4 --- /dev/null +++ b/challenges/medium/121_fused_logit_penalties/starter/starter.cute.py @@ -0,0 +1,21 @@ +import cutlass +import cutlass.cute as cute + + +# logits, prompt_tokens, output_tokens, presence_penalty, frequency_penalty, repetition_penalty, +# output are tensors on the GPU +@cute.jit +def solve( + logits: cute.Tensor, + prompt_tokens: cute.Tensor, + output_tokens: cute.Tensor, + presence_penalty: cute.Tensor, + frequency_penalty: cute.Tensor, + repetition_penalty: cute.Tensor, + output: cute.Tensor, + B: cute.Int32, + V: cute.Int32, + P: cute.Int32, + G: cute.Int32, +): + pass diff --git a/challenges/medium/121_fused_logit_penalties/starter/starter.jax.py b/challenges/medium/121_fused_logit_penalties/starter/starter.jax.py new file mode 100644 index 00000000..c6f7a470 --- /dev/null +++ b/challenges/medium/121_fused_logit_penalties/starter/starter.jax.py @@ -0,0 +1,21 @@ +import jax +import jax.numpy as jnp + + +# logits, prompt_tokens, output_tokens, presence_penalty, frequency_penalty, repetition_penalty +# are tensors on device +@jax.jit +def solve( + logits: jax.Array, + prompt_tokens: jax.Array, + output_tokens: jax.Array, + presence_penalty: jax.Array, + frequency_penalty: jax.Array, + repetition_penalty: jax.Array, + B: int, + V: int, + P: int, + G: int, +) -> jax.Array: + # return output tensor directly + pass diff --git a/challenges/medium/121_fused_logit_penalties/starter/starter.mojo b/challenges/medium/121_fused_logit_penalties/starter/starter.mojo new file mode 100644 index 00000000..a2e2a883 --- /dev/null +++ b/challenges/medium/121_fused_logit_penalties/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 + + +# logits, prompt_tokens, output_tokens, presence_penalty, frequency_penalty, repetition_penalty, output are device pointers +@export +def solve( + logits: UnsafePointer[Float32, MutExternalOrigin], + prompt_tokens: UnsafePointer[Int32, MutExternalOrigin], + output_tokens: UnsafePointer[Int32, MutExternalOrigin], + presence_penalty: UnsafePointer[Float32, MutExternalOrigin], + frequency_penalty: UnsafePointer[Float32, MutExternalOrigin], + repetition_penalty: UnsafePointer[Float32, MutExternalOrigin], + output: UnsafePointer[Float32, MutExternalOrigin], + B: Int32, + V: Int32, + P: Int32, + G: Int32, +) raises: + pass diff --git a/challenges/medium/121_fused_logit_penalties/starter/starter.pytorch.py b/challenges/medium/121_fused_logit_penalties/starter/starter.pytorch.py new file mode 100644 index 00000000..445c6bc2 --- /dev/null +++ b/challenges/medium/121_fused_logit_penalties/starter/starter.pytorch.py @@ -0,0 +1,19 @@ +import torch + + +# logits, prompt_tokens, output_tokens, presence_penalty, frequency_penalty, repetition_penalty, +# output are tensors on the GPU +def solve( + logits: torch.Tensor, + prompt_tokens: torch.Tensor, + output_tokens: torch.Tensor, + presence_penalty: torch.Tensor, + frequency_penalty: torch.Tensor, + repetition_penalty: torch.Tensor, + output: torch.Tensor, + B: int, + V: int, + P: int, + G: int, +): + pass diff --git a/challenges/medium/121_fused_logit_penalties/starter/starter.triton.py b/challenges/medium/121_fused_logit_penalties/starter/starter.triton.py new file mode 100644 index 00000000..27bee0c0 --- /dev/null +++ b/challenges/medium/121_fused_logit_penalties/starter/starter.triton.py @@ -0,0 +1,21 @@ +import torch +import triton +import triton.language as tl + + +# logits, prompt_tokens, output_tokens, presence_penalty, frequency_penalty, repetition_penalty, +# output are tensors on the GPU +def solve( + logits: torch.Tensor, + prompt_tokens: torch.Tensor, + output_tokens: torch.Tensor, + presence_penalty: torch.Tensor, + frequency_penalty: torch.Tensor, + repetition_penalty: torch.Tensor, + output: torch.Tensor, + B: int, + V: int, + P: int, + G: int, +): + pass