Skip to content

Commit 6e149e3

Browse files
[OPTIMIZATION] save some compute in dsv3 (#7785)
1 parent c2df4c6 commit 6e149e3

4 files changed

Lines changed: 262 additions & 100 deletions

File tree

fastdeploy/model_executor/layers/attention/mla_attention_backend.py

Lines changed: 16 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -326,6 +326,7 @@ def extract_decoder_token_from_q(
326326
assert len(cu_seqlens_q.shape) == 1
327327
assert len(seq_lens_encoder.shape) == 1
328328
assert len(seq_lens_decoder.shape) == 1
329+
assert seq_lens_encoder.shape == seq_lens_decoder.shape
329330

330331
max_bsz = seq_lens_decoder.shape[0]
331332

@@ -398,7 +399,7 @@ def insert_decoder_result_back(
398399
max_bsz = seq_lens_encoder.shape[0]
399400

400401
hidden_dim = decoder_result.shape[-2] * decoder_result.shape[-1]
401-
out = paddle.zeros([mixed_token_num, hidden_dim], dtype=decoder_result.dtype)
402+
out = paddle.empty([mixed_token_num, hidden_dim], dtype=decoder_result.dtype)
402403

403404
BLOCK_SIZE = triton.next_power_of_2(hidden_dim)
404405

@@ -525,6 +526,7 @@ def __init__(
525526
self.useless_tensor = paddle.randn([1]).cast("int32")
526527
prop = paddle.device.cuda.get_device_properties()
527528
cc = prop.major * 10 + prop.minor
529+
self.prop = prop
528530
self.is_blackwell = cc >= 100
529531

530532
if self.flash_attn_func is None:
@@ -813,7 +815,7 @@ def forward_mixed(
813815
self.max_seq_len,
814816
)
815817

816-
if self.is_blackwell:
818+
if self.prop.major == 10:
817819
# TODO support FA4
818820
fmha_out = MLAAttentionBackend.mha_baseline(
819821
q,
@@ -857,7 +859,7 @@ def forward_mixed(
857859
speculate_decoder,
858860
)
859861

860-
if int(os.getenv("USE_FLASH_MLA", "0")) == 0:
862+
if int(os.getenv("USE_FLASH_MLA", "0")) == 0 and self.prop.major == 9:
861863
assert self.num_heads <= 64, "paddle mla attention support failed"
862864
if self.heads_need_padding:
863865
q = paddle.nn.functional.pad(
@@ -910,17 +912,7 @@ def forward_mixed(
910912

911913
return fmha_out
912914
else:
913-
import flash_mla
914-
915-
decoder_q, cache_seqlens = extract_decoder_token_from_q(
916-
q,
917-
forward_meta.cu_seqlens_q,
918-
forward_meta.seq_lens_encoder,
919-
forward_meta.seq_lens_decoder,
920-
)
921-
922-
tile_scheduler_metadata, num_splits = flash_mla.get_mla_metadata()
923-
token_num = q.shape[0]
915+
decoder_q = q
924916
decoder_q.reshape_([-1, 1, self.num_heads, 576])
925917
if self.heads_need_padding:
926918
padded_q = paddle.zeros(
@@ -933,22 +925,28 @@ def forward_mixed(
933925
assert new_cache_shape[1] == 1
934926
new_cache_shape[1], new_cache_shape[2] = new_cache_shape[2], new_cache_shape[1]
935927

936-
if self.is_blackwell:
928+
if self.prop.major == 10:
929+
# blackwell
937930
decoder_res = MLAAttentionBackend.mla_blackwell(
938931
decoder_q,
939932
latent_cache,
940933
metadata.block_tables,
941-
cache_seqlens,
934+
forward_meta.cache_seqlens,
942935
attn_softmax_scale=self.attn_softmax_scale,
943936
)
944937
else:
938+
939+
import flash_mla
940+
941+
tile_scheduler_metadata, num_splits = flash_mla.get_mla_metadata()
942+
945943
decoder_res, _ = flash_mla.flash_mla_with_kvcache(
946944
decoder_q,
947945
# 外面的开源仓库的kv cache存储格式和FD的不同
948946
# 幸好这里缓存的头是1,直接view即可,否则上上下下要改很多!
949947
latent_cache.view(new_cache_shape),
950948
metadata.block_tables,
951-
cache_seqlens,
949+
forward_meta.cache_seqlens,
952950
512, # t.dv,
953951
tile_scheduler_metadata,
954952
num_splits,
@@ -958,15 +956,7 @@ def forward_mixed(
958956
if self.heads_need_padding:
959957
decoder_res = decoder_res[:, :, : self.num_heads, :].contiguous()
960958

961-
final_res = insert_decoder_result_back(
962-
decoder_res,
963-
forward_meta.cu_seqlens_q,
964-
forward_meta.seq_lens_encoder,
965-
forward_meta.seq_lens_decoder,
966-
token_num,
967-
)
968-
969-
return final_res
959+
return decoder_res
970960

971961
@staticmethod
972962
def mla_blackwell(decoder_q, latent_cache, block_table, cache_seqlens, attn_softmax_scale):
@@ -1016,11 +1006,6 @@ def mla_blackwell(decoder_q, latent_cache, block_table, cache_seqlens, attn_soft
10161006
softmax_scale = attn_softmax_scale
10171007
output_scale = 1.0
10181008

1019-
import sys
1020-
1021-
sys.path.insert(
1022-
0, "/root/paddlejob/workspace/env_run/output/zkk/cutlass/examples/python/CuTeDSL/blackwell/mla"
1023-
)
10241009
from mla_decode_fp16 import BlackwellMultiHeadLatentAttentionForwardFP16
10251010

10261011
mla = BlackwellMultiHeadLatentAttentionForwardFP16(

fastdeploy/model_executor/models/deepseek_v3.py

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
from __future__ import annotations
1818

1919
import math
20+
import os
2021
import re
2122
from typing import Dict
2223

@@ -344,6 +345,9 @@ def __init__(self, fd_config: FDConfig, layer_id: int, prefix: str = "") -> None
344345

345346
self.prefix = prefix
346347

348+
prop = paddle.device.cuda.get_device_properties()
349+
self.prop = prop
350+
347351
@staticmethod
348352
def yarn_get_mscale(scale=1, mscale=1):
349353
""" """
@@ -362,6 +366,8 @@ def forward(
362366
fused_read_cache_and_interleave,
363367
)
364368

369+
q_total_token_num = hidden_states.shape[0]
370+
365371
attn_out = None
366372
if self.use_gated_attn:
367373
gate_out = self.gate(hidden_states)
@@ -439,6 +445,36 @@ def forward(
439445
attn_out = fmha_out
440446

441447
if need_do_decode: # max_dec_len_this_time
448+
449+
if int(os.getenv("USE_FLASH_MLA", "0")) == 0 and self.prop.major == 9:
450+
pass
451+
else:
452+
from fastdeploy.model_executor.layers.attention.mla_attention_backend import (
453+
extract_decoder_token_from_q,
454+
insert_decoder_result_back,
455+
)
456+
457+
decoder_query_nope, cache_seqlens = extract_decoder_token_from_q(
458+
query_nope.reshape([0, -1]),
459+
forward_meta.cu_seqlens_q,
460+
forward_meta.seq_lens_encoder,
461+
forward_meta.seq_lens_decoder,
462+
)
463+
464+
decoder_query_pe, cache_seqlens = extract_decoder_token_from_q(
465+
query_pe.reshape([0, -1]),
466+
forward_meta.cu_seqlens_q,
467+
forward_meta.seq_lens_encoder,
468+
forward_meta.seq_lens_decoder,
469+
)
470+
assert decoder_query_nope.shape[0] == forward_meta.seq_lens_encoder.shape[0]
471+
assert decoder_query_pe.shape[0] == forward_meta.seq_lens_encoder.shape[0]
472+
473+
forward_meta.cache_seqlens = cache_seqlens
474+
475+
query_nope = decoder_query_nope.reshape([0, -1, self.qk_nope_head_dim])
476+
query_pe = decoder_query_pe.reshape([0, -1, self.qk_rope_head_dim])
477+
442478
q_nope_out = self.kv_b_proj_bmm(query_nope.transpose([1, 0, 2]), proj_type="k").transpose([1, 0, 2])
443479

444480
q_input = paddle.concat([q_nope_out, query_pe], axis=-1)
@@ -467,6 +503,17 @@ def forward(
467503
.reshape_([-1, self.num_attention_heads_tp * self.v_head_dim])
468504
)
469505

506+
if int(os.getenv("USE_FLASH_MLA", "0")) == 0 and self.prop.major == 9:
507+
pass
508+
else:
509+
fmqa_out = insert_decoder_result_back(
510+
fmqa_out.reshape([0, 1, self.num_attention_heads_tp, self.v_head_dim]),
511+
forward_meta.cu_seqlens_q,
512+
forward_meta.seq_lens_encoder,
513+
forward_meta.seq_lens_decoder,
514+
q_total_token_num,
515+
)
516+
470517
if need_do_prefill:
471518
merge_prefill_decode_output(
472519
attn_out,

0 commit comments

Comments
 (0)