File size: 21,091 Bytes
b025706
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
# SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
#
# SPDX-License-Identifier: Apache-2.0
"""
On-device penalties module with persistent buffers, mirroring TTSampling.
"""

from __future__ import annotations

from dataclasses import dataclass
from typing import Any, List, Optional

import torch

import ttnn
from models.common.lightweightmodule import LightweightModule


@dataclass
class PenaltyContext:
    prompt_mask: ttnn.Tensor
    output_mask: ttnn.Tensor
    output_counts: ttnn.Tensor
    output_counts_gathered: ttnn.Tensor
    presence_penalties: ttnn.Tensor
    frequency_penalties: ttnn.Tensor
    repetition_penalties: ttnn.Tensor
    inverse_repetition_penalties: ttnn.Tensor
    sub_core_grids: Any | None = None


def apply_penalties(logits: ttnn.Tensor, context: Optional[PenaltyContext]) -> ttnn.Tensor:
    if context is None:
        return logits

    op_kwargs = {"sub_core_grids": context.sub_core_grids} if context.sub_core_grids else {}
    # presence
    presence_term = ttnn.multiply(
        ttnn.typecast(context.output_mask, ttnn.bfloat16, **op_kwargs), context.presence_penalties, **op_kwargs
    )
    presence_term_bf16 = ttnn.typecast(presence_term, ttnn.bfloat16, **op_kwargs)
    logits = ttnn.subtract(logits, presence_term_bf16, output_tensor=logits, **op_kwargs)
    presence_term_bf16.deallocate()

    # frequency
    output_counts_bf16 = ttnn.typecast(context.output_counts, ttnn.bfloat16, **op_kwargs)

    freq_term = ttnn.multiply(output_counts_bf16, context.frequency_penalties, **op_kwargs)

    freq_term_bf16 = ttnn.typecast(freq_term, ttnn.bfloat16, **op_kwargs)
    logits = ttnn.subtract(logits, freq_term_bf16, output_tensor=logits, **op_kwargs)
    freq_term_bf16.deallocate()

    # repetition

    # If token appears in prompt or output, apply, otherwise use 1.0 for no-op.

    combined_mask_int32 = ttnn.add(context.prompt_mask, context.output_mask, **op_kwargs)
    combined_mask = ttnn.typecast(combined_mask_int32, ttnn.bfloat16, **op_kwargs)
    combined_mask_int32.deallocate()
    penalties = ttnn.where(combined_mask, context.repetition_penalties, 1.0, **op_kwargs)
    inverse_penalties = ttnn.where(combined_mask, context.inverse_repetition_penalties, 1.0, **op_kwargs)
    combined_mask.deallocate()

    # If logits are >0, divide by penalty, otherwise multiply by penalty.
    logits_bf16 = ttnn.typecast(logits, ttnn.bfloat16, **op_kwargs)
    logits_gt1 = ttnn.gt(logits_bf16, 0, **op_kwargs)
    scaling = ttnn.where(logits_gt1, inverse_penalties, penalties, **op_kwargs)
    logits_gt1.deallocate()
    penalties.deallocate()
    inverse_penalties.deallocate()
    logits = ttnn.multiply(logits, scaling, output_tensor=logits, **op_kwargs)
    scaling.deallocate()

    return logits


class TTPenalties(LightweightModule):
    """
    Penalty module with persistent device tensors, similar to TTSampling.
    """

    def __init__(self, mesh_device, args):
        super().__init__()
        self.mesh_device = mesh_device
        self.cluster_shape = mesh_device.shape
        # Floor at 32 so that ROW_MAJOR [batch, vocab] buffers passed to
        # ttnn.tilize always have physical_volume divisible by TILE_HW
        # (32*32 = 1024).  32 * V is 1024-aligned for any 32-aligned V.
        self.max_batch_size = max(getattr(args, "max_batch_size", 32), 32)

        padded_vocab_size = getattr(args, "padded_vocab_size", None)
        self.vocab_size = padded_vocab_size if padded_vocab_size is not None else args.vocab_size

        self.sub_core_grids = getattr(args, "sub_core_grids", None)
        self._op_kwargs = {"sub_core_grids": self.sub_core_grids} if self.sub_core_grids else {}

        # sampling_dp > 1 when multiple mesh rows each sample independently
        # (e.g. GPT-OSS on [4,8] Galaxy: 4 rows × 32 users = 128 total)
        self._sampling_dp = getattr(args, "sampling_dp", 1)

        # When rows are used for data parallelism (sampling_dp > 1), vocab
        # must be sharded along columns; otherwise pick the larger dimension.
        if self._sampling_dp > 1:
            num_devices = mesh_device.shape[-1]
        else:
            num_devices = max(mesh_device.shape[-1], mesh_device.shape[-2])
        self.num_devices = num_devices
        # Total batch across all rows. Host tensors use this size; after
        # (0, ...) sharding each row gets max_batch_size entries.
        self._total_batch = self.max_batch_size * self._sampling_dp

        # shard vocab size over larger cluster dim
        if mesh_device.shape[-1] == self.num_devices:
            shard_dims = (None, 1)
            shard_dims_slice = (None, 0)
        else:
            shard_dims = (1, None)
            shard_dims_slice = (0, None)

        # For row-sharded mode (sampling_dp > 1), also shard the batch dimension
        # across mesh rows so each row gets its own per-user penalty state.
        if self._sampling_dp > 1:
            assert (
                mesh_device.shape[-1] == self.num_devices
            ), "Row-sharded penalties require vocab sharding along mesh columns"
            shard_dims = (0, 1)  # batch across rows, vocab across cols
            shard_dims_gathered = (0, None)  # batch across rows, vocab replicated
            shard_dims_bf16 = (0, None)  # per-row penalty params
            per_row_batch = self.max_batch_size  # NOT divided: each row gets max_batch_size
        else:
            shard_dims_gathered = (None, None)
            shard_dims_bf16 = None
            per_row_batch = self.max_batch_size

        self.per_row_batch_size = per_row_batch
        self._shard_dims_gathered = shard_dims_gathered

        self.prompt_mask = self._alloc_int_buffer(shard_dims=shard_dims)
        # Host shadow of the per-slot prompt tokens, so a partial update keeps other rows' masks.
        self._prompt_tokens_host = None
        self.output_mask = self._alloc_int_buffer(shard_dims=shard_dims)
        self.output_counts_gathered = self._alloc_int_buffer(shard_dims=shard_dims_gathered)
        self.output_counts = self._alloc_int_buffer(shard_dims=shard_dims)
        self._shard_dims_mask = shard_dims
        self.decode_src = self._alloc_int_buffer(
            host=torch.ones(self._total_batch, 1), shard_dims=shard_dims_gathered, layout=ttnn.ROW_MAJOR_LAYOUT
        )
        self.zeros = self._alloc_int_buffer(shard_dims=shard_dims_gathered, layout=ttnn.ROW_MAJOR_LAYOUT)
        self.presence_penalties = self._alloc_bf16_buffer(shard_dims=shard_dims_bf16)
        self.frequency_penalties = self._alloc_bf16_buffer(shard_dims=shard_dims_bf16)
        self.repetition_penalties = self._alloc_bf16_buffer(shard_dims=shard_dims_bf16)
        self.inverse_repetition_penalties = self._alloc_bf16_buffer(shard_dims=shard_dims_bf16)

        vocab_per_dev = self.vocab_size // self.num_devices
        d = torch.arange(self.num_devices, dtype=torch.int32)

        # [0, 0, 0, vocab_per_dev, 0, 2*vocab_per_dev, ...]
        start_1d = torch.empty(2 * self.num_devices, dtype=torch.int32)
        start_1d[0::2] = 0
        start_1d[1::2] = d * vocab_per_dev

        # [batch, vocab_per_dev, batch, 2*vocab_per_dev, ...]
        end_1d = torch.empty(2 * self.num_devices, dtype=torch.int32)
        end_1d[0::2] = per_row_batch  # per-row batch size, exclusive
        end_1d[1::2] = (d + 1) * vocab_per_dev  # exclusive

        self.slice_start = ttnn.from_torch(
            start_1d,
            device=self.mesh_device,
            mesh_mapper=ttnn.ShardTensor2dMesh(self.mesh_device, dims=shard_dims_slice, mesh_shape=self.cluster_shape),
        )
        self.slice_end = ttnn.from_torch(
            end_1d,
            device=self.mesh_device,
            mesh_mapper=ttnn.ShardTensor2dMesh(self.mesh_device, dims=shard_dims_slice, mesh_shape=self.cluster_shape),
        )

    def _alloc_int_buffer(self, shard_dims, host=None, layout=ttnn.TILE_LAYOUT):
        if host is None:
            host = torch.zeros((self._total_batch, self.vocab_size), dtype=torch.int32)
        return ttnn.from_torch(
            host,
            dtype=ttnn.int32,
            layout=layout,
            device=self.mesh_device,
            mesh_mapper=ttnn.ShardTensor2dMesh(self.mesh_device, dims=shard_dims, mesh_shape=self.cluster_shape),
            memory_config=ttnn.DRAM_MEMORY_CONFIG,
        )

    def _alloc_bf16_buffer(self, shard_dims=None):
        host = torch.zeros((self._total_batch, 1), dtype=torch.float32)
        if shard_dims is not None:
            return ttnn.from_torch(
                host,
                dtype=ttnn.bfloat16,
                layout=ttnn.TILE_LAYOUT,
                device=self.mesh_device,
                mesh_mapper=ttnn.ShardTensor2dMesh(self.mesh_device, dims=shard_dims, mesh_shape=self.cluster_shape),
            )
        return ttnn.from_torch(host, dtype=ttnn.bfloat16, layout=ttnn.TILE_LAYOUT, device=self.mesh_device)

    def _copy_host_to_device(self, dst: ttnn.Tensor, src: torch.Tensor):
        if self._sampling_dp > 1:
            # For row-sharded buffers, create a properly sharded host tensor
            # so copy_host_to_device_tensor writes per-row shards correctly.
            mapper = ttnn.ShardTensor2dMesh(self.mesh_device, dims=(0, None), mesh_shape=self.cluster_shape)
            src_tt = ttnn.from_torch(src, dtype=dst.dtype, layout=ttnn.TILE_LAYOUT, device=None, mesh_mapper=mapper)
        else:
            src_tt = ttnn.from_torch(src, dtype=dst.dtype, layout=ttnn.TILE_LAYOUT, device=None)
        ttnn.copy_host_to_device_tensor(src_tt, dst)

    def _copy_int_host_to_device(self, dst: ttnn.Tensor, src: torch.Tensor, shard_dims):
        mapper = ttnn.ShardTensor2dMesh(self.mesh_device, dims=shard_dims, mesh_shape=self.cluster_shape)
        src_tt = ttnn.from_torch(src, dtype=dst.dtype, layout=ttnn.TILE_LAYOUT, device=None, mesh_mapper=mapper)
        ttnn.copy_host_to_device_tensor(src_tt, dst)

    def _token_counts_host(self, tokens_2d: torch.Tensor) -> torch.Tensor:
        valid = (tokens_2d >= 0) & (tokens_2d < self.vocab_size)
        token_ids = torch.where(valid, tokens_2d, torch.zeros_like(tokens_2d)).to(torch.int64)
        counts = torch.zeros((self._total_batch, self.vocab_size), dtype=torch.int32)
        counts.scatter_add_(1, token_ids, valid.to(torch.int32))
        return counts

    def reset_params(self, presence: List[float], frequency: List[float], repetition: List[float]):
        presence_tensor = self._pad_params(presence)
        frequency_tensor = self._pad_params(frequency)
        repetition_tensor = self._pad_params(repetition)
        inverse_repetition_tensor = 1 / repetition_tensor

        self._copy_host_to_device(self.presence_penalties, presence_tensor)
        self._copy_host_to_device(self.frequency_penalties, frequency_tensor)
        self._copy_host_to_device(self.repetition_penalties, repetition_tensor)
        self._copy_host_to_device(self.inverse_repetition_penalties, inverse_repetition_tensor)

    def _pad_params(self, values: List[float]) -> torch.Tensor:
        tensor = torch.tensor(values, dtype=torch.float32)
        if tensor.numel() < self._total_batch:
            pad_value = tensor[-1] if tensor.numel() > 0 else torch.tensor(0.0)
            pad = pad_value.repeat(self._total_batch - tensor.numel())
            tensor = torch.cat([tensor, pad])
        elif tensor.numel() > self._total_batch:
            tensor = tensor[: self._total_batch]
        return tensor.view(self._total_batch, 1)

    def _pad_batch_to_max(self, tokens_2d: torch.Tensor, pad_value: int) -> torch.Tensor:
        """Pad/truncate first dim to _total_batch."""
        if tokens_2d.dim() != 2:
            raise ValueError(f"Expected 2D tensor [B, S], got {tokens_2d.shape}")
        B, S = tokens_2d.shape
        if B < self._total_batch:
            pad = torch.full((self._total_batch - B, S), pad_value, dtype=tokens_2d.dtype)
            return torch.cat([tokens_2d, pad], dim=0)
        if B > self._total_batch:
            return tokens_2d[: self._total_batch]
        return tokens_2d

    def reset_prompt_tokens(self, prompt_tokens: torch.Tensor, slots: list[int] | None = None):
        """Rebuild the prompt mask. With ``slots``, only those rows are taken from
        ``prompt_tokens``; every other row keeps the prompt it was last given.

        The device buffer covers all rows at once, so a caller that only knows about the requests it
        is prefilling used to zero everyone else's mask: rows outside the call arrive as the -1
        padding and hash to an empty mask. repetition_penalty is the only consumer of prompt_mask, so
        a live request silently stopped penalising its own prompt until something refreshed it.
        A later full sampling-state reset can hide this bug. A demo without that reset keeps the
        wiped mask for the rest of the generation.
        """
        prompt_tokens_2d = prompt_tokens.reshape(-1, prompt_tokens.shape[-1])
        prompt_tokens_2d = self._pad_batch_to_max(prompt_tokens_2d, pad_value=-1)

        if slots is None:
            self._prompt_tokens_host = prompt_tokens_2d.clone()
        else:
            shadow = getattr(self, "_prompt_tokens_host", None)
            width = max(prompt_tokens_2d.shape[-1], shadow.shape[-1] if shadow is not None else 0)
            merged = torch.full((self._total_batch, width), -1, dtype=prompt_tokens_2d.dtype)
            if shadow is not None:
                merged[:, : shadow.shape[-1]] = shadow
            for slot in slots:
                slot = int(slot)
                if 0 <= slot < self._total_batch:
                    merged[slot, :] = -1
                    merged[slot, : prompt_tokens_2d.shape[-1]] = prompt_tokens_2d[slot]
            self._prompt_tokens_host = merged
            prompt_tokens_2d = merged

        # Build reset masks on host to avoid device scatter_add races on
        # duplicate prompt token ids (common in penalty tests/prompts).
        prompt_counts = self._token_counts_host(prompt_tokens_2d)
        prompt_mask = (prompt_counts > 0).to(torch.int32)
        self._copy_int_host_to_device(self.prompt_mask, prompt_mask, self._shard_dims_mask)

    def reset_output_tokens(self, tokens=None, slots: list[int] | None = None):
        if slots is not None:
            slots = sorted({int(slot) for slot in slots})
            if any(slot < 0 or slot >= self._total_batch for slot in slots):
                raise ValueError(f"Output reset slots must be in [0, {self._total_batch}), got {slots}")
            if not slots:
                return

            # Clear only the admitted slots. A [batch, 1] device mask is
            # replicated across mesh columns and broadcast across vocabulary,
            # so continuing requests keep their accumulated device counts.
            keep_rows = torch.ones((self._total_batch, 1), dtype=torch.int32)
            keep_rows[slots] = 0
            keep_rows_tt = self._alloc_int_buffer(
                host=keep_rows,
                shard_dims=self._shard_dims_gathered,
            )
            self.output_mask = ttnn.mul(
                self.output_mask, keep_rows_tt, output_tensor=self.output_mask, **self._op_kwargs
            )
            self.output_counts = ttnn.mul(
                self.output_counts, keep_rows_tt, output_tensor=self.output_counts, **self._op_kwargs
            )
            self.output_counts_gathered = ttnn.mul(
                self.output_counts_gathered,
                keep_rows_tt,
                output_tensor=self.output_counts_gathered,
                **self._op_kwargs,
            )
            keep_rows_tt.deallocate()

            if tokens is None:
                return

            # Restore any supplied history for the reset slots. Rows outside
            # ``slots`` are zero here, so adding cannot change live requests.
            tokens_2d = tokens.reshape(-1, tokens.shape[-1])
            tokens_2d = self._pad_batch_to_max(tokens_2d, pad_value=-1)
            output_counts = self._token_counts_host(tokens_2d)
            reset_rows = torch.zeros((self._total_batch, 1), dtype=torch.int32)
            reset_rows[slots] = 1
            output_counts *= reset_rows
            output_mask = (output_counts > 0).to(torch.int32)
            updates = (
                (self.output_counts_gathered, output_counts, self._shard_dims_gathered),
                (self.output_counts, output_counts, self._shard_dims_mask),
                (self.output_mask, output_mask, self._shard_dims_mask),
            )
            for destination, host_update, shard_dims in updates:
                update_tt = self._alloc_int_buffer(host=host_update, shard_dims=shard_dims)
                ttnn.add(destination, update_tt, output_tensor=destination, **self._op_kwargs)
                update_tt.deallocate()
            return

        # ALWAYS reset output buffers to zero first (this is the core accuracy fix from issue #35731)
        # This ensures penalty statistics are cleared between prefill and decode phases
        self.output_mask = ttnn.mul(self.output_mask, 0, output_tensor=self.output_mask, **self._op_kwargs)
        self.output_counts = ttnn.mul(self.output_counts, 0, output_tensor=self.output_counts, **self._op_kwargs)
        self.output_counts_gathered = ttnn.mul(
            self.output_counts_gathered, 0, output_tensor=self.output_counts_gathered, **self._op_kwargs
        )

        # THEN optionally repopulate if tokens are provided
        if tokens is not None:
            tokens_2d = tokens.reshape(-1, tokens.shape[-1])
            tokens_2d = self._pad_batch_to_max(tokens_2d, pad_value=-1)
            output_counts = self._token_counts_host(tokens_2d)
            output_mask = (output_counts > 0).to(torch.int32)
            self._copy_int_host_to_device(self.output_counts_gathered, output_counts, self._shard_dims_gathered)
            self._copy_int_host_to_device(self.output_counts, output_counts, self._shard_dims_mask)
            self._copy_int_host_to_device(self.output_mask, output_mask, self._shard_dims_mask)

    def update_output_tokens(self, new_tokens):
        # Reshape decode token to [batch, 1] for scatter_add.
        # Non-row-sharded: token shape is [1,1,1,batch] → shape[-1]==batch, shape[-2]==1
        # Row-sharded:     token shape is [1,1,batch,1] → shape[-2]==batch, shape[-1]==1
        batch = self.per_row_batch_size
        fast_path = (new_tokens.shape[-1] == batch and new_tokens.shape[-2] == 1) or (
            new_tokens.shape[-2] == batch and new_tokens.shape[-1] == 1
        )
        if fast_path:
            new_tokens = ttnn.reshape(new_tokens, [batch, 1], **self._op_kwargs)
            src = self.decode_src
        else:
            src = self._alloc_int_buffer(
                host=torch.ones(self._total_batch, new_tokens.shape[-1]),
                shard_dims=self._shard_dims_gathered,
                layout=ttnn.ROW_MAJOR_LAYOUT,
            )
        self.token_bin_counts_and_mask(
            new_tokens=new_tokens,
            counts=self.output_counts_gathered,
            src=src,
            counts_sliced=self.output_counts,
            mask=self.output_mask,
        )

    def token_bin_counts_and_mask(self, new_tokens, src, counts=None, mask=None, counts_sliced=None):
        counts_new = ttnn.scatter_add(self.zeros, 1, new_tokens, src, **self._op_kwargs)

        new_tokens.deallocate()
        # need to use use_low_perf because llama galaxy runs out of L1 otherwise
        counts_new = ttnn.tilize(
            counts_new, **self._op_kwargs, use_low_perf=True if self.sub_core_grids is not None else False
        )
        if counts:
            counts = ttnn.add(counts, counts_new, output_tensor=counts, **self._op_kwargs)
        else:
            counts = counts_new
        counts_sliced = ttnn.slice(
            counts,
            self.slice_start,
            self.slice_end,
            output_tensor=counts_sliced,
            slice_dim=1,
            num_devices=self.num_devices,
            **self._op_kwargs,
        )

        mask = ttnn.gt(counts_sliced, 0, output_tensor=mask, **self._op_kwargs)
        return counts, mask

    def apply(self, tt_logits: ttnn.Tensor) -> ttnn.Tensor:
        if tt_logits is None:
            return tt_logits
        context = PenaltyContext(
            prompt_mask=self.prompt_mask,
            output_mask=self.output_mask,
            output_counts=self.output_counts,
            output_counts_gathered=self.output_counts_gathered,
            presence_penalties=self.presence_penalties,
            frequency_penalties=self.frequency_penalties,
            repetition_penalties=self.repetition_penalties,
            inverse_repetition_penalties=self.inverse_repetition_penalties,
            sub_core_grids=self.sub_core_grids,
        )
        original_shape = tt_logits.shape
        reshaped = ttnn.reshape(tt_logits, (-1, original_shape[-1]))
        apply_penalties(reshaped, context)
        return ttnn.reshape(reshaped, original_shape)