Download code/models/common/sampling/tt_penalties.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 21.1 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/sampling/tt_penalties.py
- Command line
-
hf download hf://tt-hous/clef/code/models/common/sampling/tt_penalties.py
-
curl -L -o tt_penalties.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/sampling/tt_penalties.py
21.1 kB
| # 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 | |
| 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) | |