Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
104 changes: 67 additions & 37 deletions examples/model/qwen3_14b/runner/npu_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -144,7 +144,7 @@ class _DecodeKernelInputs:
seq_lens: torch.Tensor
block_table: torch.Tensor
slot_mapping: torch.Tensor
logits: torch.Tensor
logits: torch.Tensor | DeviceTensor


@dataclass
Expand Down Expand Up @@ -180,6 +180,7 @@ def __init__(
self._device_id = device_id
self._l3_worker: Any | None = None
self._l3_static_tensors: dict[tuple[int, tuple[int, ...], torch.dtype], object] = {}
self._l3_output_tensors: dict[tuple[str, tuple[int, ...], torch.dtype], DeviceTensor] = {}
self._static_args: _StaticKernelArgs | None = None
self._pending_kv_cache_specs: dict[str, tuple[ModelConfig, RuntimeConfig]] = {}
if compiled is not None:
Expand Down Expand Up @@ -582,7 +583,11 @@ def run_prefill(self, model: RuntimeModel, batch: PrefillBatch) -> PrefillResult
compiled = self._compiled
prefill_inputs = self._prepare_prefill_inputs(model, batch)

logits_padded = compiled.prefill_logits_buffer
logits_padded = (
self._output_kernel_arg("prefill_logits", compiled.prefill_logits_buffer)
if batch.allow_device_greedy_sampling
else compiled.prefill_logits_buffer
)

kv_cache = self._materialize_kv_cache(model)
k_cache = kv_cache.key_pages
Expand All @@ -597,16 +602,16 @@ def run_prefill(self, model: RuntimeModel, batch: PrefillBatch) -> PrefillResult
for batch_idx, alloc in enumerate(batch.kv_allocations):
seq_len = int(batch.seq_lens[batch_idx].item())
alloc.tokens_used = max(alloc.tokens_used, seq_len)
sampled_ids, next_hidden = self._maybe_run_sample_embed(
sampled_ids, next_hidden = self._device_sampling_outputs(
logits_padded,
compiled.prefill_sampled_ids_buffer,
compiled.prefill_next_hidden_buffer,
None,
prefill_inputs.actual_batch,
allow=batch.allow_device_greedy_sampling,
)
return PrefillResult(
Comment thread
zmnobug marked this conversation as resolved.
last_hidden=None,
logits=logits_padded[: prefill_inputs.actual_batch, : model.config.vocab_size],
logits=self._result_logits(model, prefill_inputs.actual_batch, logits_padded),
sampled_token_ids=sampled_ids,
next_hidden_states=next_hidden,
)
Expand Down Expand Up @@ -637,7 +642,11 @@ def run_decode(self, model: RuntimeModel, batch: DecodeBatch) -> DecodeResult:
k_cache = kv_cache.key_pages
v_cache = kv_cache.value_pages

kernel_inputs = self._pad_decode_inputs(model, decode_inputs)
kernel_inputs = self._pad_decode_inputs(
model,
decode_inputs,
device_outputs=batch.allow_device_greedy_sampling,
)

# Padded block_table / slot_mapping only ever reference row 0's
# already-valid pages, so bound-check exactly what the kernel will read.
Expand All @@ -649,7 +658,8 @@ def run_decode(self, model: RuntimeModel, batch: DecodeBatch) -> DecodeResult:
)
for batch_idx, alloc in enumerate(batch.kv_allocations):
alloc.tokens_used = max(alloc.tokens_used, int(batch.seq_lens[batch_idx].item()))
sampled_ids, next_hidden = self._integrated_sample_result(
sampled_ids, next_hidden = self._device_sampling_outputs(
None,
compiled.decode_sampled_ids_buffer,
# decode_fwd's next_hidden output is the embedding for sampled_ids_in
# used by this decode step. The newly sampled token is embedded at the
Expand All @@ -661,24 +671,31 @@ def run_decode(self, model: RuntimeModel, batch: DecodeBatch) -> DecodeResult:
)
return DecodeResult(
hidden_states=decode_inputs.hidden.float(),
logits=kernel_inputs.logits[: kernel_inputs.actual_batch, : model.config.vocab_size].to(
logits=self._result_logits(model, kernel_inputs.actual_batch, kernel_inputs.logits).to(
decode_inputs.hidden.device
),
sampled_token_ids=sampled_ids,
next_hidden_states=next_hidden,
)

@staticmethod
def _integrated_sample_result(
def _device_sampling_outputs(
self,
logits: torch.Tensor | DeviceTensor | None,
sampled_ids_buffer: torch.Tensor,
next_hidden_buffer: torch.Tensor | None,
actual_batch: int,
*,
allow: bool,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
"""Read device sampling output and optional precomputed next hidden rows."""
"""Run or read device greedy sampling outputs."""
if not allow:
return None, None
if logits is not None:
self._run_distributed_program(
self._compiled.greedy_sample,
logits,
sampled_ids_buffer,
)
next_hidden = (
next_hidden_buffer[:actual_batch].clone()
if next_hidden_buffer is not None
Expand All @@ -689,35 +706,12 @@ def _integrated_sample_result(
next_hidden,
)

def _maybe_run_sample_embed(
self,
logits: torch.Tensor,
sampled_ids_buffer: torch.Tensor,
next_hidden_buffer: torch.Tensor,
actual_batch: int,
*,
allow: bool,
) -> tuple[torch.Tensor | None, torch.Tensor | None]:
"""Run device greedy sampling when the request is greedy."""
if not allow:
return None, None
compiled = self._compiled
self._run_distributed_program(
compiled.greedy_sample,
logits,
sampled_ids_buffer,
)
return (
sampled_ids_buffer[:actual_batch, :1].clone(),
None,
)

def _prefill_kernel_args(
self,
inputs: _PrefillInputs,
k_cache: DeviceTensor,
v_cache: DeviceTensor,
logits: torch.Tensor,
logits: torch.Tensor | DeviceTensor,
) -> tuple[Any, ...]:
"""Return arguments in ``qwen3_prefill_host`` signature order."""
static = self._require_static_args()
Expand Down Expand Up @@ -787,7 +781,38 @@ def _decode_kernel_args(
self._compiled.decode_next_hidden_buffer,
)

def _pad_decode_inputs(self, model: RuntimeModel, inputs: _DecodeInputs) -> _DecodeKernelInputs:
def _output_kernel_arg(self, name: str, host_buffer: torch.Tensor) -> DeviceTensor:
"""Allocate a reusable worker-resident buffer for a large kernel output."""
key = (name, tuple(host_buffer.shape), host_buffer.dtype)
cached = self._l3_output_tensors.get(key)
if cached is not None:
return cached
dev = self._shared_l3_worker().alloc_tensor(tuple(host_buffer.shape), host_buffer.dtype)
self._l3_output_tensors[key] = dev
return dev

@staticmethod
def _result_logits(
model: RuntimeModel,
actual_batch: int,
logits: torch.Tensor | DeviceTensor,
) -> torch.Tensor:
"""Return host logits when available, otherwise a greedy-only empty placeholder."""
if isinstance(logits, DeviceTensor):
return torch.empty(
(actual_batch, 0),
dtype=torch.float32,
device="cpu",
)
Comment thread
zmnobug marked this conversation as resolved.
return logits[:actual_batch, : model.config.vocab_size]

def _pad_decode_inputs(
self,
model: RuntimeModel,
inputs: _DecodeInputs,
*,
device_outputs: bool = False,
) -> _DecodeKernelInputs:
"""Pad active decode rows to the fixed kernel batch.

The fused decode kernel computes all ``max_batch_size`` rows. Inactive
Expand Down Expand Up @@ -850,7 +875,11 @@ def _pad_decode_inputs(self, model: RuntimeModel, inputs: _DecodeInputs) -> _Dec
kernel_batch,
rows_each=1,
),
logits=compiled.decode_logits_buffer,
logits=(
self._output_kernel_arg("decode_logits", compiled.decode_logits_buffer)
if device_outputs
else compiled.decode_logits_buffer
),
)

def _run_distributed_program(self, callable_spec: _L3Callable, *args: Any) -> Any:
Expand Down Expand Up @@ -967,6 +996,7 @@ def close(self) -> None:
finally:
self._l3_worker = None
self._l3_static_tensors.clear()
self._l3_output_tensors.clear()

def _prepare_prefill_inputs(
self,
Expand Down
27 changes: 11 additions & 16 deletions python/core/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -236,19 +236,12 @@ def _generate_batch_impl(
runtime_model,
prefill_batch,
)
prefill_logits = prefill_result.logits
prefill_sampled_token_ids = (
prefill_result.sampled_token_ids
if allow_device_greedy_sampling
else None
)

sampling_params = self._sampler.from_generate_config(generate_config)
current_tokens = self._sample_batch_rows(
prefill_logits,
current_tokens = self._sample_result_rows(
prefill_result,
sampling_params,
len(requests),
prefill_sampled_token_ids,
allow_device_greedy_sampling,
)
active_indices = list(range(len(requests)))
finish_reasons = ["length"] * len(requests)
Expand Down Expand Up @@ -313,11 +306,11 @@ def _generate_batch_impl(
kv_allocations=active_allocations,
),
)
decoded_tokens = self._sample_batch_rows(
decode_result.logits,
decoded_tokens = self._sample_result_rows(
decode_result,
sampling_params,
len(next_active),
decode_result.sampled_token_ids if allow_device_greedy_sampling else None,
allow_device_greedy_sampling,
)
for row_idx, request_idx in enumerate(next_active):
current_tokens[request_idx] = decoded_tokens[row_idx]
Expand Down Expand Up @@ -433,21 +426,23 @@ def _generate_result(self, model_id: str, prompt: str, config: GenerateConfig) -
"""Generate one result by reusing the batch path."""
return self.generate_batch(model_id, [prompt], config)[0]

def _sample_batch_rows(
def _sample_result_rows(
self,
logits: torch.Tensor,
result,
sampling_params,
row_count: int,
sampled_token_ids: torch.Tensor | None = None,
allow_device_sampled: bool,
) -> list[int]:
"""Return sampled token IDs, preferring executor-provided device samples."""
sampled_token_ids = result.sampled_token_ids if allow_device_sampled else None
if sampled_token_ids is not None:
flat_ids = sampled_token_ids.view(-1)
if flat_ids.numel() < row_count:
raise ValueError(
f"sampled_token_ids has {flat_ids.numel()} rows, expected at least {row_count}"
)
return [int(flat_ids[idx].item()) for idx in range(row_count)]
logits = result.logits
return [
self._sampler.sample(
self._select_batch_row(logits, row_idx),
Expand Down
14 changes: 1 addition & 13 deletions python/core/serving_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -255,19 +255,13 @@ def _batch_prefill(
request = sr.request
will_be_computed = sr.num_computed_tokens + sr.num_new_tokens
if will_be_computed >= request.num_prompt_tokens:
logits = (
prefill_result.logits[i]
if prefill_result.logits.dim() > 1
else prefill_result.logits
)
params = SamplingParams(
temperature=request.temperature,
top_p=request.top_p,
top_k=request.top_k,
)
token_id = self._sample_result_row(
prefill_result,
logits,
params,
i,
allow_device_greedy_sampling,
Expand Down Expand Up @@ -330,19 +324,13 @@ def _batch_decode(

for i, sr in enumerate(scheduled):
request = sr.request
logits = (
decode_result.logits[i]
if decode_result.logits.dim() > 1
else decode_result.logits
)
params = SamplingParams(
temperature=request.temperature,
top_p=request.top_p,
top_k=request.top_k,
)
token_id = self._sample_result_row(
decode_result,
logits,
params,
i,
allow_device_greedy_sampling,
Expand All @@ -352,7 +340,6 @@ def _batch_decode(
def _sample_result_row(
self,
result,
logits: torch.Tensor,
params: SamplingParams,
row_idx: int,
allow_device_sampled: bool,
Expand All @@ -366,6 +353,7 @@ def _sample_result_row(
f"sampled_token_ids has {flat.numel()} rows, expected row {row_idx}"
)
return int(flat[row_idx].item())
logits = result.logits[row_idx] if result.logits.dim() > 1 else result.logits
return self.sampler.sample(logits, params)

def _worker_entry(
Expand Down
2 changes: 1 addition & 1 deletion tests/test_batching.py
Original file line number Diff line number Diff line change
Expand Up @@ -547,7 +547,7 @@ def test_pypto_executor_uses_cached_kernel_weights_after_registration(monkeypatc
compiled=compiled,
)
monkeypatch.setattr(runner, "_shared_l3_worker", lambda: _FakeWorker())
monkeypatch.setattr(runner, "_compute_kv_cache_pages", lambda config, runtime: 1)
monkeypatch.setattr(runner, "_compute_kv_cache_pages", lambda config, runtime, device_id=0: 1)
monkeypatch.setattr(runner, "_print_memory_breakdown", lambda *a, **kw: None)
runner.init_kv_cache(model.config.model_id, model.config, model.runtime)
monkeypatch.setattr(runner, "_static_device_tensor", lambda tensor: tensor)
Expand Down
47 changes: 46 additions & 1 deletion tests/test_device_sampling_submission.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,8 +9,10 @@
from __future__ import annotations

from pathlib import Path
from types import SimpleNamespace

import pytest
import torch


ROOT = Path(__file__).resolve().parents[1]
Expand Down Expand Up @@ -65,7 +67,50 @@ def test_prefill_keeps_sampling_in_standalone_device_kernel() -> None:
assert "_greedy_sample_inline" not in prefill
assert "_token_embed_inline" not in prefill
assert "compiled.greedy_sample" in runner
assert "_maybe_run_sample_embed(" in runner
assert "_device_sampling_outputs(" in runner


def test_device_greedy_keeps_large_outputs_worker_resident(monkeypatch) -> None:
from examples.model.qwen3_14b.runner.npu_runner import Qwen314BModelRunner
from pypto.runtime import DeviceTensor

class _FakeWorker:
def __init__(self) -> None:
self.alloc_calls = 0

def alloc_tensor(self, shape, dtype):
self.alloc_calls += 1
return DeviceTensor(self.alloc_calls, tuple(shape), dtype)

runner = object.__new__(Qwen314BModelRunner)
runner._l3_output_tensors = {}
worker = _FakeWorker()
monkeypatch.setattr(runner, "_shared_l3_worker", lambda: worker)

host_buffer = torch.empty((2, 4), dtype=torch.float32)
first = runner._output_kernel_arg("decode_logits", host_buffer)
second = runner._output_kernel_arg("decode_logits", host_buffer)
other = runner._output_kernel_arg("prefill_logits", host_buffer)

assert first is second
assert other is not first
assert worker.alloc_calls == 2

model = SimpleNamespace(
runtime=SimpleNamespace(device=torch.device("meta")),
config=SimpleNamespace(vocab_size=3),
)
host_logits = torch.arange(8, dtype=torch.float32).reshape(2, 4)
assert torch.equal(
runner._result_logits(model, 1, host_logits),
host_logits[:1, :3],
)

device_logits = DeviceTensor(99, (2, 4), torch.float32)
placeholder = runner._result_logits(model, 2, device_logits)
assert placeholder.shape == (2, 0)
assert placeholder.dtype == torch.float32
assert placeholder.device.type == "cpu"


def _device_greedy_argmax_with_clamp(logits):
Expand Down
Loading
Loading