Download code/models/common/llm_runtime/prefill/runtime.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 20.8 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/llm_runtime/prefill/runtime.py
- Command line
-
hf download hf://tt-hous/clef/code/models/common/llm_runtime/prefill/runtime.py
-
curl -L -o runtime.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/llm_runtime/prefill/runtime.py
20.8 kB
| # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC | |
| # SPDX-License-Identifier: Apache-2.0 | |
| """Public prefill orchestration facade and transient cleanup.""" | |
| from __future__ import annotations | |
| import os | |
| from typing import Any, Iterable, Sequence | |
| import torch | |
| from models.common.llm_runtime.config import PageTableLayout | |
| from models.common.llm_runtime.prefill import postprocess as prefill_postprocess | |
| from models.common.llm_runtime.prefill import result_collector as prefill_result_collector | |
| from models.common.llm_runtime.prefill import sequence_runner as prefill_sequence_runner | |
| from models.common.llm_runtime.prefill import trace as prefill_trace | |
| from models.common.llm_runtime.prefill.config import PrefillRuntimeConfig | |
| from models.common.llm_runtime.prefill.inputs import ( | |
| PrefillDeviceInputs, | |
| PrefillInputStager, | |
| PrefillPositionInputs, | |
| allocate_device_tensors, | |
| copy_into_device_tensors, | |
| ) | |
| from models.common.llm_runtime.prefill.plan import ( | |
| PrefillChunk, | |
| PrefillRequest, | |
| _max_prefill_chunk_size, | |
| _padded_prefill_length, | |
| _plan_prefill_requests, | |
| ) | |
| from models.common.llm_runtime.prefill.sampling_helpers import _slice_sampling_params | |
| from models.common.llm_runtime.prefill.signatures import ( | |
| PreparedPrefill, | |
| build_program_signatures, | |
| build_trace_signature, | |
| ) | |
| from models.common.llm_runtime.tensor_resources import ( | |
| TensorResourceOrphan, | |
| best_effort_deallocate_owned_tensors, | |
| raise_cleanup_failures, | |
| release_orphans, | |
| ) | |
| from models.common.modules.sampling.params import prepare_sampling_params | |
| from models.common.sampling.sampling_params import SamplingParams | |
| class PrefillRuntime: | |
| """Plan, execute, trace, and assemble prefill for one execution lane. | |
| The normal eager call chain is | |
| `EagerExecutor.prefill_forward()` → `prepare` → `invoke` → | |
| `assemble`. Trace warmup uses `capture_plan`; replay uses | |
| `refresh_trace` and `finish_trace` before the same | |
| `assemble` step. Callers pass host request values and never invoke | |
| the private chunk-sequence, staging, or sampling helpers directly. | |
| The runtime borrows the model, mesh, and output reader. It owns staged | |
| prefill tensors and retains failed releases for retry by `cleanup`. | |
| """ | |
| def __init__(self, config: PrefillRuntimeConfig) -> None: | |
| if not isinstance(config, PrefillRuntimeConfig): | |
| raise TypeError("config must be a PrefillRuntimeConfig") | |
| self.config = config | |
| self._sampling_state_controller = config.sampling_state_controller | |
| self._sampling_state = config.sampling_state | |
| self._transient_orphans: list[TensorResourceOrphan] = [] | |
| self.inputs = PrefillInputStager( | |
| model=config.model, | |
| mesh_device=config.mesh_device, | |
| release_transient=self._release_or_retain_transient, | |
| ) | |
| self.postprocessor = prefill_postprocess.PrefillPostprocessor( | |
| config, | |
| allocate_device_tensors=lambda values: allocate_device_tensors( | |
| values, | |
| mesh_device=self.config.mesh_device, | |
| ), | |
| copy_into_device_tensors=copy_into_device_tensors, | |
| ) | |
| self.assembler = prefill_result_collector.PrefillResultAssembler( | |
| config, | |
| postprocessor=self.postprocessor, | |
| release_transient=lambda values: self._release_or_retain_transient(values), | |
| ) | |
| self.sequence_runner = prefill_sequence_runner.PrefillSequenceRunner( | |
| input_stager=self.inputs, | |
| postprocessor=self.postprocessor, | |
| run_hidden_body=lambda *args, **kwargs: self._run_hidden_body(*args, **kwargs), | |
| run_chunk_body=lambda *args, **kwargs: self._run_chunk_body(*args, **kwargs), | |
| release_transient=lambda values: self._release_or_retain_transient(values), | |
| ) | |
| self.trace = prefill_trace.PrefillTraceLifecycle( | |
| hooks=prefill_trace.PrefillTraceHooks( | |
| input_stager=self.inputs, | |
| postprocessor=self.postprocessor, | |
| run_hidden_body=lambda *args, **kwargs: self._run_hidden_body(*args, **kwargs), | |
| run_chunk_hidden_body=lambda *args, **kwargs: self._run_chunk_hidden_body(*args, **kwargs), | |
| release_transient=lambda values: self._release_or_retain_transient(values), | |
| trace_capture_prime_sequence_lengths=config.trace_capture_prime_sequence_lengths, | |
| ) | |
| ) | |
| # 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 KV geometry before allocation or execution.""" | |
| self.config = self.config.with_page_table_layout(layout) | |
| self.postprocessor.configure(self.config) | |
| self.assembler.configure(self.config) | |
| def can_trace( | |
| self, | |
| *, | |
| tokens: torch.Tensor, # ↓ Core request | |
| prompt_lens: torch.Tensor | None = None, # ↓ Sequence metadata | |
| start_pos: torch.Tensor | None = None, | |
| ) -> bool: | |
| """Classify trace applicability without allocating planned request tensors.""" | |
| if not isinstance(tokens, torch.Tensor) or tokens.ndim != 2 or int(tokens.shape[0]) == 0: | |
| return False | |
| batch_size, token_width = map(int, tokens.shape) | |
| if prompt_lens is not None and (not isinstance(prompt_lens, torch.Tensor) or prompt_lens.ndim != 1): | |
| return False | |
| if start_pos is not None and (not isinstance(start_pos, torch.Tensor) or start_pos.ndim != 1): | |
| return False | |
| lengths = [token_width] * batch_size if prompt_lens is None else [int(value) for value in prompt_lens] | |
| cached = [0] * batch_size if start_pos is None else [int(value) for value in start_pos] | |
| if len(lengths) != batch_size or len(cached) != batch_size: | |
| return False | |
| for length, num_cached_tokens in zip(lengths, cached): | |
| if ( | |
| num_cached_tokens < 0 | |
| or num_cached_tokens % self.config.page_table_layout.block_size | |
| or length <= num_cached_tokens | |
| or length > token_width | |
| ): | |
| return False | |
| padded_length = _padded_prefill_length(length - num_cached_tokens) | |
| invocation_length = ( | |
| _max_prefill_chunk_size(padded_length, self.config.max_prefill_chunk_size) | |
| if padded_length > self.config.max_prefill_chunk_size | |
| else padded_length | |
| ) | |
| # Cached/chunk starts are runtime tensors. Static trace capability | |
| # is therefore checked against invocation geometry, not the | |
| # request's current cached offset. | |
| if not self.config.can_enable_trace(invocation_length, 0): | |
| return False | |
| return True | |
| def prepare( | |
| self, | |
| *, | |
| tokens: torch.Tensor, # ↓ Core request | |
| page_table: torch.Tensor, | |
| prompt_lens: torch.Tensor | None = None, # ↓ Sequence metadata | |
| start_pos: torch.Tensor | None = None, | |
| empty_slots: Sequence[int] | None = None, # ↓ Lane routing | |
| sampling_params: SamplingParams | None = None, # ↓ Sampling | |
| prompt_tokens: Any = None, # ↓ Request-owned sampling state | |
| output_tokens: Any = None, | |
| slot_remap: Any = None, | |
| ) -> tuple[PreparedPrefill, ...]: | |
| """Plan host inputs once and return immutable requests for execution. | |
| One public request may produce several prepared requests when batching, | |
| prefix caching, or chunking requires distinct program invocations. | |
| """ | |
| self._ensure_usable() | |
| self.postprocessor.validate_sampling_request(sampling_params) | |
| layout = self.config.page_table_layout | |
| requests = _plan_prefill_requests( | |
| tokens=tokens, | |
| page_table=page_table, | |
| prompt_lens=prompt_lens, | |
| start_pos=start_pos, | |
| empty_slots=empty_slots, | |
| block_size=layout.block_size, | |
| max_batch_size=self.config.max_batch_size, | |
| max_prefill_chunk_size=self.config.max_prefill_chunk_size, | |
| supports_batched_prefill=self.config.supports_batched_prefill, | |
| # Batched prefill makes prefill logits batch-variant (numerics | |
| # depend on wave composition). On multi-chip Blackhole this is | |
| # unverified: seeded token-accuracy/eval gates may break. Measure | |
| # batched on vs off per BH SKU before enabling it there, and check | |
| # whether https://github.com/tenstorrent/tt-metal/issues/47238 | |
| # (batch-invariant kernel fix) has landed. | |
| disable_batched_prefill=( | |
| self.config.disable_batched_prefill | |
| or bool(os.environ.get("DISABLE_BATCHED_PREFILL")) | |
| or (sampling_params is not None and not self.config.batched_prefill_batched_extract) | |
| ), | |
| max_prefill_batch_size=self.config.max_prefill_batch_size, | |
| max_actual_page_table_width=layout.raw_capacity_width, | |
| canonical_page_table_width=layout.prefill_width, | |
| ) | |
| prepared = [] | |
| pending_slot_remap = slot_remap | |
| input_slots = list(range(int(tokens.shape[0]))) if empty_slots is None else [int(slot) for slot in empty_slots] | |
| fallback_prompt_tokens = _prompt_history_from_prefill_tokens(tokens, prompt_lens) | |
| input_prompt_tokens = _select_prefill_state_rows( | |
| fallback_prompt_tokens if prompt_tokens is None else prompt_tokens, | |
| input_slots=input_slots, | |
| input_batch_size=int(tokens.shape[0]), | |
| lane_capacity=self.config.max_batch_size, | |
| ) | |
| input_output_tokens = _select_prefill_state_rows( | |
| output_tokens, | |
| input_slots=input_slots, | |
| input_batch_size=int(tokens.shape[0]), | |
| lane_capacity=self.config.max_batch_size, | |
| ) | |
| for request in requests: | |
| request_sampling = _slice_sampling_params(sampling_params, request.source_rows) | |
| request_prompt_tokens = _select_rows(input_prompt_tokens, request.source_rows) | |
| request_output_tokens = _select_rows(input_output_tokens, request.source_rows) | |
| prepared_sampling = ( | |
| None | |
| if request_sampling is None | |
| else 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, | |
| prompt_tokens=request_prompt_tokens, | |
| output_tokens=request_output_tokens, | |
| slot_remap=pending_slot_remap, | |
| ) | |
| ) | |
| pending_slot_remap = None | |
| sampling_path = self.postprocessor.classify_sampling_path(request, prepared_sampling) | |
| penalties_enabled = prepared_sampling.penalties_enabled if prepared_sampling is not None else False | |
| logprobs_enabled = prepared_sampling.log_probs_enabled if prepared_sampling is not None else False | |
| signatures = build_program_signatures( | |
| request, | |
| sampling_path, | |
| static_q128_topk_supported=self.config.static_q128_topk_supported, | |
| penalties_enabled=penalties_enabled, | |
| logprobs_enabled=logprobs_enabled, | |
| ) | |
| trace_signature = build_trace_signature( | |
| request, | |
| trace_enabled=self.config.can_enable_trace(request.chunks[0].chunk_size, 0), | |
| sampling_path=sampling_path, | |
| penalties_enabled=penalties_enabled, | |
| logprobs_enabled=logprobs_enabled, | |
| ) | |
| prepared.append( | |
| PreparedPrefill( | |
| request=request, | |
| sampling_params=request_sampling, | |
| prepared_sampling=prepared_sampling, | |
| sampling_path=sampling_path, | |
| program_signatures=signatures, | |
| trace_signature=trace_signature, | |
| ) | |
| ) | |
| return tuple(prepared) | |
| def invoke( | |
| self, | |
| prepared: PreparedPrefill, | |
| *, | |
| count_tokens: bool = True, | |
| ) -> prefill_result_collector.InvocationResult: | |
| """Run a prepared request eagerly without replanning or reclassification.""" | |
| self._ensure_usable() | |
| compile_only_state = self._sampling_state_controller is not None and not count_tokens | |
| self._prepare_sampling_state(prepared, count_tokens=count_tokens) | |
| try: | |
| return self.sequence_runner.run(prepared, count_tokens=count_tokens) | |
| finally: | |
| if compile_only_state: | |
| self._sampling_state_controller.reset(self._sampling_state) | |
| def capture_plan(self, prepared: PreparedPrefill) -> prefill_trace.PrefillCapturePlan: | |
| """Describe persistent inputs and capture work for one eligible request.""" | |
| self._ensure_usable() | |
| return self.trace.capture_plan(prepared) | |
| def refresh_trace( | |
| self, | |
| prepared: PreparedPrefill, | |
| persistent: prefill_trace.PrefillHiddenPersistentInputs, | |
| workspace: prefill_trace.PrefillReplayState, | |
| chunk: PrefillChunk | None = None, | |
| ) -> None: | |
| """Refresh borrowed persistent inputs for one replay.""" | |
| self.trace.refresh(prepared, persistent, workspace, chunk) | |
| def finish_trace( | |
| self, | |
| prepared: PreparedPrefill, | |
| hidden: Any, | |
| workspace: prefill_trace.PrefillReplayState, | |
| ) -> prefill_result_collector.InvocationResult: | |
| """Post-process a replayed hidden-state tensor into a normal result.""" | |
| self._prepare_sampling_state(prepared, count_tokens=True) | |
| return self.trace.finish(prepared, hidden, workspace) | |
| def assemble( | |
| self, | |
| prepared_results: Iterable[tuple[PreparedPrefill, prefill_result_collector.InvocationResult]], | |
| *, | |
| batch_size: int, | |
| sampling_params: SamplingParams | None = None, | |
| ) -> torch.Tensor | tuple[torch.Tensor, Any]: | |
| """Read phase outputs, restore source-row order, and release transients.""" | |
| return self.assembler.assemble( | |
| prepared_results, | |
| batch_size=batch_size, | |
| sampling_params=sampling_params, | |
| ) | |
| def cleanup(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 _prepare_sampling_state(self, prepared: PreparedPrefill, *, count_tokens: bool) -> None: | |
| controller = self._sampling_state_controller | |
| if controller is None: | |
| return | |
| sampling = self.postprocessor.prepared_sampling(prepared) | |
| if sampling is None: | |
| return | |
| if not count_tokens: | |
| controller.reset(self._sampling_state, sampling) | |
| return | |
| controller.admit_prefill( | |
| self._sampling_state, | |
| sampling, | |
| slots=prepared.request.slots, | |
| positions=prepared.request.last_token_indices, | |
| ) | |
| def _run_chunk_body( | |
| self, | |
| prepared: PreparedPrefill, | |
| chunk: PrefillChunk, | |
| device_inputs: PrefillDeviceInputs, | |
| position_inputs: PrefillPositionInputs, | |
| ) -> Any: | |
| return self.config.model.prefill_forward( | |
| self.config.model.embed_prefill(device_inputs.tokens), | |
| [device_inputs.rotary_cos, device_inputs.rotary_sin], | |
| user_id=0, | |
| page_table=device_inputs.page_table, | |
| chunk_page_table=device_inputs.chunk_page_table, | |
| chunk_start_idx=chunk.chunk_start_idx, | |
| get_last_token=-1, | |
| chunk_start_idx_tensor=device_inputs.chunk_start_idx, | |
| last_token_slice=(position_inputs.slice_start, position_inputs.slice_end), | |
| last_token_index=(position_inputs.row_index if prepared.sampling_params is not None else None), | |
| ) | |
| def _run_chunk_hidden_body( | |
| self, | |
| prepared: PreparedPrefill, | |
| chunk: PrefillChunk, | |
| device_inputs: PrefillDeviceInputs, | |
| ) -> Any: | |
| """Run only the shared model body; postprocessing remains alias-local.""" | |
| return self.config.model.prefill_forward( | |
| self.config.model.embed_prefill(device_inputs.tokens), | |
| [device_inputs.rotary_cos, device_inputs.rotary_sin], | |
| user_id=0, | |
| page_table=device_inputs.page_table, | |
| chunk_page_table=device_inputs.chunk_page_table, | |
| chunk_start_idx=None, | |
| get_last_token=-1, | |
| chunk_start_idx_tensor=device_inputs.chunk_start_idx, | |
| last_token_slice=None, | |
| last_token_index=None, | |
| ) | |
| def _run_hidden_body( | |
| self, | |
| request: PrefillRequest, | |
| device_inputs: PrefillDeviceInputs, | |
| *, | |
| fill_rows: int | None = None, | |
| ) -> Any: | |
| if fill_rows is None: | |
| fill_rows = len(request.source_rows) | |
| if fill_rows < len(request.source_rows) or fill_rows > request.padded_batch_size: | |
| raise ValueError("fill_rows must cover active rows without exceeding padded batch size") | |
| return self.config.model.prefill_forward( | |
| self.config.model.embed_prefill(device_inputs.tokens), | |
| [device_inputs.rotary_cos, device_inputs.rotary_sin], | |
| user_id=list(range(fill_rows)) if request.kind == "batched" else 0, | |
| page_table=device_inputs.page_table, | |
| chunk_page_table=device_inputs.chunk_page_table, | |
| get_last_token=-1, | |
| batch_size=request.padded_batch_size, | |
| chunk_start_idx_tensor=device_inputs.chunk_start_idx, | |
| ) | |
| def _release_or_retain_transient(self, values: Any) -> list[BaseException]: | |
| 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 _ensure_usable(self) -> None: | |
| if self._transient_orphans: | |
| raise RuntimeError("PrefillRuntime has unreleased transient resources; cleanup is required") | |
| def _select_prefill_state_rows(value: Any, *, input_slots: list[int], input_batch_size: int, lane_capacity: int): | |
| if value is None: | |
| return None | |
| length = _leading_length(value) | |
| if length == lane_capacity: | |
| return _select_rows(value, input_slots) | |
| if length == input_batch_size: | |
| return value | |
| if length == 1: | |
| return _select_rows(value, [0] * input_batch_size) | |
| raise ValueError( | |
| f"prefill sampling state has {length} rows, expected 1, request batch {input_batch_size}, " | |
| f"or lane capacity {lane_capacity}" | |
| ) | |
| def _prompt_history_from_prefill_tokens( | |
| tokens: torch.Tensor, | |
| prompt_lens: torch.Tensor | None, | |
| ) -> torch.Tensor: | |
| history = tokens.clone() | |
| width = int(tokens.shape[1]) | |
| lengths = [width] * int(tokens.shape[0]) if prompt_lens is None else [int(value) for value in prompt_lens] | |
| for row, length in enumerate(lengths): | |
| if length < 0 or length > width: | |
| raise ValueError("prompt_lens must fit the prefill token width") | |
| history[row, length:] = -1 | |
| return history | |
| def _select_rows(value: Any, rows: Sequence[int]): | |
| if value is None: | |
| return None | |
| if isinstance(value, torch.Tensor): | |
| indices = torch.tensor(tuple(int(row) for row in rows), dtype=torch.long, device=value.device) | |
| return value.index_select(0, indices) | |
| if isinstance(value, list): | |
| return [value[int(row)] for row in rows] | |
| if isinstance(value, tuple): | |
| return tuple(value[int(row)] for row in rows) | |
| raise TypeError(f"request-owned sampling state must be a tensor or sequence, got {type(value).__name__}") | |
| def _leading_length(value: Any) -> int: | |
| if isinstance(value, torch.Tensor): | |
| if value.ndim == 0: | |
| return 1 | |
| return int(value.shape[0]) | |
| if isinstance(value, (list, tuple)): | |
| return len(value) | |
| raise TypeError(f"request-owned sampling state must be a tensor or sequence, got {type(value).__name__}") | |