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