Download code/models/common/llm_runtime/prefill/postprocess.py from tt-hous/clef: direct link, hf CLI and curl.
- Browser
- Download file 15.9 kB
-
https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/llm_runtime/prefill/postprocess.py
- Command line
-
hf download hf://tt-hous/clef/code/models/common/llm_runtime/prefill/postprocess.py
-
curl -L -o postprocess.py https://huggingface.co/tt-hous/clef/resolve/main/code/models/common/llm_runtime/prefill/postprocess.py
15.9 kB
| # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC | |
| # SPDX-License-Identifier: Apache-2.0 | |
| """Prefill postprocessing and device-sampling behavior.""" | |
| from __future__ import annotations | |
| from typing import Any, Callable, Iterable | |
| import torch | |
| import ttnn | |
| from models.common.llm_runtime.prefill.config import PrefillRuntimeConfig | |
| from models.common.llm_runtime.prefill.inputs import PrefillPositionInputs | |
| from models.common.llm_runtime.prefill.plan import PrefillRequest | |
| from models.common.llm_runtime.prefill.sampling_helpers import _TILE_SIZE, SamplingPath | |
| from models.common.llm_runtime.prefill.signatures import PreparedPrefill | |
| from models.common.modules.sampling.params import PreparedSamplingParams, prepare_sampling_params | |
| from models.common.sampling.sampling_params import SamplingParams | |
| KPTSignature = tuple[tuple[int, ...], tuple[float, ...], tuple[float, ...]] | None | |
| class PrefillPostprocessor: | |
| """Own prefill output selection, sampling tensors, and finalization.""" | |
| def __init__( | |
| self, | |
| config: PrefillRuntimeConfig, | |
| *, | |
| allocate_device_tensors: Callable[[Any], Any], | |
| copy_into_device_tensors: Callable[[Any, Any], Any], | |
| ) -> None: | |
| self.config = config | |
| self._allocate_device_tensors = allocate_device_tensors | |
| self._copy_into_device_tensors = copy_into_device_tensors | |
| def configure(self, config: PrefillRuntimeConfig) -> None: | |
| self.config = config | |
| def prepared_sampling(self, prepared: Any) -> PreparedSamplingParams | None: | |
| """Return native prepared state, adapting older test collaborators.""" | |
| value = getattr(prepared, "prepared_sampling", None) | |
| if value is not None: | |
| return value | |
| raw = getattr(prepared, "sampling_params", None) | |
| if raw is None: | |
| return None | |
| return prepare_sampling_params( | |
| raw, | |
| self.sampling_output_rows(prepared), | |
| max_device_top_k=self.config.max_device_top_k, | |
| allow_force_argmax=self.config.allow_force_argmax, | |
| ) | |
| 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, | |
| request: PrefillRequest, | |
| sampling_params: PreparedSamplingParams | None, | |
| ) -> SamplingPath: | |
| if sampling_params is None: | |
| return "logits" | |
| if request.kind != "single": | |
| return "topk" | |
| return sampling_params.sampling_path | |
| def sampling_batch_size(self, request: PrefillRequest) -> int: | |
| if self.config.device_sampling_enabled: | |
| return self.config.sampling_batch_size | |
| return request.padded_batch_size | |
| def sampling_output_rows(self, prepared: PreparedPrefill) -> int: | |
| # TT sampling validates K/P/T against the physical logits row count. | |
| # The static Q128 path retains one complete tile and selects the exact | |
| # logical row on the host, so its sampling tensors must span that tile. | |
| if self.uses_static_q128_topk(prepared.request, prepared.sampling_path): | |
| return _TILE_SIZE | |
| return self.sampling_batch_size(prepared.request) | |
| def uses_static_q128_topk( | |
| self, | |
| request: PrefillRequest, | |
| sampling_path: SamplingPath, | |
| ) -> bool: | |
| return ( | |
| self.config.static_q128_topk_supported | |
| and sampling_path == "topk" | |
| and request.kind == "single" | |
| and not request.uses_chunked_prefill | |
| and request.padded_sequence_length == 128 | |
| ) | |
| def make_device_kpt( | |
| self, | |
| sampling_params: PreparedSamplingParams | None, | |
| batch_size: int, | |
| force_topk: bool, | |
| ) -> tuple[Any, Any, Any] | None: | |
| host = self.make_host_kpt(sampling_params, batch_size, force_topk) | |
| if host is None: | |
| return None | |
| return tuple(self._allocate_device_tensors(host)) | |
| def make_host_kpt( | |
| self, | |
| sampling_params: PreparedSamplingParams | None, | |
| batch_size: int, | |
| force_topk: bool, | |
| ) -> tuple[Any, Any, Any] | None: | |
| if sampling_params is None: | |
| return None | |
| if not force_topk and sampling_params.sampling_path == "argmax": | |
| return None | |
| if batch_size > sampling_params.batch_size: | |
| raise ValueError("prepared sampling parameters do not cover the physical prefill batch") | |
| k, p, temperature = self.kpt_values(sampling_params, batch_size) | |
| 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 kpt_values(sampling_params: PreparedSamplingParams, batch_size: int): | |
| active_rows = tuple(row for row, active in enumerate(sampling_params.active_mask) if active) | |
| if len(active_rows) == 1 and batch_size > 1: | |
| row = active_rows[0] | |
| return ( | |
| (sampling_params.top_k[row],) * batch_size, | |
| (sampling_params.top_p[row],) * batch_size, | |
| (sampling_params.temperature[row],) * batch_size, | |
| ) | |
| return ( | |
| sampling_params.top_k[:batch_size], | |
| sampling_params.top_p[:batch_size], | |
| sampling_params.temperature[:batch_size], | |
| ) | |
| def refresh_kpt( | |
| self, | |
| device_kpt: tuple[Any, Any, Any] | None, | |
| sampling_params: PreparedSamplingParams | None, | |
| batch_size: int, | |
| force_topk: bool, | |
| ) -> None: | |
| host_kpt = self.make_host_kpt(sampling_params, batch_size, force_topk) | |
| if (host_kpt is None) != (device_kpt is None): | |
| raise RuntimeError("sampling parameters changed the compiled sampling path") | |
| if host_kpt is not None: | |
| self._copy_into_device_tensors(host_kpt, device_kpt) | |
| def refresh_workspace_sampling( | |
| self, | |
| prepared: PreparedPrefill, | |
| *, | |
| kpt: tuple[Any, Any, Any] | None, | |
| kpt_signature: KPTSignature, | |
| ) -> KPTSignature: | |
| if prepared.sampling_path != "topk": | |
| return kpt_signature | |
| sampling_batch_size = self.sampling_output_rows(prepared) | |
| prepared_sampling = self.prepared_sampling(prepared) | |
| if prepared_sampling is None: | |
| kpt_value = None | |
| else: | |
| sampling = prepared_sampling | |
| kpt_value = self.kpt_values(sampling, sampling_batch_size) | |
| if kpt_signature != kpt_value: | |
| self.refresh_kpt( | |
| kpt, | |
| prepared_sampling, | |
| sampling_batch_size, | |
| force_topk=True, | |
| ) | |
| return kpt_value | |
| return kpt_signature | |
| def make_sampling_output(self, batch_size: int) -> Any: | |
| return ttnn.from_torch( | |
| torch.zeros((1, 1, 1, int(batch_size)), dtype=torch.int32), | |
| device=self.config.mesh_device, | |
| dtype=ttnn.uint32, | |
| layout=ttnn.ROW_MAJOR_LAYOUT, | |
| memory_config=ttnn.DRAM_MEMORY_CONFIG, | |
| mesh_mapper=ttnn.ReplicateTensorToMesh(self.config.mesh_device), | |
| ) | |
| def sample_device( | |
| self, | |
| logits: Any, | |
| kpt: tuple[Any, Any, Any] | None, | |
| sampling: PreparedSamplingParams, | |
| sampled_output: Any | None = None, | |
| *, | |
| count_tokens: bool = True, | |
| ) -> Any: | |
| controller = self.config.sampling_state_controller | |
| if controller is not None: | |
| return controller.prefill_forward( | |
| logits, | |
| self.config.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], | |
| tt_out_tok=sampled_output, | |
| count_tokens=count_tokens, | |
| ) | |
| if kpt is None: | |
| return self.config.model.sampling.decode_forward( | |
| logits, | |
| tt_out_tok=sampled_output, | |
| 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=sampled_output, | |
| enable_log_probs=sampling.enable_log_probs, | |
| ) | |
| def finish_regular_prefill( | |
| self, | |
| prepared: PreparedPrefill, | |
| hidden: Any, | |
| kpt: tuple[Any, Any, Any] | None, | |
| position_inputs: PrefillPositionInputs, | |
| *, | |
| sampled_output: Any | None = None, | |
| owned: list[Any] | None = None, | |
| count_tokens: bool = True, | |
| ) -> Any: | |
| request = prepared.request | |
| relative_last = [last - cached for last, cached in zip(request.last_token_indices, request.cached_tokens)] | |
| if request.kind == "batched" and not self.config.batched_prefill_batched_extract: | |
| hidden = ttnn.reshape( | |
| hidden, | |
| [request.padded_batch_size, 1, request.padded_sequence_length, int(hidden.shape[-1])], | |
| ) | |
| outputs = [] | |
| for local_row, last_token in enumerate(relative_last): | |
| logits = self.config.model.post_process_prefill_output( | |
| hidden[local_row : local_row + 1], | |
| last_token, | |
| ) | |
| retain_owned(owned, logits) | |
| output = ttnn.untilize(logits, use_multicore=True) | |
| retain_owned(owned, output) | |
| outputs.append(output) | |
| return outputs | |
| if request.kind == "batched": | |
| padded_last = list(relative_last) + [0] * (request.padded_batch_size - len(relative_last)) | |
| logits = self.config.model.post_process_batched_prefill_output( | |
| hidden, | |
| padded_last, | |
| request.padded_batch_size, | |
| request.padded_sequence_length, | |
| ) | |
| elif self.uses_static_q128_topk(request, prepared.sampling_path): | |
| logits = self.config.model.post_process_prefill_output(hidden, relative_last[0]) | |
| else: | |
| # Runtime last-token bounds and row index for host sampling too. The static path keyed | |
| # both its 32-row tile slice and the final 1-row slice on the prompt offset, so every | |
| # new offset compiled fresh programs on the first real request - after warmup had | |
| # recorded the decode traces (measured: 2 buffers left live across replays per model). | |
| # With the device-side row pick the logits already hold the last token in row 0. | |
| logits = self.config.model.post_process_prefill_output( | |
| hidden, | |
| relative_last[0], | |
| last_token_slice=(position_inputs.slice_start, position_inputs.slice_end), | |
| last_token_index=position_inputs.row_index, | |
| ) | |
| retain_owned(owned, logits) | |
| if prepared.sampling_params is not None: | |
| selected = fit_prefill_sampling_logits(logits, self.sampling_output_rows(prepared)) | |
| retain_owned(owned, selected) | |
| output = self.sample_device( | |
| selected, | |
| kpt, | |
| self.prepared_sampling(prepared), | |
| sampled_output, | |
| count_tokens=count_tokens, | |
| ) | |
| else: | |
| output = ttnn.untilize(logits, use_multicore=True) | |
| if request.kind == "single" and not request.uses_chunked_prefill and int(output.shape[2]) > 1: | |
| # Host sampling uses runtime row selection above. If the model returns a | |
| # padded tile, the selected token is in row 0; discard the extra rows. | |
| retain_owned(owned, output) | |
| output = ttnn.slice(output, (0, 0, 0, 0), (1, 1, 1, int(output.shape[-1]))) | |
| retain_owned(owned, output) | |
| return output | |
| def finish_prefill_sequence( | |
| self, | |
| prepared: PreparedPrefill, | |
| final_step_output: Any, | |
| kpt: tuple[Any, Any, Any] | None, | |
| position_inputs: PrefillPositionInputs, | |
| *, | |
| sampled_output: Any | None, | |
| owned: list[Any], | |
| count_tokens: bool = True, | |
| ) -> Any: | |
| if not prepared.request.uses_chunked_prefill: | |
| return self.finish_regular_prefill( | |
| prepared, | |
| final_step_output, | |
| kpt, | |
| position_inputs, | |
| sampled_output=sampled_output, | |
| owned=owned, | |
| count_tokens=count_tokens, | |
| ) | |
| if prepared.sampling_params is not None: | |
| selected = fit_prefill_sampling_logits(final_step_output, self.sampling_output_rows(prepared)) | |
| retain_owned(owned, selected) | |
| output = self.sample_device( | |
| selected, | |
| kpt, | |
| self.prepared_sampling(prepared), | |
| sampled_output, | |
| count_tokens=count_tokens, | |
| ) | |
| else: | |
| output = ttnn.untilize(final_step_output, use_multicore=True) | |
| retain_owned(owned, output) | |
| return output | |
| def retain_owned(owned: list[Any] | None, value: Any) -> None: | |
| if owned is None or value is None or any(existing is value for existing in owned): | |
| return | |
| owned.append(value) | |
| def without_borrowed(values: Iterable[Any], borrowed: Iterable[Any]) -> tuple[Any, ...]: | |
| """Remove trace/workspace-owned leaves from replay-local ownership.""" | |
| borrowed_ids = {id(value) for value in borrowed if value is not None} | |
| def prune(value): | |
| if value is None or id(value) in borrowed_ids: | |
| return None | |
| if isinstance(value, tuple): | |
| kept = tuple(item for item in (prune(item) for item in value) if item is not None) | |
| return kept or None | |
| if isinstance(value, list): | |
| kept = [item for item in (prune(item) for item in value) if item is not None] | |
| return kept or None | |
| return value | |
| return tuple(item for item in (prune(value) for value in values) if item is not None) | |
| def new_logprob_output(output: Any, persistent_sampled_output: Any | None) -> Any | None: | |
| if persistent_sampled_output is None: | |
| return None | |
| if not isinstance(output, tuple) or len(output) != 2: | |
| raise TypeError("sampled prefill output must contain (tokens, log_probs)") | |
| return output[1] | |
| def fit_prefill_sampling_logits(logits, target_batch: int): | |
| target_batch = int(target_batch) | |
| if target_batch <= 0: | |
| raise ValueError("prefill sampling target batch must be positive") | |
| current_batch = int(logits.shape[2]) | |
| if current_batch == target_batch: | |
| return logits | |
| if current_batch > target_batch: | |
| return ttnn.slice( | |
| logits, | |
| (0, 0, 0, 0), | |
| (logits.shape[0], logits.shape[1], target_batch, logits.shape[3]), | |
| ) | |
| return ttnn.pad(logits, [(0, 0), (0, 0), (0, target_batch - current_batch), (0, 0)], value=0.0) | |