-
Notifications
You must be signed in to change notification settings - Fork 1.2k
Fix Qwen3.5 static-shape prefill (#18832) #21873
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -56,6 +56,7 @@ def __init__( | |
| max_batch_size: int, | ||
| use_kv_cache: bool, | ||
| vocab_size: int, | ||
| enable_dynamic_shape: bool = True, | ||
| device: str = "cpu", | ||
| ): | ||
| """ | ||
|
|
@@ -67,11 +68,13 @@ def __init__( | |
| max_batch_size: max batch size. | ||
| use_kv_cache: whether to use a KV cache. | ||
| vocab_size: number of items in the vocab. | ||
| enable_dynamic_shape: whether the model accepts multi-token prefill. | ||
| device: device to run the runner on. | ||
| """ | ||
| self.max_seq_len = max_seq_len | ||
| self.max_batch_size = max_batch_size | ||
| self.use_kv_cache = use_kv_cache | ||
| self.enable_dynamic_shape = enable_dynamic_shape | ||
| self.tokenizer = get_tokenizer(tokenizer_path, tokenizer_config_path) | ||
| self.device = device | ||
| # For some models like qwen, mismatch is acceptable: https://github.com/QwenLM/Qwen3/issues/466#issuecomment-2146759706 | ||
|
|
@@ -88,6 +91,50 @@ def forward( | |
| ) -> torch.Tensor: | ||
| pass | ||
|
|
||
| def _prefill_chunk( | ||
| self, | ||
| prompt_tokens: List[int], | ||
| start_pos: int, | ||
| ) -> torch.Tensor: | ||
| # Parallel prefill processes the whole prompt chunk in one call when dynamic | ||
| # shapes are enabled or the KV cache is disabled. | ||
| if self.enable_dynamic_shape or not self.use_kv_cache: | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. why the kv_cache condition? Also in the if body we already know its none. |
||
| return self.forward( | ||
| tokens=torch.tensor( | ||
| [prompt_tokens], dtype=torch.long, device=self.device | ||
| ), | ||
| input_pos=( | ||
| torch.tensor([start_pos], dtype=torch.long, device=self.device) | ||
| if self.use_kv_cache | ||
| else None | ||
| ), | ||
| ) | ||
| else: | ||
| # Sequential prefill processes one token per call and uses the KV cache | ||
| # to preserve context across calls. | ||
| logits = self.forward( | ||
| tokens=torch.tensor( | ||
| [[prompt_tokens[0]]], dtype=torch.long, device=self.device | ||
| ), | ||
| input_pos=torch.tensor( | ||
| [start_pos], | ||
| dtype=torch.long, | ||
| device=self.device, | ||
| ), | ||
| ) | ||
| for prompt_pos, prompt_token in enumerate(prompt_tokens[1:], start=1): | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. why not do the start=0 also in the loop? what is the perf impact of this vs. chunking? |
||
| logits = self.forward( | ||
| tokens=torch.tensor( | ||
| [[prompt_token]], dtype=torch.long, device=self.device | ||
| ), | ||
| input_pos=torch.tensor( | ||
| [start_pos + prompt_pos], | ||
| dtype=torch.long, | ||
| device=self.device, | ||
| ), | ||
| ) | ||
| return logits | ||
|
|
||
| def generate( # noqa: C901 | ||
| self, | ||
| prompt_tokens: List[int], | ||
|
|
@@ -99,14 +146,7 @@ def generate( # noqa: C901 | |
| ) -> List[int]: | ||
| # Prefill | ||
| prefill_start = time.time() | ||
| logits = self.forward( | ||
| tokens=torch.tensor([prompt_tokens], dtype=torch.long, device=self.device), | ||
| input_pos=( | ||
| torch.tensor([pos_base], dtype=torch.long, device=self.device) | ||
| if self.use_kv_cache | ||
| else None | ||
| ), | ||
| ) | ||
| logits = self._prefill_chunk(prompt_tokens, pos_base) | ||
| prefill_time = time.time() - prefill_start | ||
|
|
||
| current_token = next_token(logits, temperature, top_p) | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,113 @@ | ||
| # Copyright (c) Meta Platforms, Inc. and affiliates. | ||
| # All rights reserved. | ||
| # | ||
| # This source code is licensed under the BSD-style license found in the | ||
| # LICENSE file in the root directory of this source tree. | ||
|
|
||
| import unittest | ||
| from unittest.mock import patch | ||
|
|
||
| import torch | ||
| from executorch.examples.models.llama.runner.generation import LlamaRunner | ||
|
|
||
|
|
||
| class _Tokenizer: | ||
| n_words = 100 | ||
| eos_id = 99 | ||
|
|
||
| def decode_token(self, token: int) -> str: | ||
| return str(token) | ||
|
|
||
|
|
||
| class _RecordingRunner(LlamaRunner): | ||
| def __init__( | ||
| self, | ||
| *, | ||
| use_kv_cache: bool, | ||
| enable_dynamic_shape: bool, | ||
| ): | ||
| with patch( | ||
| "executorch.examples.models.llama.runner.generation.get_tokenizer", | ||
| return_value=_Tokenizer(), | ||
| ): | ||
| super().__init__( | ||
| tokenizer_path="unused", | ||
| max_seq_len=5, | ||
| max_batch_size=1, | ||
| use_kv_cache=use_kv_cache, | ||
| vocab_size=100, | ||
| enable_dynamic_shape=enable_dynamic_shape, | ||
| ) | ||
| self.calls = [] | ||
|
|
||
| def forward(self, tokens, input_pos=None): | ||
| if ( | ||
| self.use_kv_cache | ||
| and not self.enable_dynamic_shape | ||
| and tokens.shape != (1, 1) | ||
| ): | ||
| raise RuntimeError( | ||
| f"static input requires shape (1, 1), got {tokens.shape}" | ||
| ) | ||
| self.calls.append( | ||
| ( | ||
| tokens.flatten().tolist(), | ||
| None if input_pos is None else input_pos.item(), | ||
| ) | ||
| ) | ||
| last_token = tokens[0, -1].item() | ||
| sampled_token = 20 if last_token == 12 else 99 if last_token == 20 else 0 | ||
| logits = torch.full((1, 100), -1.0) | ||
| logits[0, sampled_token] = 1.0 | ||
| return logits | ||
|
|
||
|
|
||
| class GenerationTest(unittest.TestCase): | ||
| def test_static_kv_cache_prefills_one_token_at_a_time(self): | ||
| runner = _RecordingRunner(use_kv_cache=True, enable_dynamic_shape=False) | ||
|
|
||
| generated = runner.generate( | ||
| prompt_tokens=[10, 11, 12], | ||
| max_seq_len=5, | ||
| temperature=0, | ||
| pos_base=7, | ||
| ) | ||
|
|
||
| self.assertEqual(generated, [20, 99]) | ||
| self.assertEqual( | ||
| runner.calls, | ||
| [([10], 7), ([11], 8), ([12], 9), ([20], 10)], | ||
| ) | ||
|
|
||
| def test_dynamic_kv_cache_preserves_parallel_prefill(self): | ||
| runner = _RecordingRunner(use_kv_cache=True, enable_dynamic_shape=True) | ||
|
|
||
| generated = runner.generate( | ||
| prompt_tokens=[10, 11, 12], | ||
| max_seq_len=5, | ||
| temperature=0, | ||
| pos_base=7, | ||
| ) | ||
|
|
||
| self.assertEqual(generated, [20, 99]) | ||
| self.assertEqual(runner.calls, [([10, 11, 12], 7), ([20], 10)]) | ||
|
|
||
| def test_static_non_kv_cache_preserves_full_sequence_calls(self): | ||
| runner = _RecordingRunner(use_kv_cache=False, enable_dynamic_shape=False) | ||
|
|
||
| generated = runner.generate( | ||
| prompt_tokens=[10, 11, 12], | ||
| max_seq_len=5, | ||
| temperature=0, | ||
| pos_base=7, | ||
| ) | ||
|
|
||
| self.assertEqual(generated, [20, 99]) | ||
| self.assertEqual( | ||
| runner.calls, | ||
| [([10, 11, 12], None), ([10, 11, 12, 20], None)], | ||
| ) | ||
|
|
||
|
|
||
| if __name__ == "__main__": | ||
| unittest.main() |
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
when dynamic shape is not enabled, why fall back to a single token? Curious - also why is it called prefill_chunk?