# 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