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
17 changes: 11 additions & 6 deletions vllm_ascend/attention/sfa_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -356,8 +356,9 @@ class AscendSFAImpl(MLAAttentionImpl):
# Supports forward using the all-gather o_proj weight for decode requests when Sharded CP is enabled.
o_proj_full_pool: torch.Tensor | None = None

# qk_hadamard tensor shared when dsa c8 enabled
qk_hadamard: torch.Tensor | None = None
# q_hadamard and k_hadamard tensor shared when dsa c8 enabled
q_hadamard: torch.Tensor | None = None
k_hadamard: torch.Tensor | None = None

def __init__(
self,
Expand Down Expand Up @@ -525,8 +526,12 @@ def process_weights_after_loading(self, act_dtype: torch.dtype):
# if mlapo, W_UK_T can't trans nz
self.W_UK_T = maybe_trans_nz(self.W_UK_T)

if self.use_sparse_c8_indexer and AscendSFAImpl.qk_hadamard is None:
AscendSFAImpl.qk_hadamard = torch.tensor(scipy.linalg.hadamard(128), dtype=torch.bfloat16, device="npu") / (
if self.use_sparse_c8_indexer and AscendSFAImpl.q_hadamard is None:
AscendSFAImpl.q_hadamard = torch.tensor(scipy.linalg.hadamard(128), dtype=torch.bfloat16, device="npu") / (
128**0.5
)
if self.use_sparse_c8_indexer and AscendSFAImpl.k_hadamard is None:
AscendSFAImpl.k_hadamard = torch.tensor(scipy.linalg.hadamard(128), dtype=torch.bfloat16, device="npu") / (
128**0.5
)
Comment on lines +529 to 536

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

The Hadamard matrix is being created twice, which is inefficient. Since q_hadamard and k_hadamard are initialized with the same value, you can create the tensor once and assign it to both, avoiding redundant computation and memory allocation.

        if self.use_sparse_c8_indexer and (AscendSFAImpl.q_hadamard is None or AscendSFAImpl.k_hadamard is None):
            hadamard_matrix = torch.tensor(scipy.linalg.hadamard(128), dtype=torch.bfloat16, device="npu") / (128**0.5)
            if AscendSFAImpl.q_hadamard is None:
                AscendSFAImpl.q_hadamard = hadamard_matrix
            if AscendSFAImpl.k_hadamard is None:
                AscendSFAImpl.k_hadamard = hadamard_matrix


Expand Down Expand Up @@ -890,7 +895,7 @@ def indexer_select_pre_process(
k_li = torch.cat([k_li_pe, k_li_nope], dim=-1) # [b*s,128]

if self.use_sparse_c8_indexer:
k_li = k_li @ AscendSFAImpl.qk_hadamard
k_li = k_li @ AscendSFAImpl.k_hadamard
k_li, k_li_scale = torch_npu.npu_dynamic_quant(k_li.view(-1, self.head_dim), dst_type=self.c8_k_cache_dtype)
k_li_scale = k_li_scale.to(self.c8_k_scale_cache_dtype) # [b*s,]
k_li_scale = k_li_scale.unsqueeze(-1) # [b*s,1]
Expand Down Expand Up @@ -930,7 +935,7 @@ def indexer_select_post_process(

if self.use_sparse_c8_indexer:
q_li_shape_ori = q_li.shape
q_li = q_li @ AscendSFAImpl.qk_hadamard
q_li = q_li @ AscendSFAImpl.q_hadamard
q_li, q_li_scale = torch_npu.npu_dynamic_quant(q_li.view(-1, self.head_dim), dst_type=self.c8_k_cache_dtype)
q_li_scale = q_li_scale.to(self.c8_k_scale_cache_dtype)

Expand Down
11 changes: 10 additions & 1 deletion vllm_ascend/patch/worker/patch_weight_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,16 @@ def patch_deepseek(module):
def new_remap(name: str, params_dict: dict):
name = ori_maybe_remap_kv_scale_name(name, params_dict)

replace_scale_names = ["fa_q.scale", "fa_k.scale", "fa_v.scale", "fa_q.offset", "fa_k.offset", "fa_v.offset"]
replace_scale_names = [
"fa_q.scale",
"fa_k.scale",
"fa_v.scale",
"fa_q.offset",
"fa_k.offset",
"fa_v.offset",
"indexer.q_rot",
"indexer.k_rot",
]

for scale_name in replace_scale_names:
if name.endswith(scale_name):
Expand Down
24 changes: 24 additions & 0 deletions vllm_ascend/quantization/methods/kv_c8.py
Original file line number Diff line number Diff line change
Expand Up @@ -63,3 +63,27 @@ def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
repeated_quant_kscale = fa_k_scale.repeat(self.kv_lora_rank)
layer.quant_kscale = repeated_quant_kscale.view(1, self.kv_lora_rank)
layer.quant_kscale = 1.0 / torch.nn.Parameter(layer.quant_kscale.to(torch.float), requires_grad=False)


@register_scheme("INT8_DYNAMIC", "attention")
class AscendSFAQuantAttentionMethod:
def __init__(self):
vllm_config = get_current_vllm_config()
config = vllm_config.model_config.hf_config
self.index_head_dim = config.index_head_dim

def create_weights(self, layer: torch.nn.Module) -> None:
extra_module_names = ["indexer"]
for name in extra_module_names:
setattr(layer, name, torch.nn.Module())
params_dict = {}
params_dict["indexer.q_rot"] = torch.empty((self.index_head_dim, self.index_head_dim), dtype=torch.float32)
params_dict["indexer.k_rot"] = torch.empty((self.index_head_dim, self.index_head_dim), dtype=torch.float32)
for name, weight in params_dict.items():
module_name, weight_name = name.split(".")
module = getattr(layer, module_name)
weight_param = torch.nn.Parameter(weight, requires_grad=False)
module.register_parameter(weight_name, weight_param)

def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
pass
22 changes: 19 additions & 3 deletions vllm_ascend/quantization/modelslim_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -379,6 +379,8 @@ def get_quant_type_for_layer(
# Attention
if layer_type == "attention" and "fa_quant_type" in quant_description:
return quant_description["fa_quant_type"]
if layer_type == "attention" and "indexer_quant_type" in quant_description:
return quant_description["indexer_quant_type"]
# Linear / MoE
return get_linear_quant_type(quant_description, prefix, packed_modules_mapping)

Expand Down Expand Up @@ -582,7 +584,9 @@ def get_quant_method(self, layer: torch.nn.Module, prefix: str) -> Optional["Qua
return AscendUnquantizedLinearMethod()
scheme = create_scheme_for_layer(self.quant_description, prefix, "linear", self.packed_modules_mapping)
return AscendLinearMethod(scheme)
elif isinstance(layer, AttentionLayerBase) and self.is_fa_quant_layer(prefix):
elif isinstance(layer, AttentionLayerBase) and (
self.is_fa_quant_layer(prefix) or self.is_indexer_quant_layer(prefix)
):
scheme = create_scheme_for_layer(self.quant_description, prefix, "attention", self.packed_modules_mapping)
return AscendKVCacheMethod(scheme)
elif isinstance(layer, FusedMoE):
Expand Down Expand Up @@ -636,6 +640,13 @@ def is_fa_quant_layer(self, prefix):
return True
return False

def is_indexer_quant_layer(self, prefix):
if self.enable_indexer_quant:
layer_id_str = "".join(re.findall(r"\.(\d+)\.", prefix))
if layer_id_str.isdigit() and int(layer_id_str) in self.indexer_quant_layers:
return True
return False

def enabling_fa_quant(self, vllm_config, layer_name) -> bool:
is_decode_instance = (
vllm_config.kv_transfer_config is not None
Expand Down Expand Up @@ -773,8 +784,13 @@ def _add_kvcache_quant_metadata(self):
fa_quant_type = self.quant_description.get("fa_quant_type", "")
self.enable_fa_quant = fa_quant_type != ""
self.kvcache_quant_layers = []
if self.enable_fa_quant:
indexer_quant_type = self.quant_description.get("indexer_quant_type", "")
self.enable_indexer_quant = indexer_quant_type != ""
self.indexer_quant_layers = []
if self.enable_fa_quant or self.enable_indexer_quant:
for key in self.quant_description:
_id = "".join(re.findall(r"\.(\d+)\.", key))
if "fa_k.scale" in key:
_id = "".join(re.findall(r"\.(\d+)\.", key))
self.kvcache_quant_layers.append(int(_id))
if "indexer.quant_type" in key:
self.indexer_quant_layers.append(int(_id))
Loading