diff --git a/examples/model/qwen3_14b/runner/npu_runner.py b/examples/model/qwen3_14b/runner/npu_runner.py index a372c437..329decb4 100644 --- a/examples/model/qwen3_14b/runner/npu_runner.py +++ b/examples/model/qwen3_14b/runner/npu_runner.py @@ -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 @@ -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: @@ -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 @@ -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( 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, ) @@ -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. @@ -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 @@ -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 @@ -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() @@ -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", + ) + 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 @@ -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: @@ -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, diff --git a/python/core/engine.py b/python/core/engine.py index d5084713..6dcfb5fd 100644 --- a/python/core/engine.py +++ b/python/core/engine.py @@ -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) @@ -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] @@ -433,14 +426,15 @@ 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: @@ -448,6 +442,7 @@ def _sample_batch_rows( 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), diff --git a/python/core/serving_worker.py b/python/core/serving_worker.py index 9576d3b6..d21674a0 100644 --- a/python/core/serving_worker.py +++ b/python/core/serving_worker.py @@ -255,11 +255,6 @@ 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, @@ -267,7 +262,6 @@ def _batch_prefill( ) token_id = self._sample_result_row( prefill_result, - logits, params, i, allow_device_greedy_sampling, @@ -330,11 +324,6 @@ 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, @@ -342,7 +331,6 @@ def _batch_decode( ) token_id = self._sample_result_row( decode_result, - logits, params, i, allow_device_greedy_sampling, @@ -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, @@ -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( diff --git a/tests/test_batching.py b/tests/test_batching.py index 43c51544..281c3e5a 100644 --- a/tests/test_batching.py +++ b/tests/test_batching.py @@ -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) diff --git a/tests/test_device_sampling_submission.py b/tests/test_device_sampling_submission.py index 056ddf3c..447e81fe 100644 --- a/tests/test_device_sampling_submission.py +++ b/tests/test_device_sampling_submission.py @@ -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] @@ -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): diff --git a/tests/test_qwen3_serving.py b/tests/test_qwen3_serving.py index 61145c39..d9e0d45f 100644 --- a/tests/test_qwen3_serving.py +++ b/tests/test_qwen3_serving.py @@ -140,6 +140,7 @@ def harness(): device="cpu", kv_dtype="bfloat16", weight_dtype="float32", + max_num_batched_tokens=MAX_SEQ_LEN, ), max_num_running_reqs=MAX_BATCH_SIZE, long_prefill_token_threshold=default_threshold,