clef / code /models /common /modules /sampling /sampling_1d.py
tt-hous's picture
Add files using upload-large-folder tool
2415c4c verified
Raw History Blame Contribute Delete
36.2 kB
# SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
"""
Sampling1D: Top-k/top-p/temperature sampling for 1D mesh topologies.
TTTv2 module — declarative config, lazy buffer allocation, k/p/temp as per-call args.
No mutable sampling state stored on the module.
See also: models/common/sampling/tt_sampling.py (TTTv1 source)
"""
from __future__ import annotations
import inspect
import sys
from dataclasses import dataclass, replace
from typing import Any, Optional
from loguru import logger
import ttnn
from models.common.lightweightmodule import LightweightModule
from models.common.modules.lazy_buffer import LazyBuffer, resolve_lazy_buffer
from models.common.modules.tt_ccl import get_tt_ccl
from models.common.sampling.vocab_padding import (
build_invalid_vocab_mask,
build_tail_invalid_vocab_mask,
get_vocab_shard_dims,
)
# ---------------------------------------------------------------------------
# Power-of-2 helpers (local copies; keep this TTTv2 module self-contained)
# ---------------------------------------------------------------------------
def _is_power_of_2(n: int) -> bool:
return n > 0 and (n & (n - 1)) == 0
def _upper_power_of_2(n: int) -> int:
if n <= 1:
return 1
return 1 << (n - 1).bit_length()
# ---------------------------------------------------------------------------
# Config
# ---------------------------------------------------------------------------
@dataclass
class Sampling1DConfig:
"""Declarative config for Sampling1D.
Buffer fields use the triple-type pattern: ``LazyBuffer | ttnn.Tensor | None``.
k/p/temp are NOT stored here — they are per-call forward args.
"""
vocab_size: int # Required. Padded logits width; caller pre-pads to be divisible by num_devices.
valid_vocab_size: Optional[int] = None # Real token vocabulary size. Defaults to vocab_size.
mesh_device: Optional[ttnn.MeshDevice] = None # None → GetDefaultDevice()
tt_ccl: Any = None # None → get_tt_ccl(mesh_device) if multi-device
max_batch_size: int = 32
max_top_k: int = 32
sub_core_grids: Any = None
sub_core_grid_topk: Any = None
start_core: Optional[ttnn.CoreCoord] = None # None → CoreCoord(0,0)
num_gather_links: int = 1
sampling_memory_config: Optional[ttnn.MemoryConfig] = None # None → DRAM_MEMORY_CONFIG
allow_force_argmax: bool = False
num_argmax_gather_links: Optional[int] = None # None → same as num_gather_links
ag_topology: Optional[ttnn.Topology] = None # None → Topology.Linear
argmax_chunks_per_sync: int = 10
argmax_num_workers_per_link: int = 1
# Pad each per-device logit shard up to the next power of 2 before ttnn.topk. Big device-perf
# win for non-power-of-2 vocab on the multi-device path (TTTv1's pad_logits_to_power_of_2).
# Strict TTTv1 parity: only the multi-device path is padded — the 1×1 multi_step split path
# is left unpadded.
pad_to_power_of_2: bool = False
# --- Persistent buffer specs (LazyBuffer | ttnn.Tensor | None) ---
# Static index buffers (computed from vocab_size + num_devices, never mutated)
index_offsets: LazyBuffer | ttnn.Tensor | None = None # [1,1,32,max_top_k*num_devices], int32, TILE
invalid_vocab_mask: LazyBuffer | ttnn.Tensor | None = None # Full fallback mask, sharded like logits
invalid_vocab_tail_mask: LazyBuffer | ttnn.Tensor | None = None # Compact [1,1,32,tail] local mask
invalid_vocab_tail_width: int = 0
# Seed/ID buffers (seeds mutable via LazyBuffer.update(), user_ids static)
seeds: LazyBuffer | ttnn.Tensor | None = None # [32], uint32, ROW_MAJOR
user_ids: LazyBuffer | ttnn.Tensor | None = None # [32], uint32, ROW_MAJOR
@staticmethod
def _buf_resolved(buf) -> bool:
if buf is None:
return False
if isinstance(buf, ttnn.Tensor):
return True
return buf.is_resolved()
def is_resolved(self) -> bool:
if self.mesh_device is None:
return False
if self.mesh_device.get_num_devices() > 1 and self.tt_ccl is None:
return False
required_buffers = ["index_offsets", "seeds", "user_ids"]
if not all(self._buf_resolved(getattr(self, f)) for f in required_buffers):
return False
if self.valid_vocab_size is not None and self.valid_vocab_size < self.vocab_size:
return self._buf_resolved(self.invalid_vocab_mask) or self._buf_resolved(self.invalid_vocab_tail_mask)
return True
# ---------------------------------------------------------------------------
# Module
# ---------------------------------------------------------------------------
class Sampling1D(LightweightModule):
"""Top-k/top-p/temperature sampling for 1D mesh topologies.
k/p/temp are per-call forward args — NOT stored on the module. This eliminates
mutable sampling state.
"""
def __init__(self, vocab_size: int, mesh_device: ttnn.MeshDevice | None = None, **kwargs):
"""Happy path — minimal required args."""
super().__init__()
self.config = _resolve_sampling1d_config(
Sampling1DConfig(vocab_size=vocab_size, mesh_device=mesh_device, **kwargs)
)
self._device_buffers_loaded = False
self._bind_strategy()
@classmethod
def from_config(cls, config: Sampling1DConfig) -> Sampling1D:
"""Power path — fully custom config."""
instance = object.__new__(cls)
super(Sampling1D, instance).__init__()
instance.config = _resolve_sampling1d_config(config)
instance._device_buffers_loaded = False
instance._bind_strategy()
return instance
# -- Strategy binding (TTTv2: no if-else in forward) ----------------------
def _bind_strategy(self):
"""Bind self._topk to the correct strategy based on mesh topology."""
cluster_shape = self.config.mesh_device.shape
self._multi_step_reduction = list(cluster_shape) == [1, 1]
if self._multi_step_reduction:
self._topk = self._topk_single_device
else:
self._topk = self._topk_multi_device
# Argmax strategy: single vs multi-device
num_devices = self.config.mesh_device.get_num_devices()
if num_devices > 1:
self._pre_argmax_gather = self._argmax_all_gather
else:
self._pre_argmax_gather = self._argmax_noop
# Memory config strategy for top-k post-processing
cfg = self.config
if cfg.sampling_memory_config is not None and cfg.sampling_memory_config != ttnn.DRAM_MEMORY_CONFIG:
self._prepare_topk_memory = self._topk_memory_sharded_roundtrip
else:
self._prepare_topk_memory = self._topk_memory_noop
# CCL introspection (port from TTSampling.__init__ lines 77-91)
self._line_all_gather = getattr(self.config.tt_ccl, "line_all_gather", None) if self.config.tt_ccl else None
self._line_all_gather_supports_buffer_key = False
if callable(self._line_all_gather):
try:
sig = inspect.signature(self._line_all_gather)
params = sig.parameters
self._line_all_gather_supports_buffer_key = "buffer_key" in params or any(
p.kind == inspect.Parameter.VAR_KEYWORD for p in params.values()
)
except (TypeError, ValueError):
logger.warning("Unable to inspect line_all_gather signature")
# -- Device buffers (idempotent) ------------------------------------------
def load_device_buffers(self):
"""Materialize all LazyBuffer fields from resolved config."""
if self._device_buffers_loaded:
return
assert self.config.is_resolved(), "config must be resolved before loading device buffers!"
cfg = self.config
try:
self._index_offsets = _materialize(cfg.index_offsets)
self._invalid_vocab_mask = (
_materialize(cfg.invalid_vocab_mask) if cfg.invalid_vocab_mask is not None else None
)
self._invalid_vocab_tail_mask = (
_materialize(cfg.invalid_vocab_tail_mask) if cfg.invalid_vocab_tail_mask is not None else None
)
self._invalid_vocab_tail_width = cfg.invalid_vocab_tail_width
self._seeds = _materialize(cfg.seeds)
self._user_ids = _materialize(cfg.user_ids)
from models.common.utils import LogProbsCalculator # lazy: transitively imports torch
self._log_probs_calculator = LogProbsCalculator(cfg.mesh_device, cfg.sub_core_grids, cfg.tt_ccl)
# Pre-compute static sub_core_grids for ttnn.sampling()
self._sampling_sub_core_grids = (
ttnn.num_cores_to_corerangeset_in_subcoregrids(
cfg.start_core, cfg.max_batch_size, cfg.sub_core_grids, row_wise=True
)
if cfg.sub_core_grids is not None
else None
)
except BaseException as primary:
try:
self.release()
except BaseException as cleanup_error:
failures = (cleanup_error,) + tuple(getattr(cleanup_error, "cleanup_failures", ()))
_attach_cleanup_failures(primary, failures)
raise
self._device_buffers_loaded = True
def release(self) -> None:
"""Release model-owned buffers idempotently and permit rematerialization."""
failures = []
calculator = getattr(self, "_log_probs_calculator", None)
if calculator is not None:
try:
calculator.release()
except BaseException as error:
failures.append(error)
else:
self._log_probs_calculator = None
for name in (
"index_offsets",
"invalid_vocab_mask",
"invalid_vocab_tail_mask",
"seeds",
"user_ids",
):
specification = getattr(self.config, name, None)
if isinstance(specification, LazyBuffer):
try:
specification.release()
except BaseException as error:
failures.append(error)
if not isinstance(specification, LazyBuffer) or specification._value is None:
setattr(self, f"_{name}", None)
self._sampling_sub_core_grids = None
self._device_buffers_loaded = False
if failures:
_raise_cleanup_failures(failures)
# -- Forward methods ------------------------------------------------------
def decode_forward(
self,
logits: ttnn.Tensor,
*,
k: ttnn.Tensor | None = None,
p: ttnn.Tensor | None = None,
temp: ttnn.Tensor | None = None,
seeds: ttnn.Tensor | None = None,
tt_out_tok: ttnn.Tensor | None = None,
enable_log_probs: bool | list[bool] = False,
):
"""Sample tokens from logits.
Args:
logits: Input logits tensor (sharded across devices)
k, p, temp: Per-call sampling parameters. All None + allow_force_argmax → argmax path.
All provided → top-k sampling path.
seeds: Optional per-call seed override. If None, uses config seeds buffer.
tt_out_tok: Optional output tensor to write results to.
enable_log_probs: Per-call logprobs toggle (bool, or per-user list of bool — if any
user is enabled the whole batch computes logprobs). Refreshed every call, so no
mutable logprobs state is stored on the module. Only the top-k path emits logprobs
(sampled-token logprob); the argmax path never does (see ``_sample_argmax``). The
calculator additionally requires a multi-device shard (T3K 1×8) — it returns ``None``
on 1×1/1×2 even when enabled.
Returns:
(token_ids, log_probs_or_none)
"""
self.load_device_buffers()
cfg = self.config
# Per-call logprobs mode — refreshed each forward from the arg, never persisted on the
# module (TTTv2 principle: no mutable sampling state). Mirrors main's reset_params path.
self._log_probs_calculator.set_log_probs_mode(enable_log_probs)
# Route: argmax or top-k
if k is None and p is None and temp is None:
if cfg.allow_force_argmax:
return self._sample_argmax(logits, tt_out_tok)
else:
raise ValueError("k/p/temp are all None but allow_force_argmax is False")
if k is None or p is None or temp is None:
raise ValueError("k, p, temp must all be provided, or all be None (for argmax)")
return self._sample_topk(logits, k, p, temp, seeds, tt_out_tok)
def forward(self, logits, **kwargs):
"""Dispatcher."""
return self.decode_forward(logits, **kwargs)
# -- Argmax path (port of the argmax branch of TTSampling.forward) --------
def _sample_argmax(self, logits, tt_out_tok):
slice_valid_vocab = self._can_slice_valid_vocab_for_argmax()
if not slice_valid_vocab:
logits = self._mask_invalid_vocab_logits(logits)
logits = self._pre_argmax_gather(logits)
if slice_valid_vocab:
logits = self._slice_valid_vocab_for_argmax(logits)
x_untilized = ttnn.untilize(logits, use_multicore=True)
tt_out_tok = ttnn.argmax(
x_untilized,
dim=-1,
output_tensor=tt_out_tok,
keepdim=False,
)
# Argmax path never emits logprobs (main's contract: force-argmax is disabled whenever
# logprobs are requested). Return None unconditionally — do not call the calculator.
return tt_out_tok, None
def _get_argmax_all_gather_config(self, cluster_axis):
"""Clamp the tuned all-gather config to what the actual submesh supports.
Port of main's ``_get_force_argmax_all_gather_config`` (#44246): bound num_links to the
links available on the submesh, and force Linear topology below 8 devices (Ring routing
like D0→D12 only wraps cleanly on T3K-class 8-device groups).
"""
cfg = self.config
num_links = cfg.num_argmax_gather_links
if hasattr(cfg.tt_ccl, "get_num_links"):
num_links = min(num_links, cfg.tt_ccl.get_num_links(cluster_axis))
topology = cfg.ag_topology
if cfg.mesh_device.get_num_devices() < 8:
topology = ttnn.Topology.Linear
return max(1, num_links), topology
def _argmax_all_gather(self, logits):
"""Multi-device: all-gather logits before argmax.
On ring-capable meshes (e.g. T3K 1×8) use Ring topology with no barrier semaphore to match the
model's logits gather and avoid trace-capture issues seen with some barrier-based configurations.
For other meshes, fall back to the clamped Linear+barrier path.
"""
cfg = self.config
cluster_axis = None if 1 in cfg.mesh_device.shape else 1
num_links, topology = self._get_argmax_all_gather_config(cluster_axis)
kwargs = {}
if cluster_axis is not None:
kwargs["cluster_axis"] = cluster_axis
return ttnn.experimental.all_gather_async(
logits,
persistent_output_buffer=None,
dim=3,
multi_device_global_semaphore=cfg.tt_ccl.get_and_cycle_ag_semaphore_handles(cluster_axis),
num_links=num_links,
memory_config=logits.memory_config(),
topology=topology,
barrier_semaphore=cfg.tt_ccl.get_and_cycle_barrier_semaphore_handle(cluster_axis),
chunks_per_sync=cfg.argmax_chunks_per_sync,
num_workers_per_link=cfg.argmax_num_workers_per_link,
num_buffers_per_channel=2,
**kwargs,
)
@staticmethod
def _argmax_noop(logits):
"""Single-device: no gather needed."""
return logits
# -- Top-k memory strategies (bound at init, no if-else in forward) -------
def _topk_memory_sharded_roundtrip(self, topk_values, topk_indices_int32):
"""Non-DRAM sampling_memory_config: round-trip through sharded memory."""
cfg = self.config
topk_values_sharded = ttnn.to_memory_config(
topk_values, memory_config=cfg.sampling_memory_config, dtype=ttnn.bfloat16
)
topk_values = ttnn.to_memory_config(topk_values_sharded, memory_config=ttnn.DRAM_MEMORY_CONFIG)
ttnn.deallocate(topk_values_sharded)
topk_indices_int32 = ttnn.to_memory_config(topk_indices_int32, cfg.sampling_memory_config)
return topk_values, topk_indices_int32
@staticmethod
def _topk_memory_noop(topk_values, topk_indices_int32):
"""DRAM memory config: no extra round-trip needed."""
return topk_values, topk_indices_int32
# -- Top-k sampling (port of the top-k branch of TTSampling.forward) ------
def _sample_topk(self, logits, k, p, temp, seeds, tt_out_tok):
cfg = self.config
x_bf16 = ttnn.typecast(logits, dtype=ttnn.bfloat16, sub_core_grids=cfg.sub_core_grids)
x_bf16 = self._mask_invalid_vocab_logits(x_bf16)
# Strategy-dispatched top-k
topk_values, topk_indices = self._topk(x_bf16)
# Convert indices to int32
topk_indices_int32 = ttnn.typecast(topk_indices, dtype=ttnn.int32, sub_core_grids=cfg.sub_core_grids)
topk_values, topk_indices_int32 = self._prepare_topk_memory(topk_values, topk_indices_int32)
# Add device offsets for global vocabulary indices
index_offsets, sliced_offsets = self._slice_user_rows(
self._index_offsets, int(topk_indices_int32.shape[2]), cfg.sampling_memory_config
)
topk_global_indices = ttnn.add(
index_offsets,
topk_indices_int32,
dtype=ttnn.int32,
memory_config=cfg.sampling_memory_config,
sub_core_grids=cfg.sub_core_grids,
)
ttnn.deallocate(topk_indices_int32)
if sliced_offsets:
ttnn.deallocate(index_offsets)
# Use distinct names so we can free the interleaved intermediate after untilize
topk_global_indices_interleaved = ttnn.to_memory_config(topk_global_indices, ttnn.DRAM_MEMORY_CONFIG)
topk_global_indices = ttnn.untilize(
topk_global_indices_interleaved, use_multicore=True, sub_core_grids=cfg.sub_core_grids
)
ttnn.deallocate(topk_global_indices_interleaved)
# Seed the RNG
seeds_tensor = seeds if seeds is not None else self._seeds
ttnn.manual_seed(
seeds=seeds_tensor,
user_ids=self._user_ids,
sub_core_grids=self._sampling_sub_core_grids,
)
# Sample
tt_out_tok = ttnn.sampling(
topk_values,
topk_global_indices,
k=k,
p=p,
temp=temp,
sub_core_grids=self._sampling_sub_core_grids,
output_tensor=tt_out_tok,
)
ttnn.deallocate(topk_values)
ttnn.deallocate(topk_global_indices)
# Logprobs (old path: single sampled-token logprob). Gated on enable_log_probs so the
# disabled case incurs zero extra ops. The calculator itself returns None unless the mesh
# is multi-device with num_devices ∈ {8, 32} (T3K 1×8).
if self._log_probs_calculator.enable_log_probs:
log_probs = self._log_probs_calculator.calculate_log_probs(logits, tt_out_tok)
else:
log_probs = None
return tt_out_tok, log_probs
def _mask_invalid_vocab_logits(self, logits):
if self._invalid_vocab_tail_mask is not None:
return self._mask_invalid_vocab_tail_logits(logits)
if self._invalid_vocab_mask is None:
return logits
return ttnn.add(
logits,
self._invalid_vocab_mask,
memory_config=logits.memory_config(),
sub_core_grids=self.config.sub_core_grids,
)
def _mask_invalid_vocab_tail_logits(self, logits):
cfg = self.config
tail_width = self._invalid_vocab_tail_width
local_width = logits.shape[-1]
valid_width = local_width - tail_width
if tail_width <= 0 or valid_width < 0:
return self._mask_invalid_vocab_logits_fallback(logits)
if valid_width == 0:
return ttnn.add(
logits,
self._invalid_vocab_tail_mask,
memory_config=logits.memory_config(),
sub_core_grids=cfg.sub_core_grids,
)
valid_logits = ttnn.slice(
logits,
[0, 0, 0, 0],
[logits.shape[0], logits.shape[1], logits.shape[2], valid_width],
memory_config=logits.memory_config(),
sub_core_grids=cfg.sub_core_grids,
)
tail_logits = ttnn.slice(
logits,
[0, 0, 0, valid_width],
[logits.shape[0], logits.shape[1], logits.shape[2], local_width],
memory_config=logits.memory_config(),
sub_core_grids=cfg.sub_core_grids,
)
masked_tail_logits = ttnn.add(
tail_logits,
self._invalid_vocab_tail_mask,
memory_config=logits.memory_config(),
sub_core_grids=cfg.sub_core_grids,
)
masked_logits = ttnn.concat(
[valid_logits, masked_tail_logits],
dim=3,
memory_config=logits.memory_config(),
sub_core_grids=cfg.sub_core_grids,
)
ttnn.deallocate(valid_logits)
ttnn.deallocate(tail_logits)
ttnn.deallocate(masked_tail_logits)
return masked_logits
def _mask_invalid_vocab_logits_fallback(self, logits):
if self._invalid_vocab_mask is None:
return logits
return ttnn.add(
logits,
self._invalid_vocab_mask,
memory_config=logits.memory_config(),
sub_core_grids=self.config.sub_core_grids,
)
def _can_slice_valid_vocab_for_argmax(self):
cfg = self.config
return cfg.valid_vocab_size < cfg.vocab_size and cfg.valid_vocab_size % ttnn.TILE_SIZE == 0
def _slice_valid_vocab_for_argmax(self, logits):
cfg = self.config
if not self._can_slice_valid_vocab_for_argmax() or logits.shape[-1] != cfg.vocab_size:
return logits
return ttnn.slice(
logits,
[0, 0, 0, 0],
[logits.shape[0], logits.shape[1], logits.shape[2], cfg.valid_vocab_size],
memory_config=logits.memory_config(),
sub_core_grids=cfg.sub_core_grids,
)
def _slice_user_rows(self, tensor, active_batch, memory_config):
if int(tensor.shape[2]) == active_batch:
return tensor, False
return (
ttnn.slice(
tensor,
[0, 0, 0, 0],
[tensor.shape[0], tensor.shape[1], active_batch, tensor.shape[3]],
memory_config=memory_config,
sub_core_grids=self.config.sub_core_grids,
),
True,
)
# -- Top-k strategies (bound at init, no if-else in forward) --------------
def _topk_single_device(self, x_bf16):
"""Split vocab in half → two topk → concat. Port of tt_sampling.py:346-371."""
cfg = self.config
x_list = ttnn.split(x_bf16, x_bf16.shape[-1] // 2, dim=3)
values_parts = []
indices_parts = []
for i in range(len(x_list)):
vals, idxs = ttnn.topk(
x_list[i],
k=cfg.max_top_k,
dim=-1,
sub_core_grids=cfg.sub_core_grid_topk,
)
values_parts.append(vals)
indices_parts.append(idxs)
x_list[i].deallocate()
gathered_values = ttnn.concat(values_parts, dim=3)
gathered_indices = ttnn.concat(indices_parts, dim=3)
for v, i in zip(values_parts, indices_parts):
ttnn.deallocate(v)
ttnn.deallocate(i)
return gathered_values, gathered_indices
def _topk_multi_device(self, x_bf16):
"""Local topk → all_gather across devices. Port of tt_sampling.py:372-421."""
cfg = self.config
cluster_shape = cfg.mesh_device.shape
# Pad the per-device shard up to the next power of 2 so ttnn.topk hits its fast path.
# Padded entries get -inf so they are never selected. Mirrors the padding in TTSampling.forward.
if cfg.pad_to_power_of_2 and not _is_power_of_2(x_bf16.shape[-1]):
padded_width = _upper_power_of_2(x_bf16.shape[-1])
x_bf16 = ttnn.pad(
x_bf16,
[(0, 0), (0, 0), (0, 0), (0, padded_width - x_bf16.shape[-1])],
value=-sys.float_info.max,
sub_core_grids=cfg.sub_core_grids,
)
topk_values, topk_indices = ttnn.topk(
x_bf16,
k=cfg.max_top_k,
dim=-1,
sub_core_grids=cfg.sub_core_grid_topk,
)
# For 1D meshes use cluster_axis=None
sampling_cluster_axis = None if 1 in cluster_shape else 0
# Gather values
gathered_values = self._perform_all_gather(
topk_values,
dim=3,
cluster_axis=sampling_cluster_axis,
memory_config=ttnn.DRAM_MEMORY_CONFIG,
num_links=cfg.num_gather_links,
buffer_key="SAMPLING_VALUES",
)
ttnn.deallocate(topk_values)
# Gather indices
gathered_indices = self._perform_all_gather(
topk_indices,
dim=3,
cluster_axis=sampling_cluster_axis,
memory_config=ttnn.DRAM_MEMORY_CONFIG,
num_links=cfg.num_gather_links,
buffer_key="SAMPLING_INDICES",
)
ttnn.deallocate(topk_indices)
return gathered_values, gathered_indices
# -- CCL helper -----------------------------------------------------------
def _perform_all_gather(self, tensor, dim, cluster_axis, memory_config, num_links, buffer_key=None):
"""Flexible all-gather: prefer line_all_gather if available, else ttnn.all_gather.
Port of TTSampling._perform_all_gather.
"""
if callable(self._line_all_gather):
kwargs = {
"dim": dim,
"cluster_axis": cluster_axis,
"memory_config": memory_config,
"num_links": num_links,
}
if self._line_all_gather_supports_buffer_key and buffer_key is not None:
kwargs["buffer_key"] = buffer_key
return self._line_all_gather(tensor, **kwargs)
return ttnn.all_gather(
tensor,
dim=dim,
num_links=num_links,
memory_config=memory_config,
cluster_axis=cluster_axis,
topology=ttnn.Topology.Linear,
)
# -- (Backward compat) Model args factory ----------------------------------
@classmethod
def from_model_args(cls, mesh_device, tt_ccl, args, model_config=None) -> Sampling1D:
"""Backward compat factory for TTTv1 model args."""
cluster_shape = mesh_device.shape
if min(cluster_shape) > 1:
raise ValueError(
f"Sampling1D only supports 1D mesh topologies, got shape {cluster_shape}. "
"Use TTSampling for Galaxy (2D) topologies."
)
padded_vocab_size = getattr(args, "padded_vocab_size", None)
vocab_size = padded_vocab_size if padded_vocab_size is not None else args.vocab_size
valid_vocab_size = getattr(args, "vocab_size", vocab_size)
# Extract config from model_config dict
mc = model_config or getattr(args, "model_config", {})
num_gather_links = 1
if "GALAXY_NUM_LINKS" in mc:
max_links = mc["GALAXY_NUM_LINKS"]
max_top_k = getattr(args, "max_top_k", 32)
num_gather_links = min(max_top_k // 32, max_links) if max_top_k // 32 <= max_links else max_links
sampling_memory_config = mc.get("DECODE_SAMPLING_INPUT_MEMCFG", ttnn.DRAM_MEMORY_CONFIG)
allow_force_argmax = False
num_argmax_gather_links = num_gather_links
ag_topology = ttnn.Topology.Linear
argmax_chunks_per_sync = 10
argmax_num_workers_per_link = 1
if "SAMPLING_AG_CONFIG" in mc:
ag_cfg = mc["SAMPLING_AG_CONFIG"]
allow_force_argmax = ag_cfg.get("allow_force_argmax", False)
num_argmax_gather_links = ag_cfg.get("num_links", num_gather_links)
ag_topology = ag_cfg.get("topology", ttnn.Topology.Linear)
argmax_chunks_per_sync = ag_cfg.get("chunks_per_sync", 10)
config = Sampling1DConfig(
vocab_size=vocab_size,
valid_vocab_size=valid_vocab_size,
mesh_device=mesh_device,
tt_ccl=tt_ccl,
max_batch_size=getattr(args, "max_batch_size", 32),
max_top_k=getattr(args, "max_top_k", 32),
sub_core_grids=getattr(args, "sub_core_grids", None),
sub_core_grid_topk=getattr(args, "sub_core_grid_topk", None),
start_core=getattr(args, "start_core", ttnn.CoreCoord(0, 0)),
num_gather_links=num_gather_links,
sampling_memory_config=sampling_memory_config,
allow_force_argmax=allow_force_argmax,
num_argmax_gather_links=num_argmax_gather_links,
ag_topology=ag_topology,
argmax_chunks_per_sync=argmax_chunks_per_sync,
argmax_num_workers_per_link=argmax_num_workers_per_link,
pad_to_power_of_2=getattr(args, "pad_logits_to_power_of_2", False),
)
return cls.from_config(config)
# ---------------------------------------------------------------------------
# Config resolution
# ---------------------------------------------------------------------------
def _resolve_sampling1d_config(config: Sampling1DConfig) -> Sampling1DConfig:
"""Fill None fields with topology-aware defaults."""
import torch
to_set: dict = {}
# Phase 1: Device and CCL
mesh_device = config.mesh_device or ttnn.GetDefaultDevice()
to_set["mesh_device"] = mesh_device
cluster_shape = mesh_device.shape
num_devices = mesh_device.get_num_devices()
multi_step_reduction = list(cluster_shape) == [1, 1]
if num_devices > 1 and config.tt_ccl is None:
to_set["tt_ccl"] = get_tt_ccl(mesh_device)
# Phase 2: Scalar config defaults
valid_vocab_size = config.valid_vocab_size if config.valid_vocab_size is not None else config.vocab_size
if valid_vocab_size > config.vocab_size:
raise ValueError(f"valid_vocab_size ({valid_vocab_size}) must be <= vocab_size ({config.vocab_size})")
to_set["valid_vocab_size"] = valid_vocab_size
if config.start_core is None:
to_set["start_core"] = ttnn.CoreCoord(0, 0)
if config.sampling_memory_config is None:
to_set["sampling_memory_config"] = ttnn.DRAM_MEMORY_CONFIG
if config.num_argmax_gather_links is None:
to_set["num_argmax_gather_links"] = config.num_gather_links
if config.ag_topology is None:
to_set["ag_topology"] = ttnn.Topology.Linear
# Phase 3: Buffer specs
B = config.max_batch_size
K = config.max_top_k
V = config.vocab_size
replicate_mapper = ttnn.ShardTensor2dMesh(mesh_device, dims=(None, None), mesh_shape=cluster_shape)
# num_devices_in_mesh for index computation
if multi_step_reduction:
num_devices_in_mesh = 2
else:
num_devices_in_mesh = max(cluster_shape[0], cluster_shape[1])
per_device_vocab = V // num_devices_in_mesh
def _resolve_buf(field_val, defaults, source_factory):
if field_val is None:
return LazyBuffer(source=source_factory(), **defaults)
if isinstance(field_val, ttnn.Tensor):
return field_val
return resolve_lazy_buffer(field_val, **defaults)
# index_offsets: [1, 1, B, K * num_devices_in_mesh]
def _make_index_offsets():
offsets = torch.ones(1, 1, B, K * num_devices_in_mesh, dtype=torch.int64)
for device_id in range(num_devices_in_mesh):
offsets[:, :, :, device_id * K : (device_id + 1) * K] = device_id * per_device_vocab
return offsets
idx_defaults = dict(
dtype=ttnn.int32,
layout=ttnn.TILE_LAYOUT,
device=mesh_device,
mesh_mapper=replicate_mapper,
memory_config=ttnn.DRAM_MEMORY_CONFIG,
)
to_set["index_offsets"] = _resolve_buf(config.index_offsets, idx_defaults, _make_index_offsets)
vocab_shard_dims = get_vocab_shard_dims(cluster_shape)
invalid_vocab_defaults = dict(
dtype=ttnn.bfloat16,
layout=ttnn.TILE_LAYOUT,
device=mesh_device,
mesh_mapper=ttnn.ShardTensor2dMesh(mesh_device, dims=vocab_shard_dims, mesh_shape=cluster_shape),
memory_config=ttnn.DRAM_MEMORY_CONFIG,
)
invalid_vocab_tail_mask = build_tail_invalid_vocab_mask(
valid_vocab_size,
V,
B,
cluster_shape,
tile_size=ttnn.TILE_SIZE,
)
if config.invalid_vocab_mask is None and (
invalid_vocab_tail_mask is not None or config.invalid_vocab_tail_mask is not None
):
to_set["invalid_vocab_tail_width"] = (
invalid_vocab_tail_mask.tail_width
if invalid_vocab_tail_mask is not None
else config.invalid_vocab_tail_width
)
to_set["invalid_vocab_tail_mask"] = _resolve_buf(
config.invalid_vocab_tail_mask,
invalid_vocab_defaults,
lambda: invalid_vocab_tail_mask.mask,
)
else:
invalid_vocab_mask = build_invalid_vocab_mask(valid_vocab_size, V, B)
if invalid_vocab_mask is not None or config.invalid_vocab_mask is not None:
to_set["invalid_vocab_mask"] = _resolve_buf(
config.invalid_vocab_mask,
invalid_vocab_defaults,
lambda: invalid_vocab_mask,
)
# seeds and user_ids: [B], uint32, ROW_MAJOR
seed_defaults = dict(
dtype=ttnn.uint32,
layout=ttnn.ROW_MAJOR_LAYOUT,
device=mesh_device,
memory_config=ttnn.DRAM_MEMORY_CONFIG,
)
to_set["seeds"] = _resolve_buf(
config.seeds, seed_defaults, lambda: torch.arange(B, dtype=torch.int64).to(torch.int32)
)
to_set["user_ids"] = _resolve_buf(
config.user_ids, seed_defaults, lambda: torch.arange(B, dtype=torch.int64).to(torch.int32)
)
resolved = replace(config, **to_set)
assert resolved.is_resolved(), "Config not fully resolved after _resolve_sampling1d_config"
return resolved
# -- Helper functions -------------------------------------------------------
def _materialize(buf):
"""Materialize a buffer field: ttnn.Tensor passthrough, LazyBuffer → get_device_buffer()."""
if isinstance(buf, ttnn.Tensor):
return buf
return buf.get_device_buffer()
def _attach_cleanup_failures(primary, failures):
if not failures:
return
previous = tuple(getattr(primary, "cleanup_failures", ()))
primary.cleanup_failures = previous + tuple(failures)
add_note = getattr(primary, "add_note", None)
if callable(add_note):
add_note(f"cleanup also encountered {len(failures)} failure(s)")
def _raise_cleanup_failures(failures):
primary = failures[0]
if len(failures) > 1:
previous = tuple(getattr(primary, "cleanup_failures", ()))
primary.cleanup_failures = previous + tuple(failures[1:])
add_note = getattr(primary, "add_note", None)
if callable(add_note):
add_note(f"cleanup also encountered {len(failures) - 1} additional failure(s)")
raise primary