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 :
+
+
+
+ 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]\).
+
+
+ Frequency (scales with the generated count):
+ \(z \leftarrow z - \texttt{frequency_penalty}[b] \cdot c_{b,v}\).
+
+
+ Presence (flat, applied once per generated token):
+ \(z \leftarrow z - \texttt{presence_penalty}[b]\) when \(c_{b,v} > 0\).
+
+
+
+ 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
+
+ Implement solve(logits, prompt_tokens, output_tokens, presence_penalty, frequency_penalty, repetition_penalty, output, B, V, P, G); do not change the signature or use external libraries beyond the standard GPU frameworks.
+ Write the penalized logits into the provided output buffer; leave logits unmodified.
+ Each request has its own presence_penalty, frequency_penalty and repetition_penalty value.
+ Apply the penalties in the order listed above; the repetition penalty must be applied before the additive penalties.
+ Token ids equal to -1 are padding slots and contribute to no count; padding may appear at any position in a row.
+
+
+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
+
+ 1 ≤ B ≤ 256
+ 1 ≤ V ≤ 131,072
+ 1 ≤ P ≤ 1,024
+ 1 ≤ G ≤ 512
+ Token ids are -1 (padding) or in the range \([0, V)\)
+ -50.0 ≤ logits[b][v] ≤ 50.0
+ 0.0 ≤ presence_penalty[b], frequency_penalty[b] ≤ 2.0
+ 0.5 ≤ repetition_penalty[b] ≤ 2.0
+ logits, the penalty vectors and output are float32; prompt_tokens and output_tokens are int32
+ Performance is measured with B = 256, V = 128,256, P = 1,024, G = 512
+
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