Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 25 additions & 1 deletion custom_ops/xpu_ops/src/ops/block_attn.cc
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,8 @@ std::vector<paddle::Tensor> BlockAttnKernel(
const paddle::optional<paddle::Tensor>& v_zeros,
const paddle::optional<paddle::Tensor>& shift,
const paddle::optional<paddle::Tensor>& smooth,
const paddle::optional<paddle::Tensor>& q_norm_weight,
const paddle::optional<paddle::Tensor>& k_norm_weight,
const paddle::optional<paddle::Tensor>& kv_signal_data_cpu,
const paddle::optional<paddle::Tensor>& cachekv_signal_thread_cpu,
const bool use_neox_rotary_style,
Expand Down Expand Up @@ -197,6 +199,16 @@ std::vector<paddle::Tensor> BlockAttnKernel(
const_cast<float*>(v_scales_inv.get().data<float>()));
}
}
const float *q_norm_weight_data{nullptr}, *k_norm_weight_data{nullptr};
if (q_norm_weight) {
q_norm_weight_data = q_norm_weight.get().data<float>();
}
if (k_norm_weight) {
k_norm_weight_data = k_norm_weight.get().data<float>();
}
PD_CHECK(!(pos_emb_type == "NEOX" && q_norm_weight_data != nullptr),
"split_neox_cache_kv_encoder not support q/k norm weight");

int ret = 0;
if (enc_batch > 0) {
xftblock::TransformerParam param;
Expand Down Expand Up @@ -383,6 +395,8 @@ std::vector<paddle::Tensor> BlockAttnKernel(
quant_v_scale, // intx_v_pc_scale
quant_k_zp, // intx_k_pc_zero
quant_v_zp, // intx_v_pc_zero
q_norm_weight_data,
k_norm_weight_data,
rope_3d);
PD_CHECK(ret == api::SUCCESS, "split_rope_cache_kv_encoder failed.");
}
Expand Down Expand Up @@ -632,6 +646,8 @@ std::vector<paddle::Tensor> BlockAttnKernel(
quant_v_scale, // intx_v_pc_scale
quant_k_zp, // intx_k_pc_zero
quant_v_zp, // intx_v_pc_zero
q_norm_weight_data,
k_norm_weight_data,
rope_3d);
PD_CHECK(ret == api::SUCCESS, "split_rope_cache_kv_encoder failed.");
}
Expand Down Expand Up @@ -858,7 +874,9 @@ std::vector<paddle::Tensor> BlockAttnKernel(
reinterpret_cast<D_Scale*>(quant_v_scale), // v_cache_scale_inv
reinterpret_cast<D_Scale*>(quant_k_zp), // k_cache_zp
reinterpret_cast<D_Scale*>(quant_v_zp), // v_cache_zp
is_cache_int8, // bool b_c8_pc
q_norm_weight_data,
k_norm_weight_data,
is_cache_int8, // bool b_c8_pc
rope_3d);
PD_CHECK(ret == api::SUCCESS, "split_rope_cache_kv_decoder failed.");
}
Expand Down Expand Up @@ -1003,6 +1021,8 @@ std::vector<paddle::Tensor> BlockAttn(
const paddle::optional<paddle::Tensor>& v_zeros,
const paddle::optional<paddle::Tensor>& shift,
const paddle::optional<paddle::Tensor>& smooth,
const paddle::optional<paddle::Tensor>& q_norm_weight,
const paddle::optional<paddle::Tensor>& k_norm_weight,
const paddle::optional<paddle::Tensor>& kv_signal_data_cpu,
const paddle::optional<paddle::Tensor>& cachekv_signal_thread_cpu,
const bool use_neox_rotary_style,
Expand Down Expand Up @@ -1032,6 +1052,8 @@ std::vector<paddle::Tensor> BlockAttn(
v_zeros, \
shift, \
smooth, \
q_norm_weight, \
k_norm_weight, \
kv_signal_data_cpu, \
cachekv_signal_thread_cpu, \
use_neox_rotary_style, \
Expand Down Expand Up @@ -1098,6 +1120,8 @@ PD_BUILD_STATIC_OP(block_attn)
paddle::Optional("v_zeros"),
paddle::Optional("shift"),
paddle::Optional("smooth"),
paddle::Optional("q_norm_weight"),
paddle::Optional("k_norm_weight"),
paddle::Optional("kv_signal_data_cpu"),
paddle::Optional("cachekv_signal_thread_cpu")})
.Attrs({"use_neox_rotary_style:bool", "rope_3d:bool"})
Expand Down
4 changes: 4 additions & 0 deletions custom_ops/xpu_ops/src/ops/pybind/pybind.cc
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,8 @@ std::vector<paddle::Tensor> BlockAttn(
const paddle::optional<paddle::Tensor>& v_zeros,
const paddle::optional<paddle::Tensor>& shift,
const paddle::optional<paddle::Tensor>& smooth,
const paddle::optional<paddle::Tensor>& q_norm_weight,
const paddle::optional<paddle::Tensor>& k_norm_weight,
const paddle::optional<paddle::Tensor>& kv_signal_data_cpu,
const paddle::optional<paddle::Tensor>& cachekv_signal_thread_cpu,
const bool use_neox_rotary_style,
Expand Down Expand Up @@ -640,6 +642,8 @@ PYBIND11_MODULE(fastdeploy_ops, m) {
py::arg("v_zeros"),
py::arg("shift"),
py::arg("smooth"),
py::arg("q_norm_weight"),
py::arg("k_norm_weight"),
py::arg("kv_signal_data_cpu"),
py::arg("cachekv_signal_thread_cpu"),
py::arg("use_neox_rotary_style"),
Expand Down
28 changes: 15 additions & 13 deletions fastdeploy/model_executor/layers/backends/xpu/attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
init_signal_layerwise,
open_shm_and_get_meta_signal,
)
from fastdeploy.model_executor.ops.xpu import block_attn

if TYPE_CHECKING:
from fastdeploy.model_executor.forward_meta import ForwardMeta
Expand Down Expand Up @@ -175,16 +176,15 @@ def forward_mixed(
layer.layer_id + self.start_layer_index,
)

k_quant_scale = getattr(layer, "cache_k_scale", None)
v_quant_scale = getattr(layer, "cache_v_scale", None)

cache_k_scale = getattr(layer, "cache_k_scale", None)
cache_v_scale = getattr(layer, "cache_v_scale", None)
cache_k_out_scale = getattr(layer, "cache_k_out_scale", None)
cache_v_out_scale = getattr(layer, "cache_v_out_scale", None)
cache_k_zp = getattr(self, "cache_k_zp", None)
cache_v_zp = getattr(self, "cache_v_zp", None)

k_zp = getattr(self, "cache_k_zp", None)
v_zp = getattr(self, "cache_v_zp", None)

from fastdeploy.model_executor.ops.xpu import block_attn
q_norm_weight = getattr(layer, "q_norm_weight", None)
k_norm_weight = getattr(layer, "k_norm_weight", None)

res = block_attn(
qkv,
Expand All @@ -203,16 +203,18 @@ def forward_mixed(
forward_meta.decoder_context_len_cache_cpu,
forward_meta.decoder_batch_map_cpu,
forward_meta.prefix_len_cpu,
k_quant_scale,
v_quant_scale,
cache_k_scale,
cache_v_scale,
cache_k_out_scale,
cache_v_out_scale,
k_zp, # zero_point_quant_scale
v_zp, # zero_point_quant_scale
cache_k_zp,
cache_v_zp,
None, # shift
None, # smooth
metadata.kv_signal_data_list[layer.layer_id], # kv_signal_data
forward_meta.kv_signal_sender, # kv_signal_sender
q_norm_weight,
k_norm_weight,
metadata.kv_signal_data_list[layer.layer_id],
forward_meta.kv_signal_sender,
layer.use_neox_rotary_style,
self.rope_3d,
)
Expand Down
2 changes: 1 addition & 1 deletion fastdeploy/model_executor/xpu_pre_and_post_process.py
Original file line number Diff line number Diff line change
Expand Up @@ -328,7 +328,7 @@ def xpu_post_process_normal(

# 2. Update the input buffer of the model
with paddle.framework._no_check_dy2st_diff():
if envs.ENABLE_V1_KVCACHE_SCHEDULER and not skip_save_output:
if envs.ENABLE_V1_KVCACHE_SCHEDULER:
update_inputs_v1(
model_output.stop_flags,
model_output.not_need_stop,
Expand Down
7 changes: 6 additions & 1 deletion fastdeploy/worker/xpu_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,14 +24,19 @@
from fastdeploy.config import FDConfig
from fastdeploy.engine.request import Request
from fastdeploy.platforms import current_platform
from fastdeploy.plugins.model_runner import load_model_runner_plugins
from fastdeploy.usage.usage_lib import report_usage_stats
from fastdeploy.utils import get_logger, set_random_seed
from fastdeploy.worker.output import ModelRunnerOutput
from fastdeploy.worker.worker_base import WorkerBase
from fastdeploy.worker.xpu_model_runner import XPUModelRunner

logger = get_logger("xpu_worker", "xpu_worker.log")

try:
XPUModelRunner = load_model_runner_plugins()
except:
from fastdeploy.worker.xpu_model_runner import XPUModelRunner


class XpuWorker(WorkerBase):
""" """
Expand Down
Loading