from typing import Literal, get_args from trainer.layers.quantization.base_config import QuantizationConfig QuantizationMethods = Literal[None] QUANTIZATION_METHODS: list[str] = list(get_args(QuantizationMethods)) # The customized quantization methods which will be added to this dict. _CUSTOMIZED_METHOD_TO_QUANT_CONFIG = {} def register_quantization_config(quantization: str): """Register a customized vllm quantization config. When a quantization method is not supported by vllm, you can register a customized quantization config to support it. Args: quantization (str): The quantization method name. Examples: >>> from trainer.layers.quantization import register_quantization_config >>> from trainer.layers.quantization import get_quantization_config >>> from trainer.layers.quantization.base_config import QuantizationConfig >>> >>> @register_quantization_config("my_quant") ... class MyQuantConfig(QuantizationConfig): ... pass >>> >>> get_quantization_config("my_quant") """ # noqa: E501 def _wrapper(quant_config_cls): if quantization in QUANTIZATION_METHODS: raise ValueError( f"The quantization method `{quantization}` is already exists.") if not issubclass(quant_config_cls, QuantizationConfig): raise ValueError("The quantization config must be a subclass of " "`QuantizationConfig`.") _CUSTOMIZED_METHOD_TO_QUANT_CONFIG[quantization] = quant_config_cls QUANTIZATION_METHODS.append(quantization) return quant_config_cls return _wrapper def get_quantization_config(quantization: str) -> type[QuantizationConfig]: if quantization not in QUANTIZATION_METHODS: raise ValueError(f"Invalid quantization method: {quantization}") method_to_config: dict[str, type[QuantizationConfig]] = {} # Update the `method_to_config` with customized quantization methods. method_to_config.update(_CUSTOMIZED_METHOD_TO_QUANT_CONFIG) return method_to_config[quantization] all = [ "QuantizationMethods", "QuantizationConfig", "get_quantization_config", "QUANTIZATION_METHODS", ]