neural_compressor.jax.quantization.config

The configs of algorithms for JAX.

Classes

OperatorConfig

Configuration pairing a quantization config with supported operators.

JaxBaseConfig

Shared base for JAX quant configs.

DynamicQuantConfig

Config class for JAX Dynamic quantization.

StaticQuantConfig

Config class for JAX Static quantization.

JaxComposableConfig

JAX composable config that is both a ComposableConfig and a JaxBaseConfig.

Functions

get_all_registered_configs(→ Dict[str, ...)

Get all registered configs for JAX framework.

get_default_dynamic_config(→ DynamicQuantConfig)

Generate the default Dynamic quantization config.

get_default_static_config(→ StaticQuantConfig)

Generate the default Static quantization config.

Module Contents

class neural_compressor.jax.quantization.config.OperatorConfig[source]

Configuration pairing a quantization config with supported operators.

class neural_compressor.jax.quantization.config.JaxBaseConfig(weight_dtype: str = 'fp8_e4m3', activation_dtype: str = 'fp8_e4m3', const_scale: bool = False, const_weight: bool = False, weight_scale_granularity: str = 'per_tensor', white_list: List[neural_compressor.common.base_config.OP_NAME_OR_MODULE_TYPE] | None = DEFAULT_WHITE_LIST, exclude_list: List[str] | None = None)[source]

Shared base for JAX quant configs.

Provides the common white_list / exclude_list selection, serialization and op-mapping behavior shared by DynamicQuantConfig and StaticQuantConfig. Subclasses supply their name, supported configs, tuning set, from_dict and the quantizable layer mapping.

class neural_compressor.jax.quantization.config.DynamicQuantConfig(weight_dtype: str = 'fp8_e4m3', activation_dtype: str = 'fp8_e4m3', const_scale: bool = False, const_weight: bool = False, weight_scale_granularity: str = 'per_tensor', white_list: List[neural_compressor.common.base_config.OP_NAME_OR_MODULE_TYPE] | None = DEFAULT_WHITE_LIST, exclude_list: List[str] | None = None)[source]

Config class for JAX Dynamic quantization.

Dynamic quantization applies quantization to both weights and activations during runtime. This configuration supports various data types for flexible quantization strategies.

Supported dtypes:
  • “fp8”: 8-bit floating-point quantization (uses ml_dtypes.float8_e4m3 by default)

  • “int8”: 8-bit integer quantization

FP8 formats available:
  • “fp8_e4m3”: 4 exponent bits, 3 mantissa bits (default for “fp8”)

  • “fp8_e5m2”: 5 exponent bits, 2 mantissa bits

class neural_compressor.jax.quantization.config.StaticQuantConfig(weight_dtype: str = 'fp8_e4m3', activation_dtype: str = 'fp8_e4m3', const_scale: bool = False, const_weight: bool = False, weight_scale_granularity: str = 'per_tensor', white_list: List[neural_compressor.common.base_config.OP_NAME_OR_MODULE_TYPE] | None = DEFAULT_WHITE_LIST, exclude_list: List[str] | None = None)[source]

Config class for JAX Static quantization.

Static quantization applies quantization to weights offline and activations during runtime using pre-computed calibration data. This configuration supports various data types for flexible quantization strategies.

Supported dtypes:
  • “fp8”: 8-bit floating-point quantization (uses ml_dtypes.float8_e4m3 by default)

  • “int8”: 8-bit integer quantization

FP8 formats available:
  • “fp8_e4m3”: 4 exponent bits, 3 mantissa bits (default for “fp8”)

  • “fp8_e5m2”: 5 exponent bits, 2 mantissa bits

class neural_compressor.jax.quantization.config.JaxComposableConfig(configs: List[BaseConfig])[source]

JAX composable config that is both a ComposableConfig and a JaxBaseConfig.

Composing behavior (__init__, __add__, serialization and the tuning stubs) is inherited from the common ComposableConfig, while JaxBaseConfig is also a base so every JAX config shares one ancestor. Base order matters: keeping ComposableConfig first ensures its composable __init__ / __add__ win over JaxBaseConfig’s single-config versions.

Only op selection is JAX-specific. Unlike the common composable – which keys model_info by config name and so collides when the same config type appears more than once (e.g. static + dynamic + static) – this subclass keeps each sub-config’s model_info positionally aligned and delegates op selection to the sub-config’s own (JAX-tweaked) to_config_mapping, merging results with last-config-wins on overlap.

neural_compressor.jax.quantization.config.get_all_registered_configs() Dict[str, neural_compressor.common.base_config.BaseConfig][source]

Get all registered configs for JAX framework.

Returns:

Mapping of config names to config classes.

Return type:

Dict[str, BaseConfig]

neural_compressor.jax.quantization.config.get_default_dynamic_config() DynamicQuantConfig[source]

Generate the default Dynamic quantization config.

Returns:

The default JAX Dynamic quantization config.

Return type:

DynamicQuantConfig

neural_compressor.jax.quantization.config.get_default_static_config() StaticQuantConfig[source]

Generate the default Static quantization config.

Returns:

The default JAX Static quantization config.

Return type:

StaticQuantConfig