From 7de3c9f41c4ff7c200802fc63fb5f38efe46ed82 Mon Sep 17 00:00:00 2001 From: usimha <135899523+u-simha@users.noreply.github.com> Date: Fri, 28 Aug 2026 16:15:23 -0700 Subject: [PATCH 1/2] Add joint post-training quantization/palettization + sparsity support Adds a hidden, settable-but-not-public `_sparsity` field to QuantizationSpec and PalettizationSpec: weights are pre-sparsified before fake-quant/palettize, and finalize() inserts sparse_to_dense (plus lut_to_dense for palettization) in the correct op order. Validation of unsupported combinations (asymmetric quant, unsigned/FP4 dtypes, quantized LUTs, per-channel scale, non-per-tensor granularity) now happens at spec-construction time via pydantic. --- .../kmeans/_prepare_for_export.py | 49 +++++++- .../palettization/kmeans/palettizer.py | 1 + .../palettization/spec/fake_palettize.py | 10 ++ src/coreai_opt/palettization/spec/spec.py | 21 ++++ .../_graph/_prepare_for_export.py | 31 ++++- src/coreai_opt/quantization/spec/factory.py | 2 + .../quantization/spec/fake_quantize.py | 10 ++ src/coreai_opt/quantization/spec/spec.py | 26 ++++ tests/export/test_joint_sparsity.py | 119 ++++++++++++++++++ 9 files changed, 262 insertions(+), 7 deletions(-) create mode 100644 tests/export/test_joint_sparsity.py diff --git a/src/coreai_opt/palettization/kmeans/_prepare_for_export.py b/src/coreai_opt/palettization/kmeans/_prepare_for_export.py index 678188d..81b9c74 100644 --- a/src/coreai_opt/palettization/kmeans/_prepare_for_export.py +++ b/src/coreai_opt/palettization/kmeans/_prepare_for_export.py @@ -6,6 +6,7 @@ from dataclasses import dataclass from os import PathLike from pathlib import Path +from typing import Any import torch import torch.nn as nn @@ -46,6 +47,32 @@ class PalettizationInfo: lut_quantization: LUTQuantizationInfo | None = None +class _SparsePalettizeReconstruction(nn.Module): + """Parametrization module inserted to reconstruct a sparse-palettized weight. + + Traces ``coreai.lut_to_dense`` and ``coreai.sparse_to_dense`` in that order. + """ + + def __init__( + self, + nonzero_indices: torch.Tensor, + lut: torch.Tensor, + mask: torch.Tensor, + vector_axis: int | None, + ) -> None: + super().__init__() + self.register_buffer("nonzero_indices", nonzero_indices) + self.register_buffer("lut", lut.reshape(1, lut.shape[-2], lut.shape[-1])) + self.register_buffer("mask", mask) + self.vector_axis = 0 if vector_axis is None else vector_axis + + def forward(self, _: Any) -> torch.Tensor: + nonzero_values = torch.ops.coreai.lut_to_dense( + self.nonzero_indices, self.lut, self.vector_axis + ) + return torch.ops.coreai.sparse_to_dense(nonzero_values, self.mask) + + def _expand_rank( tensor: torch.Tensor, target_rank: int, @@ -210,6 +237,7 @@ def _insert_mlir_custom_op( module_name: str, param_name: str, palett_info: PalettizationInfo, + fake_palett_mod: _FakePalettizeImplBase, fake_palett_idx: int, mmap_dir: str | PathLike[str] | None, ) -> None: @@ -227,6 +255,10 @@ def _insert_mlir_custom_op( 4. Both: lut_to_dense(int LUT) + constexpr_blockwise_shift_scale(fused_scale) where fused_scale = lut_scale * per_channel_scale + When ``fake_palett_mod.sparsity`` is set, the LUT lookup runs on the + nonzero-only indices and the result is packed via ``coreai::sparse_to_dense`` + instead of installing a plain Palettize/ScaledPalettize parametrization. + When ``mmap_dir`` is provided, the new MLIR module is serialized to a safetensors file under that directory and reloaded via mmap before being swapped in. @@ -261,7 +293,20 @@ def _import_coreai_torch_modules(): vector_axis = _DEFAULT_VECTOR_AXIS if palett_info.cluster_dim > 1 else None - if needs_scale: + if fake_palett_mod.sparsity is not None: + # Reuses the mask from prepare()'s forward pass. needs_scale is always + # False here: PalettizationSpec rejects lut_qspec/enable_per_channel_scale + # combined with sparsity, since both are position-dependent and would + # be scrambled by flattening to the nonzero-only indices below. + mask = fake_palett_mod._sparsity_mask.to(torch.bool) + nonzero_indices = palett_info.indices[mask] + mlir_palett_mod = _SparsePalettizeReconstruction( + nonzero_indices=nonzero_indices, + lut=palett_info.lut, + mask=mask, + vector_axis=vector_axis, + ) + elif needs_scale: lut, scale, zero_point = _resolve_mlir_lut_and_scale(palett_info) mlir_palett_mod = ScaledPalettizeParametrization( indices=palett_info.indices, @@ -336,7 +381,7 @@ def _process_palettized_parameter( _register_mil_compression_metadata(module, param_name, palett_info) elif backend == ExportBackend.CoreAI: _insert_mlir_custom_op( - module, module_name, param_name, palett_info, fake_palett_idx, mmap_dir + module, module_name, param_name, palett_info, fake_palett_mod, fake_palett_idx, mmap_dir ) diff --git a/src/coreai_opt/palettization/kmeans/palettizer.py b/src/coreai_opt/palettization/kmeans/palettizer.py index dee63ab..b36e529 100644 --- a/src/coreai_opt/palettization/kmeans/palettizer.py +++ b/src/coreai_opt/palettization/kmeans/palettizer.py @@ -432,6 +432,7 @@ def _spec_to_partial( # Serialize the spec, then layer in the owning module's compressor-specific # settings (e.g. enable_fast_kmeans_mode, rounding_precision). args = spec.model_dump_preserve_objects() + args["sparsity"] = spec._sparsity args.update(module_config._get_compressor_specific_settings()) return _KMeansFakePalettize.with_args(**args) diff --git a/src/coreai_opt/palettization/spec/fake_palettize.py b/src/coreai_opt/palettization/spec/fake_palettize.py index ee4f7e3..c5538d8 100644 --- a/src/coreai_opt/palettization/spec/fake_palettize.py +++ b/src/coreai_opt/palettization/spec/fake_palettize.py @@ -23,6 +23,7 @@ _IncompatibleClusterDimError, _IncompatibleGranularityError, ) +from coreai_opt.pruning.spec import PruneImplBase, Unstructured from coreai_opt.quantization.spec import QuantizationSpec logger = logging.getLogger(__name__) @@ -49,6 +50,7 @@ def __init__( granularity: PalettizationGranularity, cluster_dim: int, enable_per_channel_scale: bool, + sparsity: float | None = None, **kwargs, ): super().__init__(**kwargs) @@ -57,10 +59,12 @@ def __init__( self.granularity = granularity self.cluster_dim = cluster_dim self.enable_per_channel_scale = enable_per_channel_scale + self.sparsity = sparsity self.register_buffer("fake_palett_enabled", torch.tensor([1], dtype=torch.uint8)) self.register_buffer("observer_enabled", torch.tensor([1], dtype=torch.uint8)) self._disabled = False + self.register_buffer("_sparsity_mask", None, persistent=False) self.register_buffer("lut", None) self.register_buffer("indices", None) @@ -84,6 +88,12 @@ def forward(self, tensor: torch.Tensor) -> torch.Tensor: if self._disabled: return tensor + if self.sparsity is not None: + self._sparsity_mask = PruneImplBase.resolve("default").compute_mask( + tensor, self.sparsity, Unstructured() + ) + tensor = tensor * self._sparsity_mask + if self.observer_enabled[0] == 1: # Cluster weights try: diff --git a/src/coreai_opt/palettization/spec/spec.py b/src/coreai_opt/palettization/spec/spec.py index d3b4734..d095c38 100644 --- a/src/coreai_opt/palettization/spec/spec.py +++ b/src/coreai_opt/palettization/spec/spec.py @@ -95,6 +95,27 @@ class PalettizationSpec(CompressionSpec): # Private attribute for compression type _compression_type: CompressionType = PrivateAttr(default=CompressionType.PALETTIZATION) + # Sparsity level, in [0, 1]. Set via the `_sparsity` constructor/dict key. + _sparsity: float | None = PrivateAttr(default=None) + + def __init__(self, **data: Any) -> None: + sparsity = data.pop("_sparsity", None) + super().__init__(**data) + if sparsity is not None: + if not (0.0 <= sparsity <= 1.0): + raise ValueError(f"_sparsity must be in [0, 1], got {sparsity}") + self._validate_sparsity(sparsity) + self._sparsity = sparsity + + def _validate_sparsity(self, sparsity: float) -> None: + """Reject sparsity combined with a position-dependent LUT/scale mapping.""" + if self.lut_qspec is not None: + raise ValueError("lut_qspec not supported for joint sparsity.") + if not isinstance(self.granularity, PerTensorGranularity): + raise ValueError(f"granularity={self.granularity} not supported for joint sparsity.") + if self.enable_per_channel_scale: + raise ValueError("enable_per_channel_scale not supported for joint sparsity.") + @model_validator(mode="after") def validate_lut_qspec(self) -> "PalettizationSpec": """Validate that lut_qspec only uses supported configurations.""" diff --git a/src/coreai_opt/quantization/_graph/_prepare_for_export.py b/src/coreai_opt/quantization/_graph/_prepare_for_export.py index 60f1545..55e6bda 100644 --- a/src/coreai_opt/quantization/_graph/_prepare_for_export.py +++ b/src/coreai_opt/quantization/_graph/_prepare_for_export.py @@ -272,14 +272,20 @@ def _import_coreai_custom_ops(): # Extract and prepare quantization parameters scale, zero_point, minval = _extract_quantization_params(fake_quant_mod) + # Cast scale and minval to appropriate dtype for MLIR backend inference _compute_dtype_for_export = fake_quant_mod.qparams_calculator._compute_dtype_for_export scale = scale.to(dtype=_compute_dtype_for_export) if minval is not None: minval = minval.to(dtype=_compute_dtype_for_export) - # Construct quantized weights + # Construct quantized weights, reusing the mask computed during + # prepare()'s forward pass if sparsity is set. dense_weight = resolve_attr(model, input_node.target).data + mask: torch.Tensor | None = None + if fake_quant_mod.sparsity is not None: + mask = fake_quant_mod._sparsity_mask.to(torch.bool) + dense_weight = dense_weight * mask quantized_data = fake_quant_mod.quantize(dense_weight, scale, zero_point, minval) # Drop one of the offsets so that the export @@ -295,13 +301,28 @@ def _import_coreai_custom_ops(): # Register buffers and get buffer names param_name = str(input_node.target).replace(".", "_") - buffer_names = _register_quantization_buffers( - model, param_name, scale, zero_point, quantized_data, minval - ) + if mask is not None: + nonzero_data = quantized_data[mask] + buffer_names = _register_quantization_buffers( + model, param_name, scale, zero_point, minval=minval + ) + model.register_buffer(f"{param_name}_nonzero", nonzero_data) + model.register_buffer(f"{param_name}_mask", mask) + else: + buffer_names = _register_quantization_buffers( + model, param_name, scale, zero_point, quantized_data, minval + ) # Create graph nodes and replace fake quantization with model.graph.inserting_before(node): - quantized_data_node = model.graph.get_attr(buffer_names["quantized_data"]) + if mask is not None: + nonzero_node = model.graph.get_attr(f"{param_name}_nonzero") + mask_node = model.graph.get_attr(f"{param_name}_mask") + quantized_data_node = model.graph.call_function( + coreai.sparse_to_dense, (nonzero_node, mask_node) + ) + else: + quantized_data_node = model.graph.get_attr(buffer_names["quantized_data"]) scale_node = model.graph.get_attr(buffer_names["scale"]) if zero_point is not None: diff --git a/src/coreai_opt/quantization/spec/factory.py b/src/coreai_opt/quantization/spec/factory.py index 9a449b0..928f944 100644 --- a/src/coreai_opt/quantization/spec/factory.py +++ b/src/coreai_opt/quantization/spec/factory.py @@ -203,6 +203,7 @@ def create_fake_quantizer( "qparams_calculator": qparams_calculator, "quantization_target": quantization_target, "n_bits": spec.n_bits, + "sparsity": spec._sparsity, } # Automatically detect and include any extra arguments @@ -244,6 +245,7 @@ def create_fake_quantizer_partial( "quant_max": spec.quant_max, "quantization_target": quantization_target, "n_bits": spec.n_bits, + "sparsity": spec._sparsity, } # Automatically detect and include any extra arguments diff --git a/src/coreai_opt/quantization/spec/fake_quantize.py b/src/coreai_opt/quantization/spec/fake_quantize.py index aafd476..5d50b1e 100644 --- a/src/coreai_opt/quantization/spec/fake_quantize.py +++ b/src/coreai_opt/quantization/spec/fake_quantize.py @@ -26,6 +26,7 @@ is_float_quant_dtype as _is_float_quant_dtype, ) from coreai_opt.config.spec import CompressionSimulatorBase, CompressionTargetTensor +from coreai_opt.pruning.spec import PruneImplBase, Unstructured from coreai_opt.quantization._utils import get_quantization_shapes as _get_quantization_shapes from coreai_opt.quantization.spec.errors import _BlockSizeMismatchError @@ -56,6 +57,7 @@ def __init__( qparams_calculator: QParamsCalculatorBase, quantization_target: CompressionTargetTensor, n_bits: int | None = None, + sparsity: float | None = None, **kwargs, ): super().__init__() @@ -68,7 +70,9 @@ def __init__( self.quant_max = quant_max self.qparams_calculator = qparams_calculator self.quantization_target = quantization_target + self.sparsity = sparsity self.register_buffer("_disabled", torch.tensor(False)) + self.register_buffer("_sparsity_mask", None, persistent=False) # Infer n_bits from dtype if not provided if n_bits is None: @@ -140,6 +144,12 @@ def forward(self, tensor: torch.Tensor) -> torch.Tensor: if self._disabled.item(): return tensor + if self.sparsity is not None: + self._sparsity_mask = PruneImplBase.resolve("default").compute_mask( + tensor, self.sparsity, Unstructured() + ) + tensor = tensor * self._sparsity_mask + if self.observer_enabled[0] == 1: # Call the forward function of the qparams_calculator # to collect observer statistics when the observer is diff --git a/src/coreai_opt/quantization/spec/spec.py b/src/coreai_opt/quantization/spec/spec.py index 52a8372..b724fd1 100644 --- a/src/coreai_opt/quantization/spec/spec.py +++ b/src/coreai_opt/quantization/spec/spec.py @@ -372,6 +372,18 @@ class type: MinMaxRangeCalculator or custom registered class type # Private attribute for compression type _compression_type: CompressionType = PrivateAttr(default=CompressionType.QUANTIZATION) + # Sparsity level, in [0, 1]. Set via the `_sparsity` constructor/dict key. + _sparsity: float | None = PrivateAttr(default=None) + + def __init__(self, **data: Any) -> None: + sparsity = data.pop("_sparsity", None) + super().__init__(**data) + if sparsity is not None: + if not (0.0 <= sparsity <= 1.0): + raise ValueError(f"_sparsity must be in [0, 1], got {sparsity}") + self._validate_sparsity_zero_preserving(sparsity) + self._sparsity = sparsity + # Supported dtypes for quantization (class attribute for testing extensibility) SUPPORTED_DTYPES: ClassVar[set[torch.dtype]] = { # Signed integer types @@ -555,6 +567,20 @@ def validate_scale_dtype(self) -> QuantizationSpec: return self + def _validate_sparsity_zero_preserving(self, sparsity: float) -> None: + """Reject sparsity unless a raw 0 dequantizes to exactly 0.0.""" + if _is_float4_dtype(self.dtype): + raise ValueError("FP4 dtype not supported for joint sparsity.") + if self.dtype.is_floating_point: + return + + if self.qformulation != QuantizationFormulation.ZP: + raise ValueError(f"qformulation={self.qformulation} not supported for joint sparsity.") + if self.qscheme == QuantizationScheme.ASYMMETRIC: + raise ValueError(f"qscheme={self.qscheme} not supported for joint sparsity.") + if not self.dtype.is_signed: + raise ValueError(f"unsigned dtype={self.dtype} not supported for joint sparsity.") + def get_extra_args(self) -> dict[str, Any]: """ Automatically detect and return fields beyond base QuantizationSpec. diff --git a/tests/export/test_joint_sparsity.py b/tests/export/test_joint_sparsity.py new file mode 100644 index 0000000..5957d59 --- /dev/null +++ b/tests/export/test_joint_sparsity.py @@ -0,0 +1,119 @@ +# Copyright 2026 Apple Inc. +# +# Use of this source code is governed by a BSD-3-Clause license that can +# be found in the LICENSE file or at https://opensource.org/licenses/BSD-3-Clause + +"""End-to-end export tests for joint post-training quantization/palettization + sparsity.""" + +import pytest +import torch +import torch.nn as nn + +from coreai_opt import ExportBackend +from coreai_opt.palettization import ( + KMeansPalettizer, + KMeansPalettizerConfig, + ModuleKMeansPalettizerConfig, + PalettizationSpec, +) +from coreai_opt.quantization import ModuleQuantizerConfig, Quantizer, QuantizerConfig +from coreai_opt.quantization.config import ExecutionMode +from coreai_opt.quantization.spec import PerTensorGranularity, QuantizationScheme, QuantizationSpec + +from . import export_utils + + +class TestJointSparsityExport: + """PTQ/PTP + PTS (post-training quantization/palettization + sparsity), end to end.""" + + @staticmethod + def _run_quant_sparsity_export( + model: nn.Module, input_data: torch.Tensor, expected_count: int + ) -> None: + model.eval() + config = QuantizerConfig( + global_config=ModuleQuantizerConfig( + op_state_spec={ + "weight": QuantizationSpec( + dtype=torch.int8, + qscheme=QuantizationScheme.SYMMETRIC, + granularity=PerTensorGranularity(), + _sparsity=0.5, + ) + }, + op_input_spec=None, + op_output_spec=None, + ), + execution_mode=ExecutionMode.GRAPH, + ) + + quantizer = Quantizer(model, config) + prepared_model = quantizer.prepare((input_data,)) + + with torch.no_grad(): + prepared_model_output = prepared_model(input_data) + + finalized_model = quantizer.finalize(backend=ExportBackend.CoreAI) + + export_utils.convert_and_verify( + finalized_model=finalized_model, + input_data=input_data, + expected_ops={ + "sparse_to_dense": expected_count, + "constexpr_blockwise_shift_scale": expected_count, + }, + export_backend=ExportBackend.CoreAI, + prepared_model_output=prepared_model_output, + ) + + @staticmethod + def _run_palettization_sparsity_export( + model: nn.Module, input_data: torch.Tensor, expected_count: int + ) -> None: + model.eval() + config = KMeansPalettizerConfig( + global_config=ModuleKMeansPalettizerConfig( + op_state_spec={"weight": PalettizationSpec(n_bits=8, _sparsity=0.5)} + ) + ) + + palettizer = KMeansPalettizer(model, config) + prepared_model = palettizer.prepare((input_data,)) + + with torch.no_grad(): + prepared_model_output = prepared_model(input_data) + + finalized_model = palettizer.finalize(backend=ExportBackend.CoreAI) + + export_utils.convert_and_verify( + finalized_model=finalized_model, + input_data=input_data, + expected_ops={ + "lut_to_dense": expected_count, + "sparse_to_dense": expected_count, + }, + export_backend=ExportBackend.CoreAI, + prepared_model_output=prepared_model_output, + ) + + def test_quant_sparsity_mnist_export(self, custom_test_mnist_model, mnist_example_input): + self._run_quant_sparsity_export( + custom_test_mnist_model, mnist_example_input, expected_count=6 + ) + + @pytest.mark.slow + def test_quant_sparsity_resnet_export(self, resnet50_model, resnet_example_input): + self._run_quant_sparsity_export(resnet50_model, resnet_example_input, expected_count=54) + + def test_palettization_sparsity_mnist_export( + self, custom_test_mnist_model, mnist_example_input + ): + self._run_palettization_sparsity_export( + custom_test_mnist_model, mnist_example_input, expected_count=6 + ) + + @pytest.mark.slow + def test_palettization_sparsity_resnet_export(self, resnet50_model, resnet_example_input): + self._run_palettization_sparsity_export( + resnet50_model, resnet_example_input, expected_count=54 + ) From 7ce24def1431d2db51d9c9f947b5b88a0ccec06e Mon Sep 17 00:00:00 2001 From: usimha <135899523+u-simha@users.noreply.github.com> Date: Mon, 31 Aug 2026 14:43:45 -0700 Subject: [PATCH 2/2] Trim sparsity docstrings/comments and consolidate range validation Moves the [0, 1] range check for _sparsity into the existing _validate_sparsity(_zero_preserving) methods so each spec has a single validation entrypoint, and shortens a couple of over-long comments. --- .../palettization/kmeans/_prepare_for_export.py | 8 ++++---- src/coreai_opt/palettization/spec/spec.py | 4 ++-- src/coreai_opt/quantization/spec/spec.py | 4 ++-- 3 files changed, 8 insertions(+), 8 deletions(-) diff --git a/src/coreai_opt/palettization/kmeans/_prepare_for_export.py b/src/coreai_opt/palettization/kmeans/_prepare_for_export.py index 81b9c74..20bc2f7 100644 --- a/src/coreai_opt/palettization/kmeans/_prepare_for_export.py +++ b/src/coreai_opt/palettization/kmeans/_prepare_for_export.py @@ -62,6 +62,8 @@ def __init__( ) -> None: super().__init__() self.register_buffer("nonzero_indices", nonzero_indices) + # nonzero_indices is rank 1 (flattened by masking), so lut must be + # reshaped to rank 3 (lut_to_dense requires lut.rank == indices.rank + 2). self.register_buffer("lut", lut.reshape(1, lut.shape[-2], lut.shape[-1])) self.register_buffer("mask", mask) self.vector_axis = 0 if vector_axis is None else vector_axis @@ -294,10 +296,8 @@ def _import_coreai_torch_modules(): vector_axis = _DEFAULT_VECTOR_AXIS if palett_info.cluster_dim > 1 else None if fake_palett_mod.sparsity is not None: - # Reuses the mask from prepare()'s forward pass. needs_scale is always - # False here: PalettizationSpec rejects lut_qspec/enable_per_channel_scale - # combined with sparsity, since both are position-dependent and would - # be scrambled by flattening to the nonzero-only indices below. + # Reuses the mask from prepare()'s forward pass. needs_scale is always False here: + # PalettizationSpec rejects lut_qspec/enable_per_channel_scale combined with sparsity. mask = fake_palett_mod._sparsity_mask.to(torch.bool) nonzero_indices = palett_info.indices[mask] mlir_palett_mod = _SparsePalettizeReconstruction( diff --git a/src/coreai_opt/palettization/spec/spec.py b/src/coreai_opt/palettization/spec/spec.py index d095c38..798c0f7 100644 --- a/src/coreai_opt/palettization/spec/spec.py +++ b/src/coreai_opt/palettization/spec/spec.py @@ -102,13 +102,13 @@ def __init__(self, **data: Any) -> None: sparsity = data.pop("_sparsity", None) super().__init__(**data) if sparsity is not None: - if not (0.0 <= sparsity <= 1.0): - raise ValueError(f"_sparsity must be in [0, 1], got {sparsity}") self._validate_sparsity(sparsity) self._sparsity = sparsity def _validate_sparsity(self, sparsity: float) -> None: """Reject sparsity combined with a position-dependent LUT/scale mapping.""" + if not (0.0 <= sparsity <= 1.0): + raise ValueError(f"_sparsity must be in [0, 1], got {sparsity}") if self.lut_qspec is not None: raise ValueError("lut_qspec not supported for joint sparsity.") if not isinstance(self.granularity, PerTensorGranularity): diff --git a/src/coreai_opt/quantization/spec/spec.py b/src/coreai_opt/quantization/spec/spec.py index b724fd1..e776c21 100644 --- a/src/coreai_opt/quantization/spec/spec.py +++ b/src/coreai_opt/quantization/spec/spec.py @@ -379,8 +379,6 @@ def __init__(self, **data: Any) -> None: sparsity = data.pop("_sparsity", None) super().__init__(**data) if sparsity is not None: - if not (0.0 <= sparsity <= 1.0): - raise ValueError(f"_sparsity must be in [0, 1], got {sparsity}") self._validate_sparsity_zero_preserving(sparsity) self._sparsity = sparsity @@ -569,6 +567,8 @@ def validate_scale_dtype(self) -> QuantizationSpec: def _validate_sparsity_zero_preserving(self, sparsity: float) -> None: """Reject sparsity unless a raw 0 dequantizes to exactly 0.0.""" + if not (0.0 <= sparsity <= 1.0): + raise ValueError(f"_sparsity must be in [0, 1], got {sparsity}") if _is_float4_dtype(self.dtype): raise ValueError("FP4 dtype not supported for joint sparsity.") if self.dtype.is_floating_point: