diff --git a/tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_categorical_sample.py b/tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_categorical_sample.py deleted file mode 100644 index edda63ba02eb..000000000000 --- a/tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_categorical_sample.py +++ /dev/null @@ -1,768 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project - -import gc - -import pytest -import torch -import torch_npu # noqa: F401 - -from vllm_ascend.ops.triton.triton_utils import init_device_properties_triton -from vllm_ascend.ops.triton.v2.sample.categorical_sample import categorical_sample - -DEVICE = "npu" -VOCAB_SIZE = 151936 -SUPPORTED_DTYPES = [torch.float32, torch.float16, torch.bfloat16] - - -@pytest.fixture(scope="module", autouse=True) -def _npu_env(): - init_device_properties_triton() - yield - torch.npu.synchronize() - gc.collect() - torch.npu.empty_cache() - torch.npu.reset_peak_memory_stats() - - -def _seed_and_pos( - num_tokens: int, - num_reqs: int, -) -> tuple[torch.Tensor, torch.Tensor]: - seed = torch.arange(num_reqs, dtype=torch.int64, device=DEVICE) * 104729 + 17 - pos = torch.arange(num_tokens, dtype=torch.int64, device=DEVICE) + 23 - return seed, pos - - -def _sample( - logits: torch.Tensor, - expanded_idx_mapping: torch.Tensor, - temperature: torch.Tensor, - seed: torch.Tensor, - pos: torch.Tensor, - *, - apply_temperature: bool, - is_drafting: bool = False, - logits_cache: torch.Tensor | None = None, - logits_cache_col: torch.Tensor | None = None, - use_fp64: bool = False, -) -> torch.Tensor: - """Call the version-compatible public categorical entry explicitly.""" - return categorical_sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=apply_temperature, - is_drafting=is_drafting, - logits_cache=logits_cache, - logits_cache_col=logits_cache_col, - use_fp64=use_fp64, - ) - - -@pytest.mark.parametrize("num_tokens", [1, 16, 64]) -@pytest.mark.parametrize("dtype", SUPPORTED_DTYPES) -def test_categorical_sample_greedy(num_tokens, dtype): - """temperature=0 must match exact argmax on the long-vocab shape.""" - torch.manual_seed(0) - logits = torch.randn( - num_tokens, - VOCAB_SIZE, - dtype=dtype, - device=DEVICE, - ) - expanded_idx_mapping = torch.arange( - num_tokens, - dtype=torch.int32, - device=DEVICE, - ) - temperature = torch.zeros( - num_tokens, - dtype=torch.float32, - device=DEVICE, - ) - seed, pos = _seed_and_pos(num_tokens, num_tokens) - - sampled = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=False, - ) - torch.npu.synchronize() - - assert sampled.dtype == torch.int64 - assert sampled.shape == (num_tokens,) - torch.testing.assert_close( - sampled, - logits.argmax(dim=-1), - rtol=0, - atol=0, - ) - - -@pytest.mark.parametrize("dtype", SUPPORTED_DTYPES) -def test_categorical_sample_mixed_temperature_and_hierarchy_boundaries(dtype): - """Mixed greedy/random rows must cross fine/coarse/tail boundaries.""" - torch.manual_seed(1) - num_tokens = 16 - logits = torch.randn( - num_tokens, - VOCAB_SIZE, - dtype=dtype, - device=DEVICE, - ) - expanded_idx_mapping = torch.arange( - num_tokens, - dtype=torch.int32, - device=DEVICE, - ) - temperature = torch.ones( - num_tokens, - dtype=torch.float32, - device=DEVICE, - ) - temperature[:8] = 0.0 - seed, pos = _seed_and_pos(num_tokens, num_tokens) - - expected = torch.empty( - num_tokens, - dtype=torch.int64, - device=DEVICE, - ) - expected[:8] = logits[:8].argmax(dim=-1) - - support = torch.tensor( - [ - 0, - 1023, - 1024, - 8191, - 8192, - 65535, - 131071, - VOCAB_SIZE - 1, - ], - dtype=torch.int64, - device=DEVICE, - ) - logits[8:] = float("-inf") - logits[torch.arange(8, 16, device=DEVICE), support] = 3.0 - expected[8:] = support - - sampled = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - ) - torch.npu.synchronize() - - torch.testing.assert_close(sampled, expected, rtol=0, atol=0) - - -@pytest.mark.parametrize("is_drafting", [False, True]) -def test_categorical_sample_deterministic_same_seed_and_pos(is_drafting): - """Identical input, seed and position must replay deterministically.""" - torch.manual_seed(2) - num_tokens = 16 - logits = torch.randn( - num_tokens, - VOCAB_SIZE, - dtype=torch.float32, - device=DEVICE, - ) - expanded_idx_mapping = torch.arange( - num_tokens, - dtype=torch.int32, - device=DEVICE, - ) - temperature = torch.tensor( - [0.5, 1.0] * (num_tokens // 2), - dtype=torch.float32, - device=DEVICE, - ) - seed, pos = _seed_and_pos(num_tokens, num_tokens) - - sampled_1 = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - is_drafting=is_drafting, - ) - sampled_2 = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - is_drafting=is_drafting, - ) - torch.npu.synchronize() - - torch.testing.assert_close(sampled_1, sampled_2, rtol=0, atol=0) - - -def test_categorical_sample_draft_rng_uses_position_salt(): - """Draft RNG must be equivalent to the target stream at pos + 2**30.""" - torch.manual_seed(20) - num_tokens = 32 - vocab_size = 1031 - logits = torch.randn( - num_tokens, - vocab_size, - dtype=torch.float32, - device=DEVICE, - ) - expanded_idx_mapping = torch.arange( - num_tokens, - dtype=torch.int32, - device=DEVICE, - ) - temperature = torch.ones( - num_tokens, - dtype=torch.float32, - device=DEVICE, - ) - seed, pos = _seed_and_pos(num_tokens, num_tokens) - - draft = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - is_drafting=True, - ) - shifted_target = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos + (1 << 30), - apply_temperature=True, - is_drafting=False, - ) - torch.npu.synchronize() - - torch.testing.assert_close(draft, shifted_target, rtol=0, atol=0) - - -def test_categorical_sample_apply_temperature_matches_prescaled_logits(): - """In-kernel scaling must match sampling from pre-scaled FP32 logits.""" - torch.manual_seed(3) - num_tokens = 16 - logits = ( - torch.randint( - -32, - 33, - (num_tokens, VOCAB_SIZE), - dtype=torch.int32, - device=DEVICE, - ).to(torch.float32) - / 8 - ) - expanded_idx_mapping = torch.arange( - num_tokens, - dtype=torch.int32, - device=DEVICE, - ) - temperature = torch.tensor( - [0.5, 1.0, 2.0, 1.0] * (num_tokens // 4), - dtype=torch.float32, - device=DEVICE, - ) - seed, pos = _seed_and_pos(num_tokens, num_tokens) - - sampled_scaled_in_kernel = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - ) - scaled_logits = logits / temperature[:, None] - sampled_prescaled = _sample( - scaled_logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=False, - ) - torch.npu.synchronize() - - torch.testing.assert_close( - sampled_scaled_in_kernel, - sampled_prescaled, - rtol=0, - atol=0, - ) - - -@pytest.mark.parametrize("per_token_col", [False, True]) -@pytest.mark.parametrize("dtype", SUPPORTED_DTYPES) -def test_categorical_sample_logits_cache(per_token_col, dtype): - """Cache must store raw pre-temperature logits at mapped slots.""" - torch.manual_seed(4) - num_tokens = 4 - max_num_reqs = 8 - num_cols = 3 - - logits = torch.randn( - num_tokens, - VOCAB_SIZE, - dtype=dtype, - device=DEVICE, - ) - expanded_idx_mapping = torch.tensor( - [2, 5, 7, 0], - dtype=torch.int32, - device=DEVICE, - ) - temperature = torch.tensor( - [0.5, 1.0, 2.0, 0.5, 1.0, 2.0, 0.5, 1.0], - dtype=torch.float32, - device=DEVICE, - ) - seed, pos = _seed_and_pos(num_tokens, max_num_reqs) - logits_cache = torch.zeros( - max_num_reqs, - num_cols, - VOCAB_SIZE, - dtype=dtype, - device=DEVICE, - ) - - if per_token_col: - logits_cache_col = torch.tensor( - [0, 1, 2, 1], - dtype=torch.int32, - device=DEVICE, - ) - expected_cols = [0, 1, 2, 1] - else: - logits_cache_col = torch.tensor( - 1, - dtype=torch.int32, - device=DEVICE, - ) - expected_cols = [1] * num_tokens - - _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - logits_cache=logits_cache, - logits_cache_col=logits_cache_col, - ) - torch.npu.synchronize() - - used = set() - for token_idx, col in enumerate(expected_cols): - req = expanded_idx_mapping[token_idx].item() - used.add((req, col)) - torch.testing.assert_close( - logits_cache[req, col], - logits[token_idx], - rtol=0, - atol=0, - ) - - for req in range(max_num_reqs): - for col in range(num_cols): - if (req, col) not in used: - assert torch.count_nonzero(logits_cache[req, col]).item() == 0 - - -def test_categorical_sample_padding_mapping_does_not_write_cache(): - """CUDAGraph padding request index -1 must not write logits_cache.""" - torch.manual_seed(5) - num_tokens = 4 - max_num_reqs = 4 - - logits = torch.randn( - num_tokens, - VOCAB_SIZE, - dtype=torch.float32, - device=DEVICE, - ) - expanded_idx_mapping = torch.tensor( - [0, 2, -1, -1], - dtype=torch.int32, - device=DEVICE, - ) - temperature = torch.ones( - max_num_reqs, - dtype=torch.float32, - device=DEVICE, - ) - seed, pos = _seed_and_pos(num_tokens, max_num_reqs) - logits_cache = torch.zeros( - max_num_reqs, - 1, - VOCAB_SIZE, - dtype=torch.float32, - device=DEVICE, - ) - - sampled = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - logits_cache=logits_cache, - ) - torch.npu.synchronize() - - assert sampled.shape == (num_tokens,) - torch.testing.assert_close( - logits_cache[0, 0], - logits[0], - rtol=0, - atol=0, - ) - torch.testing.assert_close( - logits_cache[2, 0], - logits[1], - rtol=0, - atol=0, - ) - assert torch.count_nonzero(logits_cache[1]).item() == 0 - assert torch.count_nonzero(logits_cache[3]).item() == 0 - - -def test_categorical_sample_shared_request_mapping(): - """Rows sharing a request must share its seed and temperature stream.""" - torch.manual_seed(6) - logits_row_0 = torch.randn( - 1, - VOCAB_SIZE, - dtype=torch.float32, - device=DEVICE, - ) - logits_row_1 = torch.randn( - 1, - VOCAB_SIZE, - dtype=torch.float32, - device=DEVICE, - ) - logits = torch.cat( - [logits_row_0, logits_row_0, logits_row_1, logits_row_1], - dim=0, - ) - expanded_idx_mapping = torch.tensor( - [0, 0, 1, 1], - dtype=torch.int32, - device=DEVICE, - ) - temperature = torch.tensor( - [0.7, 1.3], - dtype=torch.float32, - device=DEVICE, - ) - seed = torch.tensor( - [12345, 67890], - dtype=torch.int64, - device=DEVICE, - ) - pos = torch.tensor( - [11, 11, 19, 19], - dtype=torch.int64, - device=DEVICE, - ) - - sampled = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - ) - torch.npu.synchronize() - - assert sampled[0].item() == sampled[1].item() - assert sampled[2].item() == sampled[3].item() - - -def test_categorical_sample_cache_does_not_change_sampled_tokens(): - """Raw-logit cache writes must not change sampled token IDs.""" - torch.manual_seed(7) - num_tokens = 16 - logits = torch.randn( - num_tokens, - VOCAB_SIZE, - dtype=torch.float32, - device=DEVICE, - ) - expanded_idx_mapping = torch.arange( - num_tokens, - dtype=torch.int32, - device=DEVICE, - ) - temperature = torch.ones( - num_tokens, - dtype=torch.float32, - device=DEVICE, - ) - seed, pos = _seed_and_pos(num_tokens, num_tokens) - logits_cache = torch.zeros( - num_tokens, - 1, - VOCAB_SIZE, - dtype=torch.float32, - device=DEVICE, - ) - - sampled_without_cache = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - ) - sampled_with_cache = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - logits_cache=logits_cache, - ) - torch.npu.synchronize() - - torch.testing.assert_close( - sampled_without_cache, - sampled_with_cache, - rtol=0, - atol=0, - ) - torch.testing.assert_close( - logits_cache[:, 0], - logits, - rtol=0, - atol=0, - ) - - -def test_categorical_sample_random_distribution_sanity(): - """Equal finite logits should reach every supported token.""" - num_tokens = 128 - support = torch.tensor( - [7, 1027, 8199, VOCAB_SIZE - 1], - dtype=torch.int64, - device=DEVICE, - ) - logits = torch.full( - (num_tokens, VOCAB_SIZE), - float("-inf"), - dtype=torch.float32, - device=DEVICE, - ) - logits[:, support] = 0.0 - expanded_idx_mapping = torch.arange( - num_tokens, - dtype=torch.int32, - device=DEVICE, - ) - temperature = torch.ones( - num_tokens, - dtype=torch.float32, - device=DEVICE, - ) - seed, pos = _seed_and_pos(num_tokens, num_tokens) - - sampled = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - ) - torch.npu.synchronize() - - sampled_cpu = sampled.cpu() - support_cpu = support.cpu() - assert set(sampled_cpu.tolist()).issubset(set(support_cpu.tolist())) - counts = torch.tensor([(sampled_cpu == token).sum().item() for token in support_cpu]) - assert (counts >= 12).all() and (counts <= 52).all(), f"unexpected counts for equal-mass support: {counts.tolist()}" - - -def test_categorical_sample_business_shape_distribution_accuracy(): - """B64/V151936 samples must match the torch.softmax reference.""" - num_tokens = 64 - num_trials = 256 - support = torch.tensor( - [ - 7, - 1023, - 1024, - 8191, - 8192, - 65535, - 131071, - VOCAB_SIZE - 1, - ], - dtype=torch.int64, - device=DEVICE, - ) - support_logits = torch.tensor( - [-1.5, -1.0, -0.5, 0.0, 0.5, 1.0, 1.5, 2.0], - dtype=torch.float32, - device=DEVICE, - ) - - logits = torch.full( - (num_tokens, VOCAB_SIZE), - float("-inf"), - dtype=torch.float32, - device=DEVICE, - ) - logits[:, support] = support_logits - expanded_idx_mapping = torch.arange( - num_tokens, - dtype=torch.int32, - device=DEVICE, - ) - temperature = torch.ones( - num_tokens, - dtype=torch.float32, - device=DEVICE, - ) - base_seed, pos = _seed_and_pos(num_tokens, num_tokens) - - counts = torch.zeros( - len(support), - dtype=torch.int64, - device=DEVICE, - ) - for trial in range(num_trials): - seed = base_seed + trial * 1000003 - sampled = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - ) - for idx, token in enumerate(support): - counts[idx] += (sampled == token).sum() - - torch.npu.synchronize() - - actual_probs = counts.to(torch.float32).cpu() / (num_tokens * num_trials) - expected_probs = torch.softmax(support_logits, dim=0).cpu() - abs_error = torch.abs(actual_probs - expected_probs) - max_abs_error = torch.max(abs_error).item() - tv_distance = 0.5 * torch.sum(abs_error).item() - - assert max_abs_error <= 0.02, ( - f"categorical probability max abs error " - f"{max_abs_error:.6f} exceeds 0.02; " - f"actual={actual_probs.tolist()}, " - f"expected={expected_probs.tolist()}" - ) - assert tv_distance <= 0.04, ( - f"categorical probability TV distance " - f"{tv_distance:.6f} exceeds 0.04; " - f"actual={actual_probs.tolist()}, " - f"expected={expected_probs.tolist()}" - ) - - -def test_categorical_sample_empty_input(): - """Empty expanded batches must return an empty int64 tensor.""" - logits = torch.empty( - 0, - 32, - dtype=torch.float32, - device=DEVICE, - ) - expanded_idx_mapping = torch.empty( - 0, - dtype=torch.int32, - device=DEVICE, - ) - temperature = torch.ones( - 1, - dtype=torch.float32, - device=DEVICE, - ) - seed = torch.zeros( - 1, - dtype=torch.int64, - device=DEVICE, - ) - pos = torch.empty( - 0, - dtype=torch.int64, - device=DEVICE, - ) - - sampled = _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - ) - - assert sampled.dtype == torch.int64 - assert sampled.shape == (0,) - - -def test_categorical_sample_use_fp64_is_not_supported(): - logits = torch.zeros( - 1, - VOCAB_SIZE, - dtype=torch.float32, - device=DEVICE, - ) - expanded_idx_mapping = torch.zeros( - 1, - dtype=torch.int32, - device=DEVICE, - ) - temperature = torch.ones( - 1, - dtype=torch.float32, - device=DEVICE, - ) - seed, pos = _seed_and_pos(1, 1) - - with pytest.raises( - NotImplementedError, - match="FP64 categorical sampling is not supported on NPU", - ): - _sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature=True, - use_fp64=True, - ) diff --git a/tests/e2e/pull_request/one_card/test_gumbel_sampling.py b/tests/e2e/pull_request/one_card/test_gumbel_sampling.py index dc99ccc8f50c..34c5acfa830a 100644 --- a/tests/e2e/pull_request/one_card/test_gumbel_sampling.py +++ b/tests/e2e/pull_request/one_card/test_gumbel_sampling.py @@ -11,8 +11,9 @@ import torch from vllm.v1.worker.gpu.spec_decode.dspark.speculator import DSparkSpeculator -from vllm_ascend.ops.triton.v2.sample.categorical_sample import categorical_sample +from vllm_ascend.utils import vllm_version_is from vllm_ascend.worker.v2.sample.gumbel import apply_temperature +from vllm_ascend.worker.v2.sample.gumbel import gumbel_sample as _sample_for_version from vllm_ascend.worker.v2.spec_decode.rejection_sampler_utils import rejection_sample DEVICE = "npu" @@ -30,8 +31,20 @@ def gumbel_sample( *, is_drafting: bool = False, ) -> torch.Tensor: - """Run sampling assertions against the Ascend categorical implementation.""" - return categorical_sample( + """Run the existing target-sampling assertions through each lane's API.""" + if vllm_version_is("0.28.0"): + return _sample_for_version( + logits, + expanded_idx_mapping, + temperature, + seed, + pos, + apply_temperature=apply_temperature, + logits_cache=logits_cache, + logits_cache_col=logits_cache_col, + is_drafting=is_drafting, + ) + return _sample_for_version( logits, expanded_idx_mapping, temperature, @@ -174,6 +187,27 @@ def test_gumbel_sample_deterministic(self, num_tokens, num_reqs, vocab_size): assert torch.equal(r1, r2), "gumbel_sample is non-deterministic with same seed" + def test_gumbel_sample_different_seeds(self): + """Different seeds must (almost surely) produce different results.""" + torch.manual_seed(8) + num_tokens, num_reqs, vocab_size = 16, 16, 32000 + logits = torch.randn(num_tokens, vocab_size, dtype=torch.float32, device=DEVICE) + expanded_idx_mapping = torch.arange(num_tokens, dtype=torch.int32, device=DEVICE) + temperature = torch.ones(num_reqs, dtype=torch.float32, device=DEVICE) * 1.0 + pos = torch.arange(num_tokens, dtype=torch.int32, device=DEVICE) + + seed1 = torch.randint(0, 2**31, (num_reqs,), dtype=torch.int64, device=DEVICE) + seed2 = torch.randint(0, 2**31, (num_reqs,), dtype=torch.int64, device=DEVICE) + # Ensure seeds differ + seed2[0] = seed1[0] + 1 + + r1 = gumbel_sample(logits, expanded_idx_mapping, temperature, seed1, pos, apply_temperature=False) + r2 = gumbel_sample(logits, expanded_idx_mapping, temperature, seed2, pos, apply_temperature=False) + torch.npu.synchronize() + + # With 16 tokens and vocab 32000 at temp=1.0, identical results are astronomically unlikely + assert not torch.equal(r1, r2), "Different seeds produced identical results" + @pytest.mark.parametrize( "num_tokens,num_reqs,vocab_size", [ @@ -200,44 +234,46 @@ def test_gumbel_sample_valid_token_ids(self, num_tokens, num_reqs, vocab_size): ) def test_gumbel_sample_temperature_affects_distribution(self): - """Higher temperature should make a peaked distribution less concentrated.""" + """Higher temperature should increase sampling entropy (less concentrated). + + Strategy: create logits with a clear winner. At low temp the winner should + be sampled most often. At high temp other tokens get more probability. + """ vocab_size = 100 num_trials = 256 + logits_base = torch.zeros(1, vocab_size, dtype=torch.float32, device=DEVICE) + logits_base[0, 0] = 10.0 # strong signal at token 0 - logits = torch.zeros(num_trials, vocab_size, dtype=torch.float32, device=DEVICE) - logits[:, 0] = 10.0 - expanded_idx_mapping = torch.arange(num_trials, dtype=torch.int32, device=DEVICE) - seed = torch.arange(num_trials, dtype=torch.int64, device=DEVICE) * 1000 + 42 - pos = torch.arange(num_trials, dtype=torch.int32, device=DEVICE) + expanded_idx_mapping = torch.zeros(1, dtype=torch.int32, device=DEVICE) - low_temp = torch.full((num_trials,), 0.1, dtype=torch.float32, device=DEVICE) - high_temp = torch.full((num_trials,), 5.0, dtype=torch.float32, device=DEVICE) + low_temp = torch.tensor([0.1], dtype=torch.float32, device=DEVICE) + high_temp = torch.tensor([5.0], dtype=torch.float32, device=DEVICE) - low_samples = gumbel_sample( - logits, - expanded_idx_mapping, - low_temp, - seed, - pos, - apply_temperature=True, - ) - high_samples = gumbel_sample( - logits, - expanded_idx_mapping, - high_temp, - seed, - pos, - apply_temperature=True, - ) - torch.npu.synchronize() + low_temp_winner_count = 0 + high_temp_winner_count = 0 - low_temp_winner_count = (low_samples == 0).sum().item() - high_temp_winner_count = (high_samples == 0).sum().item() + for i in range(num_trials): + seed = torch.tensor([i * 1000 + 42], dtype=torch.int64, device=DEVICE) + pos = torch.tensor([i], dtype=torch.int32, device=DEVICE) + s_low = gumbel_sample( + logits_base.clone(), expanded_idx_mapping, low_temp, seed, pos, apply_temperature=True + ) + s_high = gumbel_sample( + logits_base.clone(), expanded_idx_mapping, high_temp, seed, pos, apply_temperature=True + ) + if s_low.item() == 0: + low_temp_winner_count += 1 + if s_high.item() == 0: + high_temp_winner_count += 1 + + torch.npu.synchronize() + # Low temp should pick the winner much more often than high temp assert low_temp_winner_count > high_temp_winner_count, ( f"Low temp winner count ({low_temp_winner_count}) should be > " f"high temp winner count ({high_temp_winner_count})" ) + # Low temp with such a strong signal should almost always pick token 0 assert low_temp_winner_count > num_trials * 0.9, ( f"Low temp winner count ({low_temp_winner_count}/{num_trials}) should be >90%" ) @@ -595,11 +631,9 @@ def test_gumbel_sample_rejects_narrow_cache(self): def test_dspark_uses_ascend_gumbel(self): """Exercise the inherited DSpark entry point with real NPU sampling.""" # Check the installed implementation, not this test module's API wrapper. - assert DSparkSpeculator._sample_logits.__globals__["gumbel_sample"] is categorical_sample + assert DSparkSpeculator._sample_logits.__globals__["gumbel_sample"] is _sample_for_version speculator = DSparkSpeculator.__new__(DSparkSpeculator) speculator._d2t_scatter_index = None - speculator.acceptance_estimator = None - speculator.draft_watermarker = None speculator.temperature = torch.tensor([0.5, 1.5], device=DEVICE) speculator.seeds = torch.tensor([3, 7], dtype=torch.int64, device=DEVICE) speculator._step_cols = torch.arange(2, dtype=torch.int32, device=DEVICE) diff --git a/tests/ut/attention/test_attention_v1.py b/tests/ut/attention/test_attention_v1.py index 7d7319229f6d..854b3b4115d1 100644 --- a/tests/ut/attention/test_attention_v1.py +++ b/tests/ut/attention/test_attention_v1.py @@ -19,7 +19,6 @@ ) from vllm_ascend.attention.utils import ( AscendCommonAttentionMetadata, - PagedAttentionGraphParam, cache_graph_workspace, needs_layer_aware_fia_graph_replay, using_paged_attention, @@ -887,123 +886,3 @@ def test_forward_decode_only_swa_seq_len_mismatch( mock_reshape_and_cache.assert_called_once() assert output.shape == (10, 8, 64) - - @patch("vllm_ascend.attention.attention_v1.torch.npu.stream") - @patch("vllm_ascend.attention.attention_v1.torch.npu.graph_task_update_begin") - @patch("vllm_ascend.attention.attention_v1.torch.npu.graph_task_update_end") - @patch("torch_npu.npu_fused_infer_attention_score") - @patch("vllm_ascend.attention.attention_v1.get_graph_params") - @patch("vllm_ascend.attention.attention_v1._EXTRA_CTX") - @patch("vllm_ascend.attention.attention_v1.using_paged_attention", return_value=False) - @patch("vllm_ascend.attention.attention_v1.needs_layer_aware_fia_graph_replay", return_value=False) - @patch("vllm_ascend.attention.attention_v1._ATTN_KEYS_BUFFER", new=[]) - def test_update_graph_params( - self, - mock_needs_layer_aware_fia_graph_replay, - mock_using_paged_attention, - mock_EXTRA_CTX, - mock_get_graph_params, - mock_fia, - mock_graph_task_update_end, - mock_graph_task_update_begin, - mock_stream, - ): - """Test behavior when _ATTN_KEYS_BUFFER is [] after dummy_run.""" - - mock_EXTRA_CTX.sinks = False - mock_EXTRA_CTX.is_draft_model = False - - param: list[MagicMock | None] = [MagicMock()] * 22 - param[16] = None # sliding_window - param[17] = None # c8_k_aq_scale - param[21] = None # layer_name - - mock_get_graph_params.return_value.attn_params = {1: [tuple(param)] * 3} - mock_get_graph_params.return_value.handles = {1: [MagicMock()] * 3} - mock_get_graph_params.return_value.events = {1: [MagicMock()] * 3} - - attn_metadata_keys = [ - "model.layers.10.self_attn.attn", - "model.layers.2.self_attn.attn", - "model.layers.5.self_attn.attn", - ] - forward_context = MagicMock() - forward_context.attn_metadata = {key: MagicMock() for key in attn_metadata_keys} - # breakpoint() - self.impl.update_graph_params(self.mock_stream, forward_context, 1, self.mock_vllm_config) - - expected = [ - "model.layers.2.self_attn.attn", - "model.layers.5.self_attn.attn", - "model.layers.10.self_attn.attn", - ] - self.assertEqual(attn_module._ATTN_KEYS_BUFFER, expected) - self.assertEqual(mock_fia.out.call_count, 3) - - @patch("vllm_ascend.attention.attention_v1.torch.npu.stream") - @patch("vllm_ascend.attention.attention_v1.torch.npu.graph_task_update_begin") - @patch("vllm_ascend.attention.attention_v1.torch.npu.graph_task_update_end") - @patch("vllm_ascend.attention.attention_v1.torch_npu._npu_paged_attention") - @patch("vllm_ascend.attention.attention_v1.torch_npu._npu_paged_attention_get_workspace", return_value=MagicMock()) - @patch("vllm_ascend.attention.attention_v1.get_graph_params") - @patch("vllm_ascend.attention.attention_v1._EXTRA_CTX") - @patch("vllm_ascend.attention.attention_v1.using_paged_attention", return_value=True) - @patch("vllm_ascend.attention.attention_v1.needs_layer_aware_fia_graph_replay", return_value=False) - @patch("vllm_ascend.attention.attention_v1._ATTN_KEYS_BUFFER", new=[]) - def test_update_graph_params_handles_captured_paged_attention_params( - self, - mock_needs_layer_aware_fia_graph_replay, - mock_using_paged_attention, - mock_EXTRA_CTX, - mock_get_graph_params, - mock_get_workspace, - mock_paged_attention, - mock_graph_task_update_end, - mock_graph_task_update_begin, - mock_stream, - ): - mock_EXTRA_CTX.sinks = False - mock_EXTRA_CTX.is_draft_model = False - - query = MagicMock() - key_cache = MagicMock() - value_cache = MagicMock() - block_table = MagicMock() - output = MagicMock() - captured_seq_lens = MagicMock() - current_seq_lens = MagicMock() - pa_param = PagedAttentionGraphParam( - ( - query, - key_cache, - value_cache, - 8, - 8, - 1.0, - block_table, - captured_seq_lens, - output, - ), - "model.layers.0.self_attn.attn", - ) - - mock_get_graph_params.return_value.attn_params = {1: [pa_param]} - mock_get_graph_params.return_value.handles = {1: [MagicMock()]} - mock_get_graph_params.return_value.events = {1: [MagicMock()]} - - forward_context = MagicMock() - forward_context.attn_metadata = { - "model.layers.0.self_attn.attn": MagicMock( - seq_lens=current_seq_lens, - block_tables=block_table, - seq_lens_list=[10], - ), - } - - self.impl.update_graph_params(self.mock_stream, forward_context, 1, self.mock_vllm_config) - - mock_get_workspace.assert_called_once() - mock_paged_attention.assert_called_once() - self.assertEqual(mock_paged_attention.call_args.kwargs["context_lens"], current_seq_lens) - mock_graph_task_update_begin.assert_called_once() - mock_graph_task_update_end.assert_called_once() diff --git a/tests/ut/compilation/test_acl_graph.py b/tests/ut/compilation/test_acl_graph.py index 830488940089..0787fc6e75d0 100644 --- a/tests/ut/compilation/test_acl_graph.py +++ b/tests/ut/compilation/test_acl_graph.py @@ -12,6 +12,7 @@ # limitations under the License. # This file is a part of the vllm-ascend project. # +import contextlib import weakref from unittest.mock import MagicMock, Mock, patch @@ -47,7 +48,8 @@ from vllm_ascend.device_allocator.sleep_mem_optimized import AclGraphSleepWakeupManager -def test_update_full_graph_params_dispatches_draft_metadata_by_keyword(): +@patch("vllm_ascend.compilation.acl_graph.use_updatable_graph", return_value=False) +def test_update_full_graph_params_dispatches_draft_metadata_by_keyword(mock_use_updatable): impl_cls = MagicMock() attn_backend = MagicMock() attn_backend.get_impl_cls.return_value = impl_cls @@ -123,6 +125,12 @@ def setUp(self): self.addCleanup(self.get_ascend_config_patcher.stop) self.mock_get_ascend_config.return_value.ascend_compilation_config.enable_super_kernel = False + self.exit_stack = contextlib.ExitStack() + self.addCleanup(self.exit_stack.close) + self.mock_updatable_graph = self.exit_stack.enter_context( + patch("vllm_ascend.compilation.acl_graph.UpdatableGraph") + ) + # Mock VllmConfig self.mock_vllm_config = MagicMock(spec=VllmConfig) self.mock_vllm_config.compilation_config = MagicMock() @@ -284,7 +292,7 @@ def test_call_capture_graph_first_time( # Mock torch.npu.NPUGraph mock_npu_graph = MagicMock() - mock_torch.npu.NPUGraph.return_value = mock_npu_graph + self.mock_updatable_graph.return_value = mock_npu_graph # Mock torch.npu.graph context manager mock_graph_context = MagicMock() @@ -316,7 +324,7 @@ def test_call_capture_graph_first_time( # Verify graph capture happened mock_validate_cudagraph_capturing_enabled.assert_called_once() - mock_torch.npu.NPUGraph.assert_called_once() + self.mock_updatable_graph.assert_called_once() mock_torch.npu.graph.assert_called_once_with(mock_npu_graph, pool=self.mock_graph_pool) self.mock_runnable.assert_called_once_with(test_tensor, "arg2") @@ -367,7 +375,7 @@ def test_capture_respects_super_kernel_setting( self.mock_get_ascend_config.return_value.ascend_compilation_config.enable_super_kernel = enabled mock_npu_graph = MagicMock() - mock_torch.npu.NPUGraph.return_value = mock_npu_graph + self.mock_updatable_graph.return_value = mock_npu_graph mock_graph_context = MagicMock() mock_torch.npu.graph.return_value = mock_graph_context mock_graph_context.__enter__ = Mock(return_value=None) @@ -422,7 +430,7 @@ def test_call_replay_graph( # Mock torch.npu.NPUGraph mock_npu_graph = MagicMock() - mock_torch.npu.NPUGraph.return_value = mock_npu_graph + self.mock_updatable_graph.return_value = mock_npu_graph # Mock torch.npu.graph context manager mock_graph_context = MagicMock() @@ -454,7 +462,7 @@ def test_call_replay_graph( # Verify graph capture happened during first call mock_validate_cudagraph_capturing_enabled.assert_called_once() - mock_torch.npu.NPUGraph.assert_called_once() + self.mock_updatable_graph.assert_called_once() mock_torch.npu.graph.assert_called_once() # Reset mock to track second call @@ -501,7 +509,7 @@ def test_call_with_debug_mode_input_address_check( # Mock torch.npu.NPUGraph mock_npu_graph = MagicMock() - mock_torch.npu.NPUGraph.return_value = mock_npu_graph + self.mock_updatable_graph.return_value = mock_npu_graph # Mock torch.npu.graph context manager mock_graph_context = MagicMock() @@ -562,7 +570,7 @@ def test_call_with_debug_mode_input_address_mismatch( # Mock torch.npu.NPUGraph mock_npu_graph = MagicMock() - mock_torch.npu.NPUGraph.return_value = mock_npu_graph + self.mock_updatable_graph.return_value = mock_npu_graph # Mock torch.npu.graph context manager mock_graph_context = MagicMock() @@ -633,7 +641,7 @@ def test_call_capture_graph_with_gc_disable( # Mock torch.npu.NPUGraph mock_npu_graph = MagicMock() - mock_torch.npu.NPUGraph.return_value = mock_npu_graph + self.mock_updatable_graph.return_value = mock_npu_graph # Mock torch.npu.graph context manager mock_graph_context = MagicMock() @@ -675,7 +683,7 @@ def test_call_capture_graph_with_gc_disable( # Verify graph capture happened mock_validate_cudagraph_capturing_enabled.assert_called_once() - mock_torch.npu.NPUGraph.assert_called_once() + self.mock_updatable_graph.assert_called_once() mock_torch.npu.graph.assert_called_once_with(mock_npu_graph, pool=self.mock_graph_pool) # Should return the original output (not weak ref) since weak_ref_output is not enabled @@ -712,7 +720,7 @@ def test_call_capture_graph_with_weak_ref_output( # Mock torch.npu.NPUGraph mock_npu_graph = MagicMock() - mock_torch.npu.NPUGraph.return_value = mock_npu_graph + self.mock_updatable_graph.return_value = mock_npu_graph # Mock torch.npu.graph context manager mock_graph_context = MagicMock() @@ -749,7 +757,7 @@ def test_call_capture_graph_with_weak_ref_output( # Verify graph capture happened mock_validate_cudagraph_capturing_enabled.assert_called_once() - mock_torch.npu.NPUGraph.assert_called_once() + self.mock_updatable_graph.assert_called_once() mock_torch.npu.graph.assert_called_once_with(mock_npu_graph, pool=self.mock_graph_pool) # Should return the weak ref output when weak_ref_output option is enabled @@ -778,7 +786,7 @@ def test_call_capture_graph_with_debug_log( with patch("vllm_ascend.compilation.acl_graph.torch") as mock_torch: # Mock torch.npu.NPUGraph mock_npu_graph = MagicMock() - mock_torch.npu.NPUGraph.return_value = mock_npu_graph + self.mock_updatable_graph.return_value = mock_npu_graph # Mock torch.npu.graph context manager mock_graph_context = MagicMock() diff --git a/tests/ut/ops/test_fused_moe.py b/tests/ut/ops/test_fused_moe.py index 3f3d2c4cae52..85a03b782f55 100644 --- a/tests/ut/ops/test_fused_moe.py +++ b/tests/ut/ops/test_fused_moe.py @@ -12,7 +12,6 @@ from vllm_ascend.ascend_forward_context import MoECommType from vllm_ascend.device.hardware import AscendDeviceType -from vllm_ascend.ops import register_custom_ops as custom_ops from vllm_ascend.ops.fused_moe import fused_moe as fused_moe_module from vllm_ascend.ops.fused_moe import routed_experts as routed_experts_module from vllm_ascend.ops.fused_moe import shared_experts as shared_experts_module @@ -38,26 +37,6 @@ from vllm_ascend.quantization.quant_type import QuantType -@pytest.fixture -def runtime_moe_all_reduce(monkeypatch): - context = SimpleNamespace(moe_comm_type=MoECommType.ALLGATHER, no_compile_layers={}) - monkeypatch.setattr(fused_moe_module, "_EXTRA_CTX", context) - monkeypatch.setattr(custom_ops, "_EXTRA_CTX", context) - monkeypatch.setattr(custom_ops, "get_forward_context", lambda: context) - monkeypatch.setattr( - custom_ops, - "tensor_model_parallel_all_reduce", - lambda states: fused_moe_module.tensor_model_parallel_all_reduce(states), - ) - library = torch.library.Library("vllm", "IMPL", "CPU") - library.impl("maybe_all_reduce_tensor_model_parallel", custom_ops._maybe_all_reduce_tensor_model_parallel_impl) - library.impl("maybe_all_reduce_shared_expert", custom_ops._maybe_all_reduce_shared_expert_impl) - try: - yield context - finally: - library._destroy() - - def _build_weight_layer(): return SimpleNamespace( w13_weight=nn.Parameter(torch.randn(2, 3, 4)), @@ -411,18 +390,12 @@ def test_runner_reduction_contract(monkeypatch, moe_comm_type, is_sequence_paral ) def test_final_output_never_all_reduces_sequence_shards( monkeypatch, - runtime_moe_all_reduce, is_sequence_parallel, output_is_reduced, should_reduce, ): - runtime_moe_all_reduce.moe_comm_type = MoECommType.MC2 if output_is_reduced else MoECommType.ALLGATHER runner = AscendMoERunner.__new__(AscendMoERunner) - runner.layer_name = "test.reduction" - runtime_moe_all_reduce.no_compile_layers[runner.layer_name] = runner runner.moe_config = SimpleNamespace(is_sequence_parallel=is_sequence_parallel) - runner.routed_output_transform = None - runner.ascend_shared_experts = None states = torch.ones(2, 4) reduced_states = states + 1 all_reduce = MagicMock(return_value=reduced_states) @@ -458,17 +431,11 @@ def test_final_output_never_all_reduces_sequence_shards( ) def test_shared_output_reduction_depends_on_weight_layout( monkeypatch, - runtime_moe_all_reduce, mode, fused_output_is_reduced, reduce_shared, ): - runtime_moe_all_reduce.moe_comm_type = MoECommType.MC2 if fused_output_is_reduced else MoECommType.ALLGATHER runner = AscendMoERunner.__new__(AscendMoERunner) - runner.layer_name = "test.reduction" - runtime_moe_all_reduce.no_compile_layers[runner.layer_name] = runner - runner.moe_config = SimpleNamespace(is_sequence_parallel=False) - runner.routed_output_transform = None runner.ascend_shared_experts = SimpleNamespace(parallel_mode=MagicMock(return_value=mode)) shared_output = torch.ones(2, 4) reduced_output = shared_output + 1 @@ -479,9 +446,9 @@ def test_shared_output_reduction_depends_on_weight_layout( all_reduce, ) - result = runner._maybe_reduce_shared_expert_output( + result = runner._reduce_shared_output_if_needed( shared_output, - not fused_output_is_reduced, # A stale tracing-time flag must not control runtime reduction. + fused_output_is_reduced, ) if reduce_shared: @@ -502,13 +469,10 @@ def test_shared_output_reduction_depends_on_weight_layout( ) def test_local_shared_expert_dp_reduces_partial_routed_output( monkeypatch, - runtime_moe_all_reduce, mode, reduce_routed, ): runner = AscendMoERunner.__new__(AscendMoERunner) - runner.layer_name = "test.reduction" - runtime_moe_all_reduce.no_compile_layers[runner.layer_name] = runner runner.ascend_shared_experts = SimpleNamespace(parallel_mode=MagicMock(return_value=mode)) runner.routed_output_transform = None runner.moe_config = SimpleNamespace( @@ -2406,73 +2370,45 @@ def fake_linear(x, w): @pytest.mark.parametrize("initial_comm", [MoECommType.ALLGATHER, MoECommType.MC2]) -@pytest.mark.parametrize("layout", ["tp", "no_shared", "shared_dp", "sp", "latent"]) -def test_maybe_all_reduce_graph_replay_preserves_shared_and_routed_outputs( - monkeypatch, - runtime_moe_all_reduce, - initial_comm, - layout, -): +def test_compiled_moe_forward_keeps_runtime_reduction(monkeypatch, initial_comm): from torch.fx.experimental.proxy_tensor import make_fx - context = runtime_moe_all_reduce - tp_size = 8 - is_sp = layout == "sp" - has_shared = layout != "no_shared" runner = AscendMoERunner.__new__(AscendMoERunner) nn.Module.__init__(runner) - runner.layer_name = "test.runtime_maybe_all_reduce" - context.no_compile_layers[runner.layer_name] = runner - runner.moe_config = SimpleNamespace(is_sequence_parallel=is_sp, tp_size=tp_size, ep_size=tp_size) - runner.routed_experts = SimpleNamespace(quant_method=SimpleNamespace(has_unpadded_output=False)) - runner.routed_scaling_factor = 1.0 - runner.routed_output_transform = (lambda x: x.square()) if layout == "latent" else None - mode = { - "sp": SharedExpertParallelMode.SEQUENCE_PARALLEL_ONLY, - "shared_dp": SharedExpertParallelMode.SHARED_EXPERT_DATA_PARALLEL_ONLY, - }.get(layout, SharedExpertParallelMode.TENSOR_PARALLEL) - runner.ascend_shared_experts = SimpleNamespace(parallel_mode=lambda: mode, multistream_overlap=False) - runner.apply_routed_input_transform = lambda x: (x, x if has_shared else None) - runner._maybe_pad_hidden_states = lambda shared, x: (x, None, 3) - runner._encode_layer_name = lambda: runner.layer_name - runner._maybe_add_zero_expert_output = lambda x: x - monkeypatch.setattr(fused_moe_module, "tensor_model_parallel_all_reduce", lambda x: x * tp_size) - - library = torch.library.Library("moe_reduction_test", "FRAGMENT") - library.define("experts(Tensor x) -> (Tensor, Tensor)") - - def experts(x): - routed_reduced = is_sp or context.moe_comm_type != MoECommType.ALLGATHER - shared = x * (tp_size if layout in ("sp", "shared_dp") else 1) - routed = x * (2 * tp_size if routed_reduced else 2) - return shared, routed - - library.impl("experts", experts, "CPU") - torch.library.register_fake( - "moe_reduction_test::experts", - lambda x: (torch.empty_like(x), torch.empty_like(x)), - lib=library, - ) - - def forward_entry(x, *args): - shared, routed = torch.ops.moe_reduction_test.experts(x) - return (shared, routed) if has_shared else routed - - runner._forward_entry = forward_entry + runner.layer_name = "test.runtime_reduction" + runner.moe_config = SimpleNamespace(is_sequence_parallel=False) + context = SimpleNamespace(moe_comm_type=initial_comm) + monkeypatch.setattr(fused_moe_module, "_EXTRA_CTX", context) + monkeypatch.setattr( + fused_moe_module, + "get_forward_context", + lambda: SimpleNamespace(no_compile_layers={runner.layer_name: runner}), + ) + monkeypatch.setattr(fused_moe_module, "tensor_model_parallel_all_reduce", lambda states: states * 4) + + def upstream_forward(self, hidden_states, router_logits, **kwargs): + return self._maybe_reduce_final_output(hidden_states.clone(), None) + + monkeypatch.setattr(fused_moe_module.MoERunner, "forward", upstream_forward) + library = torch.library.Library("vllm", "IMPL", "CPU") + library.impl("ascend_moe_forward_complete", fused_moe_module._ascend_moe_forward_complete) try: - context.moe_comm_type = initial_comm states = torch.ones(2, 4) graph = make_fx(lambda x: runner(x, x))(states) - expected_routed = (2 * tp_size) ** 2 if layout == "latent" else 2 * tp_size - expected = expected_routed + (tp_size if has_shared else 0) - for comm in (MoECommType.ALLGATHER, MoECommType.MC2, MoECommType.ALLTOALL, MoECommType.FUSED_MC2): + # Reuse this exact graph as vLLM does, without retracing Python guards. + for comm in (MoECommType.ALLGATHER, MoECommType.MC2, MoECommType.ALLTOALL): context.moe_comm_type = comm - torch.testing.assert_close(graph(states), torch.full((2, 3), float(expected))) - if not is_sp: - assert any( - node.target == torch.ops.vllm.maybe_all_reduce_tensor_model_parallel.default - for node in graph.graph.nodes - ) - assert not any("ascend_moe_forward_complete" in str(node.target) for node in graph.graph.nodes) + expected = states * (4 if comm == MoECommType.ALLGATHER else 1) + torch.testing.assert_close(graph(states), expected) + assert any(node.target == torch.ops.vllm.ascend_moe_forward_complete.default for node in graph.graph.nodes) finally: library._destroy() + + +@pytest.mark.parametrize("shared_width", [None, 8]) +def test_complete_moe_fake_preserves_local_token_count(shared_width): + hidden = torch.empty(3, 4) + shared = torch.empty(12, shared_width) if shared_width is not None else None + result = fused_moe_module._ascend_moe_forward_complete_fake(hidden, hidden, shared, None, "test") + assert result.shape == (3, shared_width or 4) + assert result.dtype == hidden.dtype diff --git a/tests/ut/patch/worker/test_patch_v2_gumbel.py b/tests/ut/patch/worker/test_patch_v2_gumbel.py index c020407d9a8b..375ffc8e3e6d 100644 --- a/tests/ut/patch/worker/test_patch_v2_gumbel.py +++ b/tests/ut/patch/worker/test_patch_v2_gumbel.py @@ -9,23 +9,12 @@ from vllm.v1.worker.gpu.sample import gumbel, sampler from vllm.v1.worker.gpu.spec_decode import speculator as base_speculator from vllm.v1.worker.gpu.spec_decode.dspark import speculator as dspark_speculator -from vllm.v1.worker.gpu.spec_decode.eagle import speculator as eagle_speculator -from vllm_ascend.ops.triton.v2.sample.categorical_sample import categorical_sample from vllm_ascend.patch.worker.patch_v2 import patch_triton -from vllm_ascend.utils import vllm_version_is +from vllm_ascend.worker.v2.sample.gumbel import gumbel_sample -@pytest.mark.parametrize( - "consumer", - [ - gumbel, - sampler, - base_speculator, - dspark_speculator, - eagle_speculator, - ], -) +@pytest.mark.parametrize("consumer", [gumbel, sampler, base_speculator, dspark_speculator]) def test_gumbel_patch_rebinds_preimported_consumers(monkeypatch, consumer): # Simulate a consumer retaining its original `from ... import` binding. stale_gumbel = MagicMock() @@ -33,8 +22,7 @@ def test_gumbel_patch_rebinds_preimported_consumers(monkeypatch, consumer): importlib.reload(patch_triton) - # vLLM-Ascend replaces gumbel_sample with categorical_sample for NPU. - assert consumer.gumbel_sample is categorical_sample + assert consumer.gumbel_sample is gumbel_sample stale_gumbel.assert_not_called() @@ -49,25 +37,14 @@ def test_dspark_sample_logits_dispatch(monkeypatch, probabilistic): speculator._step_cols = torch.arange(2, dtype=torch.int32) speculator.draft_logits = torch.empty(2, 2, 3) if probabilistic else None speculator.use_fp64_gumbel = False - - # This fixture bypasses __init__, so provide defaults used by current main. - # They are harmless extra instance attributes on the v0.28-compatible lane. - speculator.acceptance_estimator = None - speculator.draft_watermarker = None - logits = torch.tensor([[1.0, 3.0, 2.0], [4.0, 2.0, 1.0]]) idx_mapping = torch.tensor([1, 0], dtype=torch.int32) sample_pos = torch.tensor([8, 12]) sampled = torch.tensor([2, 1]) - sample = create_autospec(categorical_sample, return_value=sampled) + sample = create_autospec(gumbel_sample, return_value=sampled) monkeypatch.setattr(dspark_speculator, "gumbel_sample", sample) - result = speculator._sample_logits( - logits, - idx_mapping, - sample_pos, - step=1, - ) + result = speculator._sample_logits(logits, idx_mapping, sample_pos, step=1) if probabilistic: assert result is sampled @@ -77,19 +54,9 @@ def test_dspark_sample_logits_dispatch(monkeypatch, probabilistic): assert args[1] is idx_mapping torch.testing.assert_close(args[4], sample_pos - 1) assert kwargs["apply_temperature"] is True - if vllm_version_is("0.28.0"): - assert "is_drafting" not in kwargs - else: - assert kwargs["is_drafting"] is True assert kwargs["logits_cache"] is speculator.draft_logits - torch.testing.assert_close( - kwargs["logits_cache_col"], - speculator._step_cols[1], - ) + torch.testing.assert_close(kwargs["logits_cache_col"], speculator._step_cols[1]) assert kwargs["use_fp64"] is False else: sample.assert_not_called() - torch.testing.assert_close( - result, - logits.argmax(dim=-1) + 10, - ) + torch.testing.assert_close(result, logits.argmax(dim=-1) + 10) diff --git a/tests/ut/spec_decode/test_dspark_proposer.py b/tests/ut/spec_decode/test_dspark_proposer.py index fa7672dd577c..bdb090184539 100644 --- a/tests/ut/spec_decode/test_dspark_proposer.py +++ b/tests/ut/spec_decode/test_dspark_proposer.py @@ -154,7 +154,9 @@ def release(): proposer.parallel_drafting = True proposer.token_indices_to_sample = torch.zeros(2, dtype=torch.int32) proposer.enable_enpu = False + proposer.draft_attn_groups = [MagicMock()] proposer._update_full_graph_params_if_needed = MagicMock() + proposer._maybe_update_metadata = MagicMock() proposer.set_inputs_first_pass = MagicMock() proposer.build_draft_attn_metadata = MagicMock() diff --git a/tests/ut/spec_decode/test_eagle_proposer.py b/tests/ut/spec_decode/test_eagle_proposer.py index 3dca82922128..1aac8e7d73f9 100644 --- a/tests/ut/spec_decode/test_eagle_proposer.py +++ b/tests/ut/spec_decode/test_eagle_proposer.py @@ -713,12 +713,18 @@ def setUp(self): self.mock_dp_group = patch("vllm_ascend.ascend_forward_context.get_dp_group", return_value=mock_dp_group) self.mock_dp_group.start() + self.mock_use_updatable_graph = patch( + "vllm_ascend.spec_decode.llm_base_proposer.use_updatable_graph", return_value=False + ) + self.mock_use_updatable_graph.start() + # Set the current vllm config set_current_vllm_config(self.vllm_config) self.proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner) self.proposer.model = MagicMock() self.proposer._runnable = MagicMock() self.proposer.update_stream = MagicMock() + self.proposer.draft_attn_groups = [MagicMock()] def tearDown(self): self.mock_get_ascend_config.stop() @@ -726,6 +732,7 @@ def tearDown(self): self.mock_supports_multimodal_inputs.stop() self.mock_tp_world_size.stop() self.mock_dp_group.stop() + self.mock_use_updatable_graph.stop() # Clear the current vllm config set_current_vllm_config(None) @@ -770,6 +777,7 @@ def test_dummy_run_in_graph_capture( mock_get_context.return_value = mock_return_context mock_get_context_2.return_value = mock_return_context self.proposer.use_cuda_graph = True + self.proposer.draft_attn_groups = [MagicMock()] # cpu does not support `torch.ops.vllm.maybe_pad_and_reduce` with set_current_vllm_config(self.vllm_config): self.proposer.dummy_run(num_tokens=64, in_graph_capturing=True, aclgraph_runtime_mode=CUDAGraphMode.FULL) @@ -964,6 +972,11 @@ def setUp_and_tearDown(self): set_current_vllm_config(self.vllm_config) self.proposer = AscendEagleProposer(vllm_config=self.vllm_config, device=self.device, runner=self.runner) + self.mock_use_updatable_graph = patch( + "vllm_ascend.spec_decode.llm_base_proposer.use_updatable_graph", return_value=False + ) + self.mock_use_updatable_graph.start() + yield self.mock_cpugpubuffer.stop() @@ -971,6 +984,7 @@ def setUp_and_tearDown(self): self.mock_tp_world_size.stop() self.mock_dp_group.stop() self.mock_get_ascend_config.stop() + self.mock_use_updatable_graph.stop() # Clear the current vllm config set_current_vllm_config(None) clear_ascend_config() diff --git a/vllm_ascend/attention/attention_v1.py b/vllm_ascend/attention/attention_v1.py index 92ea544f3138..f6943126778b 100644 --- a/vllm_ascend/attention/attention_v1.py +++ b/vllm_ascend/attention/attention_v1.py @@ -44,26 +44,20 @@ from vllm_ascend.attention.attention_mask import AttentionMaskBuilder from vllm_ascend.attention.utils import ( AscendCommonAttentionMetadata, - PagedAttentionGraphParam, - cache_graph_workspace, enable_dcp, needs_layer_aware_fia_graph_replay, notify_kv_cache_written, split_decodes_and_prefills, - update_paged_attention_graph_param, using_paged_attention, ) -from vllm_ascend.compilation.acl_graph import ( - get_draft_graph_params, - get_draft_graph_prefill_params, - get_graph_params, - update_draft_graph_params_workspaces, - update_graph_params_workspaces, +from vllm_ascend.compilation.updatable_graph import ( + get_capture_resource, + register_task, ) from vllm_ascend.device.device_op import DeviceOperator from vllm_ascend.device.hardware_profile import HardwareCapability, get_current_hardware_profile from vllm_ascend.distributed.kv_transfer.kv_pool.ascend_store.attention_fence import record_attention_compute_start -from vllm_ascend.utils import vllm_version_is, weak_ref_tensors +from vllm_ascend.utils import vllm_version_is if vllm_version_is("0.28.0"): from vllm.model_executor.layers.attention.pcp import _gather_prefill_cache_inputs # type: ignore[import-not-found] @@ -72,7 +66,9 @@ # default max value of sliding window size SWA_INT_MAX = 2147483647 -_ATTN_KEYS_BUFFER = None +_FIA_WORKSPACE_KEY = "npu_fused_infer_attention_score.workspace" +_FIA_V2_WORKSPACE_KEY = "npu_fused_infer_attention_score_v2.workspace" +_PA_WORKSPACE_KEY = "npu_paged_attention.workspace" @register_backend(AttentionBackendEnum.CUSTOM, "ASCEND") @@ -461,6 +457,52 @@ def build_for_graph_capture( return attn_metadata +@dataclass(frozen=True, slots=True) +class FIAParamProvider: + layer_name: str | None + sliding_window: int | None + is_draft_model: bool = False + + def resolve(self, attn_metadata) -> dict[str, Any]: + metadata = attn_metadata[self.layer_name] + + if self.is_draft_model or not self.sliding_window: + return { + "actual_seq_lengths": metadata.actual_seq_lengths_q, + "actual_seq_lengths_kv": metadata.seq_lens_list, + "block_table": metadata.block_tables, + } + else: + return { + "actual_seq_lengths": metadata.actual_seq_lengths_q, + "actual_seq_lengths_kv": metadata.seq_lens_list, + } + + +@dataclass(frozen=True, slots=True) +class FIAV2ParamProvider: + layer_name: str | None + + def resolve(self, attn_metadata) -> dict[str, Any]: + metadata = attn_metadata[self.layer_name] + return { + "actual_seq_qlen": metadata.actual_seq_lengths_q, + "actual_seq_kvlen": metadata.seq_lens_list, + } + + +@dataclass(frozen=True, slots=True) +class PAParamProvider: + layer_name: str | None + + def resolve(self, attn_metadata) -> dict[str, Any]: + metadata = attn_metadata[self.layer_name] + return { + "context_lens": metadata.seq_lens, + "block_table": metadata.block_tables, + } + + class AscendAttentionBackendImpl(AttentionImpl): def __init__( self, @@ -531,400 +573,7 @@ def update_graph_params( speculative_config=None, draft_attn_metadatas=None, ): - use_layer_aware_replay = needs_layer_aware_fia_graph_replay() - if using_paged_attention(num_tokens, vllm_config): - # Paged Attention update logic - if _EXTRA_CTX.is_draft_model: - if _EXTRA_CTX.is_draft_model_prefill: - graph_params = get_draft_graph_prefill_params() - else: - graph_params = get_draft_graph_params() - else: - graph_params = get_graph_params() - with torch.npu.stream(update_stream): - # The workspace size depends only on shapes and seq_lens, which - # are shared by all layers within one graph-param update pass, - # so one get_workspace call per pass is sufficient. Reuse the - # workspace across layers and re-query only when any - # size-relevant input changes. This mirrors the FIA path, which already - # reuses graph_params.workspaces[num_tokens], and removes the - # redundant per-layer get_workspace calls (each backed by a - # ~36-54MB buffer allocation) per decode step. - step_ws_key = None - workspace = None - for key, param, handle, event in zip( - forward_context.attn_metadata, - graph_params.attn_params[num_tokens], - graph_params.handles[num_tokens], - graph_params.events[num_tokens], - ): - ( - query, - key_cache, - value_cache, - num_kv_heads, - num_heads, - scale, - block_table, - seq_lens, - output, - ) = param - seq_lens = forward_context.attn_metadata[key].seq_lens - - # The key covers every size-relevant input, so models - # with heterogeneous layer configs simply trigger a - # re-query instead of reusing a wrong-sized workspace. - ws_key = ( - seq_lens.data_ptr(), - tuple(seq_lens.shape), - query.shape, - query.dtype, - key_cache.shape, - key_cache.dtype, - value_cache.shape, - value_cache.dtype, - block_table.shape if block_table is not None else None, - block_table.dtype if block_table is not None else None, - output.shape, - output.dtype, - num_kv_heads, - num_heads, - scale, - ) - if step_ws_key != ws_key: - workspace = torch_npu._npu_paged_attention_get_workspace( - query=query, - key_cache=key_cache, - value_cache=value_cache, - num_kv_heads=num_kv_heads, - num_heads=num_heads, - scale_value=scale, - block_table=block_table, - context_lens=seq_lens, - out=output, - ) - step_ws_key = ws_key - torch.npu.graph_task_update_begin(update_stream, handle) - torch_npu._npu_paged_attention( - query=query, - key_cache=key_cache, - value_cache=value_cache, - num_kv_heads=num_kv_heads, - num_heads=num_heads, - scale_value=scale, - block_table=block_table, - context_lens=seq_lens, - out=output, - workspace=workspace, - ) - torch.npu.graph_task_update_end(update_stream) - event.record(update_stream) - elif _EXTRA_CTX.sinks: - # FIA update logic - if _EXTRA_CTX.is_draft_model: - graph_params = get_draft_graph_params() - attn_metadata = draft_attn_metadatas - draft_attn_key_steps = [ - (draft_step, key) - for draft_step, per_step_metadata in enumerate(attn_metadata) - for key in per_step_metadata - ] - attn_keys = [key for _, key in draft_attn_key_steps] - else: - graph_params = get_graph_params() - attn_metadata = forward_context.attn_metadata - attn_keys = list(attn_metadata.keys()) - # For Qwen3-next, since the kv_cache_config has already categorized - # linear_attn and self_attn, the attn_metadata is first arranged with - # self_attn followed by linear_attn. Therefore, using zip directly - # filters out the update operations for linear_attn. - # TODO: We use a new variable `attn_keys` to ensure the loop count is - # correct after get by `zip` because of the new structure of the attn_metadata - # when running with the merged full eagle-graph. Should check it with Qwen3-next. - num_layers = len(attn_keys) - if num_layers == 0: - return - captured_attn_params = graph_params.attn_params[num_tokens] - handles = graph_params.handles[num_tokens] - events = graph_params.events[num_tokens] - graph_param_count = len(captured_attn_params) - workspace = graph_params.workspaces.get(num_tokens) - if _EXTRA_CTX.is_draft_model: - if graph_param_count > len(draft_attn_key_steps): - repeat_count = cdiv(graph_param_count, len(draft_attn_key_steps)) - draft_attn_key_steps = (draft_attn_key_steps * repeat_count)[:graph_param_count] - else: - draft_attn_key_steps = draft_attn_key_steps[:graph_param_count] - attn_keys = [key for _, key in draft_attn_key_steps] - elif use_layer_aware_replay: - # One graph size can contain captured FIA ops from all layers. - # Repeat attn keys to match the captured op count, then use the - # stored layer name in each op param to resolve the exact - # metadata entry during replay. - attn_keys = [attn_keys[index % num_layers] for index in range(graph_param_count)] - attn_count = 0 - with torch.npu.stream(update_stream): - for key, param, handle, event in zip( - attn_keys, - captured_attn_params, - handles, - events, - ): - ( - query, - key_cache, - value, - block_tables, - attn_mask, - block_size, - seq_lens, - num_kv_heads, - num_heads, - scale, - sliding_window, - sinks, - attn_output, - softmax_lse, - layer_name, - ) = param - - if _EXTRA_CTX.is_draft_model: - draft_step, key = draft_attn_key_steps[attn_count] - seq_lens = attn_metadata[draft_step][key].seq_lens_list - actual_seq_lengths_q = attn_metadata[draft_step][key].actual_seq_lengths_q - attn_count = attn_count + 1 - else: - metadata_key = layer_name if layer_name is not None and layer_name in attn_metadata else key - seq_lens = attn_metadata[metadata_key].seq_lens_list - actual_seq_lengths_q = attn_metadata[metadata_key].actual_seq_lengths_q - - torch.npu.graph_task_update_begin(update_stream, handle) - torch_npu.npu_fused_infer_attention_score_v2.out( - query=query, - key=key_cache, - value=value, - block_table=block_tables, - atten_mask=attn_mask, - input_layout="TND", - block_size=block_size, - actual_seq_qlen=actual_seq_lengths_q, - actual_seq_kvlen=seq_lens, - num_key_value_heads=num_kv_heads, - num_query_heads=num_heads, - sparse_mode=4 if sliding_window is not None else 3, - pre_tokens=sliding_window if sliding_window is not None else SWA_INT_MAX, - next_tokens=0, - softmax_scale=scale, - learnable_sink=sinks, - workspace=workspace, - out=[attn_output, softmax_lse], - ) - torch.npu.graph_task_update_end(update_stream) - event.record(update_stream) - else: - # FIA update logic - if _EXTRA_CTX.is_draft_model: - if _EXTRA_CTX.is_draft_model_prefill: - graph_params = get_draft_graph_prefill_params() - else: - graph_params = get_draft_graph_params() - attn_metadata = draft_attn_metadatas - draft_attn_key_steps = [ - (draft_step, key) - for draft_step, per_step_metadata in enumerate(attn_metadata) - for key in per_step_metadata - ] - attn_keys = [key for _, key in draft_attn_key_steps] - else: - graph_params = get_graph_params() - attn_metadata = forward_context.attn_metadata - # Only standard (FIA) attention layers have captured graph - # params here; linear/GDN layers (GDNAttentionMetadata) are - # updated separately by update_conv1d_graph_params. So we filter by `seq_lens_list` - attn_keys = [k for k in attn_metadata if hasattr(attn_metadata[k], "seq_lens_list")] - if not use_layer_aware_replay: - # In some speculative methods (such as DFlash), the order of - # attn_keys in the Target model will be disrupted instead of - # increasing by layer index, so need regular expressions to - # reorder the attn_keys and store the results in - # _ATTN_KEYS_BUFFER. - attn_keys_length = len(graph_params.attn_params[num_tokens]) - global _ATTN_KEYS_BUFFER - if attn_keys_length == 0: - return - if not _ATTN_KEYS_BUFFER or len(_ATTN_KEYS_BUFFER) != attn_keys_length: - import regex as re - - def extract_layer_index(key: str) -> int: - match = re.search(r"(?:^|\.)layers\.(\d+)(?:\.|$)", key) - return int(match.group(1)) if match else 0 - - def is_direct_target_attn_key(key: str) -> bool: - return ( - re.search( - r"(?:^|\.)layers\.(\d+)\.self_attn\.attn$", - key, - ) - is not None - ) - - attn_keys_to_order = attn_keys[:attn_keys_length] - if getattr(speculative_config, "method", None) == "mtp": - # Step3.5 MTP can expose draft KV-cache groups in the - # target runtime metadata. The target FULL graph only - # captures direct base-model self-attention handles, so - # select that target key domain instead of depending on - # the current draft module name. - direct_target_attn_keys = [key for key in attn_keys if is_direct_target_attn_key(key)] - if len(direct_target_attn_keys) >= attn_keys_length: - attn_keys_to_order = direct_target_attn_keys - - attn_keys_tmp = attn_keys_to_order - attn_keys_tmp.sort(key=extract_layer_index) - _ATTN_KEYS_BUFFER = attn_keys_tmp[:attn_keys_length] - attn_keys[:attn_keys_length] = _ATTN_KEYS_BUFFER - # For Qwen3-next, since the kv_cache_config has already categorized - # linear_attn and self_attn, the attn_metadata is first arranged with - # self_attn followed by linear_attn. Therefore, using zip directly - # filters out the update operations for linear_attn. - # TODO: We use a new variable `attn_keys` to ensure the loop count is - # correct after get by `zip` because of the new structure of the attn_metadata - # when running with the merged full eagle-graph. Should check it with Qwen3-next. - num_layers = len(attn_keys) - if num_layers == 0: - return - captured_attn_params = graph_params.attn_params[num_tokens] - handles = graph_params.handles[num_tokens] - events = graph_params.events[num_tokens] - graph_param_count = len(captured_attn_params) - workspace = graph_params.workspaces.get(num_tokens) - if _EXTRA_CTX.is_draft_model: - if graph_param_count > len(draft_attn_key_steps): - repeat_count = cdiv(graph_param_count, len(draft_attn_key_steps)) - draft_attn_key_steps = (draft_attn_key_steps * repeat_count)[:graph_param_count] - else: - draft_attn_key_steps = draft_attn_key_steps[:graph_param_count] - attn_keys = [key for _, key in draft_attn_key_steps] - elif use_layer_aware_replay: - # Keep the replay loop length aligned with captured FIA ops; - # layer-specific metadata lookup below prevents global/sliding - # window layers from accidentally sharing the same metadata. - attn_keys = [attn_keys[index % num_layers] for index in range(graph_param_count)] - attn_count = 0 - layer_count = 0 - with torch.npu.stream(update_stream): - for key, param, handle, event in zip( - attn_keys, - captured_attn_params, - handles, - events, - ): - if isinstance(param, PagedAttentionGraphParam): - if _EXTRA_CTX.is_draft_model: - draft_step, key = draft_attn_key_steps[attn_count] - block_table = attn_metadata[draft_step][key].block_tables - seq_lens = attn_metadata[draft_step][key].seq_lens - attn_count = attn_count + 1 - else: - layer_name = param.layer_name - metadata_key = layer_name if layer_name is not None and layer_name in attn_metadata else key - block_table = attn_metadata[metadata_key].block_tables - seq_lens = attn_metadata[metadata_key].seq_lens - update_paged_attention_graph_param( - update_stream, - handle, - event, - param, - block_table, - seq_lens, - ) - continue - ( - query, - key_cache, - value, - block_tables, - attn_mask, - block_size, - seq_lens, - query_start_loc, - num_kv_heads, - num_heads, - scale, - attn_output, - softmax_lse, - sparse_mode, - pre_tokens, - next_tokens, - sliding_window, - c8_k_aq_scale, - c8_k_aq_offset, - c8_v_aq_scale, - c8_v_aq_offset, - layer_name, - ) = param - - if _EXTRA_CTX.is_draft_model: - draft_step, key = draft_attn_key_steps[attn_count] - metadata = attn_metadata[draft_step][key] - seq_lens = metadata.seq_lens_list - actual_seq_lengths_q = metadata.actual_seq_lengths_q - block_tables = metadata.block_tables - attn_count = attn_count + 1 - if not metadata.causal: - sparse_mode = 0 - else: - metadata_key = layer_name if layer_name is not None and layer_name in attn_metadata else key - seq_lens = attn_metadata[metadata_key].seq_lens_list - actual_seq_lengths_q = attn_metadata[metadata_key].actual_seq_lengths_q - # NOTE: - # For models with sliding-window attention on the FIA full-graph replay path, - # rebinding `block_tables` to the latest metadata tensor causes corrupted / - # repeated outputs in our repro on Ascend NPU. - # - # Keep the captured block_tables tensor on this affected path. - # Non-SWA models preserve the original behavior and continue to refresh - # block_tables from attn_metadata. - if not sliding_window: - block_tables = attn_metadata[metadata_key].block_tables - layer_count += 1 - - torch.npu.graph_task_update_begin(update_stream, handle) - input_layout = "TND" - extra_args = {} - if c8_k_aq_scale is not None: - extra_args = { - "key_antiquant_scale": c8_k_aq_scale, - "value_antiquant_scale": c8_v_aq_scale, - "key_antiquant_mode": 0, - "value_antiquant_mode": 0, - "inner_precise": 1, - } - input_layout = "BNSD" - sparse_mode = 0 - torch_npu.npu_fused_infer_attention_score.out( - query=query, - key=key_cache, - value=value, - block_table=block_tables, - atten_mask=attn_mask, - input_layout=input_layout, - block_size=block_size, - actual_seq_lengths=actual_seq_lengths_q, - actual_seq_lengths_kv=seq_lens, - num_key_value_heads=num_kv_heads, - num_heads=num_heads, - scale=scale, - sparse_mode=sparse_mode, - pre_tokens=pre_tokens, - next_tokens=next_tokens, - **extra_args, - workspace=workspace, - out=[attn_output, softmax_lse], - ) - torch.npu.graph_task_update_end(update_stream) - - event.record(update_stream) + raise NotImplementedError("FIA and PA should be use UpdatableGraph.") def process_weights_after_loading(self, act_dtype: torch.dtype): super().process_weights_after_loading(act_dtype) @@ -941,13 +590,6 @@ def full_graph_fia( key, value, block_size, block_table, actual_seq_lengths_kv = self._get_fia_params(key, value, attn_metadata) num_tokens = attn_metadata.actual_seq_lengths_q[-1] - if _EXTRA_CTX.is_draft_model: - if _EXTRA_CTX.is_draft_model_prefill: - graph_params = get_draft_graph_prefill_params() - else: - graph_params = get_draft_graph_params() - else: - graph_params = get_graph_params() actual_seq_lengths_q = attn_metadata.actual_seq_lengths_q softmax_lse = torch.empty(1, dtype=query.dtype, device=query.device) input_layout = "TND" @@ -955,6 +597,7 @@ def full_graph_fia( sparse_mode = 4 if self.sliding_window else 3 if attn_metadata.causal else 0 pre_tokens = self.sliding_window or SWA_INT_MAX next_tokens = 0 if self.sliding_window else SWA_INT_MAX + output_view = output[: attn_metadata.num_actual_tokens] extra_args = {} if self.enable_c8_quant and layer is not None: @@ -974,16 +617,14 @@ def full_graph_fia( # TODO: change layerout from BNSD to TND. input_layout = "BNSD" query = query.unsqueeze(2) - output = output.unsqueeze(2) + output_view = output_view.unsqueeze(2) attn_mask = None sparse_mode = 0 + use_max_workspace = self._use_max_workspace_for_fia_graph - workspace = graph_params.workspaces.get(num_tokens) - should_update_workspace_cache = False - if use_max_workspace: - # Some models mix attention layer shapes under the same graph size. - # During capture, keep the largest required workspace for that size. - candidate_workspace = torch_npu._npu_fused_infer_attention_score_get_max_workspace( + workspace = get_capture_resource( + _FIA_WORKSPACE_KEY, + lambda: torch_npu._npu_fused_infer_attention_score_get_max_workspace( query=query, key=key, value=value, @@ -995,110 +636,38 @@ def full_graph_fia( actual_seq_lengths_kv=actual_seq_lengths_kv, num_key_value_heads=self.num_kv_heads, num_heads=self.num_heads, - sparse_mode=sparse_mode, pre_tokens=pre_tokens, next_tokens=next_tokens, scale=self.scale, - **extra_args, - ) - workspace = cache_graph_workspace( - graph_params, - num_tokens, - candidate_workspace, - use_max_workspace=use_max_workspace, - ) - should_update_workspace_cache = True - elif workspace is None: - workspace = torch_npu._npu_fused_infer_attention_score_get_max_workspace( - query=query, - key=key, - value=value, - atten_mask=attn_mask, - block_table=block_table, - input_layout=input_layout, - block_size=block_size, - actual_seq_lengths=actual_seq_lengths_q, - actual_seq_lengths_kv=actual_seq_lengths_kv, - num_key_value_heads=self.num_kv_heads, - num_heads=self.num_heads, sparse_mode=sparse_mode, - pre_tokens=pre_tokens, - next_tokens=next_tokens, - scale=self.scale, **extra_args, - ) - should_update_workspace_cache = True - if should_update_workspace_cache: - if _EXTRA_CTX.is_draft_model: - update_draft_graph_params_workspaces(num_tokens, workspace) - else: - update_graph_params_workspaces(num_tokens, workspace) - - # Handle graph capturing mode - stream = torch_npu.npu.current_stream() - - event = torch.npu.ExternalEvent() - event.wait(stream) - event.reset(stream) - graph_params.events[num_tokens].append(event) - attn_params = ( - weak_ref_tensors(query), - weak_ref_tensors(key), - weak_ref_tensors(value), - weak_ref_tensors(block_table), - weak_ref_tensors(attn_mask) if attn_mask is not None else None, - block_size, - actual_seq_lengths_kv, - actual_seq_lengths_q, - self.num_kv_heads, - self.num_heads, - self.scale, - weak_ref_tensors(output), - weak_ref_tensors(softmax_lse), - sparse_mode, - pre_tokens, - next_tokens, - self.sliding_window, + ), + use_max_workspace, ) - if self.enable_c8_quant and layer is not None: - attn_params = attn_params + ( - weak_ref_tensors(layer._c8_k_aq_scale_nz_bnsd), - None, - weak_ref_tensors(layer._c8_v_aq_scale_nz_bnsd), - None, - ) # type: ignore - else: - attn_params = attn_params + (None, None, None, None) # type: ignore - layer_name = self._graph_metadata_layer_name(layer) if self._use_layer_aware_fia_graph_replay else None - attn_params = attn_params + (layer_name,) # type: ignore - graph_params.attn_params[num_tokens].append(attn_params) - - torch.npu.graph_task_group_begin(stream) - torch_npu.npu_fused_infer_attention_score.out( - query=query, - key=key, - value=value, - atten_mask=attn_mask, - block_table=block_table, - input_layout=input_layout, - block_size=block_size, - actual_seq_lengths=actual_seq_lengths_q, - actual_seq_lengths_kv=actual_seq_lengths_kv, - num_key_value_heads=self.num_kv_heads, - num_heads=self.num_heads, - scale=self.scale, - sparse_mode=sparse_mode, - pre_tokens=pre_tokens, - next_tokens=next_tokens, - workspace=workspace, - out=[output, softmax_lse], - **extra_args, + register_task( + torch_npu.npu_fused_infer_attention_score.out, + { + "query": query, + "key": key, + "value": value, + "atten_mask": attn_mask, + "block_table": block_table, + "input_layout": input_layout, + "block_size": block_size, + "actual_seq_lengths": actual_seq_lengths_q, + "actual_seq_lengths_kv": actual_seq_lengths_kv, + "num_key_value_heads": self.num_kv_heads, + "num_heads": self.num_heads, + "pre_tokens": pre_tokens, + "next_tokens": next_tokens, + "scale": self.scale, + "sparse_mode": sparse_mode, + "workspace": workspace, + "out": [output_view, softmax_lse], + **extra_args, + }, + FIAParamProvider(self._layer_name, self.sliding_window, _EXTRA_CTX.is_draft_model), ) - - output = output.view(num_tokens, self.num_heads, self.head_size) - - handle = torch.npu.graph_task_group_end(stream) - graph_params.handles[num_tokens].append(handle) return output, num_tokens def full_graph_fia_v2( @@ -1112,46 +681,13 @@ def full_graph_fia_v2( key, value, block_size, block_table, actual_seq_lengths_kv = self._get_fia_params(key, value, attn_metadata) actual_seq_lengths_kv = attn_metadata.seq_lens num_tokens = attn_metadata.actual_seq_lengths_q[-1] - if _EXTRA_CTX.is_draft_model: - graph_params = get_draft_graph_params() - else: - graph_params = get_graph_params() - actual_seq_lengths_q = attn_metadata.actual_seq_lengths_q softmax_lse = torch.empty(1, dtype=query.dtype, device=query.device) + output_view = output[: attn_metadata.num_actual_tokens] use_max_workspace = self._use_max_workspace_for_fia_graph - workspace = graph_params.workspaces.get(num_tokens) - should_update_workspace_cache = False - if use_max_workspace: - # See full_graph_fia: this path needs the max workspace across layer - # variants sharing the same graph size. - candidate_workspace = torch_npu._npu_fused_infer_attention_score_v2_get_max_workspace( - query=query, - key=key, - value=value, - atten_mask=attn_metadata.attn_mask, - block_table=block_table, - input_layout="TND", - block_size=block_size, - actual_seq_qlen=actual_seq_lengths_q, - actual_seq_kvlen=actual_seq_lengths_kv, - num_key_value_heads=self.num_kv_heads, - softmax_scale=self.scale, - num_query_heads=self.num_heads, - sparse_mode=4 if self.sliding_window is not None else 3, - pre_tokens=self.sliding_window if self.sliding_window is not None else SWA_INT_MAX, - next_tokens=0, - learnable_sink=self.sinks, - ) - workspace = cache_graph_workspace( - graph_params, - num_tokens, - candidate_workspace, - use_max_workspace=use_max_workspace, - ) - should_update_workspace_cache = True - elif workspace is None: - workspace = torch_npu._npu_fused_infer_attention_score_v2_get_max_workspace( + workspace = get_capture_resource( + _FIA_V2_WORKSPACE_KEY, + lambda: torch_npu._npu_fused_infer_attention_score_v2_get_max_workspace( query=query, key=key, value=value, @@ -1168,63 +704,33 @@ def full_graph_fia_v2( pre_tokens=self.sliding_window if self.sliding_window is not None else SWA_INT_MAX, next_tokens=0, learnable_sink=self.sinks, - ) - should_update_workspace_cache = True - if should_update_workspace_cache: - if _EXTRA_CTX.is_draft_model: - update_draft_graph_params_workspaces(num_tokens, workspace) - else: - update_graph_params_workspaces(num_tokens, workspace) - - # Handle graph capturing mode - stream = torch_npu.npu.current_stream() - - event = torch.npu.ExternalEvent() - event.wait(stream) - event.reset(stream) - graph_params.events[num_tokens].append(event) - graph_params.attn_params[num_tokens].append( - ( - weak_ref_tensors(query), - weak_ref_tensors(key), - weak_ref_tensors(value), - weak_ref_tensors(block_table), - weak_ref_tensors(attn_metadata.attn_mask), - block_size, - actual_seq_lengths_kv, - self.num_kv_heads, - self.num_heads, - self.scale, - self.sliding_window, - self.sinks, - weak_ref_tensors(output), - weak_ref_tensors(softmax_lse), - self._graph_metadata_layer_name() if self._use_layer_aware_fia_graph_replay else None, - ) + ), + use_max_workspace, ) - torch.npu.graph_task_group_begin(stream) - torch_npu.npu_fused_infer_attention_score_v2.out( - query=query, - key=key, - value=value, - atten_mask=attn_metadata.attn_mask, - block_table=block_table, - input_layout="TND", - block_size=block_size, - actual_seq_qlen=actual_seq_lengths_q, - actual_seq_kvlen=actual_seq_lengths_kv, - num_key_value_heads=self.num_kv_heads, - num_query_heads=self.num_heads, - sparse_mode=4 if self.sliding_window is not None else 3, - pre_tokens=self.sliding_window if self.sliding_window is not None else SWA_INT_MAX, - next_tokens=0, - softmax_scale=self.scale, - learnable_sink=self.sinks, - workspace=workspace, - out=[output, softmax_lse], + register_task( + torch_npu.npu_fused_infer_attention_score_v2.out, + { + "query": query, + "key": key, + "value": value, + "atten_mask": attn_metadata.attn_mask, + "block_table": block_table, + "input_layout": "TND", + "block_size": block_size, + "actual_seq_qlen": actual_seq_lengths_q, + "actual_seq_kvlen": actual_seq_lengths_kv, + "num_key_value_heads": self.num_kv_heads, + "num_query_heads": self.num_heads, + "sparse_mode": 4 if self.sliding_window is not None else 3, + "pre_tokens": self.sliding_window if self.sliding_window is not None else SWA_INT_MAX, + "next_tokens": 0, + "softmax_scale": self.scale, + "learnable_sink": self.sinks, + "workspace": workspace, + "out": [output_view, softmax_lse], + }, + FIAV2ParamProvider(self._layer_name), ) - handle = torch.npu.graph_task_group_end(stream) - graph_params.handles[num_tokens].append(handle) return output, num_tokens def full_graph_pa( @@ -1233,51 +739,9 @@ def full_graph_pa( attn_metadata: AscendMetadata, output: torch.Tensor | None = None, ): - graph_params = get_graph_params() - num_tokens = query.shape[0] - if _EXTRA_CTX.capturing: - # Get workspace from cache or calculate it if not present. - workspace = graph_params.workspaces.get(num_tokens) - if workspace is None: - workspace = torch_npu._npu_paged_attention_get_workspace( - query=query, - key_cache=self.key_cache, - value_cache=self.value_cache, - num_kv_heads=self.num_kv_heads, - num_heads=self.num_heads, - scale_value=self.scale, - block_table=attn_metadata.block_tables, - context_lens=attn_metadata.seq_lens, - out=output, - ) - update_graph_params_workspaces(num_tokens, workspace) - - # Handle graph capturing mode - stream = torch_npu.npu.current_stream() - - event = torch.npu.ExternalEvent() - event.wait(stream) - event.reset(stream) - graph_params.events[num_tokens].append(event) - graph_params.attn_params[num_tokens].append( - PagedAttentionGraphParam( - ( - weak_ref_tensors(query), - weak_ref_tensors(self.key_cache), - weak_ref_tensors(self.value_cache), - self.num_kv_heads, - self.num_heads, - self.scale, - attn_metadata.block_tables, - attn_metadata.seq_lens, - weak_ref_tensors(output), - ), - self._graph_metadata_layer_name() if self._use_layer_aware_fia_graph_replay else None, - ) - ) - - torch.npu.graph_task_group_begin(stream) - torch_npu._npu_paged_attention( + workspace = get_capture_resource( + _PA_WORKSPACE_KEY, + lambda: torch_npu._npu_paged_attention_get_workspace( query=query, key_cache=self.key_cache, value_cache=self.value_cache, @@ -1287,11 +751,25 @@ def full_graph_pa( block_table=attn_metadata.block_tables, context_lens=attn_metadata.seq_lens, out=output, - workspace=workspace, - ) - handle = torch.npu.graph_task_group_end(stream) - graph_params.handles[num_tokens].append(handle) - return output + ), + ) + register_task( + torch_npu._npu_paged_attention, + { + "query": query, + "key_cache": self.key_cache, + "value_cache": self.value_cache, + "num_kv_heads": self.num_kv_heads, + "num_heads": self.num_heads, + "scale_value": self.scale, + "block_table": attn_metadata.block_tables, + "context_lens": attn_metadata.seq_lens, + "out": output, + "workspace": workspace, + }, + PAParamProvider(self._layer_name), + ) + return output def _get_kv_cache_view(self, key: torch.Tensor, value: torch.Tensor): if not self.use_bnsd_kv_cache: @@ -1761,8 +1239,7 @@ def forward( shape = [num_tokens, num_heads * head_size] """ assert output is not None, "Output tensor must be provided." - if self._use_layer_aware_fia_graph_replay: - self._layer_name = layer.layer_name + self._layer_name = layer.layer_name if output_scale is not None or output_block_scale is not None: raise NotImplementedError("fused output quantization is not yet supported for AscendAttentionBackendImpl") @@ -1831,8 +1308,7 @@ def forward( output_block_scale: torch.Tensor | None = None, ) -> torch.Tensor: assert output is not None, "Output tensor must be provided." - if self._use_layer_aware_fia_graph_replay: - self._layer_name = layer.layer_name + self._layer_name = layer.layer_name if output_scale is not None or output_block_scale is not None: raise NotImplementedError("fused output quantization is not yet supported for AscendC8AttentionBackendImpl") diff --git a/vllm_ascend/attention/utils.py b/vllm_ascend/attention/utils.py index 582a6fb517dc..28a247f20399 100644 --- a/vllm_ascend/attention/utils.py +++ b/vllm_ascend/attention/utils.py @@ -5,7 +5,6 @@ import torch import torch.nn.functional as F -import torch_npu from vllm.config import VllmConfig, get_current_vllm_config from vllm.distributed.kv_transfer import get_kv_transfer_group, has_kv_transfer_group, is_v1_kv_transfer_group from vllm.forward_context import ForwardContext, get_forward_context @@ -91,53 +90,6 @@ def __iter__(self): return iter(self.params) -def update_paged_attention_graph_param( - update_stream, - handle, - event, - param: PagedAttentionGraphParam, - block_table: torch.Tensor, - seq_lens: torch.Tensor, -) -> None: - ( - query, - key_cache, - value_cache, - num_kv_heads, - num_heads, - scale, - _captured_block_table, - _captured_seq_lens, - output, - ) = param.params - workspace = torch_npu._npu_paged_attention_get_workspace( - query=query, - key_cache=key_cache, - value_cache=value_cache, - num_kv_heads=num_kv_heads, - num_heads=num_heads, - scale_value=scale, - block_table=block_table, - context_lens=seq_lens, - out=output, - ) - torch.npu.graph_task_update_begin(update_stream, handle) - torch_npu._npu_paged_attention( - query=query, - key_cache=key_cache, - value_cache=value_cache, - num_kv_heads=num_kv_heads, - num_heads=num_heads, - scale_value=scale, - block_table=block_table, - context_lens=seq_lens, - out=output, - workspace=workspace, - ) - torch.npu.graph_task_update_end(update_stream) - event.record(update_stream) - - def cache_graph_workspace( graph_params, num_tokens: int, diff --git a/vllm_ascend/compilation/acl_graph.py b/vllm_ascend/compilation/acl_graph.py index e09f71d0c01a..51f047abef8d 100644 --- a/vllm_ascend/compilation/acl_graph.py +++ b/vllm_ascend/compilation/acl_graph.py @@ -23,7 +23,12 @@ from vllm_ascend.ascend_config import get_ascend_config from vllm_ascend.ascend_forward_context import _EXTRA_CTX -from ..utils import weak_ref_tensors +from ..utils import use_updatable_graph, weak_ref_tensors +from .updatable_graph import ( + ContextSource, + SharedSource, + UpdatableGraph, +) _acl_graph_wrappers: weakref.WeakSet[Any] = weakref.WeakSet() _STREAM_RESOURCE_ERROR_CODE = "207008" @@ -105,6 +110,7 @@ def __init__( *, use_eagle: bool = False, enable_enpu: bool = False, + update_stream: torch.npu.Stream | None = None, ): self.runnable = runnable self.vllm_config = vllm_config @@ -131,8 +137,21 @@ def __init__( self.concrete_aclgraph_entries: dict[BatchDescriptor, ACLGraphEntry] = {} self.enable_enpu = enable_enpu self.use_eagle = use_eagle + self.update_stream = update_stream + self.attn_backend = None + self.draft_model_metadata: list[dict[str, Any]] = [] _acl_graph_wrappers.add(self) + def set_update_stream(self, update_stream): + self.update_stream = update_stream + + def set_attn_backend(self, attn_backend): + self.attn_backend = attn_backend + + def update_draft_model_metadata(self, draft_model_metadata: list[dict[str, Any]]): + # This has been prepared for the update full graph of MRV1. + self.draft_model_metadata = draft_model_metadata + def __getattr__(self, key: str): # allow accessing the attributes of the runnable. if hasattr(self.runnable, key): @@ -179,7 +198,7 @@ def __call__(self, *args, **kwargs): input_addresses = [x.data_ptr() for x in args if isinstance(x, torch.Tensor)] entry.input_addresses = input_addresses - aclgraph = torch.npu.NPUGraph() + aclgraph = UpdatableGraph() with ExitStack() as stack: if self.aclgraph_options.gc_disable: @@ -288,9 +307,29 @@ def __call__(self, *args, **kwargs): need_sync = self.runtime_mode == CUDAGraphMode.FULL and not is_draft_eagle if not self.enable_enpu and need_sync: torch.npu.current_stream().synchronize() - entry.aclgraph.replay() + if self.runtime_mode == CUDAGraphMode.FULL and use_updatable_graph(self.attn_backend): + self._updatable_graph_replay(forward_context, entry.aclgraph) + else: + entry.aclgraph.replay() return entry.output + def _updatable_graph_replay( + self, + forward_context, + graph: UpdatableGraph, + ): + assert self.update_stream is not None + if _EXTRA_CTX.is_draft_model: + resolved_tasks = graph.resolve_tasks(SharedSource(self.draft_model_metadata)) + else: + resolved_tasks = graph.resolve_tasks(ContextSource(forward_context.attn_metadata)) + if self.enable_enpu: + graph.update(self.update_stream, resolved_tasks) + graph.replay() + else: + graph.replay() + graph.update(self.update_stream, resolved_tasks) + def weak_ref_workspaces(params): if params is None: @@ -310,6 +349,9 @@ def update_full_graph_params( speculative_config=None, draft_attn_metadatas=None, ): + if use_updatable_graph(attn_backend): + return + # vLLM >= 0.27.1 (main) makes get_current_vllm_config() raise # AssertionError outside set_current_vllm_config(); the SFA backend # resolution in get_impl_cls() needs the config. diff --git a/vllm_ascend/compilation/updatable_graph.py b/vllm_ascend/compilation/updatable_graph.py new file mode 100644 index 000000000000..566c20de525e --- /dev/null +++ b/vllm_ascend/compilation/updatable_graph.py @@ -0,0 +1,212 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from collections.abc import Callable, Hashable, Sequence +from contextvars import ContextVar, Token +from dataclasses import dataclass, replace +from typing import Any, Protocol + +import torch +import torch_npu +from vllm.logger import logger + +from vllm_ascend.utils import weak_ref_tensors + +Params = dict[str, Any] + + +class ParamProvider(Protocol): + def resolve(self, context) -> Params: ... + + +class ParamSource(Protocol): + def get( + self, + provider: ParamProvider, + ) -> Sequence[Params]: ... + + +@dataclass(frozen=True, slots=True) +class ContextSource: + context: Any + + def get( + self, + provider: ParamProvider, + ) -> Sequence[Params]: + return (provider.resolve(self.context),) + + +@dataclass(frozen=True, slots=True) +class SharedSource: + params: Sequence[Params] + + def get( + self, + _provider: ParamProvider, + ) -> Sequence[Params]: + return self.params + + +_ACTIVE_GRAPH: ContextVar["UpdatableGraph | None"] = ContextVar("capturing_updatable_graph", default=None) + + +@dataclass(slots=True) +class GraphUpdateTask: + operation: Callable[..., Any] + kwargs: dict[str, Any] + provider: ParamProvider + provider_index: int + handle: Any + event: Any + + def bind(self, params: Params) -> "GraphUpdateTask": + runtime_kwargs = {**self.kwargs, **params} + return replace(self, kwargs=runtime_kwargs) + + def apply(self, update_stream) -> None: + torch.npu.graph_task_update_begin(update_stream, self.handle) + self.operation(**self.kwargs) + torch.npu.graph_task_update_end(update_stream) + self.event.record(update_stream) + + +class UpdatableGraph(torch.npu.NPUGraph): + def __init__(self) -> None: + super().__init__() + self.tasks: list[GraphUpdateTask] = [] + self.provider_sizes: dict[ParamProvider, int] = {} + self.capture_resources: dict[Hashable, Any] = {} + self.capture_token: Token[UpdatableGraph | None] | None = None + + def capture_begin(self, pool=None, capture_error_mode: str = "global") -> None: + super().capture_begin(pool=pool, capture_error_mode=capture_error_mode) + assert self.capture_token is None + self.capture_token = _ACTIVE_GRAPH.set(self) + + def capture_end(self) -> None: + try: + super().capture_end() + finally: + assert self.capture_token is not None + _ACTIVE_GRAPH.reset(self.capture_token) + self.capture_token = None + self.capture_resources = weak_ref_tensors(self.capture_resources) + + def get_capture_resource( + self, + key: Hashable, + factory: Callable[[], Any], + use_max_workspace: bool = False, + ) -> Any: + if key not in self.capture_resources: + self.capture_resources[key] = factory() + if use_max_workspace: + # Some models mix attention layer shapes under the same graph size. + # During capture, keep the largest required workspace for that size. + candidate_workspace = factory() + if ( + candidate_workspace.numel() * candidate_workspace.element_size() + > self.capture_resources[key].numel() * self.capture_resources[key].element_size() + ): + self.capture_resources[key] = candidate_workspace + return self.capture_resources[key] + + def register_task( + self, + operation: Callable[..., Any], + kwargs: dict[str, Any], + provider: ParamProvider, + ) -> None: + stream = torch.npu.current_stream() + event = torch.npu.ExternalEvent() + event.wait(stream) + event.reset(stream) + torch.npu.graph_task_group_begin(stream) + operation(**kwargs) + handle = torch.npu.graph_task_group_end(stream) + weak_kwargs = weak_ref_tensors(kwargs) + provider_index = self.provider_sizes.get(provider, 0) + self.provider_sizes[provider] = provider_index + 1 + self.tasks.append( + GraphUpdateTask( + operation, + weak_kwargs, + provider, + provider_index, + handle, + event, + ) + ) + + def resolve_tasks( + self, + source: ParamSource, + ) -> tuple[GraphUpdateTask, ...]: + params_by_provider = {provider: source.get(provider) for provider in self.provider_sizes} + for provider, size in self.provider_sizes.items(): + assert len(params_by_provider[provider]) == size + return tuple(task.bind(params_by_provider[task.provider][task.provider_index]) for task in self.tasks) + + def update( + self, + update_stream, + resolved_tasks: tuple[GraphUpdateTask, ...], + ) -> None: + logger.debug_once("Updating host-side attention metadata with UpdatableGraph.") + with torch.npu.stream(update_stream): + # This is specially designed for PA. + ws_buffer: dict[Hashable, Any] = {} + for task in resolved_tasks: + if task.operation == torch_npu._npu_paged_attention: + ws_key = _get_ws_key(task.kwargs) + if ws_buffer.get(ws_key) is None: + ws_kwargs = task.kwargs.copy() + ws_kwargs.pop("workspace") + workspace = torch_npu._npu_paged_attention_get_workspace(**ws_kwargs) + ws_buffer[ws_key] = workspace + task.kwargs["workspace"] = ws_buffer[ws_key] + task.apply(update_stream) + + +def register_task( + operation: Callable[..., Any], + kwargs: dict[str, Any], + provider: ParamProvider, +) -> None: + graph = _ACTIVE_GRAPH.get() + if graph is None: + operation(**kwargs) + else: + graph.register_task(operation, kwargs, provider) + + +def get_capture_resource( + key: Hashable, + factory: Callable[[], Any], + use_max_workspace: bool = False, +) -> Any: + graph = _ACTIVE_GRAPH.get() + if graph is None: + return factory() + return graph.get_capture_resource(key, factory, use_max_workspace) + + +def _get_ws_key(kwargs): + def _sig(t: Any) -> tuple | None: + if t is None: + return (None, None) + return (t.shape, t.dtype) + + return ( + kwargs["context_lens"].data_ptr(), + tuple(kwargs["context_lens"].shape), + _sig(kwargs["query"]), + _sig(kwargs["key_cache"]), + _sig(kwargs["value_cache"]), + _sig(kwargs["block_table"]), + _sig(kwargs["out"]), + kwargs["num_kv_heads"], + kwargs["num_heads"], + kwargs["scale_value"], + ) diff --git a/vllm_ascend/ops/fused_moe/fused_moe.py b/vllm_ascend/ops/fused_moe/fused_moe.py index 6ce8236f713a..d303dfc4bcd4 100644 --- a/vllm_ascend/ops/fused_moe/fused_moe.py +++ b/vllm_ascend/ops/fused_moe/fused_moe.py @@ -295,43 +295,47 @@ def local_num_experts(self) -> int: def ep_rank(self) -> int: return self.moe_config.ep_rank - def _should_reduce_routed_before_combine(self) -> bool: - """Static layer policy; communication-dependent decisions stay in the op.""" - # Shared DP already produces a complete output, so reduce routed first. - if self._get_shared_expert_parallel_mode() is SharedExpertParallelMode.SHARED_EXPERT_DATA_PARALLEL_ONLY: - return True - # A routed transform must receive the complete TP result. - if self.routed_output_transform is None or self.moe_config.is_sequence_parallel: - return False - return self.moe_config.tp_size > 1 or self.moe_config.ep_size > 1 - - # Shared-expert layout-specific communication is handled by - # AscendSharedExperts, so only standard TP weights need a separate - # all-reduce when routed output has already been reduced. - def _maybe_reduce_shared_expert_output( # type: ignore[misc] + def _reduce_shared_output_if_needed( self, shared_output: torch.Tensor | None, - fused_output_is_reduced: bool | None = None, + fused_output_is_reduced: bool, ) -> torch.Tensor | None: if ( - shared_output is None - or self._get_shared_expert_parallel_mode() is not SharedExpertParallelMode.TENSOR_PARALLEL + shared_output is not None + and fused_output_is_reduced + and self._get_shared_expert_parallel_mode() is SharedExpertParallelMode.TENSOR_PARALLEL ): - return shared_output - # The upstream boolean can be specialized during tracing. Only the - # model-static early-reduction policy may be tested outside the op. - if self._should_reduce_routed_before_combine(): - return tensor_model_parallel_all_reduce(shared_output) - return torch.ops.vllm.maybe_all_reduce_shared_expert(shared_output, self.layer_name) + shared_output = tensor_model_parallel_all_reduce(shared_output) + return shared_output + + def _maybe_reduce_shared_expert_output( # type: ignore[misc] + self, + shared_output: torch.Tensor | None, + fused_output_is_reduced: bool | None = None, + ) -> torch.Tensor | None: + if fused_output_is_reduced is None: + fused_output_is_reduced = self._fused_output_is_reduced + return self._reduce_shared_output_if_needed( + shared_output, + fused_output_is_reduced, + ) def _maybe_reduce_routed_output_before_transform( self, fused_output: torch.Tensor, fused_output_is_reduced: bool, ) -> tuple[torch.Tensor, bool]: - if self._should_reduce_routed_before_combine(): - fused_output = torch.ops.vllm.maybe_all_reduce_tensor_model_parallel(fused_output, self.layer_name) - return fused_output, True + fused_output, fused_output_is_reduced = super()._maybe_reduce_routed_output_before_transform( + fused_output, + fused_output_is_reduced, + ) + + if ( + self._get_shared_expert_parallel_mode() is SharedExpertParallelMode.SHARED_EXPERT_DATA_PARALLEL_ONLY + and not fused_output_is_reduced + ): + fused_output = tensor_model_parallel_all_reduce(fused_output) + fused_output_is_reduced = True return fused_output, fused_output_is_reduced def _maybe_reduce_final_output( # type: ignore[misc] @@ -340,11 +344,13 @@ def _maybe_reduce_final_output( # type: ignore[misc] trunc_size: int | None, output_is_reduced: bool | None = None, ) -> torch.Tensor: - # Do not branch on output_is_reduced: it can describe the tracing - # batch rather than the batch replaying this graph. Early reduction - # and sequence parallelism are static properties of the layer. - if not self.moe_config.is_sequence_parallel and not self._should_reduce_routed_before_combine(): - states = torch.ops.vllm.maybe_all_reduce_tensor_model_parallel(states, self.layer_name) + if output_is_reduced is None: + output_is_reduced = self._fused_output_is_reduced + if not output_is_reduced and not self.moe_config.is_sequence_parallel: + # Use the normal TP collective when the upstream reduction + # contract requires it. Sequence-parallel outputs are token + # shards, so reducing them position-wise would corrupt the result. + states = tensor_model_parallel_all_reduce(states) if trunc_size is not None and trunc_size > 0: return states[..., :trunc_size] return states diff --git a/vllm_ascend/ops/register_custom_ops.py b/vllm_ascend/ops/register_custom_ops.py index a9d7b0d7cfda..193bfacd1f8d 100644 --- a/vllm_ascend/ops/register_custom_ops.py +++ b/vllm_ascend/ops/register_custom_ops.py @@ -9,7 +9,7 @@ from vllm.forward_context import get_forward_context from vllm.utils.torch_utils import direct_register_custom_op -from vllm_ascend.ascend_forward_context import _EXTRA_CTX, MoECommType +from vllm_ascend.ascend_forward_context import _EXTRA_CTX from vllm_ascend.ops.rotary_embedding import rope_forward_oot from vllm_ascend.ops.triton.muls_add import muls_add_triton from vllm_ascend.utils import is_vl_model @@ -124,37 +124,6 @@ def _maybe_pad_and_reduce_impl(x: torch.Tensor) -> torch.Tensor: return ep_group.reduce_scatter(padded_x.view(-1, *x.shape[1:]), 0) -def _routed_output_is_reduced(layer_name: str) -> bool: - runner = get_forward_context().no_compile_layers[layer_name] - is_sequence_parallel = runner.moe_config.is_sequence_parallel - comm = _EXTRA_CTX.moe_comm_type - return comm in { - MoECommType.MC2, - MoECommType.ALLTOALL, - MoECommType.FUSED_MC2, - } or (comm == MoECommType.ALLGATHER and is_sequence_parallel) - - -def _maybe_all_reduce_tensor_model_parallel_impl( - states: torch.Tensor, - layer_name: str, -) -> torch.Tensor: - """Reduce routed/final output only if dispatch has not already reduced it.""" - if _routed_output_is_reduced(layer_name): - return states - return tensor_model_parallel_all_reduce(states) - - -def _maybe_all_reduce_shared_expert_impl( - shared_output: torch.Tensor, - layer_name: str, -) -> torch.Tensor: - """Reduce shared TP output separately when routed output is already reduced.""" - if _routed_output_is_reduced(layer_name): - return tensor_model_parallel_all_reduce(shared_output) - return shared_output - - def _maybe_all_gather_and_maybe_unpad_fake(x: torch.Tensor) -> torch.Tensor: forward_context = get_forward_context() ep_group = get_ep_group() @@ -239,22 +208,6 @@ def _muls_add_impl_fake( dispatch_key="PrivateUse1", ) -direct_register_custom_op( - op_name="maybe_all_reduce_tensor_model_parallel", - op_func=_maybe_all_reduce_tensor_model_parallel_impl, - fake_impl=lambda states, layer_name: states, - mutates_args=[], - dispatch_key="PrivateUse1", -) - -direct_register_custom_op( - op_name="maybe_all_reduce_shared_expert", - op_func=_maybe_all_reduce_shared_expert_impl, - fake_impl=lambda shared_output, layer_name: shared_output, - mutates_args=[], - dispatch_key="PrivateUse1", -) - direct_register_custom_op( op_name="quantize", op_func=_quantize_impl, diff --git a/vllm_ascend/ops/triton/docs/v2/sample/categorical_sample.md b/vllm_ascend/ops/triton/docs/v2/sample/categorical_sample.md deleted file mode 100644 index 4f65764f8710..000000000000 --- a/vllm_ascend/ops/triton/docs/v2/sample/categorical_sample.md +++ /dev/null @@ -1,70 +0,0 @@ -# categorical_sample - -## Description - -- **Function**: Samples one token per input row from the categorical distribution represented by `logits`. It is used as the Ascend NPU replacement for vLLM `gumbel_sample`: `temperature == 0` performs greedy argmax; non-zero temperature performs random categorical sampling. The implementation keeps the existing sampling interface, including optional raw-logit caching for speculative decoding. -- **Formula**: - - Let the effective logit for token `i` be `x_i = logits_i / temperature` when `apply_temperature=True` and `temperature != 0`, otherwise `x_i = logits_i`. - - Greedy mode (`temperature == 0`): `sample = argmax_i x_i`. - - Random mode (`temperature != 0`): `P(sample=i) = exp(x_i) / sum_j exp(x_j)`. - - For numerical stability, each coarse block `b` uses `m_b = max_{i in b} x_i` and `S_b = sum_{i in b} exp(x_i - m_b)`. With `M = max_b m_b`, the global block mass is `W_b = S_b * exp(m_b - M)`. - - A stateless uniform draw `u(seed, pos)` defines `threshold = u * sum_b W_b`. The same threshold is propagated through coarse-block, fine-block, and token-level cumulative masses; no additional random draw is used. -- **Algorithm flow** (processed row by row, independently): - 1. Split each vocabulary row into 8192-element coarse blocks. Flatten `(token, coarse_block)` tasks and distribute them across at most the available Vector Cores. - 2. `_categorical_prepare_mass_kernel` loads each coarse block, optionally stores raw pre-temperature logits into `logits_cache`, applies temperature when requested, computes the block max/argmax, and computes probability masses. Each 8192-element block is further reduced into eight 1024-element fine-block masses. - 3. `_categorical_sample_kernel` processes token rows across at most the available Vector Cores. For random rows, it uses the prepared block masses to select one 8192-element coarse block, then uses the eight prepared fine-block masses to select one 1024-element fine block. - 4. Reload only the selected 1024 logits, apply the same temperature rule, compute the token masses and cumulative sum, and return the token whose cumulative mass crosses the propagated threshold. - 5. For greedy rows, use the prepared block maxima/argmax metadata to return the global argmax. -- **Supported modes**: Atlas A2, Atlas A3, and Ascend 950 - -## Parameters - -> [!NOTE] -> All parameters are required. - -| Parameter | Input/Output/Attribute | Description | Data type | Data format | -| --- | --- | --- | --- | --- | -| `logits` | Input | Logits with shape `[num_tokens, vocab_size]`; vocabulary dimension must be contiguous | fp32 / bf16 | ND | -| `expanded_idx_mapping` | Input | Maps each token row to a request-state index, shape `[num_tokens]` | int32 | ND | -| `temperature` | Input | Per-request temperature, shape `[max_num_reqs]`; `0` selects greedy mode | fp32 | ND | -| `seed` | Input | Per-request stateless RNG seed, shape `[max_num_reqs]` | int64 | ND | -| `pos` | Input | Per-token RNG position, shape `[num_tokens]`; converted to int32 inside the Triton sampling kernel | int32 / int64 | ND | -| `apply_temperature` | Attribute | If `True`, non-zero-temperature rows use `logits / temperature`; if `False`, logits are sampled as provided | bool | scalar | -| `is_drafting` | Attribute | If `True`, adds a fixed RNG-position salt to keep the optional drafting stream separate | bool | scalar | -| `logits_cache` | Input/Output | Optional cache buffer `[max_num_reqs, num_cols, cache_vocab_size]`; stores raw logits before temperature scaling | fp32 / bf16 | ND | -| `logits_cache_col` | Input | Optional cache-column selector; either a scalar or shape `[num_tokens]` | int32 | scalar / ND | -| `use_fp64` | Attribute | Compatibility argument; must be `False` on this NPU implementation | bool | scalar | -| `sampled_token_ids` | Output | Sampled token ID for each input row, shape `[num_tokens]` | int64 | ND | - -## Constraints - -- `logits` must be rank 2 with shape `[num_tokens, vocab_size]`, `vocab_size > 0`, and `logits.stride(-1) == 1`. The current single-operator tests cover fp32 and bf16. -- `expanded_idx_mapping.shape == [num_tokens]`; valid rows index `temperature` and `seed` by request state. A negative mapping is treated as an invalid/padded request by the kernel and must not be used to write `logits_cache`. -- `temperature.shape[0]` and `seed.shape[0]` must cover every non-negative request index referenced by `expanded_idx_mapping`. Temperature values are expected to be non-negative. -- `temperature == 0` is greedy mode. For non-zero temperature, random sampling is categorical. `apply_temperature=False` means the caller has already applied any required temperature scaling. -- `pos.shape == [num_tokens]`. The NPU RNG path converts positions to int32; inference positions must therefore remain within the supported int32 range. -- The vocabulary size does not need to be divisible by 8192 or 1024. Tail elements are masked. The long-vocabulary inference shape `vocab_size=151936` therefore exercises both coarse and fine tail handling. -- `logits_cache`, when provided, must satisfy `logits_cache.size(-1) >= vocab_size`. Cache values are written before temperature scaling. -- `logits_cache_col`, when provided, must be either a 0-D scalar tensor or a tensor with one column index per token. If multiple tokens map to the same `(request, cache_col)` location, their cache writes alias; callers requiring deterministic cache contents must avoid that mapping. -- `use_fp64=True` is not supported and raises `NotImplementedError`. -- Finite logits and `-inf` masking are supported. NaN and `+inf` inputs do not have an additional operator-specific normalization contract. -- The Python wrapper obtains `num_tokens` and `vocab_size` from tensor shape metadata; it does not perform a device-to-host `.item()` synchronization. Runtime task counts are passed to Triton as scalar kernel arguments. -- The implementation introduces no explicit device-to-host synchronization. The single-operator accuracy tests run eagerly; graph-capture behavior is validated by higher-level vLLM/vLLM-Ascend integration tests rather than by this single-operator test. - -## Origin and Differences - -- **Origin**: Replaces the `gumbel_sample` path from `vllm/v1/worker/gpu/sample/gumbel.py` for Ascend NPU. Gumbel-Max and direct categorical sampling represent the same categorical distribution for non-zero temperature. -- **Differences**: - - NPU adaptation for performance: replaces per-vocabulary Gumbel RNG and two logarithms with one stateless uniform draw per sampled row plus hierarchical probability-mass selection. The prepare stage uses 8192-element coarse blocks and stores eight 1024-element fine-block masses so the final sampling stage reloads only one 1024-element block. - - NPU adaptation for performance: flattens `(token, coarse_block)` work and limits launches to the available Vector Core count; the sampling stage similarly distributes token rows over the available Vector Cores. - - Modified for vllm-ascend logic: preserves the existing sampling call contract, including per-request temperature/seed, request-index mapping, optional pre-temperature `logits_cache`, scalar or per-token cache columns, and greedy rows mixed with random rows. - - Random sampling is distribution-equivalent to Gumbel-Max but does not consume random numbers in the same way. Therefore the same `seed`/`pos` is deterministic within `categorical_sample`, but is not required to return the exact same token sequence as the previous Gumbel implementation. - - FP64/fixed-point sampling is not added by this operator; `use_fp64=True` remains unsupported on NPU. - -## Test Cases - -The single-operator tests use the long-vocabulary inference shape `vocab_size=151936`, which is not divisible by either the 8192 coarse block or the 1024 fine block and therefore covers hierarchy-tail masking. Batch/token counts include `1`, `16`, and `64`, matching the operator's long-vocabulary performance validation range. Greedy and cache results are checked bit-exact. Random sampling cannot be validated by exact token equality against Gumbel-Max because the RNG consumption pattern is intentionally different; the tests instead cover deterministic replay, mixed greedy/random rows, exact single-support sampling across hierarchy boundaries, temperature-scaling equivalence, cache semantics, and an equal-mass finite-support distribution sanity check. - -```bash -pytest -sv tests/e2e/nightly/single_node/ops/singlecard_ops/triton/test_categorical_sample.py -``` diff --git a/vllm_ascend/ops/triton/v2/sample/categorical_sample.py b/vllm_ascend/ops/triton/v2/sample/categorical_sample.py deleted file mode 100644 index 4710e9a9021a..000000000000 --- a/vllm_ascend/ops/triton/v2/sample/categorical_sample.py +++ /dev/null @@ -1,443 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 -# SPDX-FileCopyrightText: Copyright contributors to the vLLM project -"""Triton categorical sampling operator for Ascend NPU.""" - -from collections.abc import Callable - -import torch -from vllm.triton_utils import tl, triton - -from vllm_ascend.ops.triton.triton_utils import get_vectorcore_num, init_device_properties_triton -from vllm_ascend.utils import vllm_version_is - -# Hierarchical sampling: prepare 8K coarse-block masses once, then sample -# from one 1K fine block. The two granularities are tuned independently. -_COARSE_BLOCK_SIZE = 8192 -_FINE_BLOCK_SIZE = 1024 -_NUM_FINE_BLOCKS = _COARSE_BLOCK_SIZE // _FINE_BLOCK_SIZE - - -def _get_vectorcore_num() -> int: - try: - return int(get_vectorcore_num()) - except AssertionError: - init_device_properties_triton() - return int(get_vectorcore_num()) - - -# Stage 1: scan logits and build hierarchical probability-mass metadata. -# Stage 2: select coarse/fine blocks from metadata and sample one token. -@triton.jit( - do_not_specialize=[ - "block_argmax_stride", - "block_max_stride", - "block_mass_stride", - "fine_mass_stride_0", - "fine_mass_stride_1", - "logits_cache_stride_0", - "logits_cache_stride_1", - "logits_stride", - "num_tokens", - "vocab_size", - "num_blocks", - ] -) -def _categorical_prepare_mass_kernel( - block_argmax_ptr, - block_argmax_stride, - block_max_ptr, - block_max_stride, - block_mass_ptr, - block_mass_stride, - fine_mass_ptr, - fine_mass_stride_0, - fine_mass_stride_1, - logits_cache_ptr, - logits_cache_stride_0, - logits_cache_stride_1, - logits_cache_col_ptr, - logits_ptr, - logits_stride, - expanded_idx_mapping_ptr, - temp_ptr, - num_tokens, - vocab_size, - num_blocks, - COARSE_BLOCK_SIZE: tl.constexpr, - FINE_BLOCK_SIZE: tl.constexpr, - NUM_FINE_BLOCKS: tl.constexpr, - APPLY_TEMPERATURE: tl.constexpr, - PER_TOKEN_COL: tl.constexpr, -): - worker_id = tl.program_id(0).to(tl.int64) - num_workers = tl.num_programs(0).to(tl.int64) - num_tokens_i64 = num_tokens.to(tl.int64) - num_blocks_i64 = num_blocks.to(tl.int64) - total_tasks = num_tokens_i64 * num_blocks_i64 - tasks_per_worker = total_tasks // num_workers - extra_tasks = total_tasks % num_workers - task_start = worker_id * tasks_per_worker + tl.minimum(worker_id, extra_tasks) - task_count = tasks_per_worker + (worker_id < extra_tasks) - lanes = tl.arange(0, COARSE_BLOCK_SIZE) - fine_block_ids = tl.arange(0, NUM_FINE_BLOCKS) - - for task_idx in tl.range(task_start, task_start + task_count): - token_idx = task_idx // num_blocks_i64 - block_idx = task_idx - token_idx * num_blocks_i64 - req_state_idx = tl.load(expanded_idx_mapping_ptr + token_idx).to(tl.int64) - is_valid_req = req_state_idx >= 0 - temp = tl.load(temp_ptr + req_state_idx, mask=is_valid_req, other=0.0).to(tl.float32) - - offsets = block_idx * COARSE_BLOCK_SIZE + lanes - mask = offsets < vocab_size - logits = tl.load(logits_ptr + token_idx * logits_stride + offsets, mask=mask, other=float("-inf")).to( - tl.float32 - ) - - if logits_cache_ptr is not None: - if logits_cache_col_ptr is not None: - if PER_TOKEN_COL: - col = tl.load(logits_cache_col_ptr + token_idx) - else: - col = tl.load(logits_cache_col_ptr) - else: - col = 0 - tl.store( - logits_cache_ptr + req_state_idx * logits_cache_stride_0 + col * logits_cache_stride_1 + offsets, - logits, - mask=mask & is_valid_req, - ) - - if temp != 0.0 and APPLY_TEMPERATURE: - logits = logits / temp - - block_max, block_argmax = tl.max(logits, axis=0, return_indices=True) - has_mass = block_max > float("-inf") - safe_block_max = tl.where(has_mass, block_max, 0.0) - - weights = tl.where(mask & (temp != 0.0) & has_mass, tl.exp(logits - safe_block_max), 0.0) - weights_2d = tl.reshape(weights, (NUM_FINE_BLOCKS, FINE_BLOCK_SIZE)) - fine_mass_values = tl.sum(weights_2d, axis=1) - block_mass = tl.sum(fine_mass_values, axis=0) - - tl.store( - block_argmax_ptr + token_idx * block_argmax_stride + block_idx, - block_idx * COARSE_BLOCK_SIZE + block_argmax, - ) - tl.store(block_max_ptr + token_idx * block_max_stride + block_idx, block_max) - tl.store(block_mass_ptr + token_idx * block_mass_stride + block_idx, block_mass) - tl.store( - fine_mass_ptr + token_idx * fine_mass_stride_0 + block_idx * fine_mass_stride_1 + fine_block_ids, - fine_mass_values, - ) - - -@triton.jit( - do_not_specialize=[ - "block_argmax_stride", - "block_max_stride", - "block_mass_stride", - "fine_mass_stride_0", - "fine_mass_stride_1", - "logits_stride", - "num_tokens", - "vocab_size", - "num_blocks", - ] -) -def _categorical_sample_kernel( - sampled_ptr, - block_argmax_ptr, - block_argmax_stride, - block_max_ptr, - block_max_stride, - block_mass_ptr, - block_mass_stride, - fine_mass_ptr, - fine_mass_stride_0, - fine_mass_stride_1, - logits_ptr, - logits_stride, - expanded_idx_mapping_ptr, - seeds_ptr, - pos_ptr, - temp_ptr, - num_tokens, - vocab_size, - num_blocks, - COARSE_BLOCK_SIZE: tl.constexpr, - FINE_BLOCK_SIZE: tl.constexpr, - NUM_FINE_BLOCKS: tl.constexpr, - PADDED_NUM_BLOCKS: tl.constexpr, - APPLY_TEMPERATURE: tl.constexpr, - IS_DRAFTING: tl.constexpr, -): - worker_id = tl.program_id(0).to(tl.int64) - num_workers = tl.num_programs(0).to(tl.int64) - num_tokens_i64 = num_tokens.to(tl.int64) - tokens_per_worker = num_tokens_i64 // num_workers - extra_tokens = num_tokens_i64 % num_workers - token_start = worker_id * tokens_per_worker + tl.minimum(worker_id, extra_tokens) - token_count = tokens_per_worker + (worker_id < extra_tokens) - - for token_idx in tl.range(token_start, token_start + token_count): - req_state_idx = tl.load(expanded_idx_mapping_ptr + token_idx).to(tl.int64) - is_valid_req = req_state_idx >= 0 - temp = tl.load(temp_ptr + req_state_idx, mask=is_valid_req, other=0.0).to(tl.float32) - is_random = temp != 0.0 - - block_ids = tl.arange(0, PADDED_NUM_BLOCKS) - valid_block_mask = block_ids < num_blocks - block_max = tl.load( - block_max_ptr + token_idx * block_max_stride + block_ids, - mask=valid_block_mask, - other=float("-inf"), - ).to(tl.float32) - greedy_block = tl.argmax(block_max, axis=0) - greedy_token = tl.load(block_argmax_ptr + token_idx * block_argmax_stride + greedy_block) - - stored_mass = tl.load( - block_mass_ptr + token_idx * block_mass_stride + block_ids, - mask=valid_block_mask & is_random, - other=0.0, - ).to(tl.float32) - global_max = tl.max( - tl.where(valid_block_mask & is_random, block_max, float("-inf")), - axis=0, - ) - safe_global_max = tl.where(global_max > float("-inf"), global_max, 0.0) - block_mass = stored_mass * tl.exp(block_max - safe_global_max) - block_mass = tl.where(valid_block_mask & is_random, block_mass, 0.0) - total_mass = tl.sum(block_mass, axis=0) - has_total_mass = total_mass > 0.0 - - seed = tl.load(seeds_ptr + req_state_idx, mask=is_valid_req & is_random, other=0) - position = tl.load(pos_ptr + token_idx).to(tl.int32) - if IS_DRAFTING: - DRAFT_NOISE_SALT: tl.constexpr = 1 << 30 - position += DRAFT_NOISE_SALT - uniform = tl.max( - tl.rand(tl.randint(seed, position), tl.arange(0, 1)).to(tl.float32), - axis=0, - ) - threshold = uniform * total_mass - - block_prefix = tl.cumsum(block_mass, axis=0) - candidate_blocks = tl.where( - (block_prefix > threshold) & valid_block_mask & has_total_mass, - block_ids, - PADDED_NUM_BLOCKS, - ) - selected_block = tl.minimum(tl.min(candidate_blocks, axis=0), num_blocks - 1) - prefix_before_block = tl.sum( - tl.where(valid_block_mask & (block_ids < selected_block), block_mass, 0.0), - axis=0, - ) - block_threshold = threshold - prefix_before_block - - # Use prepared 1K fine-block masses to avoid rescanning the selected 8K block. - fine_block_ids = tl.arange(0, NUM_FINE_BLOCKS) - stored_fine_mass = tl.load( - fine_mass_ptr + token_idx * fine_mass_stride_0 + selected_block * fine_mass_stride_1 + fine_block_ids, - mask=is_random & has_total_mass, - other=0.0, - ).to(tl.float32) - selected_block_max = tl.load(block_max_ptr + token_idx * block_max_stride + selected_block).to(tl.float32) - block_scale = tl.exp(selected_block_max - safe_global_max) - fine_mass_values = tl.where( - is_random & has_total_mass, - stored_fine_mass * block_scale, - 0.0, - ) - fine_prefix = tl.cumsum(fine_mass_values, axis=0) - candidate_fine_blocks = tl.where( - (fine_prefix > block_threshold) & is_random & has_total_mass, - fine_block_ids, - NUM_FINE_BLOCKS, - ) - selected_fine_block = tl.minimum(tl.min(candidate_fine_blocks, axis=0), NUM_FINE_BLOCKS - 1) - prefix_before_fine = tl.sum( - tl.where(fine_block_ids < selected_fine_block, fine_mass_values, 0.0), - axis=0, - ) - token_threshold = block_threshold - prefix_before_fine - - fine_offsets = tl.arange(0, FINE_BLOCK_SIZE) - fine_base = selected_block * COARSE_BLOCK_SIZE + selected_fine_block * FINE_BLOCK_SIZE - token_ids = fine_base + fine_offsets - token_mask = token_ids < vocab_size - active_token_mask = token_mask & is_random & has_total_mass - logits = tl.load( - logits_ptr + token_idx * logits_stride + token_ids, mask=active_token_mask, other=float("-inf") - ).to(tl.float32) - if APPLY_TEMPERATURE: - safe_temp = tl.where(is_random, temp, 1.0) - logits = logits / safe_temp - - token_mass = tl.where( - active_token_mask, - tl.exp(logits - safe_global_max), - 0.0, - ) - token_prefix = tl.cumsum(token_mass, axis=0) - candidate_offsets = tl.where( - (token_prefix > token_threshold) & token_mask & has_total_mass, - fine_offsets, - FINE_BLOCK_SIZE, - ) - selected_offset = tl.min(candidate_offsets, axis=0) - fallback_offset = tl.max( - tl.where(token_mask & (token_mass > 0.0), fine_offsets, 0), - axis=0, - ) - selected_offset = tl.where( - selected_offset < FINE_BLOCK_SIZE, - selected_offset, - fallback_offset, - ) - categorical_token = fine_base + selected_offset - categorical_token = tl.where(has_total_mass, categorical_token, 0) - sampled_token = tl.where(is_random, categorical_token, greedy_token) - tl.store(sampled_ptr + token_idx, sampled_token) - - -def _categorical_sample( - logits: torch.Tensor, - expanded_idx_mapping: torch.Tensor, - temperature: torch.Tensor, - seed: torch.Tensor, - pos: torch.Tensor, - apply_temperature: bool, - logits_cache: torch.Tensor | None = None, - logits_cache_col: torch.Tensor | None = None, - use_fp64: bool = False, - *, - is_drafting: bool = False, -) -> torch.Tensor: - """Sample token ids from logits with categorical sampling. - - This internal entry keeps the v0.28 positional argument order. On newer - vLLM versions the public wrapper exposes `is_drafting` immediately after - `apply_temperature`, matching the upstream gumbel_sample contract. - """ - if use_fp64: - raise NotImplementedError("FP64 categorical sampling is not supported on NPU.") - - expanded_idx_mapping = expanded_idx_mapping.contiguous() - pos = pos.contiguous() - if logits_cache_col is not None: - logits_cache_col = logits_cache_col.contiguous() - - num_tokens, vocab_size = logits.shape - if logits_cache is not None: - assert logits_cache.size(-1) >= vocab_size, ( - f"draft logits cache vocab dim ({logits_cache.size(-1)}) is narrower " - f"than the sampled logits ({vocab_size}). Cached logits would be truncated." - ) - if num_tokens == 0: - return torch.empty(0, dtype=torch.int64, device=logits.device) - - num_blocks = triton.cdiv(vocab_size, _COARSE_BLOCK_SIZE) - padded_num_blocks = triton.next_power_of_2(num_blocks) - # Metadata produced once by the prepare kernel and consumed by the sample kernel. - block_argmax_workspace = torch.empty(num_tokens, num_blocks, dtype=torch.int64, device=logits.device) - block_max_workspace = torch.empty(num_tokens, num_blocks, dtype=torch.float32, device=logits.device) - block_mass_workspace = torch.empty(num_tokens, num_blocks, dtype=torch.float32, device=logits.device) - fine_mass = torch.empty(num_tokens, num_blocks, _NUM_FINE_BLOCKS, dtype=torch.float32, device=logits.device) - sampled = torch.empty(num_tokens, dtype=torch.int64, device=logits.device) - per_token_col = logits_cache_col is not None and logits_cache_col.dim() > 0 - - total_tasks = num_tokens * num_blocks - num_workers = min(_get_vectorcore_num(), total_tasks) - - _categorical_prepare_mass_kernel[(num_workers,)]( - block_argmax_workspace, - block_argmax_workspace.stride(0), - block_max_workspace, - block_max_workspace.stride(0), - block_mass_workspace, - block_mass_workspace.stride(0), - fine_mass, - fine_mass.stride(0), - fine_mass.stride(1), - logits_cache, - logits_cache.stride(0) if logits_cache is not None else 0, - logits_cache.stride(1) if logits_cache is not None else 0, - logits_cache_col, - logits, - logits.stride(0), - expanded_idx_mapping, - temperature, - num_tokens, - vocab_size, - num_blocks, - COARSE_BLOCK_SIZE=_COARSE_BLOCK_SIZE, - FINE_BLOCK_SIZE=_FINE_BLOCK_SIZE, - NUM_FINE_BLOCKS=_NUM_FINE_BLOCKS, - APPLY_TEMPERATURE=apply_temperature, - PER_TOKEN_COL=per_token_col, - ) - - sample_workers = min(_get_vectorcore_num(), num_tokens) - _categorical_sample_kernel[(sample_workers,)]( - sampled, - block_argmax_workspace, - block_argmax_workspace.stride(0), - block_max_workspace, - block_max_workspace.stride(0), - block_mass_workspace, - block_mass_workspace.stride(0), - fine_mass, - fine_mass.stride(0), - fine_mass.stride(1), - logits, - logits.stride(0), - expanded_idx_mapping, - seed, - pos, - temperature, - num_tokens, - vocab_size, - num_blocks, - COARSE_BLOCK_SIZE=_COARSE_BLOCK_SIZE, - FINE_BLOCK_SIZE=_FINE_BLOCK_SIZE, - NUM_FINE_BLOCKS=_NUM_FINE_BLOCKS, - PADDED_NUM_BLOCKS=padded_num_blocks, - APPLY_TEMPERATURE=apply_temperature, - IS_DRAFTING=is_drafting, - ) - return sampled - - -categorical_sample: Callable[..., torch.Tensor] -if vllm_version_is("0.28.0"): - # Preserve the legacy positional order; vLLM #54282 inserted is_drafting on main. - categorical_sample = _categorical_sample -else: - - def _categorical_sample_main( - logits: torch.Tensor, - expanded_idx_mapping: torch.Tensor, - temperature: torch.Tensor, - seed: torch.Tensor, - pos: torch.Tensor, - apply_temperature: bool, - is_drafting: bool, - logits_cache: torch.Tensor | None = None, - logits_cache_col: torch.Tensor | None = None, - use_fp64: bool = False, - ) -> torch.Tensor: - return _categorical_sample( - logits, - expanded_idx_mapping, - temperature, - seed, - pos, - apply_temperature, - logits_cache, - logits_cache_col, - use_fp64, - is_drafting=is_drafting, - ) - - categorical_sample = _categorical_sample_main diff --git a/vllm_ascend/patch/worker/patch_v2/patch_triton.py b/vllm_ascend/patch/worker/patch_v2/patch_triton.py index 76a6c2edf329..431a25aaf19a 100644 --- a/vllm_ascend/patch/worker/patch_v2/patch_triton.py +++ b/vllm_ascend/patch/worker/patch_v2/patch_triton.py @@ -20,7 +20,6 @@ from vllm_ascend.ops.triton.v2.apply_grammar_bitmask import _apply_grammar_bitmask_kernel from vllm_ascend.ops.triton.v2.mamba.precopy import precopy_mamba_align_fused_kernel from vllm_ascend.ops.triton.v2.metrics.num_nans import get_num_nans -from vllm_ascend.ops.triton.v2.sample.categorical_sample import categorical_sample from vllm_ascend.ops.triton.v2.sample.fill_logprob_token_idx import _fill_logprob_token_ids_kernel from vllm_ascend.ops.triton.v2.sample.thinking_budget import ( _load_effective_token_ascend, @@ -28,7 +27,7 @@ ) from vllm_ascend.worker.v2.sample.apply_top_k_top_p import apply_top_k_top_p_npu from vllm_ascend.worker.v2.sample.bad_words import apply_bad_words -from vllm_ascend.worker.v2.sample.gumbel import apply_temperature +from vllm_ascend.worker.v2.sample.gumbel import apply_temperature, gumbel_sample from vllm_ascend.worker.v2.sample.logprob import compute_token_logprobs, compute_topk_logprobs from vllm_ascend.worker.v2.sample.min_p import apply_min_p from vllm_ascend.worker.v2.sample.penalties import apply_penalties, bincount @@ -40,12 +39,17 @@ # triton ops that need to be filed in ops/triton penalties.apply_penalties = apply_penalties # because sampler.py and speculator.py are imported before this patch, they must be overridden +sampler.gumbel_sample = gumbel_sample prompt_logprob.compute_topk_logprobs = compute_topk_logprobs sampler.compute_topk_logprobs = compute_topk_logprobs rejection_sampler.compute_topk_logprobs = compute_topk_logprobs states.apply_min_p = apply_min_p penalties.bincount = bincount +speculator.gumbel_sample = gumbel_sample +base_speculator.gumbel_sample = gumbel_sample +dspark_speculator.gumbel_sample = gumbel_sample bad_words.apply_bad_words = apply_bad_words +gumbel.gumbel_sample = gumbel_sample gumbel.apply_temperature = apply_temperature states.apply_temperature = apply_temperature logprob.compute_token_logprobs = compute_token_logprobs @@ -53,11 +57,6 @@ rejection_sampler.rejection_sample = npu_rejection_sample dflash_speculator._prepare_dflash_inputs_kernel = _prepare_dflash_inputs_kernel_ascend # triton ops that filed in ops/triton -gumbel.gumbel_sample = categorical_sample -speculator.gumbel_sample = categorical_sample -base_speculator.gumbel_sample = categorical_sample -dspark_speculator.gumbel_sample = categorical_sample -sampler.gumbel_sample = categorical_sample topk_topp_sampler.apply_top_k_top_p_triton = apply_top_k_top_p_npu structured_outputs._apply_grammar_bitmask_kernel = _apply_grammar_bitmask_kernel mamba_utils.precopy_mamba_align_fused_kernel = precopy_mamba_align_fused_kernel diff --git a/vllm_ascend/spec_decode/llm_base_proposer.py b/vllm_ascend/spec_decode/llm_base_proposer.py index 78d3b3d03680..ffa88844bd4b 100644 --- a/vllm_ascend/spec_decode/llm_base_proposer.py +++ b/vllm_ascend/spec_decode/llm_base_proposer.py @@ -64,7 +64,7 @@ _maybe_eager_context, patch_tensor_parallel_group, ) -from vllm_ascend.utils import check_gdn_layer, enable_sp, lmhead_tp_enable, vllm_version_is +from vllm_ascend.utils import check_gdn_layer, enable_sp, lmhead_tp_enable, use_updatable_graph, vllm_version_is from vllm_ascend.worker.device_metadata import DeviceMetadataTask, DeviceMetadataTaskProvider # Currently we will fix block size to a small one since `num_reqs` can't be too large @@ -277,7 +277,7 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device, pass_hidden_st # since final block table tensor is not ready in __init__, it is delayed until dummy_run self.block_table_tensor_clone: torch.Tensor | None = None - self._runnable = self._run_merged_draft + self._runnable: Any = self._run_merged_draft if self.uses_mrope: num_dims = 3 if vllm_version_is("0.28.0") else self.draft_model_config.mrope_num_dims self.mrope_positions = torch.zeros((num_dims, self.max_num_tokens + 1), dtype=torch.int32, device=device) @@ -632,6 +632,26 @@ def _maybe_share_lm_head(self, model: nn.Module) -> None: enable_enpu=self.enable_enpu, ) + def set_update_stream(self, update_stream): + if hasattr(self._runnable, "set_update_stream"): + self._runnable.set_update_stream(update_stream) + self.update_stream = update_stream + + def _maybe_update_metadata(self, att_backend, multi_steps_attn_metadata): + if use_updatable_graph(att_backend): + update_params = [] + for per_layer_metadata in multi_steps_attn_metadata: + metadata = next(iter(per_layer_metadata.values())) + update_params.append( + { + "actual_seq_lengths": metadata.actual_seq_lengths_q, + "actual_seq_lengths_kv": metadata.seq_lens_list, + "block_table": metadata.block_tables, + } + ) + self._runnable.update_draft_model_metadata(update_params) # type: ignore + self._runnable.set_attn_backend(att_backend) # type: ignore + def _maybe_share_topk_indices(self, target_language_model: nn.Module) -> None: if hasattr(target_language_model.model, "topk_indices_buffer"): if hasattr(self.model.model, "topk_indices_buffer"): @@ -827,6 +847,12 @@ def dummy_run( self.token_indices_to_sample.fill_(0) + if aclgraph_runtime_mode == CUDAGraphMode.FULL: + self._maybe_update_metadata( + self.draft_attn_groups[0].backend, + multi_steps_attn_metadata, + ) + with set_ascend_forward_context( multi_steps_attn_metadata[0] if multi_steps_attn_metadata else None, self.vllm_config, @@ -1193,6 +1219,12 @@ def _propose( self.token_indices_to_sample[:token_indices_to_sample_len].copy_(token_indices_to_sample) self.token_indices_to_sample[token_indices_to_sample_len:].fill_(0) + if aclgraph_runtime_mode == CUDAGraphMode.FULL: + self._maybe_update_metadata( + self.draft_attn_groups[0].backend, + multi_steps_attn_metadata, + ) + active_device_metadata_executor = ( getattr(self.runner, "device_metadata_executor", None) if self.method == "dspark" else None ) diff --git a/vllm_ascend/utils.py b/vllm_ascend/utils.py index b6e5682c94d0..f4636eb7f110 100644 --- a/vllm_ascend/utils.py +++ b/vllm_ascend/utils.py @@ -1028,18 +1028,16 @@ def weak_ref_tensor(tensor: Any) -> Any: The new tensor will share the same data as the original tensor, but will not keep the original tensor alive. """ - if isinstance(tensor, torch.Tensor): + if isinstance(tensor, torch.Tensor) and tensor.device.type == "npu": return torch_npu._C._weak_ref_tensor(tensor) else: return tensor -def weak_ref_tensors( - tensors: torch.Tensor | list[torch.Tensor] | tuple[torch.Tensor], -) -> torch.Tensor | list[Any] | tuple[Any] | Any: +def weak_ref_tensors(tensors: Any) -> Any: """ - Convenience function to create weak references to tensors, - for single tensor, list of tensors or tuple of tensors. + Recursively replace tensors with weak references while preserving containers + and non-tensor values. This function should be used in the following scenario: When a tensor is created during graph capture, and it's held by a method @@ -1051,14 +1049,14 @@ def weak_ref_tensors( if isinstance(tensors, torch.Tensor): return weak_ref_tensor(tensors) if isinstance(tensors, list): - return [weak_ref_tensor(t) for t in tensors] + return [weak_ref_tensors(tensor) for tensor in tensors] if isinstance(tensors, tuple): - return tuple(weak_ref_tensor(t) for t in tensors) - # For IntermediateTensors used in pipeline parallelism + return tuple(weak_ref_tensors(tensor) for tensor in tensors) + if isinstance(tensors, dict): + return {key: weak_ref_tensors(tensor) for key, tensor in tensors.items()} if isinstance(tensors, IntermediateTensors): - ret = IntermediateTensors({key: weak_ref_tensor(val) for key, val in tensors.tensors.items()}) - return ret - raise ValueError("Invalid type for tensors") + return IntermediateTensors(weak_ref_tensors(tensors.tensors)) + return tensors def npu_stream_switch(target_stream: torch.npu.Stream, *, enabled: bool = True): @@ -1732,3 +1730,11 @@ def get_rotation_matrix(rotation_path: Path | None) -> torch.Tensor: rotation_path, ) raise e + + +def use_updatable_graph( + attn_backend, +) -> bool: + from vllm_ascend.attention.attention_v1 import AscendAttentionBackend + + return attn_backend is not None and issubclass(attn_backend, AscendAttentionBackend) diff --git a/vllm_ascend/worker/model_runner_v1.py b/vllm_ascend/worker/model_runner_v1.py index 63c205950ad6..d113dca35d02 100644 --- a/vllm_ascend/worker/model_runner_v1.py +++ b/vllm_ascend/worker/model_runner_v1.py @@ -3068,6 +3068,9 @@ def _model_forward( assert self.model is not None forward_context = get_forward_context() assert forward_context is not None + if forward_context.cudagraph_runtime_mode == CUDAGraphMode.FULL: + if hasattr(self.model, "set_attn_backend"): + self.model.set_attn_backend(self.attn_backend) model_inputs: dict[str, Any] = { "input_ids": input_ids, @@ -4200,7 +4203,10 @@ def mock_pass(param1, param2): self.update_stream = torch.npu.Stream() if self.drafter is not None: - self.drafter.update_stream = self.update_stream + if hasattr(self.drafter, "set_update_stream"): + self.drafter.set_update_stream(self.update_stream) + else: + self.drafter.update_stream = self.update_stream with _torch_cuda_wrapper(): if ( @@ -4228,6 +4234,7 @@ def mock_pass(param1, param2): runtime_mode=CUDAGraphMode.FULL, use_eagle=self.use_eagle, enable_enpu=self.enable_enpu, + update_stream=self.update_stream, ) if self.compilation_config.cudagraph_mode != CUDAGraphMode.NONE: diff --git a/vllm_ascend/worker/v2/aclgraph_utils.py b/vllm_ascend/worker/v2/aclgraph_utils.py index c4579a1d068c..9563ddea28fb 100644 --- a/vllm_ascend/worker/v2/aclgraph_utils.py +++ b/vllm_ascend/worker/v2/aclgraph_utils.py @@ -38,9 +38,16 @@ from vllm.v1.worker.utils import AttentionGroup from vllm_ascend.ascend_forward_context import _EXTRA_CTX -from vllm_ascend.compilation.acl_graph import set_graph_params, update_full_graph_params +from vllm_ascend.compilation.acl_graph import ( + set_graph_params, + update_full_graph_params, +) from vllm_ascend.compilation.breakable_aclgraph import BreakableACLGraphWrapper -from vllm_ascend.utils import vllm_version_is +from vllm_ascend.compilation.updatable_graph import ( + ContextSource, + UpdatableGraph, +) +from vllm_ascend.utils import use_updatable_graph, vllm_version_is from vllm_ascend.worker.v2.input_batch import AscendInputBatch from vllm_ascend.worker.v2.utils import communicator_switch @@ -148,6 +155,17 @@ def run_fullgraph(self, desc: BatchExecutionDescriptor) -> torch.Tensor | tuple[ num_tokens = desc.num_tokens logger.info_once("run_fullgraph with num_tokens=%s", num_tokens) assert self.update_stream is not None + with set_current_vllm_config(self.vllm_config): + attn_backend = _get_graph_update_backend(self.model_runner.attn_groups) + attn_metadata = self.model_runner.model_state.attn_metadata + + if use_updatable_graph(attn_backend): + return self._updatable_graph_replay(desc, attn_metadata) + else: + # This will be removed once the refactoring is fully complete. + return self._graph_relay(attn_backend, desc, num_tokens, attn_metadata) + + def _graph_relay(self, attn_backend, desc, num_tokens, attn_metadata): self.update_stream.wait_stream(torch.npu.current_stream()) ret = super().run_fullgraph(desc) @@ -163,7 +181,7 @@ def run_fullgraph(self, desc: BatchExecutionDescriptor) -> torch.Tensor | tuple[ with ( set_current_vllm_config(self.vllm_config), set_forward_context( - self.model_runner.model_state.attn_metadata, + attn_metadata, self.vllm_config, num_tokens=num_tokens, cudagraph_runtime_mode=desc.cg_mode, @@ -173,7 +191,6 @@ def run_fullgraph(self, desc: BatchExecutionDescriptor) -> torch.Tensor | tuple[ ), ): forward_context = get_forward_context() - attn_backend = _get_graph_update_backend(self.model_runner.attn_groups) update_full_graph_params( # FIXME(Ronald1995): support hybrid attn backend attn_backend, @@ -185,6 +202,15 @@ def run_fullgraph(self, desc: BatchExecutionDescriptor) -> torch.Tensor | tuple[ ) return ret + def _updatable_graph_replay(self, desc, attn_metadata): + graph = self.graphs[desc] + assert isinstance(graph, UpdatableGraph) + resolved_tasks = graph.resolve_tasks(ContextSource(attn_metadata)) + self.update_stream.wait_stream(torch.npu.current_stream()) + ret = super().run_fullgraph(desc) + graph.update(self.update_stream, resolved_tasks) + return ret + def capture( self, model: nn.Module, diff --git a/vllm_ascend/worker/v2/spec_decode/autoregressive/aclgraph.py b/vllm_ascend/worker/v2/spec_decode/autoregressive/aclgraph.py index 5fd09b09435a..422d6aff7659 100644 --- a/vllm_ascend/worker/v2/spec_decode/autoregressive/aclgraph.py +++ b/vllm_ascend/worker/v2/spec_decode/autoregressive/aclgraph.py @@ -26,6 +26,11 @@ set_draft_graph_prefill_params, update_full_graph_params, ) +from vllm_ascend.compilation.updatable_graph import ( + SharedSource, + UpdatableGraph, +) +from vllm_ascend.utils import use_updatable_graph from vllm_ascend.worker.v2.aclgraph_utils import ( collect_sorted_captured_token_sizes, model_capture_wrapper, @@ -133,28 +138,38 @@ def run_fullgraph(self, desc: BatchExecutionDescriptor) -> torch.Tensor | tuple[ ) else: logger.info_once("AutoRegressiveAclGraphManager: draft run_fullgraph with num_tokens=%s", num_tokens) + assert self.update_stream is not None + + attn_backend = self.speculator.attn_backend + draft_vllm_config = self.speculator.draft_vllm_config + if use_updatable_graph(attn_backend): + return self._updatable_graph_replay(desc) + else: + # This will be removed once the refactoring is fully complete. + return self._graph_replay(desc, attn_backend, num_tokens, draft_vllm_config) + + def _graph_replay(self, desc, attn_backend, num_tokens, draft_vllm_config): + self.update_stream.wait_stream(torch.npu.current_stream()) + ret = super().run_fullgraph(desc) + # Mirror vLLM's DP graph-replay token-count metadata. + num_tokens_across_dp = torch.full([self.speculator.dp_size], num_tokens) + attn_metadata = self.speculator.model_state.attn_metadata draft_attn_metadatas = self.speculator.build_draft_attn_metadatas( desc.num_reqs, desc.num_tokens, self.is_draft_model_prefill, ) - self.update_stream.wait_stream(torch.npu.current_stream()) - ret = super().run_fullgraph(desc) - - # Mirror vLLM's DP graph-replay token-count metadata. - num_tokens_across_dp = torch.full([self.speculator.dp_size], num_tokens) # sfa_v1.py:AscendSFABackend.get_impl_cls reaches # sfa_cp.py:resolve_sfa_impl, whose SFA CP selector reads the current # ModelConfig. Publish the draft config because set_forward_context() # does not update it. # TODO: Remove this explicit current-config scope once ACL graph replay # passes VllmConfig directly through the graph-update interfaces. - draft_vllm_config = self.speculator.draft_vllm_config with ( set_current_vllm_config(draft_vllm_config), set_forward_context( - self.speculator.model_state.attn_metadata, + attn_metadata, draft_vllm_config, num_tokens=num_tokens, cudagraph_runtime_mode=desc.cg_mode, @@ -168,7 +183,6 @@ def run_fullgraph(self, desc: BatchExecutionDescriptor) -> torch.Tensor | tuple[ _EXTRA_CTX.is_draft_model_prefill = self.is_draft_model_prefill forward_context = get_forward_context() - attn_backend = self.speculator.attn_backend assert attn_backend is not None, "Speculator attention backend is not initialized" update_full_graph_params( # FIXME(Ronald1995): support hybrid attn backend @@ -181,3 +195,16 @@ def run_fullgraph(self, desc: BatchExecutionDescriptor) -> torch.Tensor | tuple[ draft_attn_metadatas=draft_attn_metadatas, ) return ret + + def _updatable_graph_replay(self, desc): + graph = self.graphs[desc] + assert isinstance(graph, UpdatableGraph) + fia_params = self.speculator.build_fia_params( + desc.num_reqs, + self.is_draft_model_prefill, + ) + resolved_tasks = graph.resolve_tasks(SharedSource(fia_params)) + self.update_stream.wait_stream(torch.npu.current_stream()) + ret = super().run_fullgraph(desc) + graph.update(self.update_stream, resolved_tasks) + return ret diff --git a/vllm_ascend/worker/v2/spec_decode/autoregressive/speculator.py b/vllm_ascend/worker/v2/spec_decode/autoregressive/speculator.py index 7f3304129a5d..2a2e33610a95 100644 --- a/vllm_ascend/worker/v2/spec_decode/autoregressive/speculator.py +++ b/vllm_ascend/worker/v2/spec_decode/autoregressive/speculator.py @@ -216,6 +216,8 @@ def init_cudagraph_manager(self, cudagraph_mode: CUDAGraphMode) -> None: if self.speculative_config.enforce_eager: cudagraph_mode = CUDAGraphMode.NONE super().init_cudagraph_manager(cudagraph_mode) + assert self.prefill_cudagraph_manager is not None + assert self.decode_cudagraph_manager is not None # The Ascend graph managers are patched onto the upstream module and # created by super().init_cudagraph_manager without a speculator ref. # They need this speculator to update full-graph params, so set it here. @@ -611,6 +613,46 @@ def _update_decode_attn_metadata( decode_metadata.actual_seq_lengths_q = query_lens_list metadata.seq_lens_cpu.copy_(next_seq_lens_cpu) + def build_fia_params( + self, + num_reqs_padded: int, + is_draft_model_prefill: bool, + ) -> list[dict[str, Any]]: + metadata = next( + metadata + for layer_name, metadata in self.model_state.attn_metadata.items() + if layer_name in self.draft_attn_layer_names + ) + block_table = metadata.block_tables + if block_table is not None and block_table.shape[0] < num_reqs_padded: + block_table = block_table.as_strided((num_reqs_padded, block_table.shape[1]), block_table.stride()) + + if is_draft_model_prefill: + return [ + { + "actual_seq_lengths": metadata.actual_seq_lengths_q, + "actual_seq_lengths_kv": metadata.seq_lens_list, + "block_table": block_table, + } + ] + assert self.input_batch is not None + num_reqs = self.input_batch.num_reqs + query_start_loc = list(range(1, num_reqs_padded + 1)) + fia_params: list[dict[str, Any]] = [] + for step in range(1, self.num_speculative_steps): + seq_lens = [ + min(int(seq_len) + step, self.max_model_len) for seq_len in self.input_batch.seq_lens_np[:num_reqs] + ] + seq_lens.extend([0] * (num_reqs_padded - num_reqs)) + fia_params.append( + { + "actual_seq_lengths": query_start_loc, + "actual_seq_lengths_kv": seq_lens, + "block_table": block_table, + } + ) + return fia_params + def _calc_next_seq_lens_cpu(self, seq_lens_cpu, num_reqs, num_reqs_padded, step): # NOTE(drslark) to achieve fully alignment with vllm, `num_rejected` should be subtracted from `seq_lens` # to avoid extra sync overhead, `v2` is currently aligned with NPU `v1` only diff --git a/vllm_ascend/worker/v2/spec_decode/dflash/aclgraph.py b/vllm_ascend/worker/v2/spec_decode/dflash/aclgraph.py index 4a6295f82bda..42407638c940 100644 --- a/vllm_ascend/worker/v2/spec_decode/dflash/aclgraph.py +++ b/vllm_ascend/worker/v2/spec_decode/dflash/aclgraph.py @@ -19,6 +19,11 @@ set_draft_graph_params, update_full_graph_params, ) +from vllm_ascend.compilation.updatable_graph import ( + ContextSource, + UpdatableGraph, +) +from vllm_ascend.utils import use_updatable_graph from vllm_ascend.worker.v2.aclgraph_utils import collect_sorted_captured_token_sizes, model_capture_wrapper from vllm_ascend.worker.v2.utils import communicator_switch @@ -81,14 +86,20 @@ def capture( def run_fullgraph(self, desc: BatchExecutionDescriptor) -> torch.Tensor | tuple[torch.Tensor, list[torch.Tensor]]: """Override run_fullgraph to update full graph params in run_fullgraph.""" num_tokens = desc.num_tokens - + attn_backend = list(self.speculator.attn_backends.values())[0] draft_attn_metadatas = self.speculator.build_draft_attn_metadatas( desc.num_reqs, self.speculator.input_batch.seq_lens_cpu_upper_bound, ) + if use_updatable_graph(attn_backend): + return self._updatable_graph_replay(desc, draft_attn_metadatas) + else: + # This will be removed once the refactoring is fully complete. + return self._graph_replay(desc, attn_backend, num_tokens, draft_attn_metadatas) + + def _graph_replay(self, desc, attn_backend, num_tokens, draft_attn_metadatas): self.update_stream.wait_stream(torch.npu.current_stream()) ret = super().run_fullgraph(desc) - # refer to vllm.v1.worker.gpu.dp_utils.sync_cudagraph_and_dp_padding to # calculate num_tokens_across_dp. num_tokens_across_dp = torch.full([self.speculator.dp_size], num_tokens) @@ -107,7 +118,7 @@ def run_fullgraph(self, desc: BatchExecutionDescriptor) -> torch.Tensor | tuple[ forward_context = get_forward_context() update_full_graph_params( # FIXME(Ronald1995): support hybrid attn backend - list(self.speculator.attn_backends.values())[0], + attn_backend, self.update_stream, forward_context, num_tokens, @@ -116,3 +127,12 @@ def run_fullgraph(self, desc: BatchExecutionDescriptor) -> torch.Tensor | tuple[ draft_attn_metadatas=draft_attn_metadatas, ) return ret + + def _updatable_graph_replay(self, desc, draft_attn_metadatas): + graph = self.graphs[desc] + assert isinstance(graph, UpdatableGraph) + resolved_tasks = graph.resolve_tasks(ContextSource(draft_attn_metadatas[0])) + self.update_stream.wait_stream(torch.npu.current_stream()) + ret = super().run_fullgraph(desc) + graph.update(self.update_stream, resolved_tasks) + return ret diff --git a/vllm_ascend/worker/v2/utils.py b/vllm_ascend/worker/v2/utils.py index 26643423aa56..7fd3cf7e5996 100644 --- a/vllm_ascend/worker/v2/utils.py +++ b/vllm_ascend/worker/v2/utils.py @@ -5,6 +5,7 @@ from vllm.logger import logger from vllm_ascend.compilation.acl_graph import get_draft_graph_params, get_graph_params, weak_ref_workspaces +from vllm_ascend.compilation.updatable_graph import UpdatableGraph from vllm_ascend.utils import weak_ref_tensor, weak_ref_tensors @@ -17,7 +18,7 @@ def torch_cuda_wrapper(): torch.cuda.default_stream = torch.npu.default_stream torch.cuda.current_stream = torch.npu.current_stream torch.cuda.graph_pool_handle = torch.npu.graph_pool_handle - torch.cuda.CUDAGraph = torch.npu.NPUGraph + torch.cuda.CUDAGraph = UpdatableGraph torch.cuda.graph = torch_npu_graph_wrapper torch.cuda.synchronize = torch.npu.synchronize torch.cuda.set_stream = torch.npu.set_stream