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)
|