Download code/models/common/llm_runtime/prefill/inputs.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 14 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/llm_runtime/prefill/inputs.py
- Command line
-
hf download hf://tt-hous/clef/code/models/common/llm_runtime/prefill/inputs.py
-
curl -L -o inputs.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/llm_runtime/prefill/inputs.py
14 kB
| # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC | |
| # SPDX-License-Identifier: Apache-2.0 | |
| """Prefill host/device input staging and replay refresh.""" | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from typing import Any, Callable, Sequence | |
| import torch | |
| import ttnn | |
| from models.common.llm_runtime.prefill.plan import PrefillChunk, PrefillRequest | |
| from models.common.llm_runtime.prefill.sampling_helpers import _TILE_SIZE | |
| from models.common.llm_runtime.tensor_resources import ( | |
| attach_cleanup_failures, | |
| best_effort_deallocate_owned_tensors, | |
| raise_cleanup_failures, | |
| ) | |
| class PrefillHostInputs: | |
| tokens: Any | |
| position_indices: Any | |
| page_table: Any | |
| chunk_page_table: Any | None | |
| chunk_start_idx: Any | None | |
| def values(self) -> tuple[Any, ...]: | |
| return ( | |
| self.tokens, | |
| self.position_indices, | |
| self.page_table, | |
| self.chunk_page_table, | |
| self.chunk_start_idx, | |
| ) | |
| class PrefillDeviceInputs: | |
| tokens: Any | |
| rotary_cos: Any | |
| rotary_sin: Any | |
| page_table: Any | |
| chunk_page_table: Any | None | |
| position_indices: Any | |
| chunk_start_idx: Any | None | |
| def model_values(self) -> tuple[Any, ...]: | |
| return ( | |
| self.tokens, | |
| self.rotary_cos, | |
| self.rotary_sin, | |
| self.page_table, | |
| self.chunk_page_table, | |
| self.position_indices, | |
| self.chunk_start_idx, | |
| ) | |
| def owned_tensor_values(self) -> tuple[Any, ...]: | |
| return self.model_values() | |
| class PrefillPositionInputs: | |
| slice_start: Any | |
| slice_end: Any | |
| row_index: Any | |
| def values(self) -> tuple[Any, ...]: | |
| return self.slice_start, self.slice_end, self.row_index | |
| def owned_tensor_values(self) -> tuple[Any, ...]: | |
| return self.values() | |
| class PrefillTraceInputs: | |
| """Named geometry for allocating one persistent trace input set.""" | |
| chunk: PrefillChunk | |
| relative_last: int | |
| tokens: torch.Tensor | |
| start_pos: int | |
| chunk_page_table: torch.Tensor | None | |
| chunk_start_idx: int | None | |
| sequence_length: int | |
| class PrefillInputStager: | |
| """Own TT prefill tensor geometry, allocation, and in-place refresh.""" | |
| def __init__( | |
| self, | |
| *, | |
| model: Any, | |
| mesh_device: Any, | |
| release_transient: Callable[[Any], list[BaseException]], | |
| ) -> None: | |
| self.model = model | |
| self.mesh_device = mesh_device | |
| self._release_transient = release_transient | |
| def trace_inputs(self, request: PrefillRequest) -> PrefillTraceInputs: | |
| chunk = request.chunks[0] | |
| if request.uses_chunked_prefill: | |
| final_chunk = request.chunks[-1] | |
| relative_last = (request.last_token_indices[0] - final_chunk.chunk_start_idx) % final_chunk.chunk_size | |
| tokens = request.tokens[:, chunk.token_slice] | |
| start_pos = chunk.chunk_start_idx | |
| chunk_page_table = chunk.chunk_page_table | |
| chunk_start_idx = chunk.chunk_start_idx | |
| sequence_length = chunk.chunk_size | |
| else: | |
| relative_last = max( | |
| last - cached for last, cached in zip(request.last_token_indices, request.cached_tokens) | |
| ) | |
| tokens = request.tokens | |
| start_pos = 0 | |
| chunk_page_table = None | |
| chunk_start_idx = None | |
| sequence_length = request.padded_sequence_length | |
| return PrefillTraceInputs( | |
| chunk=chunk, | |
| relative_last=relative_last, | |
| tokens=tokens, | |
| start_pos=start_pos, | |
| chunk_page_table=chunk_page_table, | |
| chunk_start_idx=chunk_start_idx, | |
| sequence_length=sequence_length, | |
| ) | |
| def stage_step( | |
| self, | |
| request: PrefillRequest, | |
| chunk: PrefillChunk, | |
| final_relative_last: int, | |
| ) -> tuple[PrefillDeviceInputs, PrefillPositionInputs]: | |
| chunked = request.uses_chunked_prefill | |
| host_inputs = self.prepare_host_inputs( | |
| request.tokens[:, chunk.token_slice], | |
| request.page_table, | |
| start_pos=chunk.chunk_start_idx if chunked else 0, | |
| chunk_page_table=chunk.chunk_page_table if chunked else None, | |
| chunk_start_idx=chunk.chunk_start_idx if chunked else None, | |
| last_token_idx=max(request.last_token_indices), | |
| ) | |
| device_inputs = None | |
| position_inputs = None | |
| try: | |
| device_inputs = self.stage_device_inputs(host_inputs) | |
| position_values = allocate_device_tensors( | |
| self.prepare_position_inputs_host(final_relative_last, chunk.chunk_size).values(), | |
| mesh_device=self.mesh_device, | |
| ) | |
| position_inputs = PrefillPositionInputs(*position_values) | |
| except BaseException as primary: | |
| failures = self._release_transient((device_inputs, position_inputs)) | |
| attach_cleanup_failures(primary, failures) | |
| raise | |
| return device_inputs, position_inputs | |
| def prepare_host_inputs( | |
| self, | |
| tokens: torch.Tensor, | |
| page_table: torch.Tensor, | |
| *, | |
| start_pos: int = 0, | |
| chunk_page_table: torch.Tensor | None = None, | |
| chunk_start_idx: int | None = None, | |
| last_token_idx: int | None = None, | |
| ) -> PrefillHostInputs: | |
| if tokens.ndim != 2: | |
| raise ValueError("prefill tokens must be rank 2") | |
| mapper = ttnn.ReplicateTensorToMesh(self.mesh_device) | |
| tokens_tt = ttnn.from_torch( | |
| tokens.reshape(1, 1, 1, -1), | |
| device=None, | |
| dtype=ttnn.uint32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=mapper, | |
| ) | |
| rope = self.model.rope_setup | |
| rope.load_device_weights() | |
| matrix_length = int(rope.cos_matrix.shape[2]) | |
| if matrix_length <= 0: | |
| raise ValueError("rotary position table must not be empty") | |
| start_pos = int(start_pos) | |
| sequence_length = int(tokens.shape[-1]) | |
| if start_pos < 0: | |
| raise ValueError("prefill start position must be nonnegative") | |
| if last_token_idx is not None and int(last_token_idx) + 1 > matrix_length: | |
| raise ValueError(f"Sequence length {int(last_token_idx) + 1} exceeds rotary capacity {matrix_length}") | |
| position_indices = torch.arange(start_pos, start_pos + sequence_length, dtype=torch.long).clamp( | |
| max=matrix_length - 1 | |
| ) | |
| position_indices_tt = ttnn.from_torch( | |
| position_indices.reshape(1, -1), | |
| device=None, | |
| dtype=ttnn.uint32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=mapper, | |
| ) | |
| page_table_tt = ttnn.from_torch( | |
| page_table, | |
| device=None, | |
| dtype=ttnn.int32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=mapper, | |
| ) | |
| chunk_tt = ( | |
| ttnn.from_torch( | |
| chunk_page_table, | |
| device=None, | |
| dtype=ttnn.int32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=mapper, | |
| ) | |
| if chunk_page_table is not None | |
| else None | |
| ) | |
| chunk_start_tt = ( | |
| ttnn.from_torch( | |
| torch.tensor([int(chunk_start_idx)], dtype=torch.int32), | |
| device=None, | |
| dtype=ttnn.int32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=mapper, | |
| ) | |
| if chunk_start_idx is not None | |
| else None | |
| ) | |
| return PrefillHostInputs(tokens_tt, position_indices_tt, page_table_tt, chunk_tt, chunk_start_tt) | |
| def prepare_position_inputs_host(self, relative_last: int, sequence_length: int) -> PrefillPositionInputs: | |
| relative_last = int(relative_last) | |
| sequence_length = int(sequence_length) | |
| if relative_last < 0 or relative_last >= sequence_length: | |
| raise ValueError("prefill last-token position must fall within the padded sequence") | |
| block_start = (relative_last // _TILE_SIZE) * _TILE_SIZE | |
| hidden_width = int(self.model.config.dim) | |
| bounds = ((0, 0, block_start, 0), (1, 1, block_start + _TILE_SIZE, hidden_width)) | |
| mapper = ttnn.ReplicateTensorToMesh(self.mesh_device) | |
| slice_bounds = tuple( | |
| ttnn.from_torch( | |
| torch.tensor(bound, dtype=torch.int32), | |
| device=None, | |
| dtype=ttnn.uint32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=mapper, | |
| ) | |
| for bound in bounds | |
| ) | |
| row_index = ttnn.from_torch( | |
| torch.tensor([[relative_last % _TILE_SIZE]], dtype=torch.int32), | |
| device=None, | |
| dtype=ttnn.uint32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=mapper, | |
| ) | |
| return PrefillPositionInputs(slice_bounds[0], slice_bounds[1], row_index) | |
| def stage_device_inputs(self, host_inputs: PrefillHostInputs) -> PrefillDeviceInputs: | |
| raw_inputs = None | |
| rot_mats = None | |
| try: | |
| raw_inputs = allocate_device_tensors(host_inputs.values(), mesh_device=self.mesh_device) | |
| prepare_rot_mats = getattr(self.model, "prepare_prefill_rot_mats", None) | |
| if not callable(prepare_rot_mats): | |
| raise TypeError("model must provide prepare_prefill_rot_mats()") | |
| rot_mats = tuple(prepare_rot_mats(raw_inputs[1])) | |
| if len(rot_mats) != 2: | |
| raise ValueError("prepare_prefill_rot_mats() must return cosine and sine tensors") | |
| except BaseException as primary: | |
| failures = self._release_transient((rot_mats, raw_inputs)) | |
| attach_cleanup_failures(primary, failures) | |
| raise | |
| return PrefillDeviceInputs( | |
| tokens=raw_inputs[0], | |
| rotary_cos=rot_mats[0], | |
| rotary_sin=rot_mats[1], | |
| page_table=raw_inputs[2], | |
| chunk_page_table=raw_inputs[3], | |
| position_indices=raw_inputs[1], | |
| chunk_start_idx=raw_inputs[4], | |
| ) | |
| def refresh_regular_device_inputs(self, request: PrefillRequest, device_inputs: PrefillDeviceInputs) -> None: | |
| mapper = ttnn.ReplicateTensorToMesh(self.mesh_device) | |
| host_tokens = ttnn.from_torch( | |
| request.tokens.reshape(1, 1, 1, -1), | |
| device=None, | |
| dtype=ttnn.uint32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=mapper, | |
| ) | |
| host_page_table = ttnn.from_torch( | |
| request.page_table, | |
| device=None, | |
| dtype=ttnn.int32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| mesh_mapper=mapper, | |
| ) | |
| copy_into_device_tensors((host_tokens, host_page_table), (device_inputs.tokens, device_inputs.page_table)) | |
| def refresh_chunk_device_inputs( | |
| self, | |
| request: PrefillRequest, | |
| chunk: PrefillChunk, | |
| device_inputs: PrefillDeviceInputs, | |
| ) -> None: | |
| host_inputs = self.prepare_host_inputs( | |
| request.tokens[:, chunk.token_slice], | |
| request.page_table, | |
| start_pos=chunk.chunk_start_idx, | |
| chunk_page_table=chunk.chunk_page_table, | |
| chunk_start_idx=chunk.chunk_start_idx, | |
| last_token_idx=max(request.last_token_indices), | |
| ) | |
| copy_into_device_tensors( | |
| host_inputs.values(), | |
| ( | |
| device_inputs.tokens, | |
| device_inputs.position_indices, | |
| device_inputs.page_table, | |
| device_inputs.chunk_page_table, | |
| device_inputs.chunk_start_idx, | |
| ), | |
| ) | |
| self.copy_rotary_inputs(device_inputs) | |
| def copy_rotary_inputs(self, device_inputs: PrefillDeviceInputs) -> None: | |
| rot_mats = None | |
| try: | |
| rot_mats = tuple(self.model.prepare_prefill_rot_mats(device_inputs.position_indices)) | |
| if len(rot_mats) != 2: | |
| raise ValueError("prepare_prefill_rot_mats() must return cosine and sine tensors") | |
| ttnn.copy(input_a=rot_mats[0], input_b=device_inputs.rotary_cos) | |
| ttnn.copy(input_a=rot_mats[1], input_b=device_inputs.rotary_sin) | |
| except BaseException as primary: | |
| failures = self._release_transient(rot_mats) | |
| attach_cleanup_failures(primary, failures) | |
| raise | |
| failures = self._release_transient(rot_mats) | |
| if failures: | |
| raise_cleanup_failures(failures) | |
| def allocate_device_tensors(host_tensors: Sequence[Any], *, mesh_device: Any) -> list[Any]: | |
| """Allocate device tensors corresponding to one host tensor structure.""" | |
| allocated = [] | |
| try: | |
| for host_tensor in host_tensors: | |
| allocated.append(ttnn.to_device(host_tensor, device=mesh_device) if host_tensor is not None else None) | |
| except BaseException as primary: | |
| failures = best_effort_deallocate_owned_tensors(allocated) | |
| attach_cleanup_failures(primary, failures) | |
| raise | |
| return allocated | |
| def copy_into_device_tensors(host_tensors: Sequence[Any], device_tensors: Sequence[Any]) -> Sequence[Any]: | |
| """Refresh an existing device tensor structure without allocating it.""" | |
| if len(host_tensors) != len(device_tensors): | |
| raise ValueError("host/device tensor structures must have equal length") | |
| if any( | |
| (host_tensor is None) != (device_tensor is None) | |
| for host_tensor, device_tensor in zip(host_tensors, device_tensors) | |
| ): | |
| raise ValueError("host/device optional tensor structure changed") | |
| for host_tensor, device_tensor in zip(host_tensors, device_tensors): | |
| if host_tensor is None: | |
| continue | |
| ttnn.copy_host_to_device_tensor(host_tensor, device_tensor) | |
| return device_tensors | |