Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
217 changes: 217 additions & 0 deletions challenges/medium/121_fused_logit_penalties/challenge.html
Original file line number Diff line number Diff line change
@@ -0,0 +1,217 @@
<p>
Implement the fused <em>logit penalty</em> 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
<code>B</code> in-flight requests, a logit row of <code>V</code> vocabulary entries per request,
the tokens each request has already seen, and per-request penalty coefficients, produce the
penalized logits in <code>output</code>. Every tensor is <code>float32</code> except the token id
tensors, which are <code>int32</code>.
</p>

<p>
For request \(b\), let \(c_{b,v}\) be how many times token \(v\) appears in
<code>output_tokens[b]</code> (the tokens generated so far), and let \(s_{b,v}\) be true when
token \(v\) appears anywhere in <code>prompt_tokens[b]</code> or <code>output_tokens[b]</code>.
Entries equal to <code>-1</code> are padding and are ignored everywhere. Starting from
\(z = \texttt{logits}[b][v]\), apply the three penalties <strong>in this order</strong>:
</p>
<ol>
<li>
<strong>Repetition</strong> (multiplicative, only for tokens that were seen):
if \(s_{b,v}\), then \(z \leftarrow z / r_b\) when \(z &gt; 0\) and \(z \leftarrow z \cdot r_b\)
otherwise, where \(r_b = \texttt{repetition_penalty}[b]\).
</li>
<li>
<strong>Frequency</strong> (scales with the generated count):
\(z \leftarrow z - \texttt{frequency_penalty}[b] \cdot c_{b,v}\).
</li>
<li>
<strong>Presence</strong> (flat, applied once per generated token):
\(z \leftarrow z - \texttt{presence_penalty}[b]\) when \(c_{b,v} &gt; 0\).
</li>
</ol>
<p>
Note the asymmetry that makes this kernel interesting: the repetition penalty keys off
<em>prompt and generated</em> tokens, while the frequency and presence penalties key off
<em>generated</em> 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.
</p>

<svg width="700" height="300" viewBox="0 0 700 300" xmlns="http://www.w3.org/2000/svg"
style="display:block; margin:20px auto; font-family:monospace;">
<rect width="700" height="300" fill="#222" rx="8"/>
<defs>
<marker id="lp_arr" markerWidth="8" markerHeight="8" refX="6" refY="3" orient="auto">
<path d="M0,0 L0,6 L8,3 z" fill="#888"/>
</marker>
</defs>

<text x="350" y="26" fill="#aaa" font-size="11" text-anchor="middle">one request (row b), V = 6</text>

<!-- token histories -->
<text x="90" y="56" fill="#cc8844" font-size="10" text-anchor="middle">prompt_tokens[b]</text>
<rect x="20" y="64" width="140" height="26" fill="#5a3a1a" stroke="#cc8844" stroke-width="1.5"/>
<line x1="66" y1="64" x2="66" y2="90" stroke="#cc8844" stroke-width="1"/>
<line x1="113" y1="64" x2="113" y2="90" stroke="#cc8844" stroke-width="1"/>
<text x="43" y="82" fill="#ccc" font-size="11" text-anchor="middle">1</text>
<text x="90" y="82" fill="#ccc" font-size="11" text-anchor="middle">3</text>
<text x="136" y="82" fill="#666" font-size="11" text-anchor="middle">-1</text>

<text x="90" y="124" fill="#5588cc" font-size="10" text-anchor="middle">output_tokens[b]</text>
<rect x="20" y="132" width="140" height="26" fill="#2a4a7f" stroke="#5588cc" stroke-width="1.5"/>
<line x1="55" y1="132" x2="55" y2="158" stroke="#5588cc" stroke-width="1"/>
<line x1="90" y1="132" x2="90" y2="158" stroke="#5588cc" stroke-width="1"/>
<line x1="125" y1="132" x2="125" y2="158" stroke="#5588cc" stroke-width="1"/>
<text x="37" y="150" fill="#ccc" font-size="11" text-anchor="middle">2</text>
<text x="72" y="150" fill="#ccc" font-size="11" text-anchor="middle">2</text>
<text x="107" y="150" fill="#ccc" font-size="11" text-anchor="middle">5</text>
<text x="142" y="150" fill="#666" font-size="11" text-anchor="middle">-1</text>
<text x="90" y="176" fill="#666" font-size="9" text-anchor="middle">-1 = padding, skipped</text>

<!-- scatter arrow -->
<line x1="166" y1="110" x2="226" y2="110" stroke="#888" stroke-width="1.5" marker-end="url(#lp_arr)"/>
<text x="196" y="102" fill="#888" font-size="8" text-anchor="middle">scatter</text>
<text x="196" y="124" fill="#888" font-size="8" text-anchor="middle">count</text>

<!-- count table -->
<text x="330" y="56" fill="#aaa" font-size="10" text-anchor="middle">per-request table over the vocabulary</text>
<text x="248" y="82" fill="#888" font-size="9" text-anchor="end">v</text>
<text x="248" y="110" fill="#5588cc" font-size="9" text-anchor="end">count c</text>
<text x="248" y="138" fill="#cc8844" font-size="9" text-anchor="end">seen s</text>
<g font-size="10" text-anchor="middle">
<text x="275" y="82" fill="#888">0</text>
<text x="330" y="82" fill="#888">1</text>
<text x="385" y="82" fill="#888">2</text>
<text x="440" y="82" fill="#888">3</text>
<text x="495" y="82" fill="#888">4</text>
<text x="550" y="82" fill="#888">5</text>
</g>
<rect x="250" y="92" width="330" height="24" fill="#2a4a7f" stroke="#5588cc" stroke-width="1.5"/>
<rect x="250" y="120" width="330" height="24" fill="#5a3a1a" stroke="#cc8844" stroke-width="1.5"/>
<g font-size="10" text-anchor="middle" fill="#ccc">
<text x="275" y="108">0</text>
<text x="330" y="108">0</text>
<text x="385" y="108">2</text>
<text x="440" y="108">0</text>
<text x="495" y="108">0</text>
<text x="550" y="108">1</text>
<text x="275" y="136">0</text>
<text x="330" y="136">1</text>
<text x="385" y="136">1</text>
<text x="440" y="136">1</text>
<text x="495" y="136">0</text>
<text x="550" y="136">1</text>
</g>
<g stroke="#5588cc" stroke-width="1">
<line x1="305" y1="92" x2="305" y2="116"/>
<line x1="360" y1="92" x2="360" y2="116"/>
<line x1="415" y1="92" x2="415" y2="116"/>
<line x1="470" y1="92" x2="470" y2="116"/>
<line x1="525" y1="92" x2="525" y2="116"/>
</g>
<g stroke="#cc8844" stroke-width="1">
<line x1="305" y1="120" x2="305" y2="144"/>
<line x1="360" y1="120" x2="360" y2="144"/>
<line x1="415" y1="120" x2="415" y2="144"/>
<line x1="470" y1="120" x2="470" y2="144"/>
<line x1="525" y1="120" x2="525" y2="144"/>
</g>

<!-- logits row -->
<text x="248" y="196" fill="#888" font-size="9" text-anchor="end">logits</text>
<rect x="250" y="180" width="330" height="24" fill="#333" stroke="#777" stroke-width="1.5"/>
<g font-size="10" text-anchor="middle" fill="#ccc">
<text x="275" y="196">1.0</text>
<text x="330" y="196">-2.0</text>
<text x="385" y="196">3.0</text>
<text x="440" y="196">0.5</text>
<text x="495" y="196">0.0</text>
<text x="550" y="196">-1.0</text>
</g>
<g stroke="#777" stroke-width="1">
<line x1="305" y1="180" x2="305" y2="204"/>
<line x1="360" y1="180" x2="360" y2="204"/>
<line x1="415" y1="180" x2="415" y2="204"/>
<line x1="470" y1="180" x2="470" y2="204"/>
<line x1="525" y1="180" x2="525" y2="204"/>
</g>

<line x1="415" y1="208" x2="415" y2="238" stroke="#888" stroke-width="1.5" marker-end="url(#lp_arr)"/>
<text x="470" y="228" fill="#888" font-size="9" text-anchor="middle">r = 2.0, freq = 0.25, pres = 0.5</text>

<!-- output row -->
<text x="248" y="258" fill="#44aa66" font-size="9" text-anchor="end">output</text>
<rect x="250" y="242" width="330" height="24" fill="#1a5a3a" stroke="#44aa66" stroke-width="1.5"/>
<g font-size="10" text-anchor="middle" fill="#ccc">
<text x="275" y="258">1.0</text>
<text x="330" y="258">-4.0</text>
<text x="385" y="258">0.5</text>
<text x="440" y="258">0.25</text>
<text x="495" y="258">0.0</text>
<text x="550" y="258">-2.75</text>
</g>
<g stroke="#44aa66" stroke-width="1">
<line x1="305" y1="242" x2="305" y2="266"/>
<line x1="360" y1="242" x2="360" y2="266"/>
<line x1="415" y1="242" x2="415" y2="266"/>
<line x1="470" y1="242" x2="470" y2="266"/>
<line x1="525" y1="242" x2="525" y2="266"/>
</g>
<text x="415" y="286" fill="#666" font-size="9" text-anchor="middle">v = 4 was never seen, so its logit passes through untouched</text>
</svg>

<h2>Implementation Requirements</h2>
<ul>
<li>Implement <code>solve(logits, prompt_tokens, output_tokens, presence_penalty, frequency_penalty, repetition_penalty, output, B, V, P, G)</code>; do not change the signature or use external libraries beyond the standard GPU frameworks.</li>
<li>Write the penalized logits into the provided <code>output</code> buffer; leave <code>logits</code> unmodified.</li>
<li>Each request has its own <code>presence_penalty</code>, <code>frequency_penalty</code> and <code>repetition_penalty</code> value.</li>
<li>Apply the penalties in the order listed above; the repetition penalty must be applied before the additive penalties.</li>
<li>Token ids equal to <code>-1</code> are padding slots and contribute to no count; padding may appear at any position in a row.</li>
</ul>

<h2>Example</h2>
<p>
With <code>B</code> = 2, <code>V</code> = 6, <code>P</code> = 3, <code>G</code> = 4:
</p>
<p>
<strong>Input</strong> <code>logits</code> (2&times;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}
\]
<code>prompt_tokens</code> (2&times;3), <code>output_tokens</code> (2&times;4):
\[
\begin{bmatrix} 1 & 3 & -1 \\ 0 & 0 & 2 \end{bmatrix}
\qquad
\begin{bmatrix} 2 & 2 & 5 & -1 \\ 4 & -1 & -1 & -1 \end{bmatrix}
\]
<code>presence_penalty</code> = \([0.5,\ 1.0]\), <code>frequency_penalty</code> = \([0.25,\ 0.5]\),
<code>repetition_penalty</code> = \([2.0,\ 1.0]\).
</p>
<p>
<strong>Output</strong> <code>output</code> (2&times;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}
\]
</p>
<p>
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\).
</p>

<h2>Constraints</h2>
<ul>
<li>1 &le; <code>B</code> &le; 256</li>
<li>1 &le; <code>V</code> &le; 131,072</li>
<li>1 &le; <code>P</code> &le; 1,024</li>
<li>1 &le; <code>G</code> &le; 512</li>
<li>Token ids are <code>-1</code> (padding) or in the range \([0, V)\)</li>
<li>-50.0 &le; <code>logits[b][v]</code> &le; 50.0</li>
<li>0.0 &le; <code>presence_penalty[b]</code>, <code>frequency_penalty[b]</code> &le; 2.0</li>
<li>0.5 &le; <code>repetition_penalty[b]</code> &le; 2.0</li>
<li><code>logits</code>, the penalty vectors and <code>output</code> are <code>float32</code>; <code>prompt_tokens</code> and <code>output_tokens</code> are <code>int32</code></li>
<li>Performance is measured with <code>B</code> = 256, <code>V</code> = 128,256, <code>P</code> = 1,024, <code>G</code> = 512</li>
</ul>
Loading
Loading