[Kernel] Dynamic Per-Token Activation Quantization (#5037)
Co-authored-by: Varun Sundar Rabindranath <varunsundar08@gmail.com> Co-authored-by: Varun Sundar Rabindranath <varun@neuralmagic.com>
This commit is contained in:
@@ -1,12 +1,16 @@
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import torch
|
||||
from pydantic import BaseModel
|
||||
|
||||
from vllm.model_executor.layers.linear import LinearBase, LinearMethodBase
|
||||
from vllm.model_executor.layers.quantization.base_config import ( # noqa: E501
|
||||
QuantizationConfig)
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.schemes import (
|
||||
CompressedTensorsScheme, CompressedTensorsW8A8StaticTensor)
|
||||
CompressedTensorsScheme, CompressedTensorsW8A8DynamicToken,
|
||||
CompressedTensorsW8A8StaticTensor)
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.utils import (
|
||||
QuantizationArgs, QuantizationStrategy, find_first_name_or_class_match)
|
||||
|
||||
|
||||
class CompressedTensorsConfig(QuantizationConfig):
|
||||
@@ -47,10 +51,12 @@ class CompressedTensorsConfig(QuantizationConfig):
|
||||
targets = quant_config.get("targets")
|
||||
for target in targets:
|
||||
layer_quant_details[target] = {}
|
||||
layer_quant_details[target]["weight"] = quant_config.get(
|
||||
"weights")
|
||||
layer_quant_details[target]["input"] = quant_config.get(
|
||||
"input_activations")
|
||||
layer_quant_details[target][
|
||||
"weight"] = QuantizationArgs.parse_obj(
|
||||
quant_config.get("weights"))
|
||||
layer_quant_details[target][
|
||||
"input"] = QuantizationArgs.parse_obj(
|
||||
quant_config.get("input_activations"))
|
||||
|
||||
return cls(layer_quant_details=layer_quant_details, ignore=ignore)
|
||||
|
||||
@@ -58,40 +64,46 @@ class CompressedTensorsConfig(QuantizationConfig):
|
||||
def get_config_filenames(cls) -> List[str]:
|
||||
return []
|
||||
|
||||
def _get_schema(self, weight_quant: Dict, input_quant: Dict):
|
||||
# TODO: Refactor as additional cases are supported
|
||||
def _is_static_tensor_w8a8(self, weight_quant: BaseModel,
|
||||
input_quant: BaseModel) -> bool:
|
||||
is_8_bits = weight_quant.num_bits == input_quant.num_bits == 8
|
||||
is_tensor = (weight_quant.strategy == input_quant.strategy ==
|
||||
QuantizationStrategy.TENSOR.value)
|
||||
is_symmetric = weight_quant.symmetric and input_quant.symmetric
|
||||
is_static = not weight_quant.dynamic and not input_quant.dynamic
|
||||
|
||||
weight_bit = weight_quant.get("num_bits")
|
||||
input_bit = input_quant.get("num_bits")
|
||||
return is_8_bits and is_tensor and is_symmetric and is_static
|
||||
|
||||
weight_strategy = weight_quant.get("strategy")
|
||||
input_strategy = input_quant.get("strategy")
|
||||
def _is_dynamic_token_w8a8(self, weight_quant: BaseModel,
|
||||
input_quant: BaseModel) -> bool:
|
||||
is_8_bits = weight_quant.num_bits == input_quant.num_bits == 8
|
||||
is_token_tensor = (weight_quant.strategy
|
||||
== QuantizationStrategy.TENSOR.value) and (
|
||||
input_quant.strategy
|
||||
== QuantizationStrategy.TOKEN.value)
|
||||
is_symmetric = weight_quant.symmetric and input_quant.symmetric
|
||||
is_dynamic = not weight_quant.dynamic and input_quant.dynamic
|
||||
|
||||
weight_symmetric = weight_quant.get("symmetric")
|
||||
input_symmetric = input_quant.get("symmetric")
|
||||
return is_8_bits and is_token_tensor and is_symmetric and is_dynamic
|
||||
|
||||
is_8_bits = weight_bit == input_bit == 8
|
||||
is_tensor = weight_strategy == input_strategy == "tensor"
|
||||
is_symmetric = weight_symmetric and input_symmetric
|
||||
|
||||
if is_8_bits and is_tensor and is_symmetric and \
|
||||
torch.cuda.is_available():
|
||||
# CompressedTensorsW8A8StaticTensor only supports CUDA path for
|
||||
# now.
|
||||
def _get_schema(self, weight_quant: BaseModel,
|
||||
input_quant: BaseModel) -> "CompressedTensorsScheme":
|
||||
if self._is_static_tensor_w8a8(weight_quant, input_quant):
|
||||
return CompressedTensorsW8A8StaticTensor()
|
||||
raise NotImplementedError(
|
||||
"Scheme not supported. Only CUDA, 8-bit static symmtetric "
|
||||
"per tensor quantization is currently supported")
|
||||
|
||||
if self._is_dynamic_token_w8a8(weight_quant, input_quant):
|
||||
return CompressedTensorsW8A8DynamicToken()
|
||||
|
||||
raise NotImplementedError("Scheme not supported.")
|
||||
|
||||
def get_scheme(self, layer: torch.nn.Module) -> "CompressedTensorsScheme":
|
||||
|
||||
# TODO: update with matching function from `compressed_tensors`
|
||||
layer_type_name = None
|
||||
layer_name_class = type(layer).__name__.lower()
|
||||
for target in self.layer_quant_details:
|
||||
if target.lower() in layer_name_class:
|
||||
layer_type_name = target
|
||||
break
|
||||
layer_type_name = find_first_name_or_class_match(
|
||||
name="",
|
||||
module=layer,
|
||||
targets=self.layer_quant_details.keys(),
|
||||
check_contains=True)
|
||||
|
||||
if layer_type_name is None:
|
||||
raise ValueError(f"Could not matching target for layer {layer}")
|
||||
|
||||
@@ -117,7 +129,9 @@ class CompressedTensorsLinearMethod(LinearMethodBase):
|
||||
**extra_weight_attrs):
|
||||
"""
|
||||
Use the CompressedTensorsScheme associated with each layer to create
|
||||
the necessary parameters for the layer.
|
||||
the necessary parameters for the layer. See LinearMethodBase for param
|
||||
details
|
||||
|
||||
"""
|
||||
weight_loader = extra_weight_attrs.get("weight_loader")
|
||||
|
||||
@@ -139,7 +153,8 @@ class CompressedTensorsLinearMethod(LinearMethodBase):
|
||||
"""
|
||||
Use the output of create_weights and the CompressedTensorsScheme
|
||||
associated with the layer to apply the forward pass with the
|
||||
layer input.
|
||||
layer input. See LinearMethodBase for param details
|
||||
|
||||
"""
|
||||
|
||||
if bias is not None:
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
from .compressed_tensors_scheme import CompressedTensorsScheme # noqa: F401
|
||||
from .compressed_tensors_unquantized import ( # noqa: F401
|
||||
CompressedTensorsUnquantized)
|
||||
from .compressed_tensors_w8a8_dynamictoken import ( # noqa: F401, E501
|
||||
CompressedTensorsW8A8DynamicToken)
|
||||
from .compressed_tensors_w8a8_statictensor import ( # noqa: F401, E501
|
||||
CompressedTensorsW8A8StaticTensor)
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
from typing import Callable, List, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch.nn import Parameter
|
||||
|
||||
from vllm import _custom_ops as custom_ops
|
||||
from vllm.model_executor.layers.quantization.compressed_tensors.schemes import (
|
||||
CompressedTensorsScheme)
|
||||
from vllm.model_executor.utils import set_weight_attrs
|
||||
|
||||
__all__ = ["CompressedTensorsW8A8DynamicToken"]
|
||||
|
||||
|
||||
class CompressedTensorsW8A8DynamicToken(CompressedTensorsScheme):
|
||||
|
||||
def _shard_id_as_int(self, shard_id: Union[str, int]) -> int:
|
||||
if isinstance(shard_id, int):
|
||||
return shard_id
|
||||
|
||||
assert isinstance(shard_id, str)
|
||||
qkv_idxs = {"q": 0, "k": 1, "v": 2}
|
||||
assert shard_id in qkv_idxs
|
||||
return qkv_idxs[shard_id]
|
||||
|
||||
def scales_shard_splitter(
|
||||
self, param: torch.Tensor, loaded_weight: torch.Tensor,
|
||||
shard_id: Union[str, int],
|
||||
logical_widths: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
shard_id = self._shard_id_as_int(shard_id)
|
||||
offset = sum(logical_widths[:shard_id])
|
||||
size = logical_widths[shard_id]
|
||||
# update loaded weight with copies for broadcast.
|
||||
loaded_weight = loaded_weight.repeat(size)
|
||||
return param[offset:offset + size], loaded_weight
|
||||
|
||||
def create_weights(self, layer: torch.nn.Module,
|
||||
output_partition_sizes: List[int],
|
||||
input_size_per_partition: int,
|
||||
params_dtype: torch.dtype, weight_loader: Callable,
|
||||
**kwargs):
|
||||
|
||||
# When the scales have a single value, it is required that they be
|
||||
# on the CPU for performance and CUDA Graphs compatibility. Please
|
||||
# refer to the comment in
|
||||
# CompressedTensorsW8A8StaticTensor::create_weights for further
|
||||
# information.
|
||||
is_tensor_partitioned = len(output_partition_sizes) != 1
|
||||
weight_scale_dim = sum(
|
||||
output_partition_sizes) if is_tensor_partitioned else 1
|
||||
|
||||
weight_zero_point = Parameter(torch.empty(1, dtype=torch.int8),
|
||||
requires_grad=False)
|
||||
|
||||
weight_scale = Parameter(torch.empty(weight_scale_dim,
|
||||
dtype=torch.float32),
|
||||
requires_grad=False)
|
||||
|
||||
weight = Parameter(torch.empty(sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=torch.int8),
|
||||
requires_grad=False)
|
||||
|
||||
layer.register_parameter("weight", weight)
|
||||
set_weight_attrs(weight, {"input_dim": 1, "output_dim": 0})
|
||||
set_weight_attrs(weight, {"weight_loader": weight_loader})
|
||||
set_weight_attrs(weight, {"logical_widths": output_partition_sizes})
|
||||
|
||||
layer.register_parameter("weight_scale", weight_scale)
|
||||
set_weight_attrs(weight_scale, {"weight_loader": weight_loader})
|
||||
set_weight_attrs(
|
||||
weight_scale, {
|
||||
"shard_splitter": self.scales_shard_splitter,
|
||||
"logical_widths": output_partition_sizes
|
||||
})
|
||||
|
||||
layer.register_parameter("weight_zero_point", weight_zero_point)
|
||||
set_weight_attrs(weight_zero_point, {"weight_loader": weight_loader})
|
||||
|
||||
def apply_weights(self, layer: torch.nn.Module, x: torch.Tensor):
|
||||
weight = layer.weight
|
||||
weight_scale = layer.weight_scale
|
||||
|
||||
x_q, input_scales = custom_ops.scaled_int8_quant(x)
|
||||
return custom_ops.cutlass_scaled_mm_dq(x_q, weight.t(), input_scales,
|
||||
weight_scale, x.dtype)
|
||||
@@ -97,7 +97,7 @@ class CompressedTensorsW8A8StaticTensor(CompressedTensorsScheme):
|
||||
act_scale = layer.input_scale
|
||||
|
||||
# Input quantize
|
||||
x_q = custom_ops.static_scaled_int8_quant(x, act_scale)
|
||||
x_q, _ = custom_ops.scaled_int8_quant(x, act_scale)
|
||||
|
||||
return custom_ops.cutlass_scaled_mm_dq(x_q, weight.t(), act_scale,
|
||||
weight_scale, x.dtype)
|
||||
|
||||
@@ -0,0 +1,114 @@
|
||||
import re
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, Iterable, Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from torch.nn import Module
|
||||
|
||||
|
||||
class QuantizationType(str, Enum):
|
||||
"""
|
||||
Enum storing quantization type options
|
||||
"""
|
||||
|
||||
INT = "int"
|
||||
FLOAT = "float"
|
||||
|
||||
|
||||
class QuantizationStrategy(str, Enum):
|
||||
"""
|
||||
Enum storing quantization strategy options
|
||||
"""
|
||||
|
||||
TENSOR = "tensor"
|
||||
CHANNEL = "channel"
|
||||
GROUP = "group"
|
||||
BLOCK = "block"
|
||||
TOKEN = "token"
|
||||
|
||||
|
||||
class QuantizationArgs(BaseModel):
|
||||
"""
|
||||
User facing arguments used to define a quantization config
|
||||
for weights or activations
|
||||
|
||||
:param num_bits: quantization bit depth
|
||||
:param type: dtype to quantized to, either int or float
|
||||
:param symmetric: whether or not quantization scale is symmetric
|
||||
:param strategy: string determining the scope of scale/zero-point to apply
|
||||
:param group_size: group length to use for the group strategy
|
||||
:param block_structure: 2d block structure to use for the block
|
||||
strategy, must be of the format "2x4", "8x16", etc.
|
||||
:param dynamic: set True to perform dynamic quantization -
|
||||
values will not be calibrated during calibration phase,
|
||||
instead during inference new quantization ranges will be
|
||||
observed with every sample. Defaults to False for static
|
||||
quantization. Note that enabling dynamic quantization
|
||||
will change the default observer to a memoryless one
|
||||
"""
|
||||
|
||||
num_bits: int = 8
|
||||
type: QuantizationType = QuantizationType.INT
|
||||
symmetric: bool = True
|
||||
group_size: Optional[int] = None
|
||||
strategy: Optional[QuantizationStrategy] = None
|
||||
block_structure: Optional[str] = None
|
||||
dynamic: bool = False
|
||||
observer: str = Field(
|
||||
default="minmax",
|
||||
description=("The class to use to compute the quantization param - "
|
||||
"scale and zero-point'"),
|
||||
)
|
||||
observer_kwargs: Dict[str, Any] = Field(
|
||||
default_factory=dict,
|
||||
description=
|
||||
("optional dict of kwargs to be passed directly to torch quantization "
|
||||
"Observers constructor excluding quantization range or symmetry"),
|
||||
)
|
||||
|
||||
|
||||
def find_first_name_or_class_match(
|
||||
name: str,
|
||||
module: Module,
|
||||
targets: Iterable[str],
|
||||
check_contains: bool = False) -> Optional[str]:
|
||||
"""
|
||||
Helper function to map the quantization details listed in the config
|
||||
for a given list of targets against each model layer. First uses the
|
||||
layer name to try and find a match. If no name match is found, uses
|
||||
the layer class name. Returns None otherwise.
|
||||
|
||||
:param name: layer name
|
||||
:param module: torch.nn.Module
|
||||
:param targets: list of targets to match the layer against
|
||||
:param check_contains: whether or not to do a substring match
|
||||
"""
|
||||
|
||||
return _find_first_match(name, targets) or _find_first_match(
|
||||
module.__class__.__name__, targets, check_contains)
|
||||
|
||||
|
||||
def _find_first_match(value: str,
|
||||
targets: Iterable[str],
|
||||
check_contains: bool = False) -> Optional[str]:
|
||||
"""
|
||||
Returns first element of target that matches value either
|
||||
exactly or as a regex after 're:'. If check_contains is set to True,
|
||||
additionally checks if the target string is contained within the value.
|
||||
|
||||
:param value: string to compare the list of targets against
|
||||
:param targets: list of targets to match the layer against
|
||||
:param check_contains: whether or not to do a substring match
|
||||
"""
|
||||
|
||||
for target in targets:
|
||||
if target.startswith("re:"):
|
||||
pattern = target[3:]
|
||||
if re.match(pattern, value):
|
||||
return target
|
||||
elif check_contains:
|
||||
if target.lower() in value.lower():
|
||||
return target
|
||||
elif target == value:
|
||||
return target
|
||||
return None
|
||||
Reference in New Issue
Block a user