File size: 32,984 Bytes
2415c4c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 | # SPDX-FileCopyrightText: © 2026 Tenstorrent AI ULC
# SPDX-License-Identifier: Apache-2.0
"""Warmup coverage planning and trace-activation coordination."""
from __future__ import annotations
from contextlib import contextmanager
from dataclasses import dataclass, replace
from typing import Any, Callable, Iterator
import torch
from loguru import logger
from models.common.llm_runtime.config import PageTableLayout, TraceConfig, WarmupConfig
from models.common.llm_runtime.decode import DecodeRuntimeConfig
from models.common.llm_runtime.prefill.config import PrefillRuntimeConfig
from models.common.llm_runtime.program_compiler import CompiledProgram
from models.common.sampling.sampling_params import SamplingParams
@dataclass(frozen=True)
class WarmupCase:
operation: str
batch_size: int
sequence_length: int | None
sampling_path: str
cached_tokens: int = 0
@dataclass(frozen=True)
class WarmupPlan:
prefill: tuple[WarmupCase, ...]
decode: tuple[WarmupCase, ...]
@dataclass(frozen=True)
class CoverageAlias:
"""One exact compiled-program association with a configured trace."""
program_signature: Any
trace_signature: Any
@dataclass(frozen=True)
class CoverageManifest:
"""Registry-authoritative operation identities sealed at activation."""
eager_program_signatures: tuple[Any, ...]
traced_source_program_signatures: tuple[Any, ...]
trace_signatures: tuple[Any, ...]
aliases: tuple[CoverageAlias, ...]
@dataclass(frozen=True)
class WarmupCoordinatorConfig:
"""Fully resolved immutable warmup policy and coverage."""
warmup: WarmupConfig
model: Any
page_table_layout: PageTableLayout # Current geometry used to build coverage plans.
page_table_layout_ceiling: PageTableLayout # Construction-time upper bound retained across replacement.
prefill_sequence_lengths: tuple[int, ...]
lane_batch_size: int
device_sampling_enabled: bool
allow_force_argmax: bool
prime_q128_tile_ends: bool
prefill_trace_enabled: bool
decode_trace_enabled: bool
eager_plan: WarmupPlan
sampled_plan: WarmupPlan
def __post_init__(self) -> None:
if not isinstance(self.warmup, WarmupConfig):
raise TypeError("warmup must be a WarmupConfig")
if self.model is None:
raise ValueError("model is required")
if not isinstance(self.page_table_layout, PageTableLayout):
raise TypeError("page_table_layout must be a PageTableLayout")
_validate_prefill_sequence_lengths(self.prefill_sequence_lengths)
_require_positive_int("lane_batch_size", self.lane_batch_size)
for name in (
"device_sampling_enabled",
"allow_force_argmax",
"prime_q128_tile_ends",
"prefill_trace_enabled",
"decode_trace_enabled",
):
if not isinstance(getattr(self, name), bool):
raise TypeError(f"{name} must be bool")
if not self.device_sampling_enabled and self.allow_force_argmax:
raise ValueError("force-argmax capability requires device sampling")
if self.prime_q128_tile_ends is not (self.device_sampling_enabled and self.lane_batch_size >= 32):
raise ValueError("prime_q128_tile_ends must match resolved sampling and lane capabilities")
if not isinstance(self.page_table_layout_ceiling, PageTableLayout):
raise TypeError("page_table_layout_ceiling must be a PageTableLayout")
if self.page_table_layout.block_size != self.page_table_layout_ceiling.block_size:
raise ValueError("page_table_layout_ceiling cannot change block_size")
if self.page_table_layout.raw_capacity_width > self.page_table_layout_ceiling.raw_capacity_width:
raise ValueError("page_table_layout_ceiling must cover page_table_layout capacity")
if (
self.page_table_layout.prefill_width > self.page_table_layout_ceiling.prefill_width
or self.page_table_layout.decode_width > self.page_table_layout_ceiling.decode_width
):
raise ValueError("page_table_layout_ceiling must cover canonical page-table geometry")
expected_eager = _build_plan(
warmup=self.warmup,
layout=self.page_table_layout,
prefill_sequence_lengths=self.prefill_sequence_lengths,
lane_batch_size=self.lane_batch_size,
allow_force_argmax=self.allow_force_argmax,
can_sample_on_device=False,
)
expected_sampled = _build_plan(
warmup=self.warmup,
layout=self.page_table_layout,
prefill_sequence_lengths=self.prefill_sequence_lengths,
lane_batch_size=self.lane_batch_size,
allow_force_argmax=self.allow_force_argmax,
can_sample_on_device=True,
)
if self.eager_plan != expected_eager or self.sampled_plan != expected_sampled:
raise ValueError("warmup plans must match resolved policy and geometry")
@classmethod
def resolve(
cls,
*,
warmup: WarmupConfig,
trace: TraceConfig,
prefill: PrefillRuntimeConfig,
decode: DecodeRuntimeConfig,
prefill_sequence_lengths: tuple[int, ...],
) -> "WarmupCoordinatorConfig":
"""Validate resolved runtimes and derive both static coverage plans."""
if not isinstance(warmup, WarmupConfig):
raise TypeError("warmup must be a WarmupConfig")
if not isinstance(trace, TraceConfig):
raise TypeError("trace must be a TraceConfig")
if not isinstance(prefill, PrefillRuntimeConfig):
raise TypeError("prefill must be a PrefillRuntimeConfig")
if not isinstance(decode, DecodeRuntimeConfig):
raise TypeError("decode must be a DecodeRuntimeConfig")
if decode.model is not prefill.model:
raise ValueError("prefill and decode configs must share one model")
if decode.page_table_layout is not prefill.page_table_layout:
raise ValueError("prefill and decode configs must share one page-table layout")
if decode.lane_capacity != prefill.max_batch_size:
raise ValueError("prefill and decode configs must share one lane capacity")
if decode.device_sampling_enabled is not prefill.device_sampling_enabled:
raise ValueError("prefill and decode configs must share device-sampling policy")
if decode.allow_force_argmax is not prefill.allow_force_argmax:
raise ValueError("prefill and decode configs must share force-argmax capability")
if decode.page_table_layout_ceiling != prefill.page_table_layout_ceiling:
raise ValueError("prefill and decode configs must share one page-table layout ceiling")
source_lengths = warmup.prefill_seq_lens
if source_lengths is None:
source_lengths = prefill_sequence_lengths
_validate_prefill_sequence_lengths(source_lengths)
lane_batch_size = prefill.max_batch_size
device_sampling_enabled = prefill.device_sampling_enabled
allow_force_argmax = prefill.allow_force_argmax
prime_q128_tile_ends = device_sampling_enabled and lane_batch_size >= 32
eager_plan = _build_plan(
warmup=warmup,
layout=prefill.page_table_layout,
prefill_sequence_lengths=source_lengths,
lane_batch_size=lane_batch_size,
allow_force_argmax=allow_force_argmax,
can_sample_on_device=False,
)
sampled_plan = _build_plan(
warmup=warmup,
layout=prefill.page_table_layout,
prefill_sequence_lengths=source_lengths,
lane_batch_size=lane_batch_size,
allow_force_argmax=allow_force_argmax,
can_sample_on_device=True,
)
return cls(
warmup=warmup,
model=prefill.model,
page_table_layout=prefill.page_table_layout,
prefill_sequence_lengths=source_lengths,
lane_batch_size=lane_batch_size,
device_sampling_enabled=device_sampling_enabled,
allow_force_argmax=allow_force_argmax,
prime_q128_tile_ends=prime_q128_tile_ends,
prefill_trace_enabled=trace.prefill_enabled,
decode_trace_enabled=trace.decode_enabled,
eager_plan=eager_plan,
sampled_plan=sampled_plan,
page_table_layout_ceiling=prefill.page_table_layout_ceiling,
)
def with_page_table_layout(self, layout: PageTableLayout) -> "WarmupCoordinatorConfig":
"""Return the same policy with final geometry within original ceilings."""
if not isinstance(layout, PageTableLayout):
raise TypeError("layout must be a PageTableLayout")
if layout.block_size != self.page_table_layout.block_size:
raise ValueError("page-table layout replacement cannot change block_size")
if layout.raw_capacity_width > self.page_table_layout_ceiling.raw_capacity_width:
raise ValueError("page-table layout replacement cannot exceed the construction-time capacity ceiling")
if (
layout.prefill_width > self.page_table_layout_ceiling.prefill_width
or layout.decode_width > self.page_table_layout_ceiling.decode_width
):
raise ValueError("page-table layout replacement cannot expand canonical geometry")
return replace(
self,
page_table_layout=layout,
eager_plan=_build_plan(
warmup=self.warmup,
layout=layout,
prefill_sequence_lengths=self.prefill_sequence_lengths,
lane_batch_size=self.lane_batch_size,
allow_force_argmax=self.allow_force_argmax,
can_sample_on_device=False,
),
sampled_plan=_build_plan(
warmup=self.warmup,
layout=layout,
prefill_sequence_lengths=self.prefill_sequence_lengths,
lane_batch_size=self.lane_batch_size,
allow_force_argmax=self.allow_force_argmax,
can_sample_on_device=True,
),
)
class WarmupCoordinator:
"""Compile configured coverage and activate traces at one shared barrier.
``Llama3Executor.warmup_model_prefill`` and ``warmup_model_decode`` call
`warmup_prefill` and `warmup_decode` in either order. Each
method compiles its required eager programs and registers trace plans.
Capture begins only after both configured operation sets are complete.
"""
def __init__(
self,
*,
config: WarmupCoordinatorConfig,
execution: Any,
ensure_sampling_buffers: Callable[[], None],
validate_bound_cache: Callable[[Any], None],
) -> None:
if not isinstance(config, WarmupCoordinatorConfig):
raise TypeError("config must be a WarmupCoordinatorConfig")
eager = getattr(execution, "eager_executor", execution)
trace_compiler = getattr(execution, "trace_compiler", None)
prefill_config = getattr(getattr(eager, "prefill", None), "config", None)
decode_config = getattr(getattr(eager, "decode", None), "config", None)
if not isinstance(prefill_config, PrefillRuntimeConfig) or not isinstance(decode_config, DecodeRuntimeConfig):
raise TypeError("execution must compose configured prefill and decode runtimes")
if prefill_config.model is not config.model or decode_config.model is not config.model:
raise ValueError("execution runtimes must use the warmup config model")
if (
prefill_config.page_table_layout is not config.page_table_layout
or decode_config.page_table_layout is not config.page_table_layout
):
raise ValueError("execution runtimes must use the warmup config page-table layout")
if (
prefill_config.max_batch_size != config.lane_batch_size
or decode_config.lane_capacity != config.lane_batch_size
):
raise ValueError("execution runtimes must use the warmup config lane capacity")
if (
prefill_config.device_sampling_enabled is not config.device_sampling_enabled
or decode_config.device_sampling_enabled is not config.device_sampling_enabled
):
raise ValueError("execution runtimes must use the warmup config sampling policy")
if not callable(ensure_sampling_buffers):
raise TypeError("ensure_sampling_buffers must be callable")
if not callable(validate_bound_cache):
raise TypeError("validate_bound_cache must be callable")
self.config = config
self.execution = execution
self.eager = eager
self.trace_compiler = trace_compiler
self._ensure_sampling_buffers = ensure_sampling_buffers
self._validate_bound_cache = validate_bound_cache
self._eager: set[WarmupCase] = set()
self._trace_registered: set[WarmupCase] = set()
self._trace_decisions: dict[str, bool] = {}
self._sampling_decisions: dict[str, bool] = {}
self._captured = False
self._coverage_manifest: CoverageManifest | None = None
self._required_program_keys: set[Any] = set()
self._required_trace_program_keys: set[Any] = set()
self._capture_deferred = False
self._capture_pending = False
self._pending_manifest: CoverageManifest | None = None
self._prefill_trace_postprocess_primed = False
self._configuration_sealed = False
# Public API
@property
def already_warmed_up_prefill(self) -> bool:
"""Whether all configured prefill programs and traces are ready."""
can_sample_on_device = self._sampling_decisions.get("prefill", self.config.device_sampling_enabled)
required = set(self._plan(can_sample_on_device=can_sample_on_device).prefill)
if not required.issubset(self._eager):
return False
if not self.config.prefill_trace_enabled or self._trace_decisions.get("prefill") is False:
return True
return required.issubset(self._trace_registered) and self._captured
@property
def coverage_manifest(self) -> CoverageManifest | None:
"""Return the immutable identities verified by successful activation."""
return self._coverage_manifest
@property
def capture_pending(self) -> bool:
"""Whether complete validated coverage is staged for activation."""
return self._capture_pending
@property
def trace_activated(self) -> bool:
"""Whether this coordinator has completed trace capture and activation."""
return self._captured
@contextmanager
def defer_capture(self) -> Iterator["WarmupCoordinator"]:
"""Stage readiness without capturing until a multi-lane barrier commits."""
if self._capture_deferred:
raise RuntimeError("trace capture deferral is already active")
if self._capture_pending:
raise RuntimeError("trace capture is already pending")
self._capture_deferred = True
try:
yield self
finally:
self._capture_deferred = False
self._capture_pending = False
self._pending_manifest = None
def activate_pending_capture(self) -> None:
"""Commit one validated capture while its deferral context is active."""
if not self._capture_deferred:
raise RuntimeError("pending trace capture can only activate inside its deferral context")
if not self._capture_pending:
raise RuntimeError("no trace capture is pending")
self._capture_now(self._pending_manifest)
self._capture_pending = False
self._pending_manifest = None
def configure_page_table_layout(self, layout: PageTableLayout) -> None:
"""Install final paged-KV geometry before warmup compiles any program."""
if self._configuration_sealed:
raise RuntimeError("page-table layout cannot change after warmup configuration is sealed")
self.config = self.config.with_page_table_layout(layout)
def seal_configuration(self) -> None:
"""Forbid geometry replacement before physical KV allocation begins."""
self._configuration_sealed = True
def warmup_prefill(
self,
*,
kv_cache: Any, # ↓ Borrowed resources
can_sample_on_device: bool, # ↓ Execution policy
enable_trace: bool,
) -> None:
"""Compile prefill coverage and capture once decode coverage is ready."""
self._validate_hints("prefill", enable_trace, can_sample_on_device)
self._validate_bound_cache(kv_cache)
self._sampling_decisions["prefill"] = bool(can_sample_on_device)
self._trace_decisions["prefill"] = bool(enable_trace)
if enable_trace and self._trace_decisions.get("decode") is False and self.config.decode_trace_enabled:
del self._trace_decisions["decode"]
self._configuration_sealed = True
if can_sample_on_device:
self._ensure_sampling_buffers()
plan = self._plan(can_sample_on_device=can_sample_on_device)
destination = self._trace_registered if enable_trace else self._eager
cases = plan.prefill
if enable_trace and can_sample_on_device:
# The hidden-body trace is sampling-independent, but its retained
# post-trace inputs must support both aliases. Register the forced
# top-k variant first so the shared artifact owns a K/P/T buffer.
cases = tuple(sorted(cases, key=lambda case: case.sampling_path != "topk"))
for case in cases:
if case in destination:
continue
sampling = None
if case.sampling_path == "argmax":
sampling = _greedy_sampling_params(case.batch_size)
elif case.sampling_path == "topk":
sampling = _topk_sampling_params(case.batch_size)
actual_uncached_lengths = (int(case.sequence_length),)
if (
case.batch_size == 1
and case.sequence_length == 128
and case.cached_tokens == 0
and (
case.sampling_path == "argmax"
or (case.sampling_path == "topk" and self.config.prime_q128_tile_ends)
)
):
# Q128 single-user sampled postprocessing has one TT slice
# program per tile start. Prime all four without expanding the
# public warmup coverage model.
actual_uncached_lengths = (32, 64, 96, 128)
for actual_uncached_length in actual_uncached_lengths:
prompt_length = case.cached_tokens + actual_uncached_length
tokens = torch.zeros((case.batch_size, prompt_length), dtype=torch.long)
prompt_lens = torch.full((case.batch_size,), prompt_length, dtype=torch.long)
width = _ceil_div(prompt_length, self.config.page_table_layout.block_size)
page_table = torch.zeros((case.batch_size, width), dtype=torch.int32)
start_pos = (
torch.full((case.batch_size,), case.cached_tokens, dtype=torch.long) if case.cached_tokens else None
)
compile_target = self.execution if enable_trace else self.eager
programs = compile_target.compile_prefill(
tokens=tokens,
page_table=page_table,
prompt_lens=prompt_lens,
start_pos=start_pos,
empty_slots=list(range(case.batch_size)),
sampling_params=sampling,
)
self._record_required_programs(programs, traced=enable_trace)
destination.add(case)
self._maybe_capture()
def warmup_decode(
self,
*,
kv_cache: Any, # ↓ Borrowed resources
max_batch_size: int, # ↓ Coverage dimensions
num_blocks: int,
can_sample_on_device: bool, # ↓ Execution policy
enable_trace: bool,
) -> None:
"""Compile decode coverage and capture once prefill coverage is ready."""
self._validate_hints("decode", enable_trace, can_sample_on_device)
self._validate_bound_cache(kv_cache)
lane_batch = self.config.lane_batch_size
if int(max_batch_size) != lane_batch:
raise ValueError(f"decode warmup batch {max_batch_size} does not match lane capacity {lane_batch}")
if int(num_blocks) <= 0:
raise ValueError("decode warmup num_blocks must be positive")
self._sampling_decisions["decode"] = bool(can_sample_on_device)
self._trace_decisions["decode"] = bool(enable_trace)
self._configuration_sealed = True
if can_sample_on_device:
self._ensure_sampling_buffers()
plan = self._plan(can_sample_on_device=can_sample_on_device)
destination = self._trace_registered if enable_trace else self._eager
for case in plan.decode:
if case in destination:
continue
sampling = None
if case.sampling_path == "argmax":
sampling = _greedy_sampling_params(lane_batch)
elif case.sampling_path == "topk":
sampling = _topk_sampling_params(lane_batch)
compile_target = self.execution if enable_trace else self.eager
program = compile_target.compile_decode(
tokens=torch.zeros(lane_batch, dtype=torch.long),
start_pos=torch.zeros(lane_batch, dtype=torch.long),
page_table=torch.zeros((lane_batch, int(num_blocks)), dtype=torch.int32),
sampling_params=sampling,
)
self._record_required_programs(program, traced=enable_trace)
if not enable_trace:
logger.info("Compiled decode")
if sampling is not None:
logger.info("Compiled on-device sampling")
destination.add(case)
self._maybe_capture()
# Private implementation
def _plan(self, *, can_sample_on_device: bool) -> WarmupPlan:
return self.config.sampled_plan if can_sample_on_device else self.config.eager_plan
def _maybe_capture(self) -> None:
if self.trace_compiler is None or self._captured:
return
required_trace: set[WarmupCase] = set()
if self.config.prefill_trace_enabled:
prefill_decision = self._trace_decisions.get("prefill")
if prefill_decision is None:
return
if prefill_decision:
prefill_plan = self._plan(can_sample_on_device=self._sampling_decisions["prefill"])
required_trace.update(prefill_plan.prefill)
if self.config.decode_trace_enabled:
decode_decision = self._trace_decisions.get("decode")
if decode_decision is None:
return
if decode_decision:
decode_plan = self._plan(can_sample_on_device=self._sampling_decisions["decode"])
required_trace.update(decode_plan.decode)
if not required_trace:
return
if not required_trace.issubset(self._trace_registered):
return
manifest = self._prepare_capture_manifest()
if self._capture_deferred:
self._pending_manifest = manifest
self._capture_pending = True
return
self._capture_now(manifest)
def _prepare_capture_manifest(self) -> CoverageManifest | None:
# WarmupCase is only an idempotency key for public warmup calls. The
# compiler registries own identity coverage: validate their exact state
# before a single-lane capture or a multi-lane barrier reports ready.
manifest = _resolve_coverage_manifest(
self.eager,
self.trace_compiler,
required_program_keys=self._required_program_keys,
required_trace_program_keys=self._required_trace_program_keys,
)
if manifest is not None and not manifest.aliases:
raise RuntimeError("Configured trace warmup registered no program-to-trace aliases")
return manifest
def _capture_now(self, manifest: CoverageManifest | None) -> None:
self.trace_compiler.capture_all()
self._captured = True
self._coverage_manifest = manifest
self._prime_prefill_trace_postprocess()
def _record_required_programs(self, programs: Any, *, traced: bool) -> None:
if programs is None:
return
if isinstance(programs, CompiledProgram):
programs = (programs,)
if not isinstance(programs, tuple) or any(not isinstance(program, CompiledProgram) for program in programs):
raise TypeError("compile targets must return CompiledProgram values")
keys = {program.key for program in programs}
self._required_program_keys.update(keys)
if traced:
self._required_trace_program_keys.update(keys)
def _prime_prefill_trace_postprocess(self) -> None:
if (
self._prefill_trace_postprocess_primed
or self._trace_decisions.get("prefill") is False
or not self.config.prefill_trace_enabled
):
return
prefill_can_sample = self._sampling_decisions.get("prefill", self.config.device_sampling_enabled)
if not prefill_can_sample or not self.config.allow_force_argmax:
self._prefill_trace_postprocess_primed = True
return
sequence_length = (
128 if 128 in self.config.prefill_sequence_lengths else int(self.config.prefill_sequence_lengths[0])
)
width = _ceil_div(sequence_length, self.config.page_table_layout.block_size)
self.execution.prefill_forward(
tokens=torch.zeros((1, sequence_length), dtype=torch.long),
page_table=torch.zeros((1, width), dtype=torch.int32),
prompt_lens=torch.full((1,), sequence_length, dtype=torch.long),
empty_slots=[0],
start_pos=None,
sampling_params=_greedy_sampling_params(1),
)
self._prefill_trace_postprocess_primed = True
def _validate_hints(self, operation: str, enable_trace: bool, can_sample_on_device: bool) -> None:
trace_enabled = (
self.config.prefill_trace_enabled if operation == "prefill" else self.config.decode_trace_enabled
)
if enable_trace and not trace_enabled:
raise ValueError(f"{operation} trace warmup exceeds the configured trace policy")
if can_sample_on_device and not self.config.device_sampling_enabled:
raise ValueError("warmup cannot enable device sampling when it is statically disabled")
def _validate_prefill_sequence_lengths(values: Any) -> None:
if not isinstance(values, tuple) or not values:
raise ValueError("prefill sequence lengths must be a non-empty tuple")
if any(not isinstance(value, int) or isinstance(value, bool) or value <= 0 for value in values):
raise ValueError("prefill sequence lengths must contain positive integers")
if len(set(values)) != len(values):
raise ValueError("prefill sequence lengths must be unique")
def _resolve_coverage_manifest(
eager: Any,
trace_compiler: Any,
*,
required_program_keys: set[Any] | None = None,
required_trace_program_keys: set[Any] | None = None,
) -> CoverageManifest | None:
"""Resolve actual registered identities when concrete registries are available."""
program_compiler = getattr(eager, "program_compiler", None)
programs = getattr(program_compiler, "compiled_programs", None)
if programs is None or not callable(getattr(trace_compiler, "trace_key_for_program", None)):
# Lightweight host-contract doubles intentionally need not reproduce
# compiler internals; production executors always expose both registries.
return None
required_program_keys = set() if required_program_keys is None else set(required_program_keys)
required_trace_program_keys = set() if required_trace_program_keys is None else set(required_trace_program_keys)
programs_by_key = {program.key: program for program in programs}
missing_programs = required_program_keys.difference(programs_by_key)
if missing_programs:
digests = sorted(key.digest for key in missing_programs)
raise RuntimeError(f"Coverage manifest is missing required compiled programs: {digests}")
missing_aliases = {key for key in required_trace_program_keys if trace_compiler.trace_key_for_program(key) is None}
if missing_aliases:
digests = sorted(key.digest for key in missing_aliases)
raise RuntimeError(f"Coverage manifest is missing required trace aliases: {digests}")
eager_signatures = []
traced_signatures = []
aliases = []
trace_signatures_by_key = {}
for program in programs:
if not isinstance(program, CompiledProgram):
raise TypeError("program compiler snapshots must contain CompiledProgram values")
trace_key = trace_compiler.trace_key_for_program(program.key)
if trace_key is None:
eager_signatures.append(program.signature)
continue
record = trace_compiler.get(trace_key)
if record is None:
raise RuntimeError(f"Trace association {trace_key.digest} has no registered trace record")
traced_signatures.append(program.signature)
aliases.append(CoverageAlias(program.signature, record.signature))
trace_signatures_by_key.setdefault(trace_key, record.signature)
return CoverageManifest(
eager_program_signatures=tuple(eager_signatures),
traced_source_program_signatures=tuple(traced_signatures),
trace_signatures=tuple(trace_signatures_by_key.values()),
aliases=tuple(aliases),
)
def _require_positive_int(name: str, value: Any) -> None:
if not isinstance(value, int) or isinstance(value, bool) or value <= 0:
raise ValueError(f"{name} must be a positive integer")
def _build_plan(
*,
warmup: WarmupConfig,
layout: PageTableLayout,
prefill_sequence_lengths: tuple[int, ...],
lane_batch_size: int,
allow_force_argmax: bool,
can_sample_on_device: bool,
) -> WarmupPlan:
sampling_paths = ["logits"]
if can_sample_on_device:
sampling_paths.append("topk")
prefill = []
for sequence_length in prefill_sequence_lengths:
batches = warmup.prefill_batch_sizes if sequence_length == 128 else (1,)
for batch_size in batches:
if batch_size <= lane_batch_size:
batch_sampling_paths = sampling_paths + (
["argmax"] if can_sample_on_device and allow_force_argmax and batch_size == 1 else []
)
prefill.extend(
WarmupCase("prefill", batch_size, sequence_length, sampling_path)
for sampling_path in batch_sampling_paths
)
cached_prompt_length = layout.block_size + sequence_length
if cached_prompt_length <= layout.raw_capacity_width * layout.block_size:
prefill.extend(
WarmupCase(
"prefill",
1,
sequence_length,
sampling_path,
cached_tokens=layout.block_size,
)
for sampling_path in sampling_paths
+ (["argmax"] if can_sample_on_device and allow_force_argmax else [])
)
decode_paths = ["logits"]
if can_sample_on_device:
if allow_force_argmax:
decode_paths.append("argmax")
if not allow_force_argmax or warmup.include_decode_top_k:
decode_paths.append("topk")
decode = tuple(WarmupCase("decode", lane_batch_size, None, sampling_path) for sampling_path in decode_paths)
return WarmupPlan(tuple(prefill), decode)
def _greedy_sampling_params(batch_size: int) -> SamplingParams:
return SamplingParams(
temperature=torch.zeros(batch_size),
top_k=torch.ones(batch_size, dtype=torch.int32),
top_p=torch.ones(batch_size),
)
def _topk_sampling_params(batch_size: int) -> SamplingParams:
return SamplingParams(
temperature=torch.ones(batch_size),
top_k=torch.full((batch_size,), 32, dtype=torch.int32),
top_p=torch.full((batch_size,), 0.08),
)
def _ceil_div(value: int, divisor: int) -> int:
return (value + divisor - 1) // divisor
|