Download code/models/common/llm_runtime/decode.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 57.3 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/llm_runtime/decode.py
- Command line
-
hf download hf://tt-hous/clef/code/models/common/llm_runtime/decode.py
-
curl -L -o decode.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/llm_runtime/decode.py
57.3 kB
| # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC | |
| # SPDX-License-Identifier: Apache-2.0 | |
| """Decode preparation, invocation, feedback, readback, and local resources.""" | |
| from __future__ import annotations | |
| import contextlib | |
| import dataclasses | |
| import functools | |
| from dataclasses import dataclass, field | |
| from typing import Any | |
| import torch | |
| from loguru import logger | |
| import ttnn | |
| from models.common.llm_runtime.config import PageTableLayout | |
| from models.common.llm_runtime.output_reader import OutputReader, PendingRead | |
| from models.common.llm_runtime.tensor_resources import ( | |
| TensorResourceOrphan, | |
| attach_cleanup_failures, | |
| best_effort_deallocate_owned_tensors, | |
| raise_cleanup_failures, | |
| release_orphans, | |
| ) | |
| from models.common.modules.sampling.params import ( | |
| PreparedSamplingParams, | |
| place_prepared_sampling_params, | |
| prepare_sampling_params, | |
| slice_sampling_params, | |
| ) | |
| from models.common.modules.sampling.seed_manager_1d import SeedManager1D | |
| from models.common.sampling.sampling_params import SamplingParams | |
| class DecodeProgramSignature: | |
| """Material identity of one decode eager-program variant.""" | |
| batch_size: int | |
| page_table_width: int | |
| sampling_path: str | |
| device_feedback: bool | |
| penalties_enabled: bool = False | |
| logprobs_enabled: bool = False | |
| def key_material(self) -> tuple[tuple[str, Any], ...]: | |
| material = [ | |
| ("operation", "decode"), | |
| ("batch_size", self.batch_size), | |
| ("page_table_width", self.page_table_width), | |
| ("sampling_path", self.sampling_path), | |
| ("device_feedback", self.device_feedback), | |
| ] | |
| if self.penalties_enabled: | |
| material.append(("penalties_enabled", True)) | |
| if self.logprobs_enabled: | |
| material.append(("logprobs_enabled", True)) | |
| return tuple(material) | |
| class DecodeTraceSignature: | |
| """Material identity of one full-step decode trace.""" | |
| batch_size: int | |
| page_table_width: int | |
| sampling_path: str | |
| device_feedback: bool | |
| penalties_enabled: bool = False | |
| logprobs_enabled: bool = False | |
| def key_material(self) -> tuple[tuple[str, Any], ...]: | |
| material = [ | |
| ("operation", "decode"), | |
| ("batch_size", self.batch_size), | |
| ("page_table_width", self.page_table_width), | |
| ("sampling_path", self.sampling_path), | |
| ("device_feedback", self.device_feedback), | |
| ] | |
| if self.penalties_enabled: | |
| material.append(("penalties_enabled", True)) | |
| if self.logprobs_enabled: | |
| material.append(("logprobs_enabled", True)) | |
| return tuple(material) | |
| class DecodeHostInputs: | |
| tokens: Any | |
| positions: Any | |
| rotary_indices: Any | |
| page_table: Any | |
| def values(self) -> tuple[Any, Any, Any, Any]: | |
| return self.tokens, self.positions, self.rotary_indices, self.page_table | |
| class DecodeDeviceInputs: | |
| tokens: Any | |
| positions: Any | |
| rotary_indices: Any | |
| page_table: Any | |
| def values(self) -> tuple[Any, Any, Any, Any]: | |
| return self.tokens, self.positions, self.rotary_indices, self.page_table | |
| def owned_tensor_values(self) -> tuple[Any, Any, Any, Any]: | |
| return self.values() | |
| class PreparedDecode: | |
| """One validated and normalized decode request, prepared exactly once.""" | |
| tokens: torch.Tensor | |
| start_pos: torch.Tensor | |
| page_table: torch.Tensor | |
| sampling_params: SamplingParams | None | |
| prepared_sampling: PreparedSamplingParams | None | |
| sampling_path: str | |
| reset_batch: bool | |
| device_feedback: bool | |
| page_table_changed: bool | |
| def sampling_values(self): | |
| """Compatibility view over the native prepared structure.""" | |
| sampling = self.prepared_sampling | |
| if sampling is None: | |
| return None | |
| return ( | |
| sampling.top_k, | |
| sampling.top_p, | |
| sampling.temperature, | |
| sampling.all_active_rows_greedy, | |
| ) | |
| def sampling_seeds(self): | |
| return None if self.prepared_sampling is None else self.prepared_sampling.seeds | |
| class InvocationResult: | |
| value: Any | |
| owned: Any | |
| is_tokens: bool | |
| class DecodeRefreshPolicy: | |
| every_replay: tuple[str, ...] = ("sampling",) | |
| full_on_batch_reset: bool = True | |
| full_on_graph_switch: bool = True | |
| full_without_device_feedback: bool = True | |
| refresh_page_table_on_change: bool = True | |
| class DecodeCapturePlan: | |
| """Operation callbacks consumed by the trace compiler by duck typing.""" | |
| prepare_inputs: Any | |
| capture: Any | |
| refresh_policy: DecodeRefreshPolicy = DecodeRefreshPolicy() | |
| class DecodePersistentInputs: | |
| device_inputs: DecodeDeviceInputs | |
| kpt: tuple[Any, Any, Any] | None | |
| kpt_signature: list[Any] | None = None | |
| seed_buffer: Any | None = None | |
| def owned_tensor_values(self) -> tuple[Any, ...]: | |
| return self.device_inputs.values(), self.kpt | |
| class DecodeOutputLease: | |
| raw_value: Any | |
| owned_values: Any | |
| host_value: Any = None | |
| pending: PendingRead | None = None | |
| released: bool = False | |
| deallocated_tensor_ids: set[int] = field(default_factory=set, repr=False) | |
| class DecodeRuntimeConfig: | |
| """Fully resolved, immutable decode policy and borrowed collaborators.""" | |
| model: Any | |
| mesh_device: Any | |
| output_reader: OutputReader | |
| lane_capacity: int | |
| page_table_layout: PageTableLayout # Current geometry; may be replaced once before execution. | |
| page_table_layout_ceiling: PageTableLayout # Construction-time upper bound retained across replacement. | |
| cluster_shape: tuple[int, int] | |
| num_devices: int | |
| vocab_size: int | |
| device_sampling_enabled: bool | |
| force_greedy_top_k: bool | |
| allow_force_argmax: bool | |
| max_device_top_k: int | |
| sampling_batch_size: int | |
| position_feedback_capable: bool | |
| sampling_state_controller: Any | |
| sampling_state: Any | |
| def __post_init__(self) -> None: | |
| _validate_resolved_decode_config(self) | |
| def resolve( | |
| cls, | |
| *, | |
| model: Any, | |
| output_reader: OutputReader, | |
| lane_capacity: int, | |
| page_table_layout: PageTableLayout, | |
| device_sampling_enabled: bool, | |
| force_greedy_top_k: bool = False, | |
| sampling_state_controller: Any = None, | |
| sampling_state: Any = None, | |
| ) -> "DecodeRuntimeConfig": | |
| if not isinstance(output_reader, OutputReader): | |
| raise TypeError("output_reader must be an OutputReader") | |
| mesh_device = output_reader.mesh_device | |
| model_mesh = getattr(getattr(model, "config", None), "mesh_device", None) | |
| if model_mesh is not None and model_mesh is not mesh_device: | |
| raise ValueError("model and decode runtime must use the same mesh_device") | |
| if not isinstance(lane_capacity, int) or isinstance(lane_capacity, bool) or lane_capacity <= 0: | |
| raise ValueError("lane_capacity must be a positive integer") | |
| if lane_capacity > 32: | |
| raise ValueError("decode token input padding supports at most 32 lane slots") | |
| if not isinstance(device_sampling_enabled, bool): | |
| raise TypeError("device_sampling_enabled must be bool") | |
| if not isinstance(force_greedy_top_k, bool): | |
| raise TypeError("force_greedy_top_k must be bool") | |
| _validate_page_table_layout(page_table_layout) | |
| try: | |
| cluster_shape = tuple(int(value) for value in mesh_device.shape) | |
| except (AttributeError, TypeError, ValueError) as error: | |
| raise TypeError("mesh_device must provide a two-dimensional shape") from error | |
| if len(cluster_shape) != 2 or any(value <= 0 for value in cluster_shape): | |
| raise ValueError("mesh_device shape must contain two positive dimensions") | |
| model_config = getattr(model, "config", None) | |
| num_devices = getattr(model_config, "num_devices", None) | |
| if num_devices is None: | |
| num_devices = getattr(model, "num_devices", cluster_shape[0] * cluster_shape[1]) | |
| if not isinstance(num_devices, int) or isinstance(num_devices, bool) or num_devices <= 0: | |
| raise ValueError("model num_devices must be a positive integer") | |
| if num_devices != cluster_shape[0] * cluster_shape[1]: | |
| raise ValueError("model num_devices must match the decode mesh shape") | |
| vocab_size = getattr(model, "vocab_size", None) | |
| if not isinstance(vocab_size, int) or isinstance(vocab_size, bool) or vocab_size <= 0: | |
| raise ValueError("model vocab_size must be a positive integer") | |
| sampling = getattr(model, "sampling", None) | |
| sampling_config = getattr(sampling, "config", None) | |
| allow_force_argmax = getattr(sampling_config, "allow_force_argmax", False) | |
| max_device_top_k = getattr(sampling_config, "max_top_k", 0) | |
| sampling_batch_size = getattr(sampling_config, "max_batch_size", lane_capacity) | |
| if device_sampling_enabled: | |
| if not callable(getattr(sampling, "decode_forward", None)): | |
| raise TypeError("device sampling requires model.sampling.decode_forward()") | |
| if not isinstance(allow_force_argmax, bool): | |
| raise TypeError("model sampling allow_force_argmax must be bool") | |
| if not isinstance(max_device_top_k, int) or isinstance(max_device_top_k, bool) or max_device_top_k <= 0: | |
| raise ValueError("model sampling max_top_k must be a positive integer") | |
| if ( | |
| not isinstance(sampling_batch_size, int) | |
| or isinstance(sampling_batch_size, bool) | |
| or sampling_batch_size < lane_capacity | |
| ): | |
| raise ValueError("model sampling max_batch_size must cover the decode lane capacity") | |
| else: | |
| allow_force_argmax = False | |
| max_device_top_k = 0 | |
| sampling_batch_size = lane_capacity | |
| return cls( | |
| model=model, | |
| mesh_device=mesh_device, | |
| output_reader=output_reader, | |
| lane_capacity=lane_capacity, | |
| page_table_layout=page_table_layout, | |
| cluster_shape=cluster_shape, | |
| num_devices=num_devices, | |
| vocab_size=vocab_size, | |
| device_sampling_enabled=device_sampling_enabled, | |
| force_greedy_top_k=force_greedy_top_k, | |
| allow_force_argmax=allow_force_argmax, | |
| max_device_top_k=max_device_top_k, | |
| sampling_batch_size=sampling_batch_size, | |
| position_feedback_capable=callable(getattr(model, "increment_positions", None)), | |
| sampling_state_controller=sampling_state_controller, | |
| sampling_state=sampling_state, | |
| page_table_layout_ceiling=page_table_layout, | |
| ) | |
| def with_page_table_layout(self, layout: PageTableLayout) -> "DecodeRuntimeConfig": | |
| """Return a validated geometry replacement within the original ceiling.""" | |
| _validate_page_table_layout(layout) | |
| if layout.block_size != self.page_table_layout.block_size: | |
| raise ValueError("replacement page-table layout cannot change block_size") | |
| if layout.raw_capacity_width > self.page_table_layout_ceiling.raw_capacity_width: | |
| raise ValueError("replacement page-table capacity exceeds the construction-time ceiling") | |
| if layout.decode_width > self.page_table_layout_ceiling.decode_width: | |
| raise ValueError("replacement decode width exceeds the construction-time ceiling") | |
| return dataclasses.replace(self, page_table_layout=layout) | |
| class DecodeRuntime: | |
| """Prepare, execute, trace, and consume decode for one execution lane. | |
| The eager call chain is | |
| `EagerExecutor.decode_forward()` → `prepare` → `invoke` → | |
| `consume`. Trace warmup uses `capture_plan`; replay calls | |
| `refresh_trace`, `note_submitted`, and `consume`. | |
| `Llama3Executor` also exposes `read_decode_output` and | |
| `process_decode_output_host` for vLLM's asynchronous output path. | |
| The model, mesh, output reader, sampler, and KV-backed page-table values are | |
| borrowed. Only staged invocation tensors, raw outputs, output leases, and | |
| retryable decode transients are released here. | |
| Explicit ``SamplingParams.seed`` routing is a decode-top-k contract. The | |
| common prefill runtime retains its existing sampling behavior and does not | |
| consume this decode RNG stream; callers that need one coherent controlled | |
| stream must keep prefill on the logits path and begin sampling in decode. | |
| """ | |
| def __init__(self, config: DecodeRuntimeConfig): | |
| if not isinstance(config, DecodeRuntimeConfig): | |
| raise TypeError("config must be a DecodeRuntimeConfig") | |
| self.config = config | |
| self._previous_page_table: torch.Tensor | None = None | |
| self._normalization_source: torch.Tensor | None = None | |
| self._normalization_copy_blocks: tuple[int, ...] | None = None | |
| self._normalization_layout: tuple[int, int, int] | None = None | |
| self._normalized_page_table: torch.Tensor | None = None | |
| self._external_by_raw_id: dict[int, DecodeOutputLease] = {} | |
| self._external_by_host_id: dict[int, DecodeOutputLease] = {} | |
| self._transient_orphans: list[TensorResourceOrphan] = [] | |
| self._sampling_state_controller = config.sampling_state_controller | |
| self._sampling_state = config.sampling_state | |
| sampling_config = getattr(getattr(config.model, "sampling", None), "config", None) | |
| seed_buffer = getattr(sampling_config, "seeds", None) | |
| has_mutable_seed_buffer = all( | |
| callable(getattr(seed_buffer, name, None)) for name in ("update", "get_device_buffer") | |
| ) and isinstance(getattr(seed_buffer, "source", None), torch.Tensor) | |
| # Converted executors without SamplingState1D still hit this fallback. | |
| # They are vLLM-facing, so keep concurrent same-seed slots unsalted. | |
| self._seed_manager = ( | |
| self._sampling_state_controller.seed_manager | |
| if self._sampling_state_controller is not None | |
| else SeedManager1D(sampling_config, salt_duplicate_seeds=False) | |
| if config.device_sampling_enabled and has_mutable_seed_buffer | |
| else None | |
| ) | |
| self._seed_state = ( | |
| self._sampling_state.seed_state | |
| if self._sampling_state is not None | |
| else self._seed_manager.create_state() | |
| if self._seed_manager is not None | |
| else None | |
| ) | |
| if self._seed_state is not None and self._seed_state.capacity < config.lane_capacity: | |
| raise ValueError("model sampling seed buffer is smaller than the decode lane capacity") | |
| # Public API | |
| def transient_orphan_count(self) -> int: | |
| """Return the number of failed transient releases awaiting cleanup.""" | |
| return len(self._transient_orphans) | |
| def configure_page_table_layout(self, layout: PageTableLayout) -> None: | |
| """Install final physical-capacity geometry before allocation.""" | |
| self.config = self.config.with_page_table_layout(layout) | |
| def prepare( | |
| self, | |
| tokens: torch.Tensor, | |
| start_pos: torch.Tensor, | |
| page_table: torch.Tensor, | |
| *, | |
| sampling_params: Any = None, # ↓ Sampling | |
| prompt_tokens: Any = None, # ↓ Request-owned sampling state | |
| output_tokens: Any = None, | |
| slot_remap: Any = None, | |
| reset_batch: bool = False, # ↓ State transition | |
| ) -> PreparedDecode: | |
| """Normalize one host decode request into an immutable prepared value.""" | |
| self._ensure_usable() | |
| self._validate_inputs(tokens, start_pos, page_table) | |
| self._validate_sampling_request(sampling_params) | |
| feedback = self._classify_feedback(sampling_params) | |
| prepared_sampling = None | |
| if sampling_params is not None: | |
| active_slots = tuple(slot for slot, position in enumerate(start_pos) if int(position) >= 0) | |
| if not active_slots: | |
| raise ValueError("decode sampling requires at least one active slot") | |
| request_sampling = slice_sampling_params(sampling_params, active_slots) | |
| prepared_sampling = prepare_sampling_params( | |
| request_sampling, | |
| self.config.sampling_batch_size, | |
| max_device_top_k=self.config.max_device_top_k, | |
| allow_force_argmax=self.config.allow_force_argmax and not self.config.force_greedy_top_k, | |
| prompt_tokens=_select_decode_request_state( | |
| prompt_tokens, | |
| active_slots=active_slots, | |
| lane_capacity=self.config.lane_capacity, | |
| ), | |
| output_tokens=_select_decode_request_state( | |
| output_tokens, | |
| active_slots=active_slots, | |
| lane_capacity=self.config.lane_capacity, | |
| ), | |
| slot_remap=_normalize_decode_slot_remap( | |
| slot_remap, | |
| lane_capacity=self.config.lane_capacity, | |
| sampling_batch_size=self.config.sampling_batch_size, | |
| ), | |
| ) | |
| prepared_sampling = place_prepared_sampling_params( | |
| prepared_sampling, | |
| active_slots, | |
| ) | |
| normalized = self._normalize_page_table( | |
| page_table, | |
| start_pos, | |
| allow_one_step_feedback_lag=feedback, | |
| ) | |
| return PreparedDecode( | |
| tokens=tokens, | |
| start_pos=start_pos, | |
| page_table=normalized, | |
| sampling_params=sampling_params, | |
| prepared_sampling=prepared_sampling, | |
| sampling_path=self._classify_sampling_path(prepared_sampling), | |
| reset_batch=bool(reset_batch), | |
| device_feedback=feedback, | |
| page_table_changed=( | |
| self._previous_page_table is None or not torch.equal(self._previous_page_table, normalized) | |
| ), | |
| ) | |
| def program_signature(self, prepared: PreparedDecode) -> DecodeProgramSignature: | |
| """Return the eager program identity for a prepared decode request.""" | |
| self._require_prepared(prepared) | |
| return self._program_signature(prepared) | |
| def trace_signature(self, prepared: PreparedDecode) -> DecodeTraceSignature: | |
| """Return the trace identity for a prepared decode request.""" | |
| self._require_prepared(prepared) | |
| program = self._program_signature(prepared) | |
| return DecodeTraceSignature( | |
| batch_size=program.batch_size, | |
| page_table_width=program.page_table_width, | |
| sampling_path=program.sampling_path, | |
| penalties_enabled=program.penalties_enabled, | |
| logprobs_enabled=program.logprobs_enabled, | |
| device_feedback=program.device_feedback, | |
| ) | |
| def invoke( | |
| self, | |
| prepared: PreparedDecode, | |
| *, | |
| device_feedback: bool = False, | |
| count_tokens: bool = True, | |
| ) -> InvocationResult: | |
| """Stage and execute one prepared request eagerly.""" | |
| self._ensure_usable() | |
| self._require_prepared(prepared) | |
| host_inputs = self._prepare_inputs_host(prepared) | |
| device_inputs, kpt = self._stage_inputs_and_kpt(host_inputs, prepared) | |
| owned = (device_inputs, kpt) | |
| compile_only_state = not count_tokens and ( | |
| self._sampling_state_controller is not None or self._seed_manager is not None | |
| ) | |
| try: | |
| if compile_only_state: | |
| sampling = prepared.prepared_sampling | |
| if self._sampling_state_controller is not None and sampling is not None: | |
| self._sampling_state_controller.reset( | |
| self._sampling_state, | |
| dataclasses.replace(sampling, slot_remap=None), | |
| ) | |
| elif self._seed_manager is not None: | |
| # Program warmup compiles sampled decode before any real | |
| # prefill/admission boundary. Give the fallback seed | |
| # manager a temporary replacement batch so its normal | |
| # strict serving checks do not mistake synthetic warmup | |
| # rows for live requests. | |
| self._refresh_sampling_seeds(dataclasses.replace(prepared, reset_batch=True)) | |
| else: | |
| self._refresh_sampling_seeds(prepared) | |
| with _validate_module_inputs(self.config.model): | |
| output = self._run_body( | |
| device_inputs, | |
| prepared, | |
| kpt, | |
| device_feedback=device_feedback and prepared.device_feedback, | |
| count_tokens=count_tokens, | |
| advance_seeds=count_tokens, | |
| ) | |
| except BaseException as primary: | |
| if compile_only_state: | |
| try: | |
| self._reset_compile_only_sampling_state() | |
| except BaseException as cleanup_error: | |
| attach_cleanup_failures(primary, (cleanup_error,)) | |
| failures = self._release_or_retain_transient(owned) | |
| attach_cleanup_failures(primary, failures) | |
| raise | |
| if compile_only_state: | |
| self._reset_compile_only_sampling_state() | |
| self._note_submitted(prepared) | |
| return InvocationResult( | |
| value=output, | |
| owned=(output, owned), | |
| is_tokens=prepared.sampling_params is not None, | |
| ) | |
| def capture_plan(self, prepared: PreparedDecode) -> DecodeCapturePlan: | |
| """Describe persistent inputs and capture work for one decode trace.""" | |
| self._require_prepared(prepared) | |
| def prepare_inputs() -> DecodePersistentInputs: | |
| if self._sampling_state_controller is not None and prepared.prepared_sampling is not None: | |
| self._sampling_state_controller.reset( | |
| self._sampling_state, | |
| dataclasses.replace(prepared.prepared_sampling, slot_remap=None), | |
| ) | |
| host_inputs = self._prepare_inputs_host(prepared) | |
| device_inputs, kpt = self._stage_inputs_and_kpt(host_inputs, prepared) | |
| sampling = prepared.prepared_sampling | |
| signature = [(sampling.top_k, sampling.top_p, sampling.temperature)] if kpt is not None else None | |
| return DecodePersistentInputs( | |
| device_inputs=device_inputs, | |
| kpt=kpt, | |
| kpt_signature=signature, | |
| seed_buffer=self._seed_device_handle(), | |
| ) | |
| def capture(persistent: Any) -> Any: | |
| values = _persistent_values(persistent) | |
| return self._run_body( | |
| values.device_inputs, | |
| prepared, | |
| values.kpt, | |
| device_feedback=prepared.device_feedback, | |
| advance_seeds=False, | |
| ) | |
| return DecodeCapturePlan(prepare_inputs=prepare_inputs, capture=capture) | |
| def refresh_trace( | |
| self, | |
| artifact: Any, | |
| prepared: PreparedDecode, | |
| decision: Any, | |
| ) -> None: | |
| self._require_prepared(prepared) | |
| values = _persistent_values(artifact) | |
| self._validate_trace_seed_handle(values) | |
| # Sampling1D's trace captures the stable model-owned seed tensor handle. | |
| # Refresh its contents before replay so the trace cannot observe stale | |
| # request state. | |
| self._refresh_sampling_seeds(prepared) | |
| if bool(decision.full): | |
| host_inputs = self._prepare_inputs_host(prepared) | |
| _copy_host_to_device(host_inputs.values(), values.device_inputs.values()) | |
| elif bool(decision.page_table): | |
| host_inputs = self._prepare_inputs_host(prepared) | |
| ttnn.copy_host_to_device_tensor(host_inputs.page_table, values.device_inputs.page_table) | |
| if prepared.sampling_path == "topk": | |
| sampling = prepared.prepared_sampling | |
| if sampling is None: | |
| raise RuntimeError("top-k decode trace is missing prepared sampling parameters") | |
| signature = sampling.top_k, sampling.top_p, sampling.temperature | |
| if values.kpt_signature is None or values.kpt_signature[0] != signature: | |
| self._refresh_kpt(values.kpt, prepared) | |
| if values.kpt_signature is not None: | |
| values.kpt_signature[0] = signature | |
| elif values.kpt is not None: | |
| raise RuntimeError("non-top-k decode trace unexpectedly owns sampling inputs") | |
| def note_submitted(self, prepared: PreparedDecode) -> None: | |
| """Advance feedback comparison state immediately after device submission.""" | |
| self._require_prepared(prepared) | |
| self._note_submitted(prepared) | |
| def consume(self, result: InvocationResult, *, read_from_device: bool = True) -> Any: | |
| """Read and normalize an invocation or transfer it to an external lease.""" | |
| if not isinstance(result, InvocationResult): | |
| raise TypeError("result must be an InvocationResult") | |
| if not read_from_device: | |
| if result.owned is not None: | |
| lease = DecodeOutputLease(raw_value=result.value, owned_values=result.owned) | |
| self._external_by_raw_id[id(result.value)] = lease | |
| return result.value | |
| try: | |
| host = self.config.output_reader.read(result.value, blocking=True) | |
| normalized = self._normalize_host_output( | |
| host, | |
| is_tokens=result.is_tokens, | |
| ) | |
| except BaseException as primary: | |
| failures = self._release_or_retain_transient(result.owned) | |
| attach_cleanup_failures(primary, failures) | |
| raise | |
| failures = self._release_or_retain_transient(result.owned) | |
| if failures: | |
| raise_cleanup_failures(failures) | |
| return normalized | |
| def read_decode_output(self, tt_out: Any, *, async_read: bool = False) -> Any: | |
| """Read a raw externally leased decode output, optionally asynchronously.""" | |
| if not async_read: | |
| host = self.config.output_reader.read(tt_out, blocking=True) | |
| self._release_external_lease(self._external_by_raw_id.get(id(tt_out))) | |
| return host | |
| pending = self.config.output_reader.submit(tt_out) | |
| lease = self._external_by_raw_id.get(id(tt_out)) | |
| if lease is not None: | |
| lease.host_value = pending.value | |
| lease.pending = pending | |
| self._external_by_host_id[id(pending.value)] = lease | |
| return pending.value, list(pending.events) | |
| def process_decode_output_host(self, tt_out: Any, *, is_tokens: bool = False) -> tuple[Any, Any]: | |
| """Complete and normalize a host value returned by async decode read.""" | |
| completed = self.config.output_reader.complete(tt_out) | |
| self._release_external_lease(self._external_by_host_id.get(id(tt_out))) | |
| return self._normalize_host_output(completed, is_tokens=is_tokens) | |
| def drain_external_outputs(self) -> None: | |
| """Synchronize and release every outstanding externally owned output.""" | |
| failures = [] | |
| for lease in tuple(self._external_by_raw_id.values()): | |
| try: | |
| if lease.pending is None: | |
| ttnn.synchronize_device(self.config.mesh_device) | |
| self._release_external_lease(lease) | |
| except BaseException as error: | |
| failures.append(error) | |
| if failures: | |
| raise_cleanup_failures(failures) | |
| def cleanup_transients(self) -> None: | |
| """Retry every transient tensor release that previously failed.""" | |
| failures = release_orphans(self._transient_orphans) | |
| if failures: | |
| raise_cleanup_failures(failures) | |
| # Private implementation | |
| def _validate_sampling_request(self, sampling_params: SamplingParams | None) -> None: | |
| if sampling_params is not None and not self.config.device_sampling_enabled: | |
| raise ValueError("sampling parameters were supplied while device sampling is disabled") | |
| def _classify_sampling_path(self, prepared_sampling: PreparedSamplingParams | None) -> str: | |
| if prepared_sampling is None: | |
| return "logits" | |
| return prepared_sampling.sampling_path | |
| def _classify_feedback(self, sampling_params: SamplingParams | None) -> bool: | |
| return sampling_params is not None and self.config.position_feedback_capable | |
| def _convert_logits(self, value: Any) -> torch.Tensor: | |
| if isinstance(value, torch.Tensor): | |
| output = value.float() | |
| elif self.config.num_devices == 1: | |
| output = ttnn.to_torch(value).float() | |
| else: | |
| output = _concat_host_output(value, self.config.cluster_shape).float() | |
| return self._slice_logits(output) | |
| def _slice_logits(self, output: torch.Tensor) -> torch.Tensor: | |
| config = self.config | |
| return output[:, :, : config.lane_capacity, : config.vocab_size].contiguous().view(config.lane_capacity, 1, -1) | |
| def _program_signature(self, prepared: PreparedDecode) -> DecodeProgramSignature: | |
| return DecodeProgramSignature( | |
| batch_size=self.config.lane_capacity, | |
| page_table_width=int(prepared.page_table.shape[-1]), | |
| sampling_path=prepared.sampling_path, | |
| penalties_enabled=( | |
| prepared.prepared_sampling.penalties_enabled if prepared.prepared_sampling is not None else False | |
| ), | |
| logprobs_enabled=( | |
| prepared.prepared_sampling.log_probs_enabled if prepared.prepared_sampling is not None else False | |
| ), | |
| device_feedback=prepared.device_feedback, | |
| ) | |
| def _note_submitted(self, prepared: PreparedDecode) -> None: | |
| self._previous_page_table = prepared.page_table.clone() | |
| def _normalize_host_output(self, host_output: Any, *, is_tokens: bool) -> tuple[Any, Any]: | |
| if isinstance(host_output, tuple): | |
| if len(host_output) != 2: | |
| raise TypeError("runtime output tuple must contain (output, log_probs)") | |
| output, log_probs = host_output | |
| else: | |
| output, log_probs = host_output, None | |
| if is_tokens: | |
| tokens = _process_output_tokens(output, self.config.lane_capacity, self.config.cluster_shape) | |
| return tokens.to(torch.int64), _process_sampled_log_probs(log_probs, self.config.lane_capacity) | |
| return self._convert_logits(output), log_probs | |
| def _normalize_page_table(self, page_table, start_pos, *, allow_one_step_feedback_lag): | |
| layout = self.config.page_table_layout | |
| raw_width = layout.raw_capacity_width | |
| decode_width = layout.decode_width | |
| block_size = layout.block_size | |
| copy_blocks_by_row = [] | |
| for row, position_value in enumerate(start_pos): | |
| position = int(position_value) | |
| used_blocks = _num_blocks(max(0, position + 1), block_size) | |
| if used_blocks > raw_width: | |
| raise ValueError("decode position exceeds the configured paged-KV capacity") | |
| if int(page_table.shape[1]) < used_blocks: | |
| raise ValueError(f"page table is too narrow for decode row {row}") | |
| copy_blocks = used_blocks | |
| if allow_one_step_feedback_lag and position >= 0 and (position + 1) % block_size == 0: | |
| copy_blocks = min(used_blocks + 1, raw_width, int(page_table.shape[1])) | |
| copy_blocks_by_row.append(copy_blocks) | |
| layout = (raw_width, decode_width, block_size) | |
| copy_blocks_by_row = tuple(copy_blocks_by_row) | |
| source = self._normalization_source | |
| if ( | |
| source is not None | |
| and self._normalization_layout == layout | |
| and self._normalization_copy_blocks == copy_blocks_by_row | |
| and source.shape == page_table.shape | |
| and source.device == page_table.device | |
| and source.dtype == page_table.dtype | |
| and torch.equal(source, page_table) | |
| ): | |
| assert self._normalized_page_table is not None | |
| return self._normalized_page_table | |
| normalized = torch.zeros((int(page_table.shape[0]), decode_width), dtype=torch.int32, device=page_table.device) | |
| for row, copy_blocks in enumerate(copy_blocks_by_row): | |
| normalized[row, :copy_blocks] = page_table[row, :copy_blocks].to(torch.int32) | |
| self._normalization_source = page_table.clone() | |
| self._normalization_copy_blocks = copy_blocks_by_row | |
| self._normalization_layout = layout | |
| self._normalized_page_table = normalized | |
| return normalized | |
| def _prepare_inputs_host(self, prepared: PreparedDecode) -> DecodeHostInputs: | |
| config = self.config | |
| padded = torch.nn.functional.pad(prepared.tokens.reshape(-1), (0, 32 - config.lane_capacity)) | |
| tokens_tt = ttnn.unsqueeze_to_4D( | |
| ttnn.from_torch( | |
| padded, | |
| device=None, | |
| dtype=ttnn.uint32, | |
| mesh_mapper=ttnn.ReplicateTensorToMesh(config.mesh_device), | |
| ) | |
| ) | |
| nonnegative = torch.maximum(prepared.start_pos, torch.zeros_like(prepared.start_pos)) | |
| rotary = config.model.rope_setup.get_rot_idxs(nonnegative, on_host=True) | |
| mapper = ttnn.ShardTensor2dMesh( | |
| config.mesh_device, | |
| dims=(None, None), | |
| mesh_shape=config.cluster_shape, | |
| ) | |
| positions = ttnn.from_torch(prepared.start_pos, device=None, dtype=ttnn.int32, mesh_mapper=mapper) | |
| page_table = ttnn.from_torch(prepared.page_table, device=None, dtype=ttnn.int32, mesh_mapper=mapper) | |
| return DecodeHostInputs(tokens_tt, positions, rotary, page_table) | |
| def _stage_inputs_and_kpt(self, host_inputs, prepared): | |
| device_inputs = None | |
| try: | |
| raw = _copy_host_to_device(host_inputs.values(), mesh_device=self.config.mesh_device) | |
| device_inputs = DecodeDeviceInputs(*raw) | |
| kpt = self._make_device_kpt(prepared) | |
| except BaseException as primary: | |
| failures = self._release_or_retain_transient(device_inputs) | |
| attach_cleanup_failures(primary, failures) | |
| raise | |
| return device_inputs, kpt | |
| def _run_body( | |
| self, | |
| inputs, | |
| prepared, | |
| kpt, | |
| *, | |
| device_feedback, | |
| count_tokens=True, | |
| advance_seeds=True, | |
| ): | |
| model = self.config.model | |
| rot_mats = model.rope_setup.get_rot_mats(inputs.rotary_indices) | |
| logits = model.decode_forward( | |
| model.embed_decode(inputs.tokens), | |
| inputs.positions, | |
| rot_mats, | |
| page_table=inputs.page_table, | |
| ) | |
| sampling = prepared.prepared_sampling | |
| if sampling is None: | |
| return model.gather_and_untilize_logits(logits), None | |
| if self._sampling_state_controller is not None: | |
| sampling = dataclasses.replace(sampling, slot_remap=None) | |
| output = self._sampling_state_controller.decode_forward( | |
| logits, | |
| self._sampling_state, | |
| sampling, | |
| k=None if kpt is None else kpt[0], | |
| p=None if kpt is None else kpt[1], | |
| temp=None if kpt is None else kpt[2], | |
| positions=prepared.start_pos, | |
| tt_out_tok=None, | |
| count_tokens=count_tokens, | |
| advance_seeds=advance_seeds, | |
| ) | |
| else: | |
| output = self._sample_device(logits, kpt, sampling) | |
| if device_feedback: | |
| sampled_tokens = ttnn.reshape(output[0], inputs.tokens.shape) | |
| ttnn.copy(input_a=sampled_tokens, input_b=inputs.tokens) | |
| model.increment_positions(inputs.positions, inputs.rotary_indices) | |
| return output | |
| def _sample_device(self, logits, kpt, sampling: PreparedSamplingParams): | |
| if kpt is None: | |
| return self.config.model.sampling.decode_forward( | |
| logits, | |
| tt_out_tok=None, | |
| enable_log_probs=sampling.enable_log_probs, | |
| ) | |
| return self.config.model.sampling.decode_forward( | |
| logits, | |
| k=kpt[0], | |
| p=kpt[1], | |
| temp=kpt[2], | |
| tt_out_tok=None, | |
| enable_log_probs=sampling.enable_log_probs, | |
| ) | |
| def _seed_device_handle(self): | |
| if self._seed_manager is None: | |
| return None | |
| return self._seed_manager.get_seed_device_buffer() | |
| def _validate_trace_seed_handle(self, persistent: DecodePersistentInputs) -> None: | |
| if self._seed_manager is None: | |
| if persistent.seed_buffer is not None: | |
| raise RuntimeError("decode trace unexpectedly captured a seed buffer") | |
| return | |
| current = self._seed_device_handle() | |
| if persistent.seed_buffer is not current: | |
| raise RuntimeError("decode trace seed buffer handle changed after capture") | |
| def _refresh_sampling_seeds(self, prepared: PreparedDecode) -> None: | |
| if self._sampling_state_controller is not None: | |
| sampling = prepared.prepared_sampling | |
| if sampling is None: | |
| self._sampling_state_controller.seed_manager.restore_defaults(self._sampling_state.seed_state) | |
| return | |
| sampling = self._sampling_state_controller.synchronize_decode( | |
| self._sampling_state, | |
| sampling, | |
| reset_batch=prepared.reset_batch, | |
| ) | |
| self._sampling_state_controller.refresh_dynamic_inputs( | |
| self._sampling_state, | |
| sampling, | |
| positions=prepared.start_pos, | |
| ) | |
| return | |
| manager = self._seed_manager | |
| state = self._seed_state | |
| sampling = prepared.prepared_sampling | |
| active_slots = [slot for slot, position in enumerate(prepared.start_pos) if int(position) >= 0] | |
| if manager is None or state is None: | |
| if sampling is not None and any(sampling.seeds[slot] is not None for slot in active_slots): | |
| raise TypeError("explicit request seeds require a native mutable Sampling1D seed buffer") | |
| return | |
| if sampling is None: | |
| manager.cleanup(state, active_slots) | |
| manager.restore_defaults(state) | |
| return | |
| if sampling.slot_remap is not None: | |
| manager.apply_slot_remap(state, sampling.slot_remap) | |
| manager.synchronize( | |
| state, | |
| sampling.seeds, | |
| active_slots, | |
| reset_batch=prepared.reset_batch, | |
| ) | |
| if prepared.sampling_path == "topk": | |
| manager.refresh(state, active_slots, positions=prepared.start_pos) | |
| else: | |
| manager.restore_defaults(state) | |
| def _reset_compile_only_sampling_state(self) -> None: | |
| """Discard synthetic compile admissions and restore seed defaults.""" | |
| if self._sampling_state_controller is not None: | |
| self._sampling_state_controller.reset(self._sampling_state) | |
| elif self._seed_manager is not None and self._seed_state is not None: | |
| self._seed_manager.reset(self._seed_state) | |
| def _make_device_kpt(self, prepared): | |
| host = self._make_host_kpt(prepared) | |
| if host is None: | |
| return None | |
| return tuple(_copy_host_to_device(host, mesh_device=self.config.mesh_device)) | |
| def _make_host_kpt(self, prepared): | |
| sampling = prepared.prepared_sampling | |
| if sampling is None or prepared.sampling_path == "argmax": | |
| return None | |
| k, p, temperature = sampling.top_k, sampling.top_p, sampling.temperature | |
| mapper = ttnn.ReplicateTensorToMesh(self.config.mesh_device) | |
| return ( | |
| ttnn.from_torch( | |
| torch.tensor(k, dtype=torch.int32), | |
| device=None, | |
| dtype=ttnn.uint32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=mapper, | |
| ), | |
| ttnn.from_torch( | |
| torch.tensor(p, dtype=torch.float32), | |
| device=None, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=mapper, | |
| ), | |
| ttnn.from_torch( | |
| torch.tensor(temperature, dtype=torch.float32), | |
| device=None, | |
| dtype=ttnn.bfloat16, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=mapper, | |
| ), | |
| ) | |
| def _refresh_kpt(self, device_kpt, prepared): | |
| host_kpt = self._make_host_kpt(prepared) | |
| if (host_kpt is None) != (device_kpt is None): | |
| raise RuntimeError("sampling parameters changed the compiled sampling path") | |
| if host_kpt is not None: | |
| _copy_host_to_device(host_kpt, device_kpt) | |
| def _validate_inputs(self, tokens, start_pos, page_table): | |
| if not isinstance(tokens, torch.Tensor) or tokens.ndim != 1: | |
| raise ValueError("decode tokens must be a rank-1 torch.Tensor") | |
| if not isinstance(start_pos, torch.Tensor) or start_pos.ndim != 1: | |
| raise ValueError("decode start_pos must be a rank-1 torch.Tensor") | |
| if not isinstance(page_table, torch.Tensor) or page_table.ndim != 2: | |
| raise ValueError("decode page_table must be a rank-2 torch.Tensor") | |
| lane_capacity = self.config.lane_capacity | |
| if int(tokens.shape[0]) != lane_capacity: | |
| raise ValueError(f"decode batch {tokens.shape[0]} must equal lane capacity {lane_capacity}") | |
| if int(start_pos.shape[0]) != lane_capacity or int(page_table.shape[0]) != lane_capacity: | |
| raise ValueError("decode tokens, start_pos, and page_table batches must match") | |
| def _require_prepared(self, prepared): | |
| if not isinstance(prepared, PreparedDecode): | |
| raise TypeError("prepared must be a PreparedDecode") | |
| def _ensure_usable(self): | |
| if self._transient_orphans: | |
| raise RuntimeError("DecodeRuntime has unreleased transient resources; clean up this runtime") | |
| def _release_external_lease(self, lease): | |
| if lease is None or lease.released: | |
| return | |
| if lease.pending is not None: | |
| self.config.output_reader.complete(lease.pending) | |
| failures = [] | |
| if lease.owned_values is not None: | |
| failures = best_effort_deallocate_owned_tensors( | |
| (lease.raw_value, lease.owned_values), | |
| lease.deallocated_tensor_ids, | |
| ) | |
| if failures: | |
| raise_cleanup_failures(failures) | |
| lease.released = True | |
| self._external_by_raw_id.pop(id(lease.raw_value), None) | |
| if lease.host_value is not None: | |
| self._external_by_host_id.pop(id(lease.host_value), None) | |
| def _release_or_retain_transient(self, values): | |
| orphan = TensorResourceOrphan(values) | |
| failures = best_effort_deallocate_owned_tensors(orphan.values, orphan.deallocated_tensor_ids) | |
| if failures: | |
| self._transient_orphans.append(orphan) | |
| return failures | |
| def _persistent_values(value: Any) -> DecodePersistentInputs: | |
| persistent = getattr(value, "persistent_inputs", value) | |
| values = getattr(persistent, "values", persistent) | |
| if isinstance(values, DecodePersistentInputs): | |
| return values | |
| if isinstance(values, dict): | |
| device = values["device_inputs"] | |
| if not isinstance(device, DecodeDeviceInputs): | |
| device = DecodeDeviceInputs(*device) | |
| return DecodePersistentInputs( | |
| device_inputs=device, | |
| kpt=values.get("kpt"), | |
| kpt_signature=values.get("kpt_signature"), | |
| seed_buffer=values.get("seed_buffer"), | |
| ) | |
| raise TypeError("decode persistent inputs have an unsupported representation") | |
| def _select_decode_request_state( | |
| value: Any, | |
| *, | |
| active_slots: tuple[int, ...], | |
| lane_capacity: int, | |
| ) -> Any: | |
| """Convert slot-indexed decode history into active request order.""" | |
| if value is None: | |
| return None | |
| if isinstance(value, torch.Tensor): | |
| if value.ndim == 0: | |
| raise ValueError("decode sampling history must have a leading request dimension") | |
| length = int(value.shape[0]) | |
| elif isinstance(value, (list, tuple)): | |
| length = len(value) | |
| else: | |
| raise TypeError("decode sampling history must be a tensor or sequence") | |
| if length in (1, len(active_slots)): | |
| return value | |
| if length < int(lane_capacity): | |
| raise ValueError( | |
| f"decode sampling history has {length} rows, expected 1, {len(active_slots)}, " | |
| f"or at least lane capacity {lane_capacity}" | |
| ) | |
| rows = ( | |
| torch.tensor(active_slots, dtype=torch.long, device=value.device) if isinstance(value, torch.Tensor) else None | |
| ) | |
| if rows is not None: | |
| return value.index_select(0, rows) | |
| selected = [value[slot] for slot in active_slots] | |
| return tuple(selected) if isinstance(value, tuple) else selected | |
| def _normalize_decode_slot_remap( | |
| value: Any, | |
| *, | |
| lane_capacity: int, | |
| sampling_batch_size: int, | |
| ) -> Any: | |
| """Extend a lane-local remap with identity rows for sampler-only padding.""" | |
| if value is None: | |
| return None | |
| if isinstance(value, torch.Tensor): | |
| if value.ndim == 0: | |
| raise ValueError("decode slot_remap must have a leading slot dimension") | |
| flat = value.reshape(-1) | |
| length = int(flat.numel()) | |
| elif isinstance(value, (list, tuple)): | |
| flat = list(value) | |
| length = len(flat) | |
| else: | |
| raise TypeError("decode slot_remap must be a tensor or sequence") | |
| if length == int(sampling_batch_size): | |
| return value | |
| if length != int(lane_capacity): | |
| raise ValueError( | |
| f"decode slot_remap has {length} rows, expected lane capacity {lane_capacity} " | |
| f"or sampler capacity {sampling_batch_size}" | |
| ) | |
| sources = [int(source) for source in flat] | |
| if any(source < 0 or source >= int(lane_capacity) for source in sources): | |
| raise ValueError("decode slot_remap contains a source outside the lane capacity") | |
| tail = list(range(int(lane_capacity), int(sampling_batch_size))) | |
| if isinstance(value, torch.Tensor): | |
| return torch.cat( | |
| [ | |
| flat, | |
| torch.tensor(tail, dtype=value.dtype, device=value.device), | |
| ] | |
| ) | |
| extended = sources + tail | |
| return tuple(extended) if isinstance(value, tuple) else extended | |
| def _validate_module_inputs(model: Any): | |
| """Instrument one decode forward pass against declared input memory configs.""" | |
| mismatches = [] | |
| originals = [] | |
| for name, module in model.iter_executor_named_modules(): | |
| config = getattr(module, "config", None) | |
| expected = getattr(config, "decode_input_memcfg", None) | |
| if expected is None: | |
| continue | |
| if not hasattr(module, "decode_forward"): | |
| raise AttributeError(f"Module {name} has decode_input_memcfg but no decode_forward method") | |
| original = module.decode_forward | |
| originals.append((module, original)) | |
| def make_wrapper(orig, module_name, expected_memcfg): | |
| def wrapper(x: Any, *args: Any, **kwargs: Any) -> Any: | |
| if isinstance(x, ttnn.Tensor) and x.is_allocated(): | |
| actual = x.spec.memory_config | |
| if actual != expected_memcfg: | |
| mismatches.append((module_name, expected_memcfg, actual)) | |
| return orig(x, *args, **kwargs) | |
| return wrapper | |
| module.decode_forward = make_wrapper(original, name, expected) | |
| try: | |
| yield | |
| finally: | |
| for module, original in originals: | |
| module.decode_forward = original | |
| for name, expected, actual in mismatches: | |
| logger.warning(f"Config mismatch at {name}: declared {expected}, actual {actual}") | |
| def _validate_page_table_layout(layout: Any) -> None: | |
| if not isinstance(layout, PageTableLayout): | |
| raise TypeError("page_table_layout must be a PageTableLayout") | |
| def _validate_resolved_decode_config(config: DecodeRuntimeConfig) -> None: | |
| if not isinstance(config.output_reader, OutputReader): | |
| raise TypeError("output_reader must be an OutputReader") | |
| if config.output_reader.mesh_device is not config.mesh_device: | |
| raise ValueError("output_reader must use the decode mesh_device") | |
| model_mesh = getattr(getattr(config.model, "config", None), "mesh_device", None) | |
| if model_mesh is not None and model_mesh is not config.mesh_device: | |
| raise ValueError("model and decode runtime must use the same mesh_device") | |
| if ( | |
| not isinstance(config.lane_capacity, int) | |
| or isinstance(config.lane_capacity, bool) | |
| or not 0 < config.lane_capacity <= 32 | |
| ): | |
| raise ValueError("lane_capacity must be an integer from 1 through 32") | |
| _validate_page_table_layout(config.page_table_layout) | |
| if ( | |
| not isinstance(config.cluster_shape, tuple) | |
| or len(config.cluster_shape) != 2 | |
| or any(not isinstance(value, int) or isinstance(value, bool) or value <= 0 for value in config.cluster_shape) | |
| ): | |
| raise ValueError("cluster_shape must contain two positive integers") | |
| if tuple(int(value) for value in config.mesh_device.shape) != config.cluster_shape: | |
| raise ValueError("cluster_shape must match mesh_device.shape") | |
| if ( | |
| not isinstance(config.num_devices, int) | |
| or isinstance(config.num_devices, bool) | |
| or config.num_devices != config.cluster_shape[0] * config.cluster_shape[1] | |
| ): | |
| raise ValueError("num_devices must match cluster_shape") | |
| model_num_devices = getattr(getattr(config.model, "config", None), "num_devices", None) | |
| if model_num_devices is None: | |
| model_num_devices = getattr(config.model, "num_devices", config.num_devices) | |
| if model_num_devices != config.num_devices: | |
| raise ValueError("num_devices must match the model") | |
| if not isinstance(config.vocab_size, int) or isinstance(config.vocab_size, bool) or config.vocab_size <= 0: | |
| raise ValueError("vocab_size must be a positive integer") | |
| if getattr(config.model, "vocab_size", None) != config.vocab_size: | |
| raise ValueError("vocab_size must match the model") | |
| for name in ( | |
| "device_sampling_enabled", | |
| "force_greedy_top_k", | |
| "allow_force_argmax", | |
| "position_feedback_capable", | |
| ): | |
| if not isinstance(getattr(config, name), bool): | |
| raise TypeError(f"{name} must be bool") | |
| sampling = getattr(config.model, "sampling", None) | |
| sampling_config = getattr(sampling, "config", None) | |
| expected_argmax = getattr(sampling_config, "allow_force_argmax", None) if config.device_sampling_enabled else False | |
| if config.device_sampling_enabled: | |
| if not callable(getattr(sampling, "decode_forward", None)): | |
| raise TypeError("device sampling requires model.sampling.decode_forward()") | |
| if not isinstance(expected_argmax, bool): | |
| raise TypeError("model sampling allow_force_argmax must be bool") | |
| if config.allow_force_argmax is not expected_argmax: | |
| raise ValueError("allow_force_argmax must match the resolved model capability") | |
| expected_top_k = getattr(sampling_config, "max_top_k", 0) if config.device_sampling_enabled else 0 | |
| if config.max_device_top_k != expected_top_k: | |
| raise ValueError("max_device_top_k must match the resolved model sampler capability") | |
| expected_sampling_batch_size = ( | |
| getattr(sampling_config, "max_batch_size", config.lane_capacity) | |
| if config.device_sampling_enabled | |
| else config.lane_capacity | |
| ) | |
| if config.sampling_batch_size != expected_sampling_batch_size: | |
| raise ValueError("sampling_batch_size must match the resolved model sampler capacity") | |
| if (config.sampling_state_controller is None) != (config.sampling_state is None): | |
| raise ValueError("sampling_state_controller and sampling_state must be supplied together") | |
| if config.sampling_state_controller is not None: | |
| if getattr(config.sampling_state_controller, "sampling", None) is not sampling: | |
| raise ValueError("sampling state controller must borrow model.sampling") | |
| if not callable(getattr(config.sampling_state_controller, "decode_forward", None)): | |
| raise TypeError("sampling state controller must provide decode_forward()") | |
| if config.position_feedback_capable != callable(getattr(config.model, "increment_positions", None)): | |
| raise ValueError("position_feedback_capable must match the resolved model capability") | |
| if not isinstance(config.page_table_layout_ceiling, PageTableLayout): | |
| raise TypeError("page_table_layout_ceiling must be a PageTableLayout") | |
| if config.page_table_layout.block_size != config.page_table_layout_ceiling.block_size: | |
| raise ValueError("page_table_layout_ceiling cannot change block_size") | |
| if config.page_table_layout.raw_capacity_width > config.page_table_layout_ceiling.raw_capacity_width: | |
| raise ValueError("page_table_layout_ceiling must cover page_table_layout capacity") | |
| if config.page_table_layout.decode_width > config.page_table_layout_ceiling.decode_width: | |
| raise ValueError("page_table_layout_ceiling must cover decode page-table geometry") | |
| def _copy_host_to_device(host_tensors, device_tensors=None, mesh_device=None): | |
| if device_tensors is None: | |
| if mesh_device is None: | |
| raise ValueError("mesh_device is required for device allocation") | |
| allocated = [] | |
| try: | |
| for host in host_tensors: | |
| allocated.append(ttnn.to_device(host, device=mesh_device) if host is not None else None) | |
| except BaseException as primary: | |
| failures = best_effort_deallocate_owned_tensors(allocated) | |
| attach_cleanup_failures(primary, failures) | |
| raise | |
| return allocated | |
| for host, device in zip(host_tensors, device_tensors): | |
| if host is None: | |
| if device is not None: | |
| raise ValueError("host/device optional tensor structure changed") | |
| continue | |
| ttnn.copy_host_to_device_tensor(host, device) | |
| return device_tensors | |
| def _formatted_sampling_values( | |
| sampling_params, | |
| batch_size, | |
| *, | |
| max_device_top_k=32, | |
| allow_force_argmax=True, | |
| ): | |
| """Compatibility test helper backed by native exact preparation.""" | |
| formatted_size = ((int(batch_size) + 31) // 32) * 32 | |
| prepared = prepare_sampling_params( | |
| sampling_params, | |
| formatted_size, | |
| max_device_top_k=max_device_top_k, | |
| allow_force_argmax=allow_force_argmax, | |
| ) | |
| return ( | |
| prepared.top_k, | |
| prepared.top_p, | |
| prepared.temperature, | |
| prepared.all_active_rows_greedy, | |
| ) | |
| def _concat_host_output(value, cluster_shape): | |
| if isinstance(value, torch.Tensor): | |
| return value | |
| tensors = [ttnn.to_torch(tensor) for tensor in ttnn.get_device_tensors(value)] | |
| rows, columns = cluster_shape | |
| mesh = [tensors[index : index + columns] for index in range(0, len(tensors), columns)] | |
| return torch.cat([torch.cat(row, dim=-1) for row in mesh], dim=1) | |
| def _process_output_tokens(value, batch_size, cluster_shape): | |
| output = _concat_host_output(value, cluster_shape) | |
| if output.ndim >= 4: | |
| if int(output.shape[2]) >= batch_size: | |
| output = output[0, 0, :batch_size, 0] | |
| elif int(output.shape[3]) >= batch_size: | |
| output = output[0, 0, 0, :batch_size] | |
| return output.reshape(-1)[:batch_size].to(torch.int64) | |
| def _process_sampled_log_probs(value, batch_size): | |
| """Normalize replicated sampled-token logprobs to one row-major tensor.""" | |
| if value is None: | |
| return None | |
| if isinstance(value, torch.Tensor): | |
| output = value | |
| elif isinstance(value, ttnn.Tensor): | |
| replicas = ttnn.get_device_tensors(value) | |
| output = ttnn.to_torch(replicas[0] if replicas else value) | |
| else: | |
| # Preserve opaque compatibility payloads used by callers that own their | |
| # own logprob representation. Native Sampling1D returns a TT tensor. | |
| return value | |
| flat = output.reshape(-1) | |
| if int(flat.numel()) < int(batch_size): | |
| raise ValueError(f"sampled-token logprobs contain {flat.numel()} rows, expected at least {batch_size}") | |
| return flat[: int(batch_size)].to(torch.float32) | |
| def _num_blocks(sequence_length, block_size): | |
| return (int(sequence_length) + int(block_size) - 1) // int(block_size) | |