Skip to content

Commit 2fa0c57

Browse files
Ascend refacter fused module. (#136)
Vendor Improvements Narrow down the patch scope for the module. Co-authored-by: cyber-pioneer <116002591+cyber-pioneer@users.noreply.github.com>
1 parent 344e42b commit 2fa0c57

30 files changed

Lines changed: 234 additions & 6070 deletions

vllm_fl/dispatch/backends/vendor/ascend/impl/attention.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -548,8 +548,9 @@ def _get_fia_params(
548548
attn_metadata: AscendMetadata,
549549
):
550550
"""Get parameters for fused_infer_attention."""
551-
block_size = 128
551+
552552
if attn_metadata.attn_state == AscendAttentionState.PrefillNoCache:
553+
block_size = 128
553554
block_table = None
554555
actual_seq_lengths_kv = attn_metadata.actual_seq_lengths_q
555556
elif attn_metadata.attn_state == AscendAttentionState.PrefillCacheHit:
@@ -560,13 +561,18 @@ def _get_fia_params(
560561
value = self.value_cache.view(num_block, block_size, -1)
561562
actual_seq_lengths_kv = attn_metadata.seq_lens_list
562563
elif attn_metadata.attn_state == AscendAttentionState.DecodeOnly:
564+
# num_block, block_size, _, _ = self.key_cache.shape
565+
# key = self.key_cache.view(num_block, block_size, -1)
566+
# value = self.value_cache.view(num_block, block_size, -1)
563567
key = self.key_cache.view(-1, block_size, 256)
564568
value = self.value_cache.view(-1, block_size, 256)
565569
block_table = attn_metadata.block_tables
566570
actual_seq_lengths_kv = attn_metadata.seq_lens_list
567571
else:
568572
# ChunkedPrefill
569573
# num_block, block_size, _, _ = self.key_cache.shape
574+
# key = self.key_cache.view(num_block, block_size, -1)
575+
# value = self.value_cache.view(num_block, block_size, -1)
570576
key = self.key_cache.view(-1, block_size, 256)
571577
value = self.value_cache.view(-1, block_size, 256)
572578
block_table = attn_metadata.block_tables
Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1 +1,6 @@
11
# Copyright (c) 2026 BAAI. All rights reserved.
2+
from vllm_fl.dispatch.backends.vendor.ascend.impl.fla.chunk import (
3+
chunk_gated_delta_rule as chunk_gated_delta_rule_npu,
4+
)
5+
6+
__all__ = ["chunk_gated_delta_rule_npu"]

vllm_fl/dispatch/backends/vendor/ascend/impl/fla/chunk.py

Lines changed: 1 addition & 57 deletions
Original file line numberDiff line numberDiff line change
@@ -13,66 +13,10 @@
1313

1414
import torch
1515
from einops import rearrange
16-
from vllm.model_executor.layers.fla.ops.utils import SUPPRESS_LEVEL
1716

18-
from .chunk_delta_h import chunk_gated_delta_rule_fwd_h
19-
from .chunk_o import chunk_fwd_o
20-
from .chunk_scaled_dot_kkt import chunk_scaled_dot_kkt_fwd
21-
from .cumsum import chunk_local_cumsum
2217
from .l2norm import l2norm_fwd
23-
from .solve_tril import solve_tril
2418
from .utils import input_guard
25-
from .wy_fast import recompute_w_u_fwd
26-
27-
28-
def chunk_gated_delta_rule_fwd(
29-
q: torch.Tensor,
30-
k: torch.Tensor,
31-
v: torch.Tensor,
32-
g: torch.Tensor,
33-
beta: torch.Tensor,
34-
scale: float,
35-
initial_state: torch.Tensor,
36-
output_final_state: bool,
37-
cu_seqlens: Optional[torch.LongTensor] = None,
38-
):
39-
g = chunk_local_cumsum(g, chunk_size=64, cu_seqlens=cu_seqlens)
40-
# obtain WY representation. u is actually the new v.
41-
A = chunk_scaled_dot_kkt_fwd(
42-
k=k, beta=beta, g_cumsum=g, cu_seqlens=cu_seqlens, output_dtype=torch.float32
43-
)
44-
A = solve_tril(A=A, cu_seqlens=cu_seqlens, output_dtype=k.dtype)
45-
w, u = recompute_w_u_fwd(
46-
k=k,
47-
v=v,
48-
beta=beta,
49-
A=A,
50-
g_cumsum=g,
51-
cu_seqlens=cu_seqlens,
52-
)
53-
h, v_new, final_state = chunk_gated_delta_rule_fwd_h(
54-
k=k,
55-
w=w,
56-
u=u,
57-
g=g,
58-
initial_state=initial_state,
59-
output_final_state=output_final_state,
60-
cu_seqlens=cu_seqlens,
61-
)
62-
o = chunk_fwd_o(
63-
q=q,
64-
k=k,
65-
v=v_new,
66-
h=h,
67-
g=g,
68-
scale=scale,
69-
cu_seqlens=cu_seqlens,
70-
)
71-
if SUPPRESS_LEVEL < 3:
72-
return g, o, A, final_state, None, None, None
73-
elif SUPPRESS_LEVEL >= 3:
74-
return g, o, A, final_state, w, h, v_new
75-
19+
from flag_gems.runtime.backend._ascend.fla import chunk_gated_delta_rule_fwd
7620

7721
class ChunkGatedDeltaRuleFunction(torch.autograd.Function):
7822
@staticmethod

vllm_fl/dispatch/backends/vendor/ascend/impl/fla/chunk_delta_h.py

Lines changed: 0 additions & 278 deletions
This file was deleted.

0 commit comments

Comments
 (0)