diff --git a/.azure-pipelines/scripts/ai_analysis/ai_analyze.py b/.azure-pipelines/scripts/ai_analysis/ai_analyze.py index a261f4a899..693361a1ca 100644 --- a/.azure-pipelines/scripts/ai_analysis/ai_analyze.py +++ b/.azure-pipelines/scripts/ai_analysis/ai_analyze.py @@ -155,7 +155,7 @@ def _call_copilot_cli(prompt: str, timeout: int, model: str, reasoning_effort: s if ai_token: env.setdefault("GITHUB_TOKEN", ai_token) env.setdefault("GH_TOKEN", ai_token) - started = datetime.datetime.now(datetime.timezone.utc) + started = datetime.datetime.now(datetime.UTC) returncode = None stderr = "" stdout = "" @@ -169,7 +169,7 @@ def _call_copilot_cli(prompt: str, timeout: int, model: str, reasoning_effort: s except (OSError, subprocess.TimeoutExpired) as e: print(f"Warning: Copilot CLI call failed: {e}", file=sys.stderr) stderr = str(e) - ended = datetime.datetime.now(datetime.timezone.utc) + ended = datetime.datetime.now(datetime.UTC) parsed = _parse_copilot_json(stdout) raw = parsed["answer"] _append_trace( diff --git a/.azure-pipelines/scripts/compat_smoke_test.py b/.azure-pipelines/scripts/compat_smoke_test.py index 3c48e701ab..5419526d49 100644 --- a/.azure-pipelines/scripts/compat_smoke_test.py +++ b/.azure-pipelines/scripts/compat_smoke_test.py @@ -21,7 +21,7 @@ print(f"auto_round {auto_round.__version__} imported successfully (AutoRound={AutoRound.__name__})") # Verify the console_scripts entry point was installed and is runnable. -result = subprocess.run(["auto-round", "--help"], capture_output=True, text=True) +result = subprocess.run(["auto-round", "--help"], capture_output=True, text=True, check=False) if result.returncode != 0: sys.stderr.write(result.stdout) sys.stderr.write(result.stderr) diff --git a/.azure-pipelines/scripts/performance/check_performance.py b/.azure-pipelines/scripts/performance/check_performance.py index bdb74f6e85..ce4940f9d2 100644 --- a/.azure-pipelines/scripts/performance/check_performance.py +++ b/.azure-pipelines/scripts/performance/check_performance.py @@ -3,9 +3,9 @@ import sys from dataclasses import dataclass from pathlib import Path -from typing import Dict, Optional logging.basicConfig(level=logging.INFO, format="%(message)s") +logger = logging.getLogger(__name__) LOG_DIR = Path("/auto-round/log_dir") OUTPUT_BASE_DIR = Path("/auto-round/.azure-pipelines/scripts/performance") @@ -13,10 +13,10 @@ @dataclass class QuantMetrics: - tuning_time_s: Optional[float] = None - peak_ram_gb: Optional[float] = None - peak_vram_gb: Optional[float] = None - output_size_gb: Optional[float] = None + tuning_time_s: float | None = None + peak_ram_gb: float | None = None + peak_vram_gb: float | None = None + output_size_gb: float | None = None def get_dir_size_gb(path: Path) -> float: @@ -31,7 +31,7 @@ def parse_log_file(log_file: Path) -> QuantMetrics: metrics = QuantMetrics() if not log_file.exists(): - logging.warning(f"Log file not found: {log_file}") + logger.warning(f"Log file not found: {log_file}") return metrics content = log_file.read_text(encoding="utf-8") @@ -50,7 +50,7 @@ def parse_log_file(log_file: Path) -> QuantMetrics: return metrics -def get_tuning_info() -> Dict[str, Dict[str, QuantMetrics]]: +def get_tuning_info() -> dict[str, dict[str, QuantMetrics]]: summary = {} model_list = ["Qwen/Qwen3-0.6B"] @@ -60,7 +60,7 @@ def get_tuning_info() -> Dict[str, Dict[str, QuantMetrics]]: log_file = LOG_DIR / f"perf_test_{test_mode}.log" output_dir = OUTPUT_BASE_DIR / test_mode - logging.info(f"Processing {log_file}...") + logger.info(f"Processing {log_file}...") metrics = parse_log_file(log_file) metrics.output_size_gb = get_dir_size_gb(output_dir) @@ -70,15 +70,13 @@ def get_tuning_info() -> Dict[str, Dict[str, QuantMetrics]]: return summary -def compare_metric( - metric_name: str, current: Optional[float], baseline: Optional[float], tolerance: float = 0.1 -) -> bool: +def compare_metric(metric_name: str, current: float | None, baseline: float | None, tolerance: float = 0.1) -> bool: if current is None or baseline is None: - logging.error(f" [-] {metric_name}: Incomplete data (Current: {current}, Baseline: {baseline})") + logger.error(f" [-] {metric_name}: Incomplete data (Current: {current}, Baseline: {baseline})") return False if baseline == 0: - logging.warning(f" [!] {metric_name}: Baseline is 0, cannot calculate ratio.") + logger.warning(f" [!] {metric_name}: Baseline is 0, cannot calculate ratio.") return False ratio = current / baseline @@ -87,10 +85,10 @@ def compare_metric( msg = f" [*] {metric_name:<20}: Current = {current:<8} | Baseline = {baseline:<8} (Diff: {diff_percent:+.2f}%)" if 1.0 - tolerance <= ratio <= 1.0 + tolerance: - logging.info(f"{msg} -> PASS") + logger.info(f"{msg} -> PASS") return True else: - logging.error(f"{msg} -> FAIL") + logger.error(f"{msg} -> FAIL") return False @@ -99,8 +97,8 @@ def check_performance(): all_passed = True for model, modes in summary.items(): - logging.info(f"\nEvaluating Model: {model}") - logging.info("-" * 60) + logger.info(f"\nEvaluating Model: {model}") + logger.info("-" * 60) current: QuantMetrics = modes.get("current", QuantMetrics()) baseline: QuantMetrics = modes.get("baseline", QuantMetrics()) @@ -117,11 +115,11 @@ def check_performance(): if not compare_metric("Output Size (GB)", current.output_size_gb, baseline.output_size_gb, tolerance=0.01): all_passed = False - logging.info("=" * 60) + logger.info("=" * 60) if all_passed: - logging.info("✅ Performance check passed: All metrics are within acceptable limits.") + logger.info("✅ Performance check passed: All metrics are within acceptable limits.") else: - logging.error("❌ Performance check failed: Current metrics exceed acceptable limits compared to baseline.") + logger.error("❌ Performance check failed: Current metrics exceed acceptable limits compared to baseline.") sys.exit(1) diff --git a/.azure-pipelines/scripts/ut/collect_result.py b/.azure-pipelines/scripts/ut/collect_result.py index 0b23ad9e23..ebc20529d0 100644 --- a/.azure-pipelines/scripts/ut/collect_result.py +++ b/.azure-pipelines/scripts/ut/collect_result.py @@ -181,8 +181,10 @@ def generate(self, test_type: str, results: list[TestResult]) -> None: self._format_subheader(), *(self._format_row(r) for r in results), self.SEPARATOR, - f"Total: {stats['total']}, Passed: {stats['passed']}, " - f"Failed: {stats['failed']}, Skipped: {stats['skipped']}", + ( + f"Total: {stats['total']}, Passed: {stats['passed']}, " + f"Failed: {stats['failed']}, Skipped: {stats['skipped']}" + ), self.SEPARATOR, "", ] diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 6b4e7d12a5..278174a6e1 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -73,7 +73,7 @@ repos: exclude: ^auto_round/export/export_to_gguf/conversion/.*\.py$ - repo: https://github.com/astral-sh/ruff-pre-commit - rev: v0.15.20 + rev: v0.16.9 hooks: - id: ruff args: [--fix, --exit-non-zero-on-fix, --no-cache] diff --git a/auto_round/__init__.py b/auto_round/__init__.py index d3f87c2c4b..fce0c0b4e5 100644 --- a/auto_round/__init__.py +++ b/auto_round/__init__.py @@ -31,20 +31,20 @@ from .version import __version__ __all__ = [ - "__version__", + "AWQConfig", + "AdamRoundConfig", "AutoRound", - "AutoRoundLLM", - "AutoRoundMLLM", "AutoRoundAdam", "AutoRoundDiffusion", + "AutoRoundLLM", + "AutoRoundMLLM", "AutoScheme", + "OptimizedRTNConfig", "QuantizationScheme", "RTNConfig", - "OptimizedRTNConfig", + "RotationConfig", "SignRoundConfig", - "AdamRoundConfig", "SignRoundV2Config", - "AWQConfig", - "RotationConfig", "SpinQuantConfig", + "__version__", ] diff --git a/auto_round/algorithms/base.py b/auto_round/algorithms/base.py index b3fd725355..9940a47ff9 100644 --- a/auto_round/algorithms/base.py +++ b/auto_round/algorithms/base.py @@ -72,7 +72,7 @@ def __init__(self, config: Any = None) -> None: self.config = config # Name-mangled so subclasses cannot accidentally overwrite the run context. self.__run_ctx: QuantizationRunContext | None = None - self.__block_forward_runner: "BlockForwardRunner | None" = None + self.__block_forward_runner: BlockForwardRunner | None = None @classmethod def from_config(cls, config: Any) -> "BaseAlgorithm": diff --git a/auto_round/algorithms/block_runner.py b/auto_round/algorithms/block_runner.py index e5b5462fe1..90de327b3e 100644 --- a/auto_round/algorithms/block_runner.py +++ b/auto_round/algorithms/block_runner.py @@ -23,7 +23,7 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, Union +from typing import TYPE_CHECKING, Any import torch @@ -115,8 +115,8 @@ def __init__( self, batch_dim: int = 0, batch_size: int = 8, - device: Union[str, "torch.device"] = "cpu", - cache_device: Union[str, "torch.device"] = "cpu", + device: str | torch.device = "cpu", + cache_device: str | torch.device = "cpu", amp: bool = True, amp_dtype: torch.dtype | None = None, is_diffusion: bool = False, @@ -144,7 +144,7 @@ def __init__( # ── Factory ────────────────────────────────────────────────────────────── @classmethod - def from_orchestrator(cls, orchestrator: "BaseOrchestrator", enable_torch_compile=True) -> "BlockForwardRunner": + def from_orchestrator(cls, orchestrator: BaseOrchestrator, enable_torch_compile=True) -> BlockForwardRunner: """Create from an orchestrator instance (called once at orchestrator init).""" model_ctx = getattr(orchestrator, "model_context", None) is_diffusion = getattr(model_ctx, "is_diffusion", False) if model_ctx else False @@ -170,7 +170,7 @@ def __call__(self, *args, **kwargs) -> torch.Tensor: def forward( self, - block: "torch.nn.Module", + block: torch.nn.Module, inputs: list[torch.Tensor] | dict, input_others: dict, indices: torch.Tensor | None = None, @@ -310,7 +310,7 @@ def _count_samples(self, inputs: Any) -> int: else: return inputs.shape[self.batch_dim] - def _normalize_output(self, output: Any, block: "torch.nn.Module" = None) -> torch.Tensor: + def _normalize_output(self, output: Any, block: torch.nn.Module = None) -> torch.Tensor: """Normalize block output to a single tensor.""" if isinstance(output, torch.Tensor): return output @@ -343,7 +343,7 @@ def _normalize_output(self, output: Any, block: "torch.nn.Module" = None) -> tor return first raise TypeError(f"Block output[0] must be tensor, got {type(first).__name__}.") - def _get_output_dict(self, output: Any, block: "torch.nn.Module" = None) -> dict[str, torch.Tensor] | None: + def _get_output_dict(self, output: Any, block: torch.nn.Module = None) -> dict[str, torch.Tensor] | None: if isinstance(output, torch.Tensor) or not isinstance(output, (tuple, list)): return None block_cls_name = block.__class__.__name__ if block is not None else None diff --git a/auto_round/algorithms/composer.py b/auto_round/algorithms/composer.py index 725ead4fc5..67f1aeb2d1 100644 --- a/auto_round/algorithms/composer.py +++ b/auto_round/algorithms/composer.py @@ -70,7 +70,7 @@ class BlockContext: ``ValueError`` with a user-readable message. """ - model: "torch.nn.Module" + model: torch.nn.Module block_names: list[str] # scheduling group; len > 1 when nblocks > 1 block_name: str # = block_names[0] for single-block; descriptive label for multi block_index: int # 0-based index within the current all_blocks group @@ -98,7 +98,7 @@ class AlgorithmComposer: composer = AlgorithmComposer(configs, compressor=self) """ - def __init__(self, configs: list, orchestrator: "BaseOrchestrator" = None) -> None: + def __init__(self, configs: list, orchestrator: BaseOrchestrator = None) -> None: """Build the pipeline from a list of algorithm config instances. Resolution rules: @@ -240,7 +240,7 @@ def __init__(self, configs: list, orchestrator: "BaseOrchestrator" = None) -> No # ── Internal hook helpers (act_max calibration) ─────────────────────────── - def _register_act_max_hooks(self, block: "torch.nn.Module") -> list: + def _register_act_max_hooks(self, block: torch.nn.Module) -> list: """Register per-module act_max tracking hooks for static activation quantization. Returns a list of hook handles that the caller must remove when done. @@ -301,7 +301,7 @@ def should_collect(name, module): handles.append(module.register_forward_hook(collect_act_max)) return handles - def _get_fp_act_hooks(self, block: "torch.nn.Module") -> list: + def _get_fp_act_hooks(self, block: torch.nn.Module) -> list: """Register FP-input act_max + quantizer forward hooks.""" if not self.need_quanted_input(): # If having q_input, act_max will be collected in q_input forward hook, @@ -312,13 +312,13 @@ def _get_fp_act_hooks(self, block: "torch.nn.Module") -> list: handles.extend(self.block_quantizer.register_fp_input_forward_hooks(block)) return handles - def _get_q_act_hooks(self, block: "torch.nn.Module") -> list: + def _get_q_act_hooks(self, block: torch.nn.Module) -> list: """Register Q-input act_max + quantizer forward hooks.""" handles = self._register_act_max_hooks(block) handles.extend(self.block_quantizer.register_qinput_forward_hooks(block)) return handles - def _attach_act_max_for_outside_layer(self, layer: "torch.nn.Module", fp_inputs, q_inputs) -> None: + def _attach_act_max_for_outside_layer(self, layer: torch.nn.Module, fp_inputs, q_inputs) -> None: """Compute and attach act_max for an outside-block layer from cached inputs. Mirrors the hook-based act_max collection done for in-block layers, but @@ -334,10 +334,10 @@ def _attach_act_max_for_outside_layer(self, layer: "torch.nn.Module", fp_inputs, from auto_round.data_type.utils import reshape_pad_tensor_by_group_size target_input = q_inputs or fp_inputs - act_group_size = getattr(layer, "act_group_size") + act_group_size = layer.act_group_size if act_group_size is None: act_group_size = layer.group_size - act_data_type = getattr(layer, "act_data_type") + act_data_type = layer.act_data_type if act_data_type is None: act_data_type = layer.data_type is_act_nv_fp_flag = is_nv_fp(act_data_type) if act_data_type else False @@ -368,9 +368,7 @@ def need_quanted_input(self): for preprocessor in self.preprocessors: if getattr(preprocessor, "enable_quanted_input", False): return True - if getattr(self.block_quantizer, "enable_quanted_input", False): - return True - return False + return bool(getattr(self.block_quantizer, "enable_quanted_input", False)) def compress_embedding_layer(self): return self.block_quantizer.quantize_embedding_layer() @@ -512,7 +510,7 @@ def compress_block( def compress_layer_outside_block( self, - layer: "torch.nn.Module", + layer: torch.nn.Module, fp_inputs=None, q_inputs=None, disable_opt_rtn=None, # TODO wenhuach rename this to search_init_scale @@ -537,7 +535,7 @@ def compress_layer_outside_block( if fp_inputs is not None: from auto_round.compressors.utils import is_nv_fp - act_data_type = getattr(layer, "act_data_type") + act_data_type = layer.act_data_type if act_data_type is None: act_data_type = "fp" act_dynamic = getattr(layer, "act_dynamic", True) @@ -569,7 +567,7 @@ def members(self) -> list: """ return list(self.preprocessors) + list(self._rotation_members) + [self.block_quantizer] - def dispatch_block(self, block: "torch.nn.Module", input_ids, input_others: dict): + def dispatch_block(self, block: torch.nn.Module, input_ids, input_others: dict): """Dispatch block to device(s) via the pipeline's algorithms. Iterates all members; if exactly one overrides the default dispatch_block, @@ -597,7 +595,7 @@ def dispatch_block(self, block: "torch.nn.Module", input_ids, input_others: dict return overriders[0].dispatch_block(block, input_ids, input_others) return self.block_quantizer.dispatch_block(block, input_ids, input_others) - def prepare_run(self, composer: "AlgorithmComposer" = None): + def prepare_run(self, composer: AlgorithmComposer = None): for alg in self.members(): alg.prepare_run(composer=self) @@ -626,7 +624,7 @@ def _resolve_rotation_data_type(self) -> str: return getattr(self.block_quantizer.config, "data_type", "mx_fp") return "mx_fp" - def apply_model_transforms(self, model: "torch.nn.Module") -> "torch.nn.Module": + def apply_model_transforms(self, model: torch.nn.Module) -> torch.nn.Module: """Apply model-level pre-quantisation transforms (rotation) to *model*. Generic entry point invoked once by the orchestrator before calibration diff --git a/auto_round/algorithms/config.py b/auto_round/algorithms/config.py index 3202a79a5d..d533b162e2 100644 --- a/auto_round/algorithms/config.py +++ b/auto_round/algorithms/config.py @@ -34,7 +34,7 @@ def dest(self) -> str: class _MutuallyExclusiveParameterRegistry: - def __init__(self, registry: "AlgorithmParameterRegistry", group_id: int) -> None: + def __init__(self, registry: AlgorithmParameterRegistry, group_id: int) -> None: self._registry = registry self._group_id = group_id diff --git a/auto_round/algorithms/quantization/adam_round/adam.py b/auto_round/algorithms/quantization/adam_round/adam.py index e30e1ab4bb..bb67e6c1ba 100644 --- a/auto_round/algorithms/quantization/adam_round/adam.py +++ b/auto_round/algorithms/quantization/adam_round/adam.py @@ -33,8 +33,6 @@ def _get_optimizer(self, optimizer): optimizer = torch.optim.AdamW elif isinstance(optimizer, str): optimizer = getattr(torch.optim, optimizer) - else: - optimizer = optimizer return optimizer def _get_scaler(self): diff --git a/auto_round/algorithms/quantization/base.py b/auto_round/algorithms/quantization/base.py index f9cfaebd4b..6aa6593b0a 100644 --- a/auto_round/algorithms/quantization/base.py +++ b/auto_round/algorithms/quantization/base.py @@ -125,12 +125,9 @@ def quantize_embedding_layer(self) -> bool: ) except torch.OutOfMemoryError: cuda_error_msg = traceback.format_exc() - try: - logger.error(cuda_error_msg) - logger.warning("falling back to CPU") - weight, scale, zp = quant_func(module.weight.to("cpu"), **quant_kwargs) - except Exception: - raise + logger.error(cuda_error_msg) + logger.warning("falling back to CPU") + weight, scale, zp = quant_func(module.weight.to("cpu"), **quant_kwargs) module.weight.data.copy_(weight.cpu()) for param_name, val in zip(["scale", "zp"], [scale, zp]): if isinstance(val, dict): @@ -238,23 +235,20 @@ def _quantize_layer_via_rtn(self, layer: "torch.nn.Module", disable_opt_rtn: "bo except torch.OutOfMemoryError: cuda_error_msg = traceback.format_exc() layer = layer.orig_layer if hasattr(layer, "orig_layer") else layer - try: - logger.error(cuda_error_msg) - logger.warning("falling back to CPU.") - layer.to("cpu") - layer = WrapperLinear( - layer, - enable_minmax_tuning=False, - enable_norm_bias_tuning=False, - enable_round_tuning=False, - enable_torch_compile=self.compress_context.enable_torch_compile, - disable_opt_rtn=disable_opt_rtn, - enable_neuqi=getattr(self.config, "enable_neuqi", False), - iters=0, - ) - layer = layer.unwrapper({}) - except Exception: - raise + logger.error(cuda_error_msg) + logger.warning("falling back to CPU.") + layer.to("cpu") + layer = WrapperLinear( + layer, + enable_minmax_tuning=False, + enable_norm_bias_tuning=False, + enable_round_tuning=False, + enable_torch_compile=self.compress_context.enable_torch_compile, + disable_opt_rtn=disable_opt_rtn, + enable_neuqi=getattr(self.config, "enable_neuqi", False), + iters=0, + ) + layer = layer.unwrapper({}) set_module(self.model, layer_name, layer) def _compute_valid_token_mask(self, input_ids: list) -> "list | None": diff --git a/auto_round/algorithms/quantization/registry.py b/auto_round/algorithms/quantization/registry.py index 82a77742c7..6c4beb090e 100644 --- a/auto_round/algorithms/quantization/registry.py +++ b/auto_round/algorithms/quantization/registry.py @@ -8,4 +8,4 @@ def register_alg(alias, factory): register_algorithm(alias, aliases=(alias,), config_factory=factory) -__all__ = ["register_alg", "resolve_alg_config", "list_registered_algorithms"] +__all__ = ["list_registered_algorithms", "register_alg", "resolve_alg_config"] diff --git a/auto_round/algorithms/quantization/rtn/config.py b/auto_round/algorithms/quantization/rtn/config.py index ac4fc376c5..ad0c2dc1d3 100644 --- a/auto_round/algorithms/quantization/rtn/config.py +++ b/auto_round/algorithms/quantization/rtn/config.py @@ -56,8 +56,8 @@ def register_args(cls, registry: AlgorithmParameterRegistry) -> None: def __init__( self, *, - disable_opt_rtn: bool = None, - enable_opt_rtn: bool = None, + disable_opt_rtn: bool | None = None, + enable_opt_rtn: bool | None = None, enable_neuqi: bool = False, **kwargs, ) -> None: diff --git a/auto_round/algorithms/quantization/sign_round/config.py b/auto_round/algorithms/quantization/sign_round/config.py index a362a34e11..32be62e946 100644 --- a/auto_round/algorithms/quantization/sign_round/config.py +++ b/auto_round/algorithms/quantization/sign_round/config.py @@ -11,7 +11,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -from typing import Callable +from collections.abc import Callable from auto_round.algorithms.config import AlgorithmParameterRegistry from auto_round.algorithms.quantization.config import QuantizationConfig diff --git a/auto_round/algorithms/quantization/sign_round/quantizer.py b/auto_round/algorithms/quantization/sign_round/quantizer.py index e5a7e0257a..7bf06c8535 100644 --- a/auto_round/algorithms/quantization/sign_round/quantizer.py +++ b/auto_round/algorithms/quantization/sign_round/quantizer.py @@ -12,8 +12,9 @@ # See the License for the specific language governing permissions and # limitations under the License. import copy +from collections.abc import Callable from contextlib import nullcontext -from typing import TYPE_CHECKING, Any, Callable, Optional, Union +from typing import TYPE_CHECKING, Any import torch from torch import autocast @@ -135,8 +136,8 @@ def _get_loss( ref_output: torch.Tensor, indices: torch.Tensor, loss_func: Callable, - device: Union[str, torch.device] = "cpu", - valid_token_mask: Optional[torch.Tensor] = None, + device: str | torch.device = "cpu", + valid_token_mask: torch.Tensor | None = None, input_ids=None, ): autocast_ctx = ( @@ -608,10 +609,10 @@ def quantize_block( def quantize_layer_outside_block( self, layer: "torch.nn.Module", - fp_inputs: Optional[list[torch.Tensor]] = None, - q_inputs: Optional[list[torch.Tensor]] = None, - disable_opt_rtn: Optional[bool] = None, - input_ids: Optional[list[torch.Tensor]] = None, + fp_inputs: list[torch.Tensor] | None = None, + q_inputs: list[torch.Tensor] | None = None, + disable_opt_rtn: bool | None = None, + input_ids: list[torch.Tensor] | None = None, ): """Quantize a single layer that lives outside a transformer block. @@ -841,7 +842,7 @@ def _count_layer_input_elements(self, input_ids, indices: list) -> int: def _get_scaler(self): """Returns scaler, in SignRound, no need to use scaler.""" - return None + return def _scale_loss_and_backward(self, scaler: Any, loss: torch.Tensor) -> torch.Tensor: """Scales the loss and performs backward pass. diff --git a/auto_round/algorithms/quantization/sign_round/sign_sgd.py b/auto_round/algorithms/quantization/sign_round/sign_sgd.py index b1fe060256..671c178f5a 100644 --- a/auto_round/algorithms/quantization/sign_round/sign_sgd.py +++ b/auto_round/algorithms/quantization/sign_round/sign_sgd.py @@ -1,5 +1,4 @@ # -# -*- coding: utf-8 -*- # # Copyright (c) 2023 Intel Corporation # @@ -16,7 +15,7 @@ # limitations under the License. -from typing import Iterable, List, Optional +from collections.abc import Iterable # From PyTorch: # @@ -102,7 +101,7 @@ __all__ = ["SignSGD", "sgd"] -class _RequiredParameter(object): +class _RequiredParameter: """Singleton class representing a required parameter for an Optimizer.""" def __repr__(self): @@ -220,15 +219,15 @@ def __init__( nesterov: bool = False, *, maximize: bool = False, - foreach: Optional[bool] = None, + foreach: bool | None = None, differentiable: bool = False, ) -> None: if lr is not required and lr < 0.0: - raise ValueError("Invalid learning rate: {}".format(lr)) + raise ValueError(f"Invalid learning rate: {lr}") if momentum < 0.0: - raise ValueError("Invalid momentum value: {}".format(momentum)) + raise ValueError(f"Invalid momentum value: {momentum}") if weight_decay < 0.0: - raise ValueError("Invalid weight_decay value: {}".format(weight_decay)) + raise ValueError(f"Invalid weight_decay value: {weight_decay}") defaults = dict( lr=lr, @@ -242,7 +241,7 @@ def __init__( ) if nesterov and (momentum <= 0 or dampening != 0): raise ValueError("Nesterov momentum requires a momentum and zero dampening") - super(SignSGD, self).__init__(params, defaults) + super().__init__(params, defaults) def __setstate__(self, state): super().__setstate__(state) @@ -307,13 +306,13 @@ def step(self, closure=None): def sgd( - params: List[Tensor], - d_p_list: List[Tensor], - momentum_buffer_list: List[Optional[Tensor]], + params: list[Tensor], + d_p_list: list[Tensor], + momentum_buffer_list: list[Tensor | None], # kwonly args with defaults are not supported by functions compiled with torchscript issue #70627 # setting this as kwarg for now as functional API is compiled by torch/distributed/optim - has_sparse_grad: bool = None, - foreach: bool = None, + has_sparse_grad: bool | None = None, + foreach: bool | None = None, *, weight_decay: float, momentum: float, @@ -354,9 +353,9 @@ def sgd( def _single_tensor_sgd( - params: List[Tensor], - d_p_list: List[Tensor], - momentum_buffer_list: List[Optional[Tensor]], + params: list[Tensor], + d_p_list: list[Tensor], + momentum_buffer_list: list[Tensor | None], *, weight_decay: float, momentum: float, diff --git a/auto_round/algorithms/quantization/sign_roundv2/quantizer.py b/auto_round/algorithms/quantization/sign_roundv2/quantizer.py index eb9c977c6f..0d4b44fc47 100644 --- a/auto_round/algorithms/quantization/sign_roundv2/quantizer.py +++ b/auto_round/algorithms/quantization/sign_roundv2/quantizer.py @@ -12,9 +12,10 @@ # See the License for the specific language governing permissions and # limitations under the License. +from collections.abc import Callable from contextlib import nullcontext from functools import partial -from typing import TYPE_CHECKING, Callable, Union +from typing import TYPE_CHECKING import torch import transformers @@ -25,6 +26,7 @@ if TYPE_CHECKING: from auto_round.algorithms.composer import AlgorithmComposer + from auto_round.algorithms.registry import register_pipeline_member from auto_round.data_type.gguf import ( double_quant_tensor_sym_rtn, @@ -154,9 +156,8 @@ def _prepare_init_scale_weight(self) -> torch.Tensor: weight_reshape = torch.clamp(weight_reshape, clip_min, clip_max) else: logger.warning_once( - "SignRoundV2: ignoring AWQ clip range with shapes %s/%s incompatible with " - "grouped weight shape %s." - % (tuple(clip_min.shape), tuple(clip_max.shape), tuple(weight_reshape.shape)) + f"SignRoundV2: ignoring AWQ clip range with shapes {tuple(clip_min.shape)}/{tuple(clip_max.shape)} " + f"incompatible with grouped weight shape {tuple(weight_reshape.shape)}." ) return weight_reshape @@ -365,7 +366,7 @@ def _get_loss( ref_output: torch.Tensor, indices: torch.Tensor, mse_loss: Callable, - device: Union[str, torch.device] = "cpu", + device: str | torch.device = "cpu", valid_token_mask: list[torch.Tensor] | None = None, ): if self._use_outlier_suppressed_loss: diff --git a/auto_round/algorithms/registry.py b/auto_round/algorithms/registry.py index 3dfffca67f..0e96bd1e4f 100644 --- a/auto_round/algorithms/registry.py +++ b/auto_round/algorithms/registry.py @@ -5,8 +5,9 @@ import copy import importlib +from collections.abc import Callable from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, Callable +from typing import TYPE_CHECKING if TYPE_CHECKING: from auto_round.algorithms.base import BaseAlgorithm @@ -23,7 +24,7 @@ class AlgRegistryEntry: _ALG_REGISTRY: dict[str, AlgRegistryEntry] = {} _ALIAS_TO_NAME: dict[str, str] = {} -_CONFIG_IMPL_REGISTRY: dict[type, type["BaseAlgorithm"]] = {} +_CONFIG_IMPL_REGISTRY: dict[type, type[BaseAlgorithm]] = {} _builtin_algorithms_registered = False _pipeline_members_registered = False _BUILTIN_ALGORITHM_ORDER = ("rtn", "auto_round", "awq", "svdquant", "hadamard", "quarot", "spinquant") @@ -171,14 +172,14 @@ def list_registered_algorithms() -> list[str]: def register_pipeline_member(config_cls: type): - def _decorator(member_cls: type["BaseAlgorithm"]) -> type["BaseAlgorithm"]: + def _decorator(member_cls: type[BaseAlgorithm]) -> type[BaseAlgorithm]: _CONFIG_IMPL_REGISTRY[config_cls] = member_cls return member_cls return _decorator -def resolve_pipeline_member(config: object) -> type["BaseAlgorithm"]: +def resolve_pipeline_member(config: object) -> type[BaseAlgorithm]: _ensure_pipeline_members_registered() config_cls = type(config) for cls in config_cls.__mro__: diff --git a/auto_round/algorithms/transforms/__init__.py b/auto_round/algorithms/transforms/__init__.py index 13b597e719..b1ac6289bb 100644 --- a/auto_round/algorithms/transforms/__init__.py +++ b/auto_round/algorithms/transforms/__init__.py @@ -61,7 +61,7 @@ RotationConfig, ) -__all__ = [ +__all__ = [ # noqa: RUF022 # Base interfaces "BasePreprocessor", "BaseWeightTransformer", # backward-compat alias diff --git a/auto_round/algorithms/transforms/awq/base.py b/auto_round/algorithms/transforms/awq/base.py index 2c8332e5eb..d2e3999927 100644 --- a/auto_round/algorithms/transforms/awq/base.py +++ b/auto_round/algorithms/transforms/awq/base.py @@ -31,6 +31,7 @@ import inspect import re +import sys from typing import TYPE_CHECKING, Any import torch @@ -278,13 +279,13 @@ def bind(self, compressor) -> None: "AWQ does not support nblocks > 1 (got nblocks=%s). ", nblocks, ) - exit(-1) + sys.exit(-1) def can_compile_block_forward(self) -> bool: """AWQ installs per-block calibration hooks that trigger Dynamo recompiles.""" return False - def prepare_run(self, composer: "AlgorithmComposer" = None) -> None: + def prepare_run(self, composer: AlgorithmComposer = None) -> None: """Resolve model-wide mappings and group them by transformer block.""" model = self.model @@ -347,7 +348,7 @@ def register_fp_input_forward_hooks(self, block) -> list: return self._register_awq_hooks(self.model_context.model, block, block_name) return [] - def pre_quantize_block(self, ctx: "BlockContext") -> None: + def pre_quantize_block(self, ctx: BlockContext) -> None: """Apply AWQ smoothing for this block and mark modified params. Called after the reference forward (activation stats collected) and @@ -382,7 +383,7 @@ def pre_quantize_block(self, ctx: "BlockContext") -> None: modified.extend(mapping.balance_names) modified.append(mapping.smooth_name) - def post_quantize_block(self, ctx: "BlockContext") -> None: + def post_quantize_block(self, ctx: BlockContext) -> None: """Release per-block AWQ caches to free memory.""" block_mappings = self._block_mappings.get(ctx.block_name, []) if not block_mappings: @@ -608,9 +609,7 @@ def _mapping_is_smoothable(self, mapping: ResolvedMapping) -> bool: """AWQ smoothing is all-or-nothing for layers sharing one smooth scale.""" if self._mapping_has_ignored_layer(mapping): return False - if self._mapping_has_mixed_quant_params(mapping): - return False - return True + return not self._mapping_has_mixed_quant_params(mapping) def _smooth_block(self, block_prefix: str, block_mappings: list) -> None: """Run grid search and apply AWQ scales for one block. diff --git a/auto_round/algorithms/transforms/awq/config.py b/auto_round/algorithms/transforms/awq/config.py index affc1bad02..92018236db 100644 --- a/auto_round/algorithms/transforms/awq/config.py +++ b/auto_round/algorithms/transforms/awq/config.py @@ -310,7 +310,7 @@ def disable_opt_rtn(self, value: bool | None) -> None: def finalize_scheme(self) -> None: """Adjust AWQ state that depends on the resolved run scheme.""" data_type = self.data_type - is_gguf_double_quant = bool(data_type) and (data_type.endswith("_dq") or data_type.endswith("float_zp")) + is_gguf_double_quant = bool(data_type) and data_type.endswith(("_dq", "float_zp")) if self.apply_clip and is_gguf_double_quant: logger.warning( "AWQ weight clipping (apply_clip=True) is not supported for GGUF " diff --git a/auto_round/algorithms/transforms/base.py b/auto_round/algorithms/transforms/base.py index d83935c556..d904584ae3 100644 --- a/auto_round/algorithms/transforms/base.py +++ b/auto_round/algorithms/transforms/base.py @@ -22,7 +22,7 @@ from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import Any, Optional +from typing import Any import torch import torch.nn as nn @@ -97,7 +97,7 @@ class BaseRotation(ABC): """ # Registry populated by subclasses via ``BaseRotation.register``. - _REGISTRY: dict[str, type["BaseRotation"]] = {} + _REGISTRY: dict[str, type[BaseRotation]] = {} def __init__(self, config: BaseRotationConfig) -> None: self.config = config @@ -144,7 +144,7 @@ def prepare_layerwise( model: torch.nn.Module, data_type: str = "mx_fp", **kwargs: Any, - ) -> "BaseRotation": + ) -> BaseRotation: """Prepare for layer-wise rotation without modifying model weights. Called once when the rotation config has ``layerwise=True``. @@ -205,7 +205,6 @@ def finalize_layerwise(self, model: torch.nn.Module) -> None: Args: model: The fully-rotated model. """ - pass # ------------------------------------------------------------------ # Factory @@ -229,7 +228,7 @@ def _decorator(subclass: type[BaseRotation]) -> type[BaseRotation]: return _decorator @classmethod - def from_config(cls, config: BaseRotationConfig) -> "BaseRotation": + def from_config(cls, config: BaseRotationConfig) -> BaseRotation: """Instantiate the correct ``BaseRotation`` subclass for *config*. The algorithm is looked up by ``config.algorithm`` in the registry. diff --git a/auto_round/algorithms/transforms/hadamard/__init__.py b/auto_round/algorithms/transforms/hadamard/__init__.py index 54b749f989..64738e6f20 100644 --- a/auto_round/algorithms/transforms/hadamard/__init__.py +++ b/auto_round/algorithms/transforms/hadamard/__init__.py @@ -28,7 +28,7 @@ build_hadamard_transform, ) -__all__ = [ +__all__ = [ # noqa: RUF022 # Algorithm class "HadamardRotation", # Config diff --git a/auto_round/algorithms/transforms/hadamard/apply.py b/auto_round/algorithms/transforms/hadamard/apply.py index 9e7d035702..c599844cc1 100644 --- a/auto_round/algorithms/transforms/hadamard/apply.py +++ b/auto_round/algorithms/transforms/hadamard/apply.py @@ -82,7 +82,7 @@ def __init__(self, config: RotationConfig) -> None: super().__init__(config) @classmethod - def from_config(cls, config: dict | RotationConfig) -> "HadamardRotation": + def from_config(cls, config: dict | RotationConfig) -> HadamardRotation: """Build a :class:`HadamardRotation` from a raw dict or :class:`RotationConfig`.""" if isinstance(config, dict): config = RotationConfig.model_validate(config) @@ -143,7 +143,7 @@ def apply_to_model( fuse_online_to_weight=fuse_online_to_weight, compute_device=compute_device, ) - setattr(model, "_rotation_config", cfg) + model._rotation_config = cfg return model # backend == "transform": original per-Linear triton-fused path. @@ -159,7 +159,7 @@ def apply_to_model( _apply_to_module(model, module, cfg, location, data_type) # Store config on model for serialisation / downstream inspection. - setattr(model, "_rotation_config", cfg) + model._rotation_config = cfg return model # ------------------------------------------------------------------ @@ -188,9 +188,7 @@ def supports_layerwise(self) -> bool: cfg = self.config backend = getattr(cfg, "backend", "auto") hadamard_type = getattr(cfg, "hadamard_type", "") or "" - if backend == "inplace" or "inplace" in hadamard_type: - return False - return True + return not (backend == "inplace" or "inplace" in hadamard_type) def prepare_layerwise( self, @@ -198,7 +196,7 @@ def prepare_layerwise( data_type: str = "mx_fp", location: str = "weight", **kwargs: Any, - ) -> "HadamardRotation": + ) -> HadamardRotation: """Prepare for per-block Hadamard rotation without touching weights. The per-Linear Hadamard needs no global pre-computation — matrices are @@ -219,7 +217,7 @@ def prepare_layerwise( ) self._layerwise_location = location self._layerwise_data_type = data_type - setattr(model, "_rotation_config", self.config) + model._rotation_config = self.config return self def rotate_layer( diff --git a/auto_round/algorithms/transforms/hadamard/config.py b/auto_round/algorithms/transforms/hadamard/config.py index f2f2bf85a3..e1893711b5 100644 --- a/auto_round/algorithms/transforms/hadamard/config.py +++ b/auto_round/algorithms/transforms/hadamard/config.py @@ -34,7 +34,7 @@ from __future__ import annotations -from typing import Any, ClassVar, Optional +from typing import Any, ClassVar from pydantic import BaseModel, Field, field_validator @@ -46,9 +46,9 @@ __all__ = [ "RotationConfig", + "dump_group_size_to_rotation_config", "normalize_rotation_config", "to_dict_rotation_config", - "dump_group_size_to_rotation_config", ] # Supported Hadamard transform types (also used by HadamardTransform registry). @@ -75,7 +75,7 @@ class RotationConfig(AlgorithmConfig, BaseModel, BaseRotationConfig): # ---- shared ---- backend: str = Field(default="auto") - block_size: Optional[int] = Field(default=None) + block_size: int | None = Field(default=None) hadamard_type: str = Field(default="hadamard") # Apply the Hadamard rotation per decoder block, in lock-step with block-wise # quantization, instead of rotating the whole model up-front. Honoured by the @@ -84,7 +84,7 @@ class RotationConfig(AlgorithmConfig, BaseModel, BaseRotationConfig): layerwise: bool = Field(default=False) # ---- inplace-only ---- - fuse_online_to_weight: Optional[bool] = Field(default=None) + fuse_online_to_weight: bool | None = Field(default=None) allow_online_rotation: bool = Field(default=True) # for random hadamard (transform path) diff --git a/auto_round/algorithms/transforms/hadamard/dispatcher.py b/auto_round/algorithms/transforms/hadamard/dispatcher.py index 6ae4fdbad3..52e96aa445 100644 --- a/auto_round/algorithms/transforms/hadamard/dispatcher.py +++ b/auto_round/algorithms/transforms/hadamard/dispatcher.py @@ -28,7 +28,7 @@ from __future__ import annotations -from typing import Any, Union +from typing import Any import torch @@ -41,7 +41,7 @@ def _to_config( - rotation_config: Union[str, dict, RotationConfig, None], + rotation_config: str | dict | RotationConfig | None, data_type: str, ) -> RotationConfig: """Normalise *rotation_config* and return a :class:`RotationConfig` instance.""" @@ -88,7 +88,7 @@ def resolve_hadamard_backend(config: RotationConfig, data_type: str) -> str: def apply_hadamard_rotation( model: torch.nn.Module, - rotation_config: Union[str, dict, RotationConfig, None], + rotation_config: str | dict | RotationConfig | None, data_type: str, compute_device: torch.device | str = None, ) -> (torch.nn.Module, Any): @@ -138,7 +138,7 @@ def apply_hadamard_rotation( compute_device=compute_device, ) # Stash config object for downstream (export / serialization). - setattr(model, "_rotation_config", config) + model._rotation_config = config return model, hooks elif backend == "transform": diff --git a/auto_round/algorithms/transforms/hadamard/inplace/apply.py b/auto_round/algorithms/transforms/hadamard/inplace/apply.py index d2adb7508a..b809830f82 100644 --- a/auto_round/algorithms/transforms/hadamard/inplace/apply.py +++ b/auto_round/algorithms/transforms/hadamard/inplace/apply.py @@ -9,7 +9,6 @@ import gc import typing -from typing import Dict, Union import torch import tqdm @@ -280,10 +279,10 @@ def _rotate_weights( model, mapping: RotationMapping, use_fast_had: bool = True, - group_size: int = None, + group_size: int | None = None, compute_device: torch.device = None, - had_dict: dict = None, - preset: str = None, + had_dict: dict | None = None, + preset: str | None = None, fuse_online_to_weight: bool = True, ) -> None: """Apply Hadamard rotation to all weights. @@ -657,9 +656,9 @@ def _register_online_hooks( mapping: RotationMapping, fp32_had: bool = False, use_fast_had: bool = True, - group_size: int = None, - had_dict: dict = None, - preset: str = None, + group_size: int | None = None, + had_dict: dict | None = None, + preset: str | None = None, fuse_online_to_weight: bool = True, ): """Register online Hadamard pre-forward hooks on ``down_proj`` and ``o_proj``. @@ -710,8 +709,8 @@ def _online_had(dim): attn_o_suffix = mapping.attn_o.split(".")[-1] # Suffixes for Q/K/V and gate/up (for online input Had hooks) - attn_qkv_suffixes = set(attr.split(".")[-1] for attr in (mapping.attn_q, mapping.attn_k, mapping.attn_v)) - mlp_in_suffixes = set(attr.split(".")[-1] for attr in mapping.mlp_in) + attn_qkv_suffixes = {attr.split(".")[-1] for attr in (mapping.attn_q, mapping.attn_k, mapping.attn_v)} + mlp_in_suffixes = {attr.split(".")[-1] for attr in mapping.mlp_in} # --- Build hook factories --- def _make_down_proj_hook(): @@ -807,12 +806,12 @@ def _make_o_proj_hook(): def apply_rotation_transform( model, - group_size: int = None, + group_size: int | None = None, allow_online_rotation: bool = True, - rotation_matrix: Union[str, torch.Tensor, Dict[int, torch.Tensor], None] = None, + rotation_matrix: str | torch.Tensor | dict[int, torch.Tensor] | None = None, compute_device: torch.device | str = None, fp32_had: bool = False, - fuse_online_to_weight: bool = None, + fuse_online_to_weight: bool | None = None, ): """Fuse layer norms, rotate weights, and register online Hadamard hooks. diff --git a/auto_round/algorithms/transforms/hadamard/inplace/hooks.py b/auto_round/algorithms/transforms/hadamard/inplace/hooks.py index 057d881d2a..9e907749b3 100644 --- a/auto_round/algorithms/transforms/hadamard/inplace/hooks.py +++ b/auto_round/algorithms/transforms/hadamard/inplace/hooks.py @@ -9,7 +9,6 @@ """ import math -from typing import Optional import torch import torch.nn as nn @@ -156,11 +155,11 @@ class FullOnlineHadamardHook(nn.Module): def __init__( self, - had_K: Optional[torch.Tensor], - K: Optional[int], + had_K: torch.Tensor | None, + K: int | None, fp32_had: bool = False, use_fast_had: bool = True, - had_matrix: Optional[torch.Tensor] = None, + had_matrix: torch.Tensor | None = None, ) -> None: super().__init__() self.custom_had = had_matrix is not None @@ -219,12 +218,12 @@ class CrossHeadOnlineHadamardHook(nn.Module): def __init__( self, - had_K: Optional[torch.Tensor], - K: Optional[int], + had_K: torch.Tensor | None, + K: int | None, head_dim: int, fp32_had: bool = False, use_fast_had: bool = True, - had_matrix: Optional[torch.Tensor] = None, + had_matrix: torch.Tensor | None = None, ) -> None: """ Args: @@ -613,7 +612,7 @@ def __init__( group_size: int, fp32_had: bool = False, use_fast_had: bool = True, - had_matrix: Optional[torch.Tensor] = None, + had_matrix: torch.Tensor | None = None, ) -> None: super().__init__() self.group_size = group_size @@ -788,15 +787,7 @@ def register_online_had_hooks_grouped(model, mapping, group_size, fp32_had=False handles = [] for name, module in model.named_modules(): - if name.endswith(mlp_out_suffix) and isinstance(module, nn.Linear): - hook = GroupOnlineHadamardHook( - group_size=group_size, - fp32_had=fp32_had, - use_fast_had=use_fast_had, - ) - h = module.register_forward_pre_hook(hook) - handles.append(h) - elif name.endswith(attn_o_suffix) and isinstance(module, nn.Linear): + if isinstance(module, nn.Linear) and name.endswith((mlp_out_suffix, attn_o_suffix)): hook = GroupOnlineHadamardHook( group_size=group_size, fp32_had=fp32_had, diff --git a/auto_round/algorithms/transforms/hadamard/inplace/model_config.py b/auto_round/algorithms/transforms/hadamard/inplace/model_config.py index 3ecbf9b696..78fd99534d 100644 --- a/auto_round/algorithms/transforms/hadamard/inplace/model_config.py +++ b/auto_round/algorithms/transforms/hadamard/inplace/model_config.py @@ -12,16 +12,15 @@ from __future__ import annotations from dataclasses import dataclass, field -from typing import Dict, List, Optional from auto_round.utils import logger __all__ = [ + "MAPPING_REGISTRY", "RotationMapping", - "register_mapping", "get_mapping", "infer_mapping_from_model", - "MAPPING_REGISTRY", + "register_mapping", ] @@ -45,7 +44,7 @@ class RotationMapping: # -- top-level modules (dot-path from model root) -- embedding: str = "model.embed_tokens" lm_head: str = "lm_head" - positional_embedding: Optional[str] = None # e.g. "model.decoder.embed_positions" for OPT + positional_embedding: str | None = None # e.g. "model.decoder.embed_positions" for OPT # -- layers container (dot-path from model root) -- layers_attr: str = "model.layers" @@ -59,14 +58,14 @@ class RotationMapping: # -- per-layer: MLP (dot-path from each layer) -- mlp_input_ln: str = "post_attention_layernorm" - mlp_in: List[str] = field(default_factory=lambda: ["mlp.up_proj", "mlp.gate_proj"]) + mlp_in: list[str] = field(default_factory=lambda: ["mlp.up_proj", "mlp.gate_proj"]) mlp_out: str = "mlp.down_proj" # -- final norm (dot-path from model root) -- pre_head_ln: str = "model.norm" # -- head dim override (None = hidden_size // num_heads) -- - attn_head_dim: Optional[int] = None + attn_head_dim: int | None = None # -- config attr names -- num_heads_attr: str = "num_attention_heads" @@ -91,7 +90,7 @@ def _resolve(root, dot_path: str): # Registry # --------------------------------------------------------------------------- -MAPPING_REGISTRY: Dict[str, RotationMapping] = {} +MAPPING_REGISTRY: dict[str, RotationMapping] = {} def register_mapping(key: str, mapping: RotationMapping) -> RotationMapping: diff --git a/auto_round/algorithms/transforms/hadamard/patch.py b/auto_round/algorithms/transforms/hadamard/patch.py index e3f5f9f95d..2e9fbe3bee 100644 --- a/auto_round/algorithms/transforms/hadamard/patch.py +++ b/auto_round/algorithms/transforms/hadamard/patch.py @@ -31,9 +31,9 @@ from auto_round.wrapper import WrapperLinear, WrapperWALayer __all__ = [ + "patch_quantlinear", "patch_wrapperlinear_to_apply_transform", "patch_wrapperwalayer_forward_to_apply_transform", - "patch_quantlinear", ] @@ -81,7 +81,7 @@ def _qdq_weight_patched(self, value, min_scale, max_scale): _orig_qdq_act = WrapperLinear._qdq_act - def _qdq_act_patched(self, x, act_min_scale=torch.tensor(1.0), act_max_scale=torch.tensor(1.0), act_max=None): + def _qdq_act_patched(self, x, act_min_scale=None, act_max_scale=None, act_max=None): x = inp_transform(x) return _orig_qdq_act(self, x, act_min_scale=act_min_scale, act_max_scale=act_max_scale, act_max=act_max) @@ -191,7 +191,6 @@ def _pack_patched( # add transform weight self.register_buffer("hadamard_matrix", w_transform.weight.to(device)) - return QuantLinear.pack = _pack_patched QuantLinear._pack_patched = True diff --git a/auto_round/algorithms/transforms/hadamard/transforms.py b/auto_round/algorithms/transforms/hadamard/transforms.py index 890fc58707..5d77fea473 100644 --- a/auto_round/algorithms/transforms/hadamard/transforms.py +++ b/auto_round/algorithms/transforms/hadamard/transforms.py @@ -22,7 +22,8 @@ import inspect import math -from typing import Any, Callable, Dict +from collections.abc import Callable +from typing import Any import torch import torch.nn as nn @@ -34,14 +35,14 @@ from auto_round.algorithms.transforms.hadamard.utils.matrix import apply_transform_weight __all__ = [ + "HADAMARDS", "HadamardTransform", "RandomHadamardTransform", - "HADAMARDS", "build_hadamard_transform", ] -def _filter_kwargs(fn: Callable, kwargs: Dict[str, Any]) -> Dict[str, Any]: +def _filter_kwargs(fn: Callable, kwargs: dict[str, Any]) -> dict[str, Any]: """Return only the keyword arguments accepted by *fn*.""" accepted = inspect.signature(fn).parameters.keys() return {k: v for k, v in kwargs.items() if k in accepted} diff --git a/auto_round/algorithms/transforms/hadamard/utils/math.py b/auto_round/algorithms/transforms/hadamard/utils/math.py index 14b15ce0de..a2b9edeb64 100644 --- a/auto_round/algorithms/transforms/hadamard/utils/math.py +++ b/auto_round/algorithms/transforms/hadamard/utils/math.py @@ -28,7 +28,7 @@ import torch from safetensors import safe_open -__all__ = ["deterministic_hadamard_matrix", "random_hadamard_matrix", "is_pow2"] +__all__ = ["deterministic_hadamard_matrix", "is_pow2", "random_hadamard_matrix"] # Precomputed Hadamard matrices for non-power-of-2 sizes. _HADAMARD_MATRICES_PATH: Path = Path(__file__).parent / "hadamards.safetensors" diff --git a/auto_round/algorithms/transforms/member.py b/auto_round/algorithms/transforms/member.py index 448e98ff37..dac931603c 100644 --- a/auto_round/algorithms/transforms/member.py +++ b/auto_round/algorithms/transforms/member.py @@ -57,7 +57,7 @@ class RotationPreprocessor(BasePreprocessor): def __init__(self, config: Any) -> None: super().__init__(config) - self._rotation: "BaseRotation | None" = None + self._rotation: BaseRotation | None = None # Set once :meth:`rotate_model` decides layer-wise preparation succeeded. self._layerwise_active: bool = False @@ -65,7 +65,7 @@ def __init__(self, config: Any) -> None: # Lazy rotation construction # ------------------------------------------------------------------ @property - def rotation(self) -> "BaseRotation": + def rotation(self) -> BaseRotation: """The concrete :class:`BaseRotation` for this member (built lazily).""" if self._rotation is None: from auto_round.algorithms.transforms import normalize_rotation_config @@ -93,9 +93,9 @@ def is_layerwise_active(self) -> bool: # ------------------------------------------------------------------ def rotate_model( self, - model: "torch.nn.Module", + model: torch.nn.Module, data_type: str = "mx_fp", - ) -> "torch.nn.Module": + ) -> torch.nn.Module: """Rotate *model* up-front, or prepare layer-wise rotation matrices. Whether to rotate per-block is taken solely from the rotation config's @@ -129,7 +129,7 @@ def rotate_model( # ------------------------------------------------------------------ # Per-block entry (compress_block step 0) # ------------------------------------------------------------------ - def on_block_ready(self, block: "torch.nn.Module", ctx: "BlockContext") -> None: + def on_block_ready(self, block: torch.nn.Module, ctx: BlockContext) -> None: """Rotate the block about to be quantized (layer-wise mode only). No-op unless :meth:`rotate_model` prepared layer-wise rotation. Iterates @@ -158,7 +158,7 @@ def finalize_run(self) -> None: # Helpers # ------------------------------------------------------------------ @staticmethod - def _iter_layers(block: "torch.nn.Module", block_index: int): + def _iter_layers(block: torch.nn.Module, block_index: int): """Yield ``(decoder_layer, layer_idx)`` for every layer inside *block*. ``block_index`` is the global index of the *first* layer in this block. diff --git a/auto_round/algorithms/transforms/spinquant/__init__.py b/auto_round/algorithms/transforms/spinquant/__init__.py index b94ef5dc19..2a93394648 100644 --- a/auto_round/algorithms/transforms/spinquant/__init__.py +++ b/auto_round/algorithms/transforms/spinquant/__init__.py @@ -106,7 +106,7 @@ save_spinquant_config, ) -__all__ = [ +__all__ = [ # noqa: RUF022 # -- Registry algorithm (unified apply_rotation() entry) -- "SpinQuantRotation", # -- Preprocessor (QuaRot, recommended) -- diff --git a/auto_round/algorithms/transforms/spinquant/apply.py b/auto_round/algorithms/transforms/spinquant/apply.py index f3c0e0c22c..c8595f4291 100644 --- a/auto_round/algorithms/transforms/spinquant/apply.py +++ b/auto_round/algorithms/transforms/spinquant/apply.py @@ -120,7 +120,7 @@ def prepare_layerwise( model: torch.nn.Module, data_type: str = "mx_fp", **kwargs: Any, - ) -> "SpinQuantRotation": + ) -> SpinQuantRotation: """Prepare for layer-wise rotation: init R matrices only. Creates a :class:`SpinQuantPreprocessor`, calls its diff --git a/auto_round/algorithms/transforms/spinquant/cayley_optimizer.py b/auto_round/algorithms/transforms/spinquant/cayley_optimizer.py index f919deaeff..d00ef56e24 100644 --- a/auto_round/algorithms/transforms/spinquant/cayley_optimizer.py +++ b/auto_round/algorithms/transforms/spinquant/cayley_optimizer.py @@ -11,7 +11,8 @@ from __future__ import annotations -from typing import Any, Callable, Iterable +from collections.abc import Callable, Iterable +from typing import Any import torch from torch.optim.optimizer import Optimizer diff --git a/auto_round/algorithms/transforms/spinquant/inplace/apply.py b/auto_round/algorithms/transforms/spinquant/inplace/apply.py index 18c8574345..b78334f197 100644 --- a/auto_round/algorithms/transforms/spinquant/inplace/apply.py +++ b/auto_round/algorithms/transforms/spinquant/inplace/apply.py @@ -19,7 +19,7 @@ from __future__ import annotations import logging -from typing import Any, Optional +from typing import Any import torch import torch.nn as nn @@ -35,7 +35,7 @@ def register_spinquant_hooks( model: nn.Module, config: Any, - compute_device: Optional[torch.device] = None, + compute_device: torch.device | None = None, head_dim: int = 0, intermediate_size: int = 0, r4_rotation_size: int = 0, @@ -262,7 +262,7 @@ def remove_spinquant_hooks(handles: list[Any]) -> None: def apply_spinquant_in_place( model: nn.Module, config: Any, - dataloader: Optional[Any] = None, + dataloader: Any | None = None, ) -> nn.Module: """Apply SpinQuant rotations to a model **in-place**. diff --git a/auto_round/algorithms/transforms/spinquant/monkeypatch.py b/auto_round/algorithms/transforms/spinquant/monkeypatch.py index 7fc3d9b463..747830d3e8 100644 --- a/auto_round/algorithms/transforms/spinquant/monkeypatch.py +++ b/auto_round/algorithms/transforms/spinquant/monkeypatch.py @@ -19,7 +19,8 @@ import copy import functools import types -from typing import Any, Callable +from collections.abc import Callable +from typing import Any import torch import torch.nn as nn diff --git a/auto_round/algorithms/transforms/spinquant/preprocessor.py b/auto_round/algorithms/transforms/spinquant/preprocessor.py index 0c41271b67..2cbec11ea0 100644 --- a/auto_round/algorithms/transforms/spinquant/preprocessor.py +++ b/auto_round/algorithms/transforms/spinquant/preprocessor.py @@ -23,7 +23,7 @@ import logging import math from dataclasses import dataclass -from typing import Any, Optional +from typing import Any import torch import torch.nn as nn @@ -92,7 +92,7 @@ class SpinQuantConfig(BaseRotationConfig): # and R4 uses rotation_size instead of intermediate_size. # R2 always uses head_dim, R3 does not support custom size. # This follows the same convention as Quark's rotation_size. - rotation_size: Optional[int] = None + rotation_size: int | None = None # Rotation matrix type for R1–R4 # - False (default): deterministic Hadamard (same matrix every time, no need to persist) @@ -130,7 +130,7 @@ class SpinQuantConfig(BaseRotationConfig): # Numerics dtype: torch.dtype = torch.float32 - device: Optional[str] = None + device: str | None = None def __post_init__(self): if self.device is None: @@ -205,7 +205,7 @@ class SpinQuantPreprocessor: original but with weight distributions better suited for quantisation. """ - def __init__(self, model: nn.Module, config: Optional[SpinQuantConfig] = None) -> None: + def __init__(self, model: nn.Module, config: SpinQuantConfig | None = None) -> None: self.model = model self.config = config or SpinQuantConfig() @@ -242,7 +242,7 @@ def __init__(self, model: nn.Module, config: Optional[SpinQuantConfig] = None) - # ------------------------------------------------------------------ # Main entry point # ------------------------------------------------------------------ - def preprocess(self, dataloader: Optional[Any] = None) -> nn.Module: + def preprocess(self, dataloader: Any | None = None) -> nn.Module: logger.info("[SpinQuant] Starting preprocessing...") logger.info( f"[SpinQuant] Model architecture info: hidden_size={self.hidden_size}, " @@ -337,7 +337,7 @@ def preprocess(self, dataloader: Optional[Any] = None) -> nn.Module: # Layer-wise API (for block-lifecycle / block-wise quantization) # ------------------------------------------------------------------ - def prepare(self, dataloader: Optional[Any] = None) -> None: + def prepare(self, dataloader: Any | None = None) -> None: """Global preparation for layer-wise (block-wise) rotation. Performs all lightweight, non-destructive steps: @@ -874,7 +874,7 @@ def _train_rotations(self, dataloader: Any) -> None: if torch.cuda.is_available(): torch.cuda.empty_cache() - def _get_embed_tokens(self) -> Optional[nn.Module]: + def _get_embed_tokens(self) -> nn.Module | None: """Get embedding module, supporting both model.embed_tokens and model.model.embed_tokens.""" for attr_path in ("embed_tokens", "model.embed_tokens"): parts = attr_path.split(".") @@ -909,7 +909,7 @@ def _get_layers(self): yield layer return - def _get_lm_head(self) -> Optional[nn.Module]: + def _get_lm_head(self) -> nn.Module | None: """Get LM head module.""" return getattr(self.model, "lm_head", None) @@ -1386,7 +1386,7 @@ def _cleanup(self) -> None: # ------------------------------------------------------------------ # Helpers # ------------------------------------------------------------------ - def _get_rotation_tensor(self, name: str) -> Optional[torch.Tensor]: + def _get_rotation_tensor(self, name: str) -> torch.Tensor | None: if hasattr(self.model, name): tensor = getattr(self.model, name) if isinstance(tensor, (nn.Parameter, torch.Tensor)): diff --git a/auto_round/algorithms/transforms/spinquant/rotation_utils.py b/auto_round/algorithms/transforms/spinquant/rotation_utils.py index 6f5cd91ae0..921abed252 100644 --- a/auto_round/algorithms/transforms/spinquant/rotation_utils.py +++ b/auto_round/algorithms/transforms/spinquant/rotation_utils.py @@ -13,7 +13,6 @@ from __future__ import annotations import math -from typing import Optional, Tuple import torch import torch.nn as nn @@ -58,7 +57,7 @@ def is_pow2(n: int) -> bool: return n > 0 and (n & (n - 1)) == 0 -def get_hadamard_K(n: int) -> Tuple[torch.Tensor, int]: +def get_hadamard_K(n: int) -> tuple[torch.Tensor, int]: """Get the Hadamard matrix and block dimension K for a given input size. For power-of-2 sizes, K=1 (full Walsh-Hadamard via butterfly). @@ -98,7 +97,7 @@ def get_hadamard_K(n: int) -> Tuple[torch.Tensor, int]: ) -def matmul_hadU(X: torch.Tensor, hadamard_K: Optional[torch.Tensor] = None, K: Optional[int] = None) -> torch.Tensor: +def matmul_hadU(X: torch.Tensor, hadamard_K: torch.Tensor | None = None, K: int | None = None) -> torch.Tensor: """Apply normalized Hadamard transform to the last dimension of X. Uses the efficient butterfly algorithm for power-of-2 dimensions, @@ -171,19 +170,19 @@ def apply_transform_weight( __all__ = [ - "rotate_in_channels_", - "rotate_out_channels_", - "fuse_rmsnorm_in_model", - "untie_word_embeddings_if_needed", + "InputRotationWrapperHadamard", + "apply_hadamard_to_linear", + "create_block_diag_from_head_matrix", "deterministic_hadamard_matrix", - "random_hadamard_matrix", - "is_pow2", + "fuse_rmsnorm_in_model", "get_hadamard_K", - "matmul_hadU", - "create_block_diag_from_head_matrix", - "apply_hadamard_to_linear", "get_model_arch_info", - "InputRotationWrapperHadamard", + "is_pow2", + "matmul_hadU", + "random_hadamard_matrix", + "rotate_in_channels_", + "rotate_out_channels_", + "untie_word_embeddings_if_needed", ] @@ -235,8 +234,8 @@ def __init__( self, original_module: nn.Linear, rotation_size: int, - hadamard_K: Optional[torch.Tensor] = None, - K: Optional[int] = None, + hadamard_K: torch.Tensor | None = None, + K: int | None = None, ) -> None: super().__init__() @@ -331,9 +330,9 @@ def __repr__(self) -> str: def rotate_in_channels_( layer: nn.Linear, - rotation_matrix: Optional[torch.Tensor] = None, - R_in: Optional[torch.Tensor] = None, - rotated_modules: Optional[set] = None, + rotation_matrix: torch.Tensor | None = None, + R_in: torch.Tensor | None = None, + rotated_modules: set | None = None, ) -> None: """Fuse an input-side rotation into a linear layer's weight. @@ -387,9 +386,9 @@ def rotate_in_channels_( def rotate_out_channels_( layer: nn.Linear, - rotation_matrix: Optional[torch.Tensor] = None, - R_out: Optional[torch.Tensor] = None, - rotated_modules: Optional[set] = None, + rotation_matrix: torch.Tensor | None = None, + R_out: torch.Tensor | None = None, + rotated_modules: set | None = None, ) -> None: """Fuse an output-side rotation into a linear layer's weight. diff --git a/auto_round/algorithms/transforms/spinquant/serialize.py b/auto_round/algorithms/transforms/spinquant/serialize.py index 634619ad8c..e88df4f092 100644 --- a/auto_round/algorithms/transforms/spinquant/serialize.py +++ b/auto_round/algorithms/transforms/spinquant/serialize.py @@ -29,7 +29,7 @@ import math import os from dataclasses import asdict -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING import torch import torch.nn as nn @@ -65,7 +65,7 @@ def inject_spinquant_buffers( model: nn.Module, - config: "SpinQuantConfig", + config: SpinQuantConfig, ) -> int: """Inject SpinQuant rotation buffers into QuantLinear modules for serialization. @@ -132,7 +132,7 @@ def inject_spinquant_buffers( def save_spinquant_config( model: nn.Module, save_dir: str, - config: "SpinQuantConfig", + config: SpinQuantConfig, ) -> None: """Save SpinQuant config into the model's config.json for load-time reconstruction. @@ -298,7 +298,7 @@ def _preregister_buffers_on_module( def rebuild_spinquant_online( model: nn.Module, - config: Optional["SpinQuantConfig"] = None, + config: SpinQuantConfig | None = None, ) -> nn.Module: """Rebuild online SpinQuant rotations after loading a quantized model. @@ -385,7 +385,6 @@ def _patch_quantlinear_forward_spinquant(model: nn.Module) -> int: Returns: Number of QuantLinear modules with spinquant buffers found. """ - global _QUANTLINEAR_PATCHED n_with_buffers = 0 quantlinear_classes = set() @@ -523,7 +522,7 @@ def _inject_rotation_buffers( rotation_size: int, random: bool, is_trained: bool, - rotation_matrix: Optional[torch.Tensor] = None, + rotation_matrix: torch.Tensor | None = None, ) -> None: """Register rotation buffers on a QuantLinear module. @@ -649,7 +648,7 @@ def _get_r4_target_names(model: nn.Module) -> set: return targets -def _get_stored_rotation(model: nn.Module, param_name: str) -> Optional[torch.Tensor]: +def _get_stored_rotation(model: nn.Module, param_name: str) -> torch.Tensor | None: """Get a stored rotation matrix/parameter from the model. During preprocessing, rotation matrices are stored as model-level @@ -687,7 +686,7 @@ def _get_intermediate_size(model: nn.Module) -> int: return 0 -def _config_to_serializable(config: "SpinQuantConfig", model: nn.Module) -> dict: +def _config_to_serializable(config: SpinQuantConfig, model: nn.Module) -> dict: """Convert SpinQuantConfig to a JSON-serializable dict with model info.""" from auto_round.algorithms.transforms.spinquant.preprocessor import SpinQuantConfig @@ -714,7 +713,7 @@ def _config_to_serializable(config: "SpinQuantConfig", model: nn.Module) -> dict def _load_config_from_model( model: nn.Module, -) -> Optional["SpinQuantConfig"]: +) -> SpinQuantConfig | None: """Try to load SpinQuantConfig from model.config.""" from auto_round.algorithms.transforms.spinquant.preprocessor import SpinQuantConfig diff --git a/auto_round/algorithms/transforms/spinquant/training.py b/auto_round/algorithms/transforms/spinquant/training.py index e617933155..eac6088e62 100644 --- a/auto_round/algorithms/transforms/spinquant/training.py +++ b/auto_round/algorithms/transforms/spinquant/training.py @@ -47,8 +47,9 @@ import copy import logging import time +from collections.abc import Callable from dataclasses import asdict, dataclass, field -from typing import Any, Callable, Optional +from typing import Any import torch import torch.nn as nn @@ -164,7 +165,7 @@ def create_dual_optimizer( model: nn.Module, lr: float = 1e-4, smooth_lr: float = 1e-3, -) -> Optional[AdamAndSGDG]: +) -> AdamAndSGDG | None: """Create the Adam (smooth) + SGDG (rotation) dual optimiser. Returns ``None`` if no trainable parameters are found. @@ -215,8 +216,8 @@ def run_training_loop( max_iters: int = 200, loss_type: str = "kl_top", kl_top_k: int = 1000, - compute_loss_fn: Optional[Callable] = None, - on_step_end: Optional[Callable[[int, float, float], None]] = None, + compute_loss_fn: Callable | None = None, + on_step_end: Callable[[int, float, float], None] | None = None, log_interval: int = 50, ) -> TrainingResult: """Run the SpinQuant rotation training loop. @@ -338,7 +339,7 @@ def __init__( self.model = model self.config = config or SpinQuantConfig() self.enabled = enabled - self.preprocessor: Optional[SpinQuantPreprocessor] = None + self.preprocessor: SpinQuantPreprocessor | None = None def preprocess(self, dataloader: Any) -> nn.Module: """Execute SpinQuant preprocessing.""" @@ -482,11 +483,11 @@ class RotationTrainerConfig: # ---------- Misc ---------- dtype: torch.dtype = torch.float32 - device: Optional[str] = None + device: str | None = None log_interval: int = 50 # print every N steps eval_interval: int = 0 # 0 = never save_interval: int = 0 # 0 = never - checkpoint_dir: Optional[str] = None + checkpoint_dir: str | None = None def __post_init__(self): if self.device is None: @@ -604,9 +605,9 @@ class RotationTrainer: def __init__( self, model: nn.Module, - config: Optional[RotationTrainerConfig] = None, - callbacks: Optional[list[RotationTrainerCallback]] = None, - compute_loss_fn: Optional[Callable[[torch.Tensor, torch.Tensor, RotationTrainerConfig], torch.Tensor]] = None, + config: RotationTrainerConfig | None = None, + callbacks: list[RotationTrainerCallback] | None = None, + compute_loss_fn: Callable[[torch.Tensor, torch.Tensor, RotationTrainerConfig], torch.Tensor] | None = None, ) -> None: from auto_round.algorithms.transforms.spinquant.preprocessor import ( SpinQuantConfig, @@ -640,7 +641,7 @@ def __init__( # Training components (created lazily) self.optimizer = None - self._original_model: Optional[nn.Module] = None + self._original_model: nn.Module | None = None self._hook_handles: list[Any] = [] self._rotated_modules: set[nn.Module] = set() self._loss_buffer: list[float] = [] @@ -756,7 +757,7 @@ def fuse(self) -> nn.Module: self._preprocessor._cleanup() return self.model - def save_checkpoint(self, path: Optional[str] = None) -> str: + def save_checkpoint(self, path: str | None = None) -> str: """Save rotation + smooth params to disk.""" if path is None: path = f"{self.config.checkpoint_dir or '.'}/spinquant_ckpt_step{self.state['step']}.pt" diff --git a/auto_round/algorithms/transforms/svdquant/apply.py b/auto_round/algorithms/transforms/svdquant/apply.py index b1a7678015..fffff0beba 100644 --- a/auto_round/algorithms/transforms/svdquant/apply.py +++ b/auto_round/algorithms/transforms/svdquant/apply.py @@ -685,11 +685,10 @@ def _is_target(self, name: str, module: torch.nn.Module) -> bool: pattern in name or pattern in full_name for pattern in self._target_modules ): return False - if self.config.exclude_modules and any( - pattern in name or pattern in full_name for pattern in self.config.exclude_modules - ): - return False - return True + return not ( + self.config.exclude_modules + and any(pattern in name or pattern in full_name for pattern in self.config.exclude_modules) + ) @staticmethod def _new_linear_like(module: torch.nn.Linear, weight: torch.Tensor, bias: torch.Tensor | None): diff --git a/auto_round/algorithms/transforms/svdquant/smooth.py b/auto_round/algorithms/transforms/svdquant/smooth.py index 41d6b4c0c4..c3d0d4d9f2 100644 --- a/auto_round/algorithms/transforms/svdquant/smooth.py +++ b/auto_round/algorithms/transforms/svdquant/smooth.py @@ -15,8 +15,9 @@ from __future__ import annotations import math +from collections.abc import Iterable from dataclasses import dataclass -from typing import Iterable, TypeVar +from typing import TypeVar import torch diff --git a/auto_round/algorithms/utils.py b/auto_round/algorithms/utils.py index 76ce0d3045..342c50cf0d 100644 --- a/auto_round/algorithms/utils.py +++ b/auto_round/algorithms/utils.py @@ -29,7 +29,7 @@ def _is_nvfp4_value(value: Any) -> bool: return "nv_fp" in value or "nvfp4" in value -def _has_nvfp4_layer(orchestrator: "BaseOrchestrator") -> bool: +def _has_nvfp4_layer(orchestrator: BaseOrchestrator) -> bool: """Whether global or per-layer config enables any NVFP4 quantization.""" if _is_nvfp4_value(getattr(orchestrator, "data_type", "")): return True diff --git a/auto_round/auto_scheme/delta_loss.py b/auto_round/auto_scheme/delta_loss.py index 8d23ff1524..c5fbfe2329 100644 --- a/auto_round/auto_scheme/delta_loss.py +++ b/auto_round/auto_scheme/delta_loss.py @@ -19,9 +19,9 @@ import math import os import time +from collections.abc import Iterable from dataclasses import asdict from functools import wraps -from typing import Iterable, Optional, Union import torch from accelerate import dispatch_model @@ -194,11 +194,11 @@ def test_multi_card(self): """ if torch.isnan(grad).any() or torch.isnan(x_diff).any(): self.act_cnt -= 1 - return None + return - self.act_score += torch.abs((grad * x_diff.to(grad.device))).sum().item() + self.act_score += torch.abs(grad * x_diff.to(grad.device)).sum().item() self.mix_score = self.weight_score + self.act_score - return None + return if qdq_x.requires_grad: qdq_x.register_hook(save_grad) @@ -210,7 +210,7 @@ def _ensure_score_cache(self): return device = self.device with torch.no_grad(): - qdq_w, _, _ = super(AutoSchemeWrapperLinear, self)._qdq_weight( + qdq_w, _, _ = super()._qdq_weight( torch.tensor(0, device=device), torch.tensor(1.0, device=device), torch.tensor(1.0, device=device) ) self._score_qdq_cpu = qdq_w.detach().to("cpu") @@ -280,7 +280,6 @@ def save_grad(grad): w_diff = weight.to(grad.device) - self._score_qdq_cpu.to(grad.device) self.weight_score += torch.abs(grad.to(w_diff.device) * w_diff).sum().item() self.mix_score = self.weight_score + self.act_score - return None qdq_w.register_hook(save_grad) return qdq_w, 1.0, None @@ -357,9 +356,8 @@ def post_init_qdqw(self, device): def save_grad(grad): """Backward hook: accumulate weight score from grad * (weight - qdq_w).""" w_diff = self.orig_layer.weight - self.qdq_w.to(self.orig_layer.weight.device) - self.weight_score += torch.abs((grad.to(torch.float32) * w_diff.to(grad.device))).sum().item() + self.weight_score += torch.abs(grad.to(torch.float32) * w_diff.to(grad.device)).sum().item() self.mix_score = self.weight_score + self.act_score - return None self.qdq_w.requires_grad_(True) self.orig_layer.weight.requires_grad_(False) @@ -395,11 +393,11 @@ def test_multi_card(self): """ if torch.isnan(grad).any() or torch.isnan(x_diff).any(): self.act_cnt -= 1 - return None + return - self.act_score += torch.abs((grad * x_diff.to(grad.device))).sum().item() + self.act_score += torch.abs(grad * x_diff.to(grad.device)).sum().item() self.mix_score = self.weight_score + self.act_score - return None + return if qdq_x.requires_grad: qdq_x.register_hook(save_grad) @@ -450,9 +448,8 @@ def save_grad(grad): """Backward hook: accumulate weight score from grad * (weight - qdq_w).""" w_diff = self.orig_layer.weight - self.qdq_w.to(self.orig_layer.weight.device) # TODO strange, grad could be in CPU - self.weight_score += torch.abs((grad.to(w_diff.device).to(torch.float32) * w_diff)).sum().item() + self.weight_score += torch.abs(grad.to(w_diff.device).to(torch.float32) * w_diff).sum().item() self.mix_score = self.weight_score + self.act_score - return None self.qdq_w.requires_grad_(True) self.orig_layer.weight.requires_grad_(False) @@ -506,9 +503,8 @@ def post_init_qdqw(self, device): # Could not place in qdq_w, otherwise vram is def save_grad(grad): """Backward hook: accumulate weight score from grad * (weight - qdq_w).""" w_diff = self.orig_layer.weight - self.qdq_w.to(self.orig_layer.weight.device) - self.weight_score += torch.abs((grad.to(torch.float32) * w_diff.to(grad.device))).sum().item() + self.weight_score += torch.abs(grad.to(torch.float32) * w_diff.to(grad.device)).sum().item() self.mix_score = self.weight_score + self.act_score - return None self.qdq_w.requires_grad_(True) self.orig_layer.weight.requires_grad_(False) @@ -620,9 +616,7 @@ def move_to_cpu_clear_memory(module, inputs, outputs): clear_memory(device_list=major_device) all_move_device_hooks = [] - i = 0 for block_name in block_names: - i += 1 block_module = get_module(model, block_name) hook_move_gpu = block_module.register_forward_pre_hook(move_to_gpu_hook) @@ -651,7 +645,7 @@ def __init__(self, message): super().__init__(message) -def prepare_model_low_gpu(model, block_inputs: dict = None, pbar=None, major_device="cpu", disk_index=None): +def prepare_model_low_gpu(model, block_inputs: dict | None = None, pbar=None, major_device="cpu", disk_index=None): """Wrap every block's forward so that, for one calibration batch, it (1) moves itself to ``major_device`` on demand, (2) records its own inputs into ``block_inputs`` (on CPU) so they can be replayed later, and (3) moves itself back to CPU once done. @@ -1199,12 +1193,12 @@ def get_score_for_scheme( low_gpu_mem_usage=True, major_device="cpu", batch_size=1, - offload_context: Optional[OffloadManager] = None, + offload_context: OffloadManager | None = None, processor=None, is_vlm: bool = False, force_mllm: bool = False, - model_name: Optional[str] = None, - scheme_tag: Optional[str] = None, + model_name: str | None = None, + scheme_tag: str | None = None, disk_index=None, ): """Wrap every quantizable layer in ``quant_layer_names`` with a scoring wrapper, run @@ -1612,7 +1606,7 @@ def _run_forward_loop(loader): return scores_dict -def choose_bits_per_layer_with_path(layers: dict, P: int, max_states: int = None): +def choose_bits_per_layer_with_path(layers: dict, P: int, max_states: int | None = None): """ Args: layers: A dict mapping each layer name to a list of candidate options. @@ -1632,7 +1626,7 @@ def choose_bits_per_layer_with_path(layers: dict, P: int, max_states: int = None # the entire path on every transition, which becomes quadratic for large # models; linked nodes keep each transition O(1) and are expanded only once. dp: dict[int, tuple[float, tuple]] = {0: (0.0, ())} - for layer_name, opts in layers.items(): + for opts in layers.values(): new_dp: dict[int, tuple[float, tuple]] = {} for cur_params, (cur_loss, cur_path) in dp.items(): for opt in opts: @@ -1675,7 +1669,7 @@ def choose_bits_per_layer_with_path(layers: dict, P: int, max_states: int = None step = (n - 1) / (max_states - 1) selected: dict[int, tuple[float, tuple]] = {} for i in range(max_states): - idx = int(round(i * step)) + idx = round(i * step) if idx >= n: idx = n - 1 k = sorted_keys[idx] @@ -2534,7 +2528,7 @@ def batch_checkpoint(batch_idx, total_batches): def _gen_layer_config( auto_scheme: AutoScheme, - model: Union[str, torch.nn.Module], + model: str | torch.nn.Module, quant_layer_names: Iterable[str], fixed_layer_scheme: dict[str, dict], min_avg_bit_scheme, @@ -2548,7 +2542,7 @@ def _gen_layer_config( processor=None, is_vlm: bool = False, disk_index=None, - export_format: str = None, + export_format: str | None = None, ): """Score every candidate scheme in ``auto_scheme.options`` against ``quant_layer_names`` and return per-layer per-scheme losses used by the caller to pick a final bit-width @@ -3607,7 +3601,7 @@ def _enforce_w8_symmetric_entries(layer_config: dict, allow_w8_asym: bool = Fals @register_scheme_methods(("default", "DeltaLoss")) def gen_layer_config( auto_scheme: AutoScheme, - model: Union[str, torch.nn.Module], + model: str | torch.nn.Module, quant_layer_names: Iterable[str], fixed_layer_scheme: dict[str, dict], dataset: str = "pile-10k", @@ -3617,7 +3611,7 @@ def gen_layer_config( low_gpu_mem_usage=True, min_avg_bit_scheme=None, processor=None, - export_format: str = None, + export_format: str | None = None, **kwargs, ): """Public AutoScheme entry. diff --git a/auto_round/auto_scheme/gen_auto_scheme.py b/auto_round/auto_scheme/gen_auto_scheme.py index d212aca52c..7e47f30a76 100644 --- a/auto_round/auto_scheme/gen_auto_scheme.py +++ b/auto_round/auto_scheme/gen_auto_scheme.py @@ -12,8 +12,9 @@ # See the License for the specific language governing permissions and # limitations under the License. +from collections.abc import Iterable from dataclasses import dataclass -from typing import Iterable, Optional, Union +from typing import Union import torch @@ -27,17 +28,17 @@ @dataclass class AutoScheme: - avg_bits: Union[float, list[float]] - options: Union[str, list[Union[QuantizationScheme, str]], tuple[Union[QuantizationScheme, str], ...]] - shared_layers: Optional[Iterable[Iterable[str]]] = None + avg_bits: float | list[float] + options: str | list[QuantizationScheme | str] | tuple[QuantizationScheme | str, ...] + shared_layers: Iterable[Iterable[str]] | None = None method: str = "default" ignore_scale_zp_bits: bool = False - batch_size: Optional[int] = None - nsamples: Optional[int] = None - seqlen: Optional[int] = None - dataset: Optional[str] = None # Import Notice no comma for each item - device_map: Optional[Union[str, torch.device, int, dict]] = None - enable_torch_compile: Optional[bool] = None + batch_size: int | None = None + nsamples: int | None = None + seqlen: int | None = None + dataset: str | None = None # Import Notice no comma for each item + device_map: str | torch.device | int | dict | None = None + enable_torch_compile: bool | None = None low_gpu_mem_usage: bool = True low_cpu_mem_usage: bool = True @@ -94,11 +95,11 @@ def __init__( quant_layer_names: Iterable[str], fixed_layer_scheme: dict[str, dict], dataset: str = "pile-10k", - device_map: Union[str, torch.device, int, dict, None] = None, + device_map: str | torch.device | int | dict | None = None, tokenizer=None, enable_torch_compile=True, processor=None, - export_format: str = None, + export_format: str | None = None, ): self.auto_scheme = auto_scheme # Export-format context for the generation-time 8-bit-asym policy: the diff --git a/auto_round/auto_scheme/utils.py b/auto_round/auto_scheme/utils.py index c8e4f5e2a6..6ffd840d9e 100644 --- a/auto_round/auto_scheme/utils.py +++ b/auto_round/auto_scheme/utils.py @@ -14,8 +14,8 @@ import logging import math import re +from collections.abc import Iterable from dataclasses import asdict, fields -from typing import Iterable, Optional, Union import torch from accelerate import dispatch_model, infer_auto_device_map @@ -44,7 +44,7 @@ def apply_quant_scheme( model: torch.nn.Module, quant_layer_names: Iterable[str], fixed_layer_scheme: dict[str, dict], - scheme: Union[str, dict], # TODO add scale_dtype + scheme: str | dict, # TODO add scale_dtype ) -> None: """Apply a quantization scheme to each quantized layer. @@ -91,7 +91,7 @@ def compute_avg_bits_for_scheme( model: torch.nn.Module, quant_layer_names: Iterable[str], fixed_layer_scheme: dict[str, dict], - scheme: Union[str, dict, None] = None, + scheme: str | dict | None = None, ignore_scale_zp_bits: bool = False, clean_scheme: bool = True, ) -> tuple[float, float]: @@ -341,7 +341,7 @@ def parse_shared_layers(model: torch.nn.Module, shared_patterns: Iterable[Iterab return matched_groups -def _expert_key_from_layer_name(layer_name: str) -> Optional[str]: +def _expert_key_from_layer_name(layer_name: str) -> str | None: """Map one MoE-related linear layer to a unique expert key. Gate/up/down projections belonging to the same expert should map to one key. @@ -489,7 +489,7 @@ def _fill_inactive_expert_scores(scores_dict: dict[str, list[float]], block_name if not active_expert_avg_losses: continue fill_value = sum(active_expert_avg_losses) / len(active_expert_avg_losses) - for _, expert_stats in expert_stats_map.items(): + for expert_stats in expert_stats_map.values(): if expert_stats["has_active"]: continue for layer_name in expert_stats["layers"]: @@ -500,8 +500,8 @@ def _log_score_summary_by_block_and_nonblock( scores_dict: dict[str, list[float]], block_names: list[str], model=None, - scheme_tag: Optional[str] = None, - summary_stage: Optional[str] = None, + scheme_tag: str | None = None, + summary_stage: str | None = None, ): """Log a per-block (and non-block) breakdown of ``scores_dict`` losses at debug level.""" if not scores_dict: diff --git a/auto_round/autoround.py b/auto_round/autoround.py index 4826d1b0cc..8ffcc85a6f 100644 --- a/auto_round/autoround.py +++ b/auto_round/autoround.py @@ -16,7 +16,7 @@ import functools import inspect -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any import torch @@ -74,7 +74,7 @@ def _get_compressor_class(model_type: str, base_cls: type) -> type: return combined -def _resolve_quant_config_for_routing(alg_configs) -> tuple[list, list, "QuantizationConfig"]: +def _resolve_quant_config_for_routing(alg_configs) -> tuple[list, list, QuantizationConfig]: from auto_round.algorithms.config_resolver import split_quantization_configs from auto_round.algorithms.quantization.config import QuantizationConfig from auto_round.algorithms.quantization.rtn.config import RTNConfig @@ -151,7 +151,7 @@ def _build_model_type_ctor_kwargs(model, base_kwargs, mllm_kwargs, diffusion_kwa return model_type, ctor_kwargs -def _select_rtn_compressor_base_cls(quant_config: "RTNConfig", scheme, format, base_kwargs) -> type: +def _select_rtn_compressor_base_cls(quant_config: RTNConfig, scheme, format, base_kwargs) -> type: from auto_round.algorithms.quantization.rtn.config import OptimizedRTNConfig, RTNConfig from auto_round.auto_scheme.gen_auto_scheme import AutoScheme from auto_round.compressors.orchestrator import CompressionOrchestrator as Compressor @@ -193,9 +193,7 @@ def _select_rtn_compressor_base_cls(quant_config: "RTNConfig", scheme, format, b # the plain min/max initialization ignores it if getattr(quant_config, "enable_neuqi", False): enable_imatrix = True - elif data_type == "int" and (bits is None or bits < 8): - enable_imatrix = True - elif is_weight_scheme(scheme): + elif (data_type == "int" and (bits is None or bits < 8)) or is_weight_scheme(scheme): enable_imatrix = True act_bits = resolved_attrs.get("act_bits") @@ -277,7 +275,7 @@ def _iter_registered_alg_configs() -> list[tuple[str, type]]: return result -@functools.lru_cache(maxsize=None) +@functools.cache def _discover_alg_config_fields(config_cls: type) -> frozenset: """Discover accepted config fields without maintaining per-algorithm lists.""" from pydantic import BaseModel @@ -508,7 +506,7 @@ def _prepare_entry_kwargs(alg_configs, direct_kwargs): return configs, runtime_kwargs -class _CompressorBuilder(object): +class _CompressorBuilder: """Algorithm-config-driven entry point (``scheme`` + ``alg_configs``). This is the internal pipeline entry: it resolves the algorithm config(s), @@ -519,7 +517,7 @@ class _CompressorBuilder(object): """ @classmethod - def _resolve_config(cls, config: Union[str, object, list]) -> Union[object, list[object]]: + def _resolve_config(cls, config: str | object | list) -> object | list[object]: """Convert string alias(es) to the corresponding config instance(s) with default parameters.""" from auto_round.algorithms.registry import resolve_alg_config @@ -531,24 +529,24 @@ def _resolve_config(cls, config: Union[str, object, list]) -> Union[object, list def __new__( cls, - model: Union[torch.nn.Module, str], + model: torch.nn.Module | str, scheme="W4A16", - alg_configs: Union[str, object, list[Union[str, object]]] = None, + alg_configs: str | object | list[str | object] = None, tokenizer=None, platform="hf", format=None, dataset="NeelNanda/pile-10k", low_gpu_mem_usage: bool = False, - device_map: Union[str, torch.device, int, dict] = 0, - iters: int = None, + device_map: str | torch.device | int | dict = 0, + iters: int | None = None, enable_torch_compile: bool = False, seed: int = 42, low_cpu_mem_usage: bool = True, layer_config=None, - nsamples: int = None, - seqlen: int = None, + nsamples: int | None = None, + seqlen: int | None = None, **kwargs, - ) -> "BaseCompressor": + ) -> BaseCompressor: from auto_round.algorithms.quantization.rtn.config import OptimizedRTNConfig, RTNConfig from auto_round.algorithms.quantization.sign_round.config import SignRoundConfig from auto_round.algorithms.registry import normalize_algorithm_config @@ -728,28 +726,28 @@ class AutoRound: def __new__( cls, - model: Union[torch.nn.Module, str], + model: torch.nn.Module | str, tokenizer=None, platform: str = "hf", - scheme: Union[str, dict, QuantizationScheme, "AutoScheme"] = "W4A16", - schemes: Union[str, list, tuple, None] = None, - bits: Union[int, float, None] = None, - layer_config: dict[str, Union[str, dict, QuantizationScheme]] = None, - dataset: Optional[Union[str, list, tuple, torch.utils.data.DataLoader]] = None, + scheme: str | dict | QuantizationScheme | AutoScheme = "W4A16", + schemes: str | list | tuple | None = None, + bits: float | None = None, + layer_config: dict[str, str | dict | QuantizationScheme] | None = None, + dataset: str | list | tuple | torch.utils.data.DataLoader | None = None, iters: int | None = None, seqlen: int = 2048, nsamples: int = 128, batch_size: int = 8, gradient_accumulate_steps: int | None = None, low_gpu_mem_usage: bool = False, - device_map: Union[str, torch.device, int, dict] = 0, - enable_torch_compile: Optional[bool] = None, + device_map: str | torch.device | int | dict = 0, + enable_torch_compile: bool | None = None, seed: int = 42, low_cpu_mem_usage: bool = True, alg_configs=None, algorithm: str | None = None, **kwargs, - ) -> "BaseCompressor": + ) -> BaseCompressor: direct_kwargs = dict(kwargs) legacy_device = direct_kwargs.pop("device", None) if legacy_device is not None: diff --git a/auto_round/calib_dataset.py b/auto_round/calib_dataset.py index 82373687c1..d2bc379fab 100644 --- a/auto_round/calib_dataset.py +++ b/auto_round/calib_dataset.py @@ -460,7 +460,7 @@ def default_tokenizer_function(examples): "💡 This dataset uses an old script-based format. To load it, please install `datasets<=3.6.0`:\n\n" ) else: - raise error + raise calib_dataset = concatenate_datasets([dataset_mit, dataset_apache]) calib_dataset = calib_dataset.shuffle(seed=seed).take(10000) ##TODO concat data'shuffle may have bugs calib_dataset = calib_dataset.map(tokenizer_function, batched=True) @@ -584,7 +584,7 @@ def get_ultrachat_dataset( split = "train_sft" all_splits = ["train_sft", "test_sft", "train_gen", "test_gen"] if split not in all_splits: - raise ValueError("split must be one of {} for ultrachat_200k ".format(all_splits)) + raise ValueError(f"split must be one of {all_splits} for ultrachat_200k ") dataset = load_dataset("HuggingFaceH4/ultrachat_200k", split=split, streaming=True, trust_remote_code=True) dataset = dataset.shuffle(seed=seed).take(20000) @@ -765,8 +765,8 @@ def get_mbpp_dataset( if isinstance(splits, str): splits = splits.split("+") - for split in splits: - dataset = load_dataset(dataset_name, split=split) + for split_name in splits: + dataset = load_dataset(dataset_name, split=split_name) for data in dataset: samples.append({"text": data["text"] + data["code"]}) random.Random(seed).shuffle(samples) @@ -1021,9 +1021,7 @@ def filter_func(example): return False input_ids = example["input_ids"][:seqlen] input_ids_list = input_ids.tolist() - if len(input_ids_list) > 1 and seqlen > 2 and input_ids_list.count(input_ids_list[-1]) > seqlen // 2: - return False - return True + return not (len(input_ids_list) > 1 and seqlen > 2 and input_ids_list.count(input_ids_list[-1]) > seqlen // 2) def concat_dataset_element(dataset): input_ids, concat_input_ids = [eg["input_ids"] for eg in dataset], [] @@ -1085,9 +1083,9 @@ def concat_dataset_element(dataset): if key == "num": data_lens[name] = int(values[0]) if key == "concat": - do_concat = False if (len(values) > 0 and values[0].lower() == "false") else True + do_concat = not (len(values) > 0 and values[0].lower() == "false") if key == "apply_chat_template": - apply_chat_template = False if (len(values) > 0 and values[0].lower() == "false") else True + apply_chat_template = not (len(values) > 0 and values[0].lower() == "false") if key == "system_prompt": system_prompt = values[0] apply_chat_template = True @@ -1095,15 +1093,15 @@ def concat_dataset_element(dataset): get_dataset = CALIB_DATASETS.get("local") else: calib_name = name - if name not in CALIB_DATASETS.keys(): + if name not in CALIB_DATASETS: calib_name = name.split("/")[-1] - for key in CALIB_DATASETS.keys(): + for key in CALIB_DATASETS: if calib_name in key: calib_name = key break get_dataset = CALIB_DATASETS.get(calib_name) if get_dataset is None: - filtered_keys = [k for k in CALIB_DATASETS.keys() if "/" not in k] + filtered_keys = [k for k in CALIB_DATASETS if "/" not in k] raise ValueError( f"Dataset '{name}' is not found. Please choose from the supported datasets: {filtered_keys}." ) diff --git a/auto_round/calibration/__init__.py b/auto_round/calibration/__init__.py index 37baed84eb..5efa7ed2de 100644 --- a/auto_round/calibration/__init__.py +++ b/auto_round/calibration/__init__.py @@ -23,9 +23,9 @@ from auto_round.calibration import diffusion as _diffusion # noqa: F401 __all__ = [ - "Calibrator", - "CalibrationContext", "CALIBRATORS", + "CalibrationContext", + "Calibrator", "get_calibrator", "register_calibrator", ] diff --git a/auto_round/calibration/base.py b/auto_round/calibration/base.py index 583a385f68..539d9f4180 100644 --- a/auto_round/calibration/base.py +++ b/auto_round/calibration/base.py @@ -26,7 +26,8 @@ """ from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Callable +from collections.abc import Callable +from typing import TYPE_CHECKING import torch diff --git a/auto_round/calibration/diffusion.py b/auto_round/calibration/diffusion.py index 29fca01414..c278872bc9 100644 --- a/auto_round/calibration/diffusion.py +++ b/auto_round/calibration/diffusion.py @@ -23,6 +23,7 @@ import inspect import os +import sys import torch from tqdm import tqdm @@ -134,7 +135,7 @@ def calib(self, nsamples: int, bs: int) -> None: "Please use model path for quantization or " "move the pipeline object to GPU/XPU before passing them into API." ) - exit(-1) + sys.exit(-1) target_device = device_manager.device self._cpu_offload_mode = _prepare_pipeline_for_calibration( @@ -189,7 +190,7 @@ def calib(self, nsamples: int, bs: int) -> None: f"no data has been cached, please provide more data with sequence length >={self.seqlen} in the " f"dataset or decease the sequence length" ) - exit(-1) + sys.exit(-1) elif total_cnt < nsamples: logger.warning( f"Insufficient number of samples collected may affect the quantization. " diff --git a/auto_round/calibration/inputs.py b/auto_round/calibration/inputs.py index ee7ede3862..5b05c91852 100644 --- a/auto_round/calibration/inputs.py +++ b/auto_round/calibration/inputs.py @@ -13,14 +13,12 @@ # limitations under the License. """Pure helpers for shaping cached block inputs.""" -from typing import Tuple - import torch from auto_round.utils import clear_memory, to_device, to_dtype from auto_round.utils.device_manager import device_manager -__all__ = ["split_inputs", "preprocess_block_inputs"] +__all__ = ["preprocess_block_inputs", "split_inputs"] def split_inputs( @@ -29,7 +27,7 @@ def split_inputs( *, is_diffusion: bool, shared_cache_keys: tuple = (), -) -> Tuple[object, dict]: +) -> tuple[object, dict]: """Split a captured ``inputs`` dict into ``(input_ids, input_others)``. Mirrors the original ``Compressor._split_inputs`` exactly: @@ -71,7 +69,7 @@ def preprocess_block_inputs( model_context, compress_context, first_input_name: str = "input_ids", -) -> Tuple[object, dict]: +) -> tuple[object, dict]: """Move/cast cached block inputs onto the calibration cache device. Mirrors the original ``Compressor._preprocess_block_inputs`` exactly. diff --git a/auto_round/calibration/llm.py b/auto_round/calibration/llm.py index 9bd45ae2d7..287d8c7695 100644 --- a/auto_round/calibration/llm.py +++ b/auto_round/calibration/llm.py @@ -18,9 +18,10 @@ ``self.compressor.X``. """ +import sys import traceback +from collections.abc import Callable from functools import partial -from typing import Callable import accelerate import torch @@ -90,7 +91,7 @@ def calibration(self, block_names, nsamples, layer_names=None, last_cache_name=N if "flash_attn::" in error_msg and "CPU" in error_msg: cannot_calibrate_on_cpu = True else: - raise error + raise if not calibrate_on_cpu or cannot_calibrate_on_cpu: try: @@ -178,7 +179,7 @@ def calibration(self, block_names, nsamples, layer_names=None, last_cache_name=N except torch.OutOfMemoryError as e: if cannot_calibrate_on_cpu: - raise e + raise cuda_error_msg = traceback.format_exc() try: logger.info("switch to cpu to cache block inputs") @@ -319,13 +320,13 @@ def calib(self, nsamples: int, bs: int) -> None: elif isinstance(data, str): if self.tokenizer is None: logger.error("please provide tokenizer for string input") - exit(-1) + sys.exit(-1) data = self.tokenizer(data, truncation=True, max_length=self.seqlen, return_tensors="pt").data data_new = {} for key in data.keys(): data_new[key] = data[key].to(self.model.device) input_ids = data_new["input_ids"] - elif isinstance(data, tuple) or isinstance(data, list): + elif isinstance(data, (tuple, list)): data_new = to_device(data, self.model.device) input_ids = data_new[0] else: @@ -406,7 +407,7 @@ def calib(self, nsamples: int, bs: int) -> None: if isinstance(data_new, torch.Tensor): self.model(data_new, **kwargs) - elif isinstance(data_new, tuple) or isinstance(data_new, list): + elif isinstance(data_new, (tuple, list)): self.model(*data_new, **kwargs) else: self.model(**data_new, **kwargs) @@ -428,9 +429,9 @@ def calib(self, nsamples: int, bs: int) -> None: "When quantization encounters tensor shape mismatch error, " "you can try to avoid it with batch_size=1" ) - raise error + raise except Exception as error: - raise error + raise total_cnt += input_ids.shape[0] if len(input_ids.shape) > 1 else 1 if total_cnt >= nsamples: @@ -440,7 +441,7 @@ def calib(self, nsamples: int, bs: int) -> None: f"no data has been cached, please provide more data with sequence length " f">={self.seqlen} in the dataset or decease the sequence length" ) - exit(-1) + sys.exit(-1) elif total_cnt < nsamples: logger.warning_once( f"An insufficient number of samples likely reduces the accuracy of the quantized model. " @@ -487,17 +488,13 @@ def forward_capture(m, hidden_states=None, *positional_inputs, **kwargs): " or try to set the `batch_size` to 1 and " "`gradient_accumulate_steps` to your current batch size." ) - exit(-1) + sys.exit(-1) if hidden_states is not None: kwargs["hidden_states"] = hidden_states - for key in kwargs.keys(): - if ( - isinstance(kwargs[key], torch.Tensor) - or isinstance(kwargs[key], list) - or isinstance(kwargs[key], tuple) - ): + for key, value in kwargs.items(): + if isinstance(value, (torch.Tensor, list, tuple)): if ( self.has_variable_block_shape and name not in self.blocks_requiring_input_ids @@ -505,7 +502,7 @@ def forward_capture(m, hidden_states=None, *positional_inputs, **kwargs): ): continue if key not in self.inputs[name].keys(): # initialization - data = to_device(kwargs[key], device=torch.device("cpu")) + data = to_device(value, device=torch.device("cpu")) if data is None or key in self.shared_cache_keys: self.inputs[name][key] = data continue @@ -518,7 +515,7 @@ def forward_capture(m, hidden_states=None, *positional_inputs, **kwargs): else: self.inputs[name][key] = [data] else: # append cache inputs - new_data = post_process_cache_data(self.batch_size, kwargs[key], key) + new_data = post_process_cache_data(self.batch_size, value, key) if new_data is None: # shareable args or NoneType if key in self.shared_cache_keys: # Shared keys are normally the same across samples. However @@ -526,7 +523,7 @@ def forward_capture(m, hidden_states=None, *positional_inputs, **kwargs): # varies per image because each image has a different patch count. # Upgrade from shared (raw value) to per-sample list storage so # each sample gets its own positional embeddings. - raw_new = to_device(kwargs[key], device=torch.device("cpu")) + raw_new = to_device(value, device=torch.device("cpu")) stored = self.inputs[name].get(key) if isinstance(stored, list): stored.append(raw_new) @@ -544,9 +541,9 @@ def forward_capture(m, hidden_states=None, *positional_inputs, **kwargs): self.inputs[name][key].extend(list(torch.split(new_data, 1, dim=self.batch_dim))) else: self.inputs[name][key].append(new_data) - elif isinstance(kwargs[key], (str, bool, type(None))): + elif isinstance(value, (str, bool, type(None))): if key not in self.inputs[name].keys(): - self.inputs[name][key] = kwargs[key] + self.inputs[name][key] = value else: # Parameters not to be cached if check_skippable_keywords(key): @@ -561,7 +558,7 @@ def forward_capture(m, hidden_states=None, *positional_inputs, **kwargs): if hidden_states is not None: kwargs.pop("hidden_states", None) if positional_inputs: - return m.orig_forward(hidden_states=hidden_states, *positional_inputs, **kwargs) + return m.orig_forward(*positional_inputs, hidden_states=hidden_states, **kwargs) else: return m.orig_forward(hidden_states, **kwargs) else: diff --git a/auto_round/calibration/mllm.py b/auto_round/calibration/mllm.py index 4eb3ea80bd..24493aa20d 100644 --- a/auto_round/calibration/mllm.py +++ b/auto_round/calibration/mllm.py @@ -22,6 +22,8 @@ ``self.compressor``. """ +import sys + import torch from auto_round.calibration.llm import LLMCalibrator @@ -65,7 +67,7 @@ def calib(self, nsamples: int, bs: int) -> None: if hasattr(self.model, "name_or_path"): name = self.model.name_or_path - if any([m in name for m in MISTRAL_3_2_MODELS]): + if any(m in name for m in MISTRAL_3_2_MODELS): self.template = "mistral3_2" template_name = self.template @@ -207,6 +209,6 @@ def calib(self, nsamples: int, bs: int) -> None: if total_cnt == 0: logger.error("no data has been cached, please provide more data") - exit(-1) + sys.exit(-1) elif total_cnt < nsamples: logger.warning(f"Insufficient number of samples: required {nsamples}, but only {total_cnt} were processed.") diff --git a/auto_round/calibration/register.py b/auto_round/calibration/register.py index cd8dd49003..d4a0ded309 100644 --- a/auto_round/calibration/register.py +++ b/auto_round/calibration/register.py @@ -18,17 +18,15 @@ decorating its class with ``@register_calibrator("my_kind")``. """ -from typing import Type - from auto_round.calibration.base import Calibrator -CALIBRATORS: dict[str, Type[Calibrator]] = {} +CALIBRATORS: dict[str, type[Calibrator]] = {} def register_calibrator(name: str): """Class decorator: register a ``Calibrator`` subclass under ``name``.""" - def _wrap(cls: Type[Calibrator]) -> Type[Calibrator]: + def _wrap(cls: type[Calibrator]) -> type[Calibrator]: if not issubclass(cls, Calibrator): raise TypeError(f"{cls.__name__} must subclass auto_round.calibration.base.Calibrator") cls.name = name @@ -40,7 +38,7 @@ def _wrap(cls: Type[Calibrator]) -> Type[Calibrator]: return _wrap -def get_calibrator(name: str) -> Type[Calibrator]: +def get_calibrator(name: str) -> type[Calibrator]: """Look up a registered calibrator class by name.""" if name not in CALIBRATORS: raise KeyError(f"No calibrator registered under '{name}'. " f"Known: {sorted(CALIBRATORS.keys())}") diff --git a/auto_round/calibration/state.py b/auto_round/calibration/state.py index d3c027027e..0984363b3a 100644 --- a/auto_round/calibration/state.py +++ b/auto_round/calibration/state.py @@ -29,7 +29,7 @@ """ from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from auto_round.compressors.base import BaseOrchestrator @@ -53,7 +53,7 @@ class CalibrationContext: # blocks_requiring_input_ids: list = field(default_factory=list) # - batch_dim: Optional[int] = None + batch_dim: int | None = None # ── Calibration parameters ───────────────────────────────────────────── batch_size: int = 8 diff --git a/auto_round/calibration/utils.py b/auto_round/calibration/utils.py index 24707564d5..113e059e80 100644 --- a/auto_round/calibration/utils.py +++ b/auto_round/calibration/utils.py @@ -45,7 +45,7 @@ def _update_inputs(inputs: dict, q_inputs: dict) -> tuple[dict, torch.Tensor]: model_context = ModelContext() if model_context.is_diffusion: - input_id_str = [key for key in inputs.keys() if "hidden_state" in key] + input_id_str = [key for key in inputs if "hidden_state" in key] if input_id_str == ["hidden_states"]: if q_inputs is not None: q_inputs = q_inputs.pop("hidden_states", None) @@ -70,7 +70,7 @@ def _update_inputs(inputs: dict, q_inputs: dict) -> tuple[dict, torch.Tensor]: def _split_inputs_diffusion(inputs: dict) -> tuple[dict, dict]: """Split inputs for diffusion models that only have hidden_states.""" - input_id_str = [key for key in inputs.keys() if "hidden_state" in key] + input_id_str = [key for key in inputs if "hidden_state" in key] if input_id_str == ["hidden_states"]: input_ids = inputs.pop("hidden_states", None) input_others = inputs diff --git a/auto_round/cli/algorithms.py b/auto_round/cli/algorithms.py index 0136a2e449..23f9dccb8c 100644 --- a/auto_round/cli/algorithms.py +++ b/auto_round/cli/algorithms.py @@ -127,7 +127,7 @@ def _merge_parameter(merged, parameter): if not overlapping_options and merged.dest != parameter.dest: return None if merged.dest != parameter.dest or _argument_compatibility_key(merged) != _argument_compatibility_key(parameter): - option = sorted(overlapping_options)[0] if overlapping_options else merged.dest + option = min(overlapping_options) if overlapping_options else merged.dest raise ValueError(f"incompatible shared CLI argument {option!r}") kwargs = dict(merged.argparse_kwargs) kwargs["default"] = argparse.SUPPRESS diff --git a/auto_round/compressors/__init__.py b/auto_round/compressors/__init__.py index 7965acf4b3..d8849f999c 100644 --- a/auto_round/compressors/__init__.py +++ b/auto_round/compressors/__init__.py @@ -22,8 +22,8 @@ from auto_round.compressors.orchestrator import CompressionOrchestrator __all__ = [ - "BaseOrchestrator", "BaseCompressor", # backward-compat alias + "BaseOrchestrator", "CompressionOrchestrator", "ModelFreeCompressor", ] diff --git a/auto_round/compressors/base.py b/auto_round/compressors/base.py index 46fd3df873..5c3d3de4da 100644 --- a/auto_round/compressors/base.py +++ b/auto_round/compressors/base.py @@ -16,7 +16,7 @@ import os import sys from dataclasses import asdict, dataclass, fields, replace -from typing import Any, Optional, Union +from typing import Any import torch from transformers import AutoConfig, set_seed @@ -91,39 +91,39 @@ @dataclass class SerializedCompressorConfig: - bits: Optional[int] = None - act_bits: Optional[int] = None - data_type: Optional[str] = None - act_data_type: Optional[str] = None - group_size: Optional[int] = None - act_group_size: Optional[int] = None - sym: Optional[bool] = None - act_sym: Optional[bool] = None - act_dynamic: Optional[bool] = None - amp: Optional[bool] = None - batch_size: Optional[int] = None - enable_minmax_tuning: Optional[bool] = True - enable_norm_bias_tuning: Optional[bool] = False - enable_quanted_input: Optional[bool] = True - gradient_accumulate_steps: Optional[int] = None - iters: Optional[int] = None - lr: Optional[float] = None - low_gpu_mem_usage: Optional[bool] = None - minmax_lr: Optional[float] = None - nsamples: Optional[int] = None - quant_block_list: Optional[list[str]] = None - regex_config: Optional[dict[str, Any]] = None - scale_dtype: Optional[str] = None - seqlen: Optional[int] = None - supported_types: Optional[list[str]] = SUPPORTED_LAYER_TYPES - static_attention_dtype: Optional[str] = None - static_kv_dtype: Optional[str] = None - static_attention_granularity: Optional[str] = "tensor" - static_kv_granularity: Optional[str] = "tensor" - super_bits: Optional[int] = None - super_group_size: Optional[int] = None - to_quant_block_names: Optional[list[str]] = None - rotation_configs: Optional[list[dict[str, Any]]] = None + bits: int | None = None + act_bits: int | None = None + data_type: str | None = None + act_data_type: str | None = None + group_size: int | None = None + act_group_size: int | None = None + sym: bool | None = None + act_sym: bool | None = None + act_dynamic: bool | None = None + amp: bool | None = None + batch_size: int | None = None + enable_minmax_tuning: bool | None = True + enable_norm_bias_tuning: bool | None = False + enable_quanted_input: bool | None = True + gradient_accumulate_steps: int | None = None + iters: int | None = None + lr: float | None = None + low_gpu_mem_usage: bool | None = None + minmax_lr: float | None = None + nsamples: int | None = None + quant_block_list: list[str] | None = None + regex_config: dict[str, Any] | None = None + scale_dtype: str | None = None + seqlen: int | None = None + supported_types: list[str] | None = SUPPORTED_LAYER_TYPES + static_attention_dtype: str | None = None + static_kv_dtype: str | None = None + static_attention_granularity: str | None = "tensor" + static_kv_granularity: str | None = "tensor" + super_bits: int | None = None + super_group_size: int | None = None + to_quant_block_names: list[str] | None = None + rotation_configs: list[dict[str, Any]] | None = None SERIALIZATION_KEYS = tuple(field.name for field in fields(SerializedCompressorConfig)) @@ -190,7 +190,7 @@ def _formats_policy_string_of(formats) -> str: return ",".join(names) -class BaseOrchestrator(object): +class BaseOrchestrator: need_calib: bool = True compress_context: CompressContext = None model_context: ModelContext = None @@ -212,13 +212,13 @@ class BaseOrchestrator(object): locals()[_scheme_field] = _make_compressor_scheme_property(_scheme_field) @staticmethod - def _preload_model_config(model: Union[torch.nn.Module, str], trust_remote_code: bool) -> Optional[AutoConfig]: + def _preload_model_config(model: torch.nn.Module | str, trust_remote_code: bool) -> AutoConfig | None: if not isinstance(model, str): return None try: return AutoConfig.from_pretrained(model, trust_remote_code=trust_remote_code) - except (OSError, EnvironmentError, ValueError) as e: + except (OSError, ValueError) as e: logger.debug( "Failed to load config via AutoConfig.from_pretrained for %s: %s. " "Proceeding without config-based checks.", @@ -229,25 +229,25 @@ def _preload_model_config(model: Union[torch.nn.Module, str], trust_remote_code: def __init__( self, - config: Union[object, list[object]], - model: Union[torch.nn.Module, str], + config: object | list[object], + model: torch.nn.Module | str, tokenizer: Any = None, platform: str = "hf", - format: Union[str, list, None] = None, - scheme: Union[str, dict, QuantizationScheme, AutoScheme] = "W4A16", + format: str | list | None = None, + scheme: str | dict | QuantizationScheme | AutoScheme = "W4A16", low_gpu_mem_usage: bool = False, - device_map: Union[str, torch.device, int, dict] = 0, - enable_torch_compile: Optional[bool] = None, + device_map: str | torch.device | int | dict = 0, + enable_torch_compile: bool | None = None, seed: int = 42, low_cpu_mem_usage: bool = True, - layer_config: Optional[dict] = None, - nsamples: int = None, - seqlen: int = None, - scale_dtype: Optional[Union[str, torch.dtype]] = None, + layer_config: dict | None = None, + nsamples: int | None = None, + seqlen: int | None = None, + scale_dtype: str | torch.dtype | None = None, ignore_layers: str = "", quant_lm_head: bool = False, - to_quant_block_names: Optional[Union[str, list[str]]] = None, - dataset: Optional[Union[str, list, tuple, torch.utils.data.DataLoader]] = None, + to_quant_block_names: str | list[str] | None = None, + dataset: str | list | tuple | torch.utils.data.DataLoader | None = None, **kwargs, ) -> None: # ``CalibrationContext`` is the single source of truth for calibration @@ -584,10 +584,7 @@ def _needs_calibration_data(self) -> bool: # Layer-level scheme overrides can request static-activation paths # (e.g., global MXFP8 + local NVFP4 experts). Those still need # calibration data even when top-level scheme looks dynamic. - if self._layer_config_needs_calibration(check_need_act_calibration): - return True - - return False + return self._layer_config_needs_calibration(check_need_act_calibration) def _layer_config_needs_calibration(self, check_need_act_calibration) -> bool: """Return True if any raw layer_config entry implies activation calibration.""" @@ -647,12 +644,12 @@ def _replace_compression_plan(self, **changes) -> None: self.__dict__["compression_plan"] = replace(plan, **changes) @property - def scheme_context(self) -> Optional[QuantizationScheme]: + def scheme_context(self) -> QuantizationScheme | None: plan = self.__dict__.get("compression_plan") return plan.scheme.value if plan is not None else self.__dict__.get("_scheme_context") @scheme_context.setter - def scheme_context(self, value: Optional[QuantizationScheme]) -> None: + def scheme_context(self, value: QuantizationScheme | None) -> None: self.__dict__["_scheme_context"] = value plan = self.__dict__.get("compression_plan") if plan is not None and value is not None: @@ -675,7 +672,7 @@ def formats(self, value) -> None: self._replace_compression_plan(formats=tuple(value)) @property - def layer_config(self) -> Optional[dict]: + def layer_config(self) -> dict | None: plan = self.__dict__.get("compression_plan") if plan is None: return self.__dict__.get("_layer_config") @@ -688,7 +685,7 @@ def layer_config(self, value) -> None: self._replace_compression_plan(layer_config=value) @property - def regex_config(self) -> Optional[dict]: + def regex_config(self) -> dict | None: plan = self.__dict__.get("compression_plan") if plan is None: return self.__dict__.get("_regex_config") @@ -738,8 +735,8 @@ def quant_block_list(self, value) -> None: def resolve_scheme( self, - model_context: Optional[ModelContext] = None, - compress_context: Optional[CompressContext] = None, + model_context: ModelContext | None = None, + compress_context: CompressContext | None = None, ) -> None: """Phase-1 init: resolve scheme and bind config attrs (no model structure needed). @@ -1238,7 +1235,7 @@ def _maybe_log_torch_compile_default_hint(self) -> None: "'enable_torch_compile' is disabled. Enabling it can reduce tuning cost by about 20%.", ) - def _torch_compile_disabled_reason(self, ignore_user_override: bool = False) -> Optional[str]: + def _torch_compile_disabled_reason(self, ignore_user_override: bool = False) -> str | None: """Return why torch.compile must stay off for the current algorithm, else None. RTN and optimized RTN quantize each layer in a single pass, and very short @@ -1287,7 +1284,7 @@ def _torch_compile_disabled_reason(self, ignore_user_override: bool = False) -> return None - def _torch_compile_unsupported_arch_reason(self) -> Optional[str]: + def _torch_compile_unsupported_arch_reason(self) -> str | None: """Return why the model *architecture* forbids ``torch.compile``, else ``None``. Rules live in :mod:`auto_round.special_model_handler` so a new architecture can @@ -1530,7 +1527,7 @@ def alg_composer(self) -> Any: return self._alg_composer @staticmethod - def _resolve_gguf_preset_string(formats: list["OutputFormat"]) -> Optional[str]: + def _resolve_gguf_preset_string(formats: list["OutputFormat"]) -> str | None: """Return the precise GGUF preset string (e.g. ``"gguf:q4_k_m"``) for the single resolved GGUF format, or ``None`` if no GGUF format is present. @@ -1775,7 +1772,7 @@ def device(self) -> str: return device_manager.device @device.setter - def device(self, value: Union[str, torch.device]) -> None: + def device(self, value: str | torch.device) -> None: device_manager.device = value @property @@ -1840,7 +1837,7 @@ def _adjust_immediate_packing_and_saving(self): logger.warning("reset low_cpu_mem_usage to False due to tied weights") return if len(tied_weight_keys) == 1: - key = list(tied_weight_keys.keys())[0] + key = next(iter(tied_weight_keys)) if "lm_head" not in key: self.compress_context.is_immediate_saving = False if self.compress_context.low_cpu_mem_usage: @@ -1907,11 +1904,11 @@ def quantize(self) -> tuple[torch.nn.Module, dict[str, Any]]: def save_quantized( self, - output_dir: str = None, - format: Union[str, list[OutputFormat]] = None, + output_dir: str | None = None, + format: str | list[OutputFormat] | None = None, inplace: bool = True, return_folders: bool = False, - max_shard_size: Union[int, str] = None, + max_shard_size: int | str | None = None, **kwargs, ) -> torch.nn.Module: """Save the quantized model to the specified output directory in the specified format. @@ -1946,9 +1943,9 @@ def save_quantized( if isinstance(self.formats, str): self.formats = self._resolve_format_string(self.formats) self.compress_context.formats = self.formats - for format in self.formats: - save_folder = _get_save_folder_name(format) - if self.act_bits <= 8 and format.is_fake(): + for output_format in self.formats: + save_folder = _get_save_folder_name(output_format) + if self.act_bits <= 8 and output_format.is_fake(): logger.warning( "Support for exporting activation quantization is limited. " "Please ensure that your configuration is supported." @@ -1997,7 +1994,7 @@ def _revert_block_name(block_name): original_block_name, reverted_block_name ) - compressed_model = format.save_quantized( + compressed_model = output_format.save_quantized( save_folder, model=self.model_context.model, layer_config=self.layer_config, @@ -2100,9 +2097,9 @@ def _assert_w8_asym_exportable(self) -> None: def quantize_and_save( self, output_dir: str = "tmp_autoround", - format: str = None, + format: str | None = None, inplace: bool = True, - max_shard_size: Union[int, str] = None, + max_shard_size: int | str | None = None, **kwargs, ) -> tuple[torch.nn.Module, dict[str, Any]]: """Quantizes the model and saves it in the specified format(s). diff --git a/auto_round/compressors/config_resolution/__init__.py b/auto_round/compressors/config_resolution/__init__.py index dcdbb3b740..44332aef02 100644 --- a/auto_round/compressors/config_resolution/__init__.py +++ b/auto_round/compressors/config_resolution/__init__.py @@ -27,11 +27,11 @@ from auto_round.compressors.config_resolution.resolve import resolve_quantization_config, resolve_scheme_value __all__ = [ - "ResolvedQuantizationConfig", + "ConfigResolutionError", "FormatCompatibilityError", "FormatResolution", "LayerConfigResolutionError", - "ConfigResolutionError", + "ResolvedQuantizationConfig", "ResolvedScheme", "SchemeResolutionError", "resolve_quantization_config", diff --git a/auto_round/compressors/config_resolution/contracts.py b/auto_round/compressors/config_resolution/contracts.py index f095dc9367..e4cd2b74c8 100644 --- a/auto_round/compressors/config_resolution/contracts.py +++ b/auto_round/compressors/config_resolution/contracts.py @@ -15,14 +15,15 @@ from __future__ import annotations import copy +from collections.abc import Mapping from dataclasses import dataclass, field from types import MappingProxyType -from typing import Any, Mapping, Optional, Tuple +from typing import Any from auto_round.schemes import QuantizationScheme LayerConfig = Mapping[str, Mapping[str, Any]] -BlockGroups = Tuple[Tuple[str, ...], ...] +BlockGroups = tuple[tuple[str, ...], ...] def _deepcopy_mapping_proxy(value: MappingProxyType, memo: dict) -> MappingProxyType: @@ -40,7 +41,7 @@ def _deepcopy_mapping_proxy(value: MappingProxyType, memo: dict) -> MappingProxy copy._deepcopy_dispatch[MappingProxyType] = _deepcopy_mapping_proxy -def freeze_mapping(value: Optional[LayerConfig]) -> LayerConfig: +def freeze_mapping(value: LayerConfig | None) -> LayerConfig: """Return an isolated, read-only snapshot of a layer configuration mapping. Per-layer configuration values are usually dicts (e.g. ``{"bits": 4}``), but some @@ -57,7 +58,7 @@ def freeze_mapping(value: Optional[LayerConfig]) -> LayerConfig: return MappingProxyType(frozen) -def thaw_mapping(value: Optional[LayerConfig]) -> dict: +def thaw_mapping(value: LayerConfig | None) -> dict: """Return a fully mutable, deep-copyable plain-dict snapshot of a frozen mapping. This is the inverse of :func:`freeze_mapping` and should be used instead of @@ -71,7 +72,7 @@ def thaw_mapping(value: Optional[LayerConfig]) -> dict: return result -def freeze_block_groups(value: Optional[Tuple[Tuple[str, ...], ...]]) -> Optional[BlockGroups]: +def freeze_block_groups(value: tuple[tuple[str, ...], ...] | None) -> BlockGroups | None: """Freeze the two-dimensional block grouping used by model traversal.""" if value is None: return None @@ -81,7 +82,7 @@ def freeze_block_groups(value: Optional[Tuple[Tuple[str, ...], ...]]) -> Optiona @dataclass(frozen=True) class ResolvedScheme: _value: QuantizationScheme - preset_name: Optional[str] = None + preset_name: str | None = None def __post_init__(self) -> None: object.__setattr__(self, "_value", copy.deepcopy(self._value)) @@ -92,17 +93,17 @@ def value(self) -> QuantizationScheme: return copy.deepcopy(self._value) @classmethod - def from_scheme(cls, value: QuantizationScheme, preset_name: Optional[str] = None) -> "ResolvedScheme": + def from_scheme(cls, value: QuantizationScheme, preset_name: str | None = None) -> ResolvedScheme: return cls(_value=value, preset_name=preset_name) @dataclass(frozen=True) class FormatResolution: - formats: Tuple[Any, ...] + formats: tuple[Any, ...] scheme: ResolvedScheme layer_config_patch: LayerConfig = field(default_factory=lambda: MappingProxyType({})) scale_dtype: Any = None - quant_block_list: Optional[BlockGroups] = None + quant_block_list: BlockGroups | None = None def __post_init__(self) -> None: object.__setattr__(self, "formats", tuple(self.formats)) @@ -113,12 +114,12 @@ def __post_init__(self) -> None: @dataclass(frozen=True) class ResolvedQuantizationConfig: scheme: ResolvedScheme - formats: Tuple[Any, ...] + formats: tuple[Any, ...] layer_config: LayerConfig regex_config: LayerConfig = field(default_factory=lambda: MappingProxyType({})) has_qlayer_outside_block: bool = False scale_dtype: Any = None - quant_block_list: Optional[BlockGroups] = None + quant_block_list: BlockGroups | None = None def __post_init__(self) -> None: object.__setattr__(self, "formats", tuple(self.formats)) diff --git a/auto_round/compressors/config_resolution/resolve.py b/auto_round/compressors/config_resolution/resolve.py index 4d501ae0cf..6e1d9882e2 100644 --- a/auto_round/compressors/config_resolution/resolve.py +++ b/auto_round/compressors/config_resolution/resolve.py @@ -15,7 +15,8 @@ from __future__ import annotations import copy -from typing import Any, Mapping +from collections.abc import Mapping +from typing import Any from auto_round.compressors.config_resolution.contracts import ( FormatResolution, diff --git a/auto_round/compressors/diffusion/dataset.py b/auto_round/compressors/diffusion/dataset.py index 0b73b177bb..16348eb533 100644 --- a/auto_round/compressors/diffusion/dataset.py +++ b/auto_round/compressors/diffusion/dataset.py @@ -18,7 +18,6 @@ from concurrent.futures import ThreadPoolExecutor from io import StringIO from pathlib import Path -from typing import Dict, Optional import pandas as pd import torch @@ -27,7 +26,7 @@ from auto_round.utils import download_audiocaps_csv, logger -DIFFUSION_DATASET: Dict[str, Dataset] = {} +DIFFUSION_DATASET: dict[str, Dataset] = {} COCO_URL = { @@ -228,7 +227,7 @@ def __init__( self, dataset_path: str, nsamples: int = 128, - dataframe: Optional[pd.DataFrame] = None, + dataframe: pd.DataFrame | None = None, ) -> None: super().__init__() self.captions = [] @@ -272,7 +271,7 @@ def __init__( def __len__(self): return len(self.captions) - def __getitem__(self, i) -> Dict[str, torch.Tensor]: + def __getitem__(self, i) -> dict[str, torch.Tensor]: if self.image_paths is not None: return self.caption_ids[i], self.captions[i], self.image_paths[i] return self.caption_ids[i], self.captions[i] @@ -312,7 +311,7 @@ def __init__(self, dataset_path: str, nsamples: int = 128) -> None: def __len__(self): return len(self.captions) - def __getitem__(self, i) -> Dict[str, torch.Tensor]: + def __getitem__(self, i) -> dict[str, torch.Tensor]: return self.caption_ids[i], self.captions[i] diff --git a/auto_round/compressors/diffusion_mixin.py b/auto_round/compressors/diffusion_mixin.py index 415bfcbc71..7716bf1eb9 100644 --- a/auto_round/compressors/diffusion_mixin.py +++ b/auto_round/compressors/diffusion_mixin.py @@ -14,7 +14,7 @@ import json import math import os -from typing import Any, Optional, Union +from typing import Any import torch @@ -63,8 +63,8 @@ def __init__( guidance_scale: float = 7.5, num_inference_steps: int = 50, calib_num_inference_steps: int = 8, - generator_seed: Optional[int] = None, - diffusion_tuning_cache_size: Union[float, str] = 0, + generator_seed: int | None = None, + diffusion_tuning_cache_size: float | str = 0, **kwargs, ) -> None: if num_inference_steps < 1: @@ -422,8 +422,8 @@ def quantize(self) -> tuple[torch.nn.Module, dict]: def save_quantized( self, - output_dir: Optional[str] = None, - format: Optional[Union[str, list]] = None, + output_dir: str | None = None, + format: str | list | None = None, inplace: bool = True, return_folders: bool = False, **kwargs, diff --git a/auto_round/compressors/layer_config_resolver.py b/auto_round/compressors/layer_config_resolver.py index f98129b80c..4a70576e22 100644 --- a/auto_round/compressors/layer_config_resolver.py +++ b/auto_round/compressors/layer_config_resolver.py @@ -417,7 +417,7 @@ def resolve_layer_config( enable_gguf_official_mixed: bool = True, is_mllm: bool = False, fill_default_value: bool = True, - format: str = None, + format: str | None = None, ) -> LayerConfig: """Resolve final per-layer configuration without writing model attributes.""" layer_config = expand_layer_config_for_weight_renames(layer_config, model=model, to_model_names=True) @@ -488,7 +488,7 @@ def extract_regex_config( inner_supported_types=None, ignore_layers: str = "", fill_default_value: bool = True, - format: str = None, + format: str | None = None, ) -> LayerConfig: """Resolve only the regex entries retained for export metadata.""" layer_config = expand_layer_config_for_weight_renames(layer_config, model=model, to_model_names=True) @@ -545,8 +545,7 @@ def apply_plan_to_model(model, plan: ResolvedQuantizationConfig) -> None: # (``AttributeError`` on the norm's next forward). Scope the reset to the # same modules that receive the plan below. is_quant_target = ( - isinstance(module, SUPPORTED_LAYER_TYPES) - or isinstance(module, torch.nn.Embedding) + isinstance(module, (SUPPORTED_LAYER_TYPES, torch.nn.Embedding)) or module.__class__.__name__ in INNER_SUPPORTED_LAYER_TYPES ) if module_name != "" and not is_quant_target: diff --git a/auto_round/compressors/mllm/__init__.py b/auto_round/compressors/mllm/__init__.py index 1abc19f37e..90fdab74f9 100644 --- a/auto_round/compressors/mllm/__init__.py +++ b/auto_round/compressors/mllm/__init__.py @@ -17,10 +17,10 @@ from auto_round.compressors.mllm.template import TEMPLATES, Template, get_template __all__ = [ - "BasicProcessor", "MLLM_DATASET", "PROCESSORS", "TEMPLATES", + "BasicProcessor", "Template", "get_mllm_dataloader", "get_template", diff --git a/auto_round/compressors/mllm/dataset.py b/auto_round/compressors/mllm/dataset.py index fec57f44a4..e5bd4b1782 100644 --- a/auto_round/compressors/mllm/dataset.py +++ b/auto_round/compressors/mllm/dataset.py @@ -14,9 +14,10 @@ import json import os +import sys from collections import OrderedDict from inspect import signature -from typing import Any, Dict, Optional +from typing import Any import torch from torch.utils.data import DataLoader, Dataset @@ -28,7 +29,7 @@ from .template import Template from .utils import _extract_data_dir -MLLM_DATASET: Dict[str, Dataset] = {} +MLLM_DATASET: dict[str, Dataset] = {} def register_dataset(name_list): @@ -81,7 +82,7 @@ def __init__( model: torch.nn.Module, tokenizer: Any, dataset_path: str, - extra_data_dir: Optional[str] = None, + extra_data_dir: str | None = None, seqlen: int = 512, padding: bool = True, truncation: bool = True, @@ -95,7 +96,8 @@ def __init__( self.tokenizer = tokenizer if os.path.exists(dataset_path): logger.info(f"use dataset {dataset_path}, loading from disk...") - self.questions = json.load(open(dataset_path, "r")) + with open(dataset_path, "r") as f: + self.questions = json.load(f) else: if dataset_path == "liuhaotian/llava": dataset_path = "llava_conv_58k" @@ -186,8 +188,7 @@ def _check(questions, min_word_len, max_word_len, nsamples): if self.IMAGE_TOKEN in text["value"]: text["value"] = self.IMAGE_TOKEN + text["value"].replace(self.IMAGE_TOKEN, "") str_len += len(text["value"].split(" ")) - if str_len > max_len: - max_len = str_len + max_len = max(max_len, str_len) if min_word_len <= str_len < max_word_len: new_questions.append(source) if len(new_questions) >= nsamples: @@ -211,7 +212,7 @@ def _check(questions, min_word_len, max_word_len, nsamples): def __len__(self): return len(self.questions) - def __getitem__(self, i) -> Dict[str, torch.Tensor]: + def __getitem__(self, i) -> dict[str, torch.Tensor]: if self.cached_data_dict is not None and i in self.cached_data_dict: self.cached_data_dict.move_to_end(i) return self.cached_data_dict[i] @@ -304,7 +305,7 @@ def get_mllm_dataloader( template, model=model, tokenizer=tokenizer, processor=processor, image_processor=image_processor ) - if os.path.isfile(dataset) or dataset in MLLM_DATASET.keys(): + if os.path.isfile(dataset) or dataset in MLLM_DATASET: if seqlen > MLLM_DATASET[dataset].MAX_SUPPORT_SEQLEN: logger.warning( f"seqlen({seqlen}) is greater than the maximum length supported by the {dataset}," @@ -341,5 +342,5 @@ def get_mllm_dataloader( "Text only dataset cannot be used for calibrating non-text modules," " switching to liuhaotian/llava_conv_58k" ) - exit(-1) + sys.exit(-1) return dataloader, bs, seqlen diff --git a/auto_round/compressors/mllm/processor.py b/auto_round/compressors/mllm/processor.py index 00ea1da259..7f91af57bb 100644 --- a/auto_round/compressors/mllm/processor.py +++ b/auto_round/compressors/mllm/processor.py @@ -565,8 +565,9 @@ def load_system_prompt(repo_id_or_path: str, filename: str) -> str: file_path = hf_hub_download(repo_id=repo_id_or_path, filename=filename) with open(file_path, "r") as file: system_prompt = file.read() - today = datetime.today().strftime("%Y-%m-%d") - yesterday = (datetime.today() - timedelta(days=1)).strftime("%Y-%m-%d") + now = datetime.now().astimezone() + today = now.strftime("%Y-%m-%d") + yesterday = (now - timedelta(days=1)).strftime("%Y-%m-%d") model_name = repo_id_or_path.split("/")[-1] return system_prompt.format(name=model_name, today=today, yesterday=yesterday) diff --git a/auto_round/compressors/mllm/template.py b/auto_round/compressors/mllm/template.py index 31a75612bf..e3ece7b14d 100644 --- a/auto_round/compressors/mllm/template.py +++ b/auto_round/compressors/mllm/template.py @@ -16,13 +16,12 @@ import os from dataclasses import dataclass from enum import Enum, unique -from typing import Dict, List, Optional from auto_round.logger import logger from .processor import PROCESSORS, BasicProcessor -TEMPLATES: Dict[str, "Template"] = {} +TEMPLATES: dict[str, "Template"] = {} def fill_content(target, **kwargs): @@ -50,7 +49,7 @@ class Template: format_observation: str format_separator: str default_system: str - replace_tokens: List[tuple] + replace_tokens: list[tuple] extra_encode: bool default_dataset: str processor: "BasicProcessor" @@ -80,16 +79,16 @@ def _encode(self, sources): def _register_template( model_type: str, - format_user: Optional[str] = None, - format_assistant: Optional[str] = None, - format_system: Optional[str] = None, - format_function: Optional[str] = None, - format_observation: Optional[str] = None, - format_separator: Optional[str] = None, + format_user: str | None = None, + format_assistant: str | None = None, + format_system: str | None = None, + format_function: str | None = None, + format_observation: str | None = None, + format_separator: str | None = None, default_system: str = "", - replace_tokens: List[tuple] = None, - extra_encode: Optional[bool] = False, - default_dataset: Optional[bool] = "NeelNanda/pile-10k", + replace_tokens: list[tuple] | None = None, + extra_encode: bool | None = False, + default_dataset: bool | None = "NeelNanda/pile-10k", processor: "BasicProcessor" = PROCESSORS["basic"], ): """Registers a chat template.""" diff --git a/auto_round/compressors/mllm/utils.py b/auto_round/compressors/mllm/utils.py index eda37e92b8..2729a4fa16 100644 --- a/auto_round/compressors/mllm/utils.py +++ b/auto_round/compressors/mllm/utils.py @@ -43,7 +43,7 @@ def _extract_data_dir(dir_path: str): def fetch_image(path_or_url): if os.path.isfile(path_or_url): image_obj = Image.open(path_or_url) - elif path_or_url.startswith("http://") or path_or_url.startswith("https://"): + elif path_or_url.startswith(("http://", "https://")): try: response = requests.get(path_or_url, stream=True, timeout=(3, 10)) response.raise_for_status() diff --git a/auto_round/compressors/mllm_mixin.py b/auto_round/compressors/mllm_mixin.py index 4f61f83d69..59581dd030 100644 --- a/auto_round/compressors/mllm_mixin.py +++ b/auto_round/compressors/mllm_mixin.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -from typing import Any, Optional, Union +from typing import Any from auto_round.logger import logger @@ -49,8 +49,8 @@ def __init__( *args, processor: Any = None, image_processor: Any = None, - template: Optional[str] = None, - extra_data_dir: Optional[str] = None, + template: str | None = None, + extra_data_dir: str | None = None, quant_nontext_module: bool = False, **kwargs, ) -> None: @@ -103,8 +103,8 @@ def _get_calibrator_kind(self) -> str: def save_quantized( self, - output_dir: Optional[str] = None, - format: Union[str, list] = "auto_round", + output_dir: str | None = None, + format: str | list = "auto_round", inplace: bool = True, **kwargs, ) -> Any: diff --git a/auto_round/compressors/model_free.py b/auto_round/compressors/model_free.py index 89ff392e79..9d5c54f4bd 100644 --- a/auto_round/compressors/model_free.py +++ b/auto_round/compressors/model_free.py @@ -110,7 +110,7 @@ import time from concurrent.futures import FIRST_COMPLETED, ProcessPoolExecutor, ThreadPoolExecutor, as_completed, wait from dataclasses import asdict, fields -from typing import Any, Optional, Union +from typing import Any import torch from safetensors import safe_open @@ -545,14 +545,14 @@ def __init__( self, model_name_or_path: str, output_dir: str, - scheme: Union[str, QuantizationScheme] = "W4A16", - layer_config: Optional[dict] = None, + scheme: str | QuantizationScheme = "W4A16", + layer_config: dict | None = None, ignore_layers: str = "", - format: Optional[str] = None, + format: str | None = None, device: str = "cpu", quant_lm_head: bool = False, quant_nontext_module: bool = False, - enable_torch_compile: Optional[bool] = None, + enable_torch_compile: bool | None = None, disable_opt_rtn: bool = False, ) -> None: # --- raw inputs --- @@ -790,7 +790,7 @@ def _check_conv1d_and_embedding(self) -> None: if incompatible: # Group by class for a cleaner warning message incompatible_layers = [] - for cls, layers in incompatible.items(): + for layers in incompatible.values(): incompatible_layers.extend(layers) summary = ", ".join(f"{cls}({len(layers)})" for cls, layers in sorted(incompatible.items())) self.ignore_patterns.extend(incompatible_layers) @@ -1758,11 +1758,11 @@ class ModelFreeCompressor(_ModelFreeCompressorCore): def __init__( self, model_name_or_path: str, - output_dir: Optional[str] = None, - scheme: Union[str, QuantizationScheme] = "W4A16", - layer_config: Optional[dict] = None, + output_dir: str | None = None, + scheme: str | QuantizationScheme = "W4A16", + layer_config: dict | None = None, ignore_layers: str = "", - format: Optional[str] = None, + format: str | None = None, device: str = "cpu", quant_lm_head: bool = False, quant_nontext_module: bool = False, @@ -1770,7 +1770,7 @@ def __init__( tokenizer: Any = None, device_map: Any = None, low_cpu_mem_usage: bool = True, - enable_torch_compile: Optional[bool] = None, + enable_torch_compile: bool | None = None, disable_opt_rtn: bool = False, **kwargs, ) -> None: @@ -1811,7 +1811,7 @@ def __init__( # Compressor-role state (mirrors BaseCompressor attributes used by # AutoRound's post-processing code) - self._output_dir_override: Optional[str] = None # set by quantize_and_save + self._output_dir_override: str | None = None # set by quantize_and_save self.model = None self.tokenizer = tokenizer self.model_free = True @@ -1857,9 +1857,9 @@ def __init__( # AutoScheme (two-phase delta-loss selection) state. self._auto_scheme_resolved = False - self._auto_scheme_family: Optional[str] = None + self._auto_scheme_family: str | None = None - def _fallback_to_base_compressor(self, save_format: Optional[str] = None): + def _fallback_to_base_compressor(self, save_format: str | None = None): from auto_round.autoround import AutoRound logger.info( diff --git a/auto_round/compressors/orchestrator.py b/auto_round/compressors/orchestrator.py index 350c7b4861..96d1a262a2 100644 --- a/auto_round/compressors/orchestrator.py +++ b/auto_round/compressors/orchestrator.py @@ -16,7 +16,7 @@ import os import time from functools import partial -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Optional import accelerate import torch @@ -69,15 +69,15 @@ class CompressionOrchestrator(BaseOrchestrator): def __init__( self, - config: Union[object, list[object]], # TODO rename this to alg_config wenhuach - model: Union[torch.nn.Module, str], + config: object | list[object], # TODO rename this to alg_config wenhuach + model: torch.nn.Module | str, tokenizer: Any = None, platform: str = "hf", - format: Union[str, list, None] = None, - dataset: Optional[Union[str, list, tuple, torch.utils.data.DataLoader]] = None, + format: str | list | None = None, + dataset: str | list | tuple | torch.utils.data.DataLoader | None = None, low_gpu_mem_usage: bool = False, - device_map: Union[str, torch.device, int, dict] = 0, - enable_torch_compile: Optional[bool] = None, + device_map: str | torch.device | int | dict = 0, + enable_torch_compile: bool | None = None, seed: int = 42, low_cpu_mem_usage: bool = True, **kwargs, @@ -126,8 +126,8 @@ def cache_data( self, block_names: list, nsamples: int, - layer_names: Optional[list] = None, - last_cache_name: Optional[str] = None, + layer_names: list | None = None, + last_cache_name: str | None = None, ) -> Any: """Thin wrapper around ``self.calibration.collect``. @@ -520,7 +520,7 @@ def _quantize_zero_shot(self) -> tuple[torch.nn.Module, dict[str, Any]]: self.alg_composer.finalize_run() remain_layer_names = [] - block_name_set = set(name for block in all_blocks for name in block) + block_name_set = {name for block in all_blocks for name in block} for n, m in self.model.named_modules(): if not check_to_quantized(m): continue @@ -605,7 +605,7 @@ def _quantize_data_driven(self) -> tuple[torch.nn.Module, dict[str, Any]]: all_q_inputs = None # Leave it to gguf itself to handle if has_gguf and self.alg_composer.need_quanted_input(): # pylint: disable=E1101 - is_quantized_embedding = self.alg_composer.compress_embedding_layer() # + is_quantized_embedding = self.alg_composer.compress_embedding_layer() clear_memory() if is_quantized_embedding: all_inputs = copy.deepcopy(self.inputs) @@ -985,8 +985,8 @@ def quantize_block( self, block: torch.nn.Module, inputs: Any, - q_input: Union[torch.Tensor, dict, None] = None, - device: Union[str, torch.device] = "cpu", + q_input: torch.Tensor | dict | None = None, + device: str | torch.device = "cpu", auto_offload: bool = True, reference_output=None, ) -> Any: diff --git a/auto_round/compressors/shard_writer.py b/auto_round/compressors/shard_writer.py index 6232b3d884..5e26844d14 100644 --- a/auto_round/compressors/shard_writer.py +++ b/auto_round/compressors/shard_writer.py @@ -16,7 +16,7 @@ import os import re from collections import OrderedDict -from typing import Optional, Union +from typing import Optional import torch @@ -59,7 +59,7 @@ def __init__( self, model: torch.nn.Module, bits: int, - max_shard_size: Optional[Union[int, str]] = None, + max_shard_size: int | str | None = None, safe_serialization: bool = True, ) -> None: if ShardWriter._initialized: @@ -218,7 +218,7 @@ def _check_safetensors(self) -> bool: logger.warning("safetensors not installed; falling back to torch.save.") return False - def save_module(self, m: torch.nn.Module, name: str = None) -> None: + def save_module(self, m: torch.nn.Module, name: str | None = None) -> None: """Extracts and accumulates tensors from a module.""" prefix = name if name is not None else getattr(m, "global_name", "model") sd = m.state_dict() @@ -483,7 +483,7 @@ def finalize(self) -> None: logger.info(f"model has been saved to {self.output_dir}") @torch.no_grad() - def write(self, m: torch.nn.Module = None, name: str = None, is_finalize: bool = False) -> None: + def write(self, m: torch.nn.Module = None, name: str | None = None, is_finalize: bool = False) -> None: if m is None and name is None and not is_finalize and not is_finalize: raise ValueError("Must specify either name or m") if m is None and name is not None: diff --git a/auto_round/compressors/utils.py b/auto_round/compressors/utils.py index a0593bab01..eb149e6e3c 100644 --- a/auto_round/compressors/utils.py +++ b/auto_round/compressors/utils.py @@ -123,7 +123,7 @@ def block_forward( amp_dtype: torch.dtype = torch.float16, device: torch.device = torch.device("cpu"), output_return_id: int = 0, -) -> Union[torch.Tensor, dict]: +) -> torch.Tensor | dict: """Performs a forward pass through a block with the given inputs. Args: @@ -178,7 +178,7 @@ def block_forward( output = block(**input_others) else: output = block(**input_others) - if isinstance(output_return_id, int) and (isinstance(output, list) or isinstance(output, tuple)): + if isinstance(output_return_id, int) and isinstance(output, (list, tuple)): output = output[output_return_id] return output @@ -195,11 +195,11 @@ def check_skippable_keywords(key): def check_need_act_calibration( - is_act_dynamic: Union[bool, None], - act_data_type: Union[str, None] = None, - act_bits: Union[int, None] = 16, - static_kv_dtype: Union[str, None] = None, - static_attention_dtype: Union[str, None] = None, + is_act_dynamic: bool | None, + act_data_type: str | None = None, + act_bits: int | None = 16, + static_kv_dtype: str | None = None, + static_attention_dtype: str | None = None, ) -> bool: if static_kv_dtype is not None or static_attention_dtype is not None: return True @@ -208,9 +208,7 @@ def check_need_act_calibration( # None is dynamic if is_act_dynamic is not None and not is_act_dynamic: return True - if act_data_type is not None and "static" in act_data_type: - return True - return False + return act_data_type is not None and "static" in act_data_type def collect_best_params(block, cache_device="cpu"): @@ -485,11 +483,7 @@ def _get_diffusion_save_folder_name(format) -> str: formats = compress_context.formats # Use a subfolder only if there are multiple formats if len(formats) > 1: - return ( - os.path.join(compress_context.output_dir, sanitized_format, "transformer") - if compress_context.is_immediate_saving - else os.path.join(compress_context.output_dir, sanitized_format, "transformer") - ) + return os.path.join(compress_context.output_dir, sanitized_format, "transformer") # if use is_immediate_saving, we need to save model in self.output_dir/transformer folder return ( diff --git a/auto_round/context/compress.py b/auto_round/context/compress.py index 6cb0bfbcae..d90036c480 100644 --- a/auto_round/context/compress.py +++ b/auto_round/context/compress.py @@ -11,7 +11,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -from typing import Callable, Optional, Union +from collections.abc import Callable import torch @@ -36,10 +36,10 @@ def __init__( enable_torch_compile: bool = True, is_immediate_packing: bool = False, is_immediate_saving: bool = False, - formats: Union[list, str] = None, + formats: list | str | None = None, output_dir: str = "./compressed_models", - static_kv_dtype: Optional[torch.dtype] = None, # TODO later this should be scheme wenhuach - static_attention_dtype: Optional[torch.dtype] = None, + static_kv_dtype: torch.dtype | None = None, # TODO later this should be scheme wenhuach + static_attention_dtype: torch.dtype | None = None, static_kv_granularity: str = "tensor", static_attention_granularity: str = "tensor", **kwargs, diff --git a/auto_round/context/model.py b/auto_round/context/model.py index 7fb5c23385..a4c6c03a44 100644 --- a/auto_round/context/model.py +++ b/auto_round/context/model.py @@ -15,7 +15,8 @@ import gc import importlib import os -from typing import Any, Callable, Optional, Union +from collections.abc import Callable +from typing import Any import torch import transformers @@ -61,12 +62,12 @@ class ModelContext(BaseContext): def __init__( self, - model: Union[torch.nn.Module, str, None] = None, + model: torch.nn.Module | str | None = None, tokenizer: Any = None, platform: str = "hf", - model_dtype: Optional[Union[str, torch.dtype]] = None, + model_dtype: str | torch.dtype | None = None, trust_remote_code: bool = True, - config: Optional[AutoConfig] = None, + config: AutoConfig | None = None, amp: bool = True, need_calib: bool = True, is_act_quantize: bool = False, @@ -219,7 +220,7 @@ def _load_model(self): if config is None: config = AutoConfig.from_pretrained(self.model, trust_remote_code=self.trust_remote_code) self._import_custom_moe_replacements(config) - except (OSError, EnvironmentError, ValueError) as e: + except (OSError, ValueError) as e: logger.debug( "Failed to load config via AutoConfig.from_pretrained for %s: %s. " "Proceeding without config-based checks.", @@ -353,7 +354,7 @@ def _should_use_meta_skeleton(self, config) -> bool: ) return True - def _resolve_local_checkpoint_dir(self) -> Optional[str]: + def _resolve_local_checkpoint_dir(self) -> str | None: """Return a local directory holding the checkpoint shards, or ``None``. ``SafetensorsIndex`` reads shards by path, so a hub repo id has to be resolved to @@ -449,7 +450,7 @@ def _patch_custom_moe_modules(self) -> None: gate = getattr(module, "gate", None) top_k = getattr(gate, "top_k", None) if top_k is not None: - setattr(module, "top_k", top_k) + module.top_k = top_k def _set_amp_dtype(self) -> None: """Sets the automatic mixed precision (AMP) data type for the model based on the device and configuration. diff --git a/auto_round/data_type/gguf.py b/auto_round/data_type/gguf.py index 2b4de2b7ee..f111ef4304 100644 --- a/auto_round/data_type/gguf.py +++ b/auto_round/data_type/gguf.py @@ -11,7 +11,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -from typing import Any, Callable, Union +from collections.abc import Callable import torch @@ -435,10 +435,10 @@ def _gather_row_pattern(row_pattern: torch.Tensor, start: int, end: int) -> torc def _imatrix_handle_zero( - imatrix: Union[torch.Tensor, float], + imatrix: torch.Tensor | float, weight: torch.Tensor, bits: int, - group_size: Union[int, None] = None, + group_size: int | None = None, raw_imatrix: torch.Tensor = None, ): if not isinstance(imatrix, torch.Tensor): diff --git a/auto_round/data_type/int.py b/auto_round/data_type/int.py index d301f0b2e9..857c83fb18 100644 --- a/auto_round/data_type/int.py +++ b/auto_round/data_type/int.py @@ -11,7 +11,6 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -from typing import Union import torch @@ -21,7 +20,7 @@ from auto_round.utils import get_reciprocal -def search_scales(data: torch.Tensor, bits: int, qw: Union[None, torch.Tensor, float] = None) -> torch.Tensor: +def search_scales(data: torch.Tensor, bits: int, qw: None | torch.Tensor | float = None) -> torch.Tensor: # Maximum absolute value for symmetric quantization nmax = int(2.0 ** (bits - 1)) diff --git a/auto_round/data_type/mxfp.py b/auto_round/data_type/mxfp.py index fd0de4758c..c9ff8b22de 100644 --- a/auto_round/data_type/mxfp.py +++ b/auto_round/data_type/mxfp.py @@ -125,9 +125,7 @@ def search_mx_scale(tensor, bits, qw=None, data_type=None): def compute_loss(qdq_tensor, out): torch.sub(qdq_tensor, tensor, out=buf) buf.pow_(2) - if qw is not None and not isinstance(qw, (int, float)): - buf.mul_(qw) - elif qw != 1.0: + if (qw is not None and not isinstance(qw, (int, float))) or qw != 1.0: buf.mul_(qw) torch.sum(buf, dim=-1, out=out) @@ -468,7 +466,7 @@ def quant_mx_rceil_v2( return tensor.to(orig_dtype), shared_exp.to(orig_dtype), None -for key in MXFP_FORMAT_CACHE.keys(): +for key in MXFP_FORMAT_CACHE: QUANT_FUNC_WITH_DTYPE[key] = quant_mx QUANT_FUNC_WITH_DTYPE[key + "_rceil"] = quant_mx_rceil QUANT_FUNC_WITH_DTYPE["opt_rtn_" + key] = quant_mx_opt_rtn diff --git a/auto_round/data_type/neuqi.py b/auto_round/data_type/neuqi.py index d83afeb6d7..4171e3ad21 100644 --- a/auto_round/data_type/neuqi.py +++ b/auto_round/data_type/neuqi.py @@ -613,11 +613,11 @@ def neuqi_search_scale_zero(data, bits, qw=None, q_scale_thresh=1e-5, coarse_n=N if coarse_n is None: coarse_n = envs.AR_NEUQI_COARSE or d_coarse if coarse_n < 2: - raise ValueError("AR_NEUQI_COARSE=%d is too small; need at least 2 candidates." % coarse_n) + raise ValueError(f"AR_NEUQI_COARSE={coarse_n} is too small; need at least 2 candidates.") if fine_n is None: fine_n = envs.AR_NEUQI_FINE or d_fine if fine_n < 2: - raise ValueError("AR_NEUQI_FINE=%d is too small; need at least 2 candidates." % fine_n) + raise ValueError(f"AR_NEUQI_FINE={fine_n} is too small; need at least 2 candidates.") _log_search_engaged(coarse_n, fine_n) if torch.is_tensor(data): ensure_sweep_warmup(data.device) @@ -1197,11 +1197,11 @@ def neuqi_search_scale_sym(data, bits, qw=None, q_scale_thresh=1e-5, coarse_n=No # default is backend-aware (wide only on the accelerated lanes) coarse_n = envs.AR_NEUQI_COARSE or d_coarse if coarse_n < 2: - raise ValueError("AR_NEUQI_COARSE=%d is too small; need at least 2 candidates." % coarse_n) + raise ValueError(f"AR_NEUQI_COARSE={coarse_n} is too small; need at least 2 candidates.") if fine_n is None: fine_n = envs.AR_NEUQI_FINE or d_fine if fine_n < 2: - raise ValueError("AR_NEUQI_FINE=%d is too small; need at least 2 candidates." % fine_n) + raise ValueError(f"AR_NEUQI_FINE={fine_n} is too small; need at least 2 candidates.") _log_sym_search_engaged(coarse_n, fine_n) ensure_sym_search_warmup(data.device) diff --git a/auto_round/data_type/utils.py b/auto_round/data_type/utils.py index 676b2822e8..99f027c723 100644 --- a/auto_round/data_type/utils.py +++ b/auto_round/data_type/utils.py @@ -15,7 +15,6 @@ import math from functools import lru_cache from math import ceil -from typing import List, Union import torch from torch.nn import Linear, Module @@ -26,7 +25,7 @@ from auto_round.utils import check_to_quantized, logger -def reshape_pad_tensor_by_group_size(data: torch.Tensor, group_size: Union[int, list], val: float = 0.0): +def reshape_pad_tensor_by_group_size(data: torch.Tensor, group_size: int | list, val: float = 0.0): """Reshapes and pads the tensor to ensure that it can be quantized in groups of `group_size`. This function adjusts the @@ -71,7 +70,7 @@ def reshape_pad_tensor_by_group_size(data: torch.Tensor, group_size: Union[int, return data_new, orig_shape, pad_len -def revert_tensor_by_pad(data: torch.Tensor, orig_shape: tuple, pad_len: Union[int, list]): +def revert_tensor_by_pad(data: torch.Tensor, orig_shape: tuple, pad_len: int | list): """Reverts the tensor to its original shape by removing padding. This function removes the padding added during reshaping and returns the tensor to @@ -467,7 +466,7 @@ def update_fused_layer_global_scales( global_scale_name = f"{base_name}_global_scale" - def _collect_scales(mods: List[Module]) -> List[torch.Tensor]: + def _collect_scales(mods: list[Module]) -> list[torch.Tensor]: """Collect valid global_scale tensors from modules.""" scales = [] for m in mods: @@ -488,7 +487,7 @@ def _is_moe_expert_module(module: Module): """Check for MoE expert naming: w1 (gate) and w3 (up).""" return all(hasattr(module, projection) for projection in ("w1", "w3")) - def _update_global_scales(modules: List[Module]): + def _update_global_scales(modules: list[Module]): """Update global scales for a list of modules.""" scales = _collect_scales(modules) if not scales: @@ -544,7 +543,7 @@ def update_block_global_scale_if_needed(block, data_type, group_size): has_nvfp = True if not hasattr(m, "weight_global_scale"): weight_global_scale = calculate_gparam(m.weight, module_group_size) - setattr(m, "weight_global_scale", weight_global_scale) + m.weight_global_scale = weight_global_scale if not has_nvfp: return diff --git a/auto_round/envs.py b/auto_round/envs.py index 368351c094..c0337fb75a 100644 --- a/auto_round/envs.py +++ b/auto_round/envs.py @@ -15,32 +15,33 @@ # For detailed usage and configuration guide, see: docs/environments.md import os -from typing import TYPE_CHECKING, Any, Callable, Optional +from collections.abc import Callable +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: AR_LOG_LEVEL: str = "INFO" AR_USE_MODELSCOPE: bool = "False" - AR_MODEL_FREE_SHARD_PARALLELISM: Optional[int] = None - AUTO_ROUND_CACHE: Optional[str] = None + AR_MODEL_FREE_SHARD_PARALLELISM: int | None = None + AUTO_ROUND_CACHE: str | None = None AUTO_ROUND_GGUF_AUTO_UPDATE: bool = False AR_DISABLE_GGUF_MTP_EXPORT: bool = False - LLAMA_CPP_ROOT: Optional[str] = None - AR_AUTO_SCHEME_NSAMPLES: Optional[int] = None - AR_AUTO_SCHEME_BATCH_SIZE: Optional[int] = None - AR_AUTO_SCHEME_CACHE: Optional[str] = None + LLAMA_CPP_ROOT: str | None = None + AR_AUTO_SCHEME_NSAMPLES: int | None = None + AR_AUTO_SCHEME_BATCH_SIZE: int | None = None + AR_AUTO_SCHEME_CACHE: str | None = None AR_AUTO_SCHEME_NO_SERIAL_FALLBACK: bool = False AR_ENABLE_AUTO_SCHEME_PARALLEL: bool = True AR_NVFP4_E5M3_CACHE_HP_WEIGHT: bool = False AR_DISK_STREAM_MODEL: bool = False AR_DISABLE_META_LOAD: bool = False - AR_RESUME_DIR: Optional[str] = None + AR_RESUME_DIR: str | None = None AR_FORCE_MOE_ROUTING_ALL_EXPERTS: bool = False AR_QUANTIZE_BAGEL_MOE_GEN: bool = False AR_NVFP4_FUSED_LAYER_GLOBAL_SCALE: bool = True AR_ALLOW_W8_ASYM: bool = False -def _get_optional_positive_int_env(name: str) -> Optional[int]: +def _get_optional_positive_int_env(name: str) -> int | None: """Read an optional env var that must be a positive integer when set.""" raw = os.getenv(name) if raw is None: diff --git a/auto_round/eval/eval_cli.py b/auto_round/eval/eval_cli.py index 27d28994f9..4433e6c295 100644 --- a/auto_round/eval/eval_cli.py +++ b/auto_round/eval/eval_cli.py @@ -241,7 +241,7 @@ def eval(args): add_bos_token=args.add_bos_token, ) print(make_table(res)) - print("evaluation running time=%ds" % (time.time() - st)) + print(f"evaluation running time={int(time.time() - st)}s") else: st = time.time() if "auto" in str(batch_size) and args.mllm: @@ -262,7 +262,7 @@ def eval(args): from lm_eval.utils import make_table # pylint: disable=E0401 print(make_table(res)) - print("evaluation running time=%ds" % (time.time() - st)) + print(f"evaluation running time={int(time.time() - st)}s") def eval_with_vllm(args): @@ -349,7 +349,7 @@ def eval_with_vllm(args): ) print(make_table(res)) - print("evaluation running time=%ds" % (time.time() - st)) + print(f"evaluation running time={int(time.time() - st)}s") def eval_task_by_task( diff --git a/auto_round/eval/evaluation.py b/auto_round/eval/evaluation.py index ab4638edef..48ff1cfd50 100644 --- a/auto_round/eval/evaluation.py +++ b/auto_round/eval/evaluation.py @@ -13,7 +13,7 @@ # limitations under the License. import os -from typing import Optional, Union +import sys from auto_round.logger import logger from auto_round.utils import dispatch_model_block_wise @@ -57,9 +57,9 @@ def _normalize_model_eval_dtype(model, eval_model_dtype): def simple_evaluate_user_model( user_model, tokenizer, - batch_size: Optional[int] = 1, - limit: Optional[Union[int, float]] = None, - max_batch_size: Optional[int] = 64, + batch_size: int | None = 1, + limit: float | None = None, + max_batch_size: int | None = 64, eval_model_dtype="auto", add_bos_token: bool = False, mllm: bool = False, @@ -98,11 +98,11 @@ def simple_evaluate_user_model( def simple_evaluate( model, - model_args: Optional[Union[str, dict]] = None, - batch_size: Optional[int] = None, - limit: Optional[Union[int, float]] = None, - max_batch_size: Optional[int] = None, - device: Optional[str] = None, + model_args: str | dict | None = None, + batch_size: int | None = None, + limit: float | None = None, + max_batch_size: int | None = None, + device: str | None = None, **kwargs, ): import lm_eval # pylint: disable=E0401 @@ -141,7 +141,7 @@ def evaluate_diffusion_model(args, autoround=None, model=None, pipe=None): logger.error( "Quantized model is meta and diffusers doesn't support loading auto-round quantized model now. Exit." ) - exit(0) + sys.exit(0) pipe = autoround.pipe pipe.to(model.dtype) pipe.transformer = model @@ -336,7 +336,7 @@ def evaluate_with_model_instance(model, tokenizer, device_str, args): fewshot_as_multiturn=getattr(args, "fewshot_as_multiturn", False), ) print(make_table(res)) - print("evaluation running time=%ds" % (time.time() - st)) + print(f"evaluation running time={int(time.time() - st)}s") def evaluate_with_model_path(eval_folder, device_str, autoround, args): @@ -412,7 +412,7 @@ def evaluate_with_model_path(eval_folder, device_str, autoround, args): fewshot_as_multiturn=getattr(args, "fewshot_as_multiturn", False), ) print(make_table(res)) - print("evaluation running time=%ds" % (time.time() - st)) + print(f"evaluation running time={int(time.time() - st)}s") def run_model_evaluation(model, tokenizer, autoround, folders, formats, args): diff --git a/auto_round/experimental/attention.py b/auto_round/experimental/attention.py index 476159505a..546e2329df 100644 --- a/auto_round/experimental/attention.py +++ b/auto_round/experimental/attention.py @@ -18,8 +18,8 @@ import contextlib import inspect +from collections.abc import Callable from functools import partial -from typing import Callable, Optional from weakref import ref import torch @@ -41,9 +41,9 @@ __all__ = [ "QuantizedAttentionImpl", + "attention_quant_ctx", "init_hooked_attention", "is_attention_calibration_tensor_name", - "attention_quant_ctx", ] diff --git a/auto_round/experimental/kv_cache.py b/auto_round/experimental/kv_cache.py index 948a48f911..83bef24f36 100644 --- a/auto_round/experimental/kv_cache.py +++ b/auto_round/experimental/kv_cache.py @@ -19,7 +19,7 @@ import contextlib from enum import Enum from functools import partial -from typing import Any, Dict, List, Optional, Tuple, Union +from typing import Any import torch from transformers.cache_utils import DynamicCache @@ -37,10 +37,10 @@ from auto_round.utils import logger __all__ = [ - "initialize_quantized_kv_cache", - "prep_attention_module_for_calibration", "freeze_module_quantization_", + "initialize_quantized_kv_cache", "kvcache_quant_context", + "prep_attention_module_for_calibration", ] @@ -71,7 +71,7 @@ class KVCacheScaleType(Enum): # NOTE: Using _ suffix to denote l is modified in place -def _pad_and_append_at_idx_(lst: List, idx: int, val: Any) -> list: +def _pad_and_append_at_idx_(lst: list, idx: int, val: Any) -> list: """ Append value val to list lst at index idx, right padding if necessary Needed because user may ignore some layers in configuration, meaning @@ -107,7 +107,7 @@ class QuantizedKVParameterCache(DynamicCache): def __new__(cls, *args, **kwargs): """Singleton""" if cls._instance is None: - cls._instance = super(QuantizedKVParameterCache, cls).__new__(cls) + cls._instance = super().__new__(cls) return cls._instance def __init__(self, dtype: torch.dtype | str = torch.float8_e4m3fn, granularity: str = "tensor"): @@ -125,10 +125,10 @@ def __init__(self, dtype: torch.dtype | str = torch.float8_e4m3fn, granularity: super().__init__() # each index corresponds to layer_idx of the attention layer - self.k_scales: List[torch.Tensor] = [] - self.v_scales: List[torch.Tensor] = [] - self.k_amax: List[float] = [] - self.v_amax: List[float] = [] + self.k_scales: list[torch.Tensor] = [] + self.v_scales: list[torch.Tensor] = [] + self.k_amax: list[float] = [] + self.v_amax: list[float] = [] self._initialized = True def update( @@ -136,8 +136,8 @@ def update( key_states: torch.Tensor, value_states: torch.Tensor, layer_idx: int, - cache_kwargs: Optional[Dict[str, Any]] = None, - ) -> Tuple[torch.Tensor, torch.Tensor]: + cache_kwargs: dict[str, Any] | None = None, + ) -> tuple[torch.Tensor, torch.Tensor]: """ Get the k_scale and v_scale and output the quant-dequant key_states and value_states """ @@ -155,7 +155,7 @@ def update( return keys_to_return, values_to_return - def get_seq_length(self, layer_idx: Optional[int] = 0) -> int: + def get_seq_length(self, layer_idx: int | None = 0) -> int: """ Returns the sequence length of the cached states. A layer index can be optionally passed. @@ -173,12 +173,12 @@ def get_seq_length(self, layer_idx: Optional[int] = 0) -> int: def reset_states(self): """reset the kv states (used in calibration)""" - self.key_cache: List[torch.Tensor] = [] - self.value_cache: List[torch.Tensor] = [] + self.key_cache: list[torch.Tensor] = [] + self.value_cache: list[torch.Tensor] = [] # Used in `generate` to keep tally of how many tokens the cache has seen self._seen_tokens = 0 - self._quantized_key_cache: List[torch.Tensor] = [] - self._quantized_value_cache: List[torch.Tensor] = [] + self._quantized_key_cache: list[torch.Tensor] = [] + self._quantized_value_cache: list[torch.Tensor] = [] def reset(self): """ @@ -253,7 +253,7 @@ def initialize_quantized_kv_cache(module: torch.nn.Module, dtype=torch.float8_e4 ) quantized_kv_cache = QuantizedKVParameterCache(dtype=dtype, granularity=granularity) - setattr(module, "kv_cache", quantized_kv_cache) + module.kv_cache = quantized_kv_cache logger.debug(f"Initialized quantized kv_cache for {module.__class__.__name__} {getattr(module, 'layer_idx', None)}") if quantized_kv_cache.is_nvfp4: # Global scales are only registered once calibration has observed KV @@ -266,14 +266,14 @@ def initialize_quantized_kv_cache(module: torch.nn.Module, dtype=torch.float8_e4 def calibrate_kv_cache_input_hook( - module: torch.nn.Module, args: Any, kwargs: Dict[str, Any] -) -> Tuple[Tuple[Any, ...], Dict[str, Any]]: + module: torch.nn.Module, args: Any, kwargs: dict[str, Any] +) -> tuple[tuple[Any, ...], dict[str, Any]]: """ Hook to update inputs to attention layers when running kv_cache quantization. Will update the passed in kv_cache to singleton QuantizedKVParameterCache. """ - kv_cache = getattr(module, "kv_cache") + kv_cache = module.kv_cache # Start from transformers 4.55.2, the `past_key_value` was renamed to `past_key_values`. # https://github.com/huggingface/transformers/blob/52c6c1bb6e27ca87c4faede34a4c2a7404c17c4d/src/transformers/models/llama/modeling_llama.py#L279-L280 if "past_key_values" in kwargs: @@ -316,7 +316,7 @@ def calibrate_kv_cache_output_hook(module: torch.nn.Module, _args: Any, _output: """ Hook to update k_scale and v_scale parameters when running kv_cache quantization. """ - kv_cache = getattr(module, "kv_cache") + kv_cache = module.kv_cache if kv_cache.is_nvfp4: layer_idx = module.layer_idx k_amax = kv_cache.k_amax[layer_idx] if layer_idx < len(kv_cache.k_amax) else 0.0 diff --git a/auto_round/experimental/qmodules/base.py b/auto_round/experimental/qmodules/base.py index 8b7a9c1383..be148110d1 100644 --- a/auto_round/experimental/qmodules/base.py +++ b/auto_round/experimental/qmodules/base.py @@ -13,7 +13,6 @@ # limitations under the License. from abc import ABC, abstractmethod -from typing import Optional, Union import torch diff --git a/auto_round/experimental/qmodules/fake.py b/auto_round/experimental/qmodules/fake.py index e7419e91d3..322ab9aeb0 100644 --- a/auto_round/experimental/qmodules/fake.py +++ b/auto_round/experimental/qmodules/fake.py @@ -12,7 +12,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -from typing import Optional import torch @@ -31,8 +30,8 @@ def __init__( in_features: int, out_features: int, config: QuantizationScheme, - weight: Optional[torch.Tensor] = None, - bias: Optional[torch.Tensor] = None, + weight: torch.Tensor | None = None, + bias: torch.Tensor | None = None, dtype: torch.dtype = torch.bfloat16, ): super().__init__() diff --git a/auto_round/experimental/qmodules/fp4_utils.py b/auto_round/experimental/qmodules/fp4_utils.py index f2dafa1c8b..a8c1f7a4ea 100644 --- a/auto_round/experimental/qmodules/fp4_utils.py +++ b/auto_round/experimental/qmodules/fp4_utils.py @@ -27,7 +27,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -from typing import Optional import torch @@ -58,9 +57,7 @@ def pack_fp4_to_uint8(x: torch.Tensor) -> torch.Tensor: return packed.reshape(rows, columns // 2) -def unpack_fp4_from_uint8( - a: torch.Tensor, m: int, n: int, dtype: Optional[torch.dtype] = torch.bfloat16 -) -> torch.Tensor: +def unpack_fp4_from_uint8(a: torch.Tensor, m: int, n: int, dtype: torch.dtype | None = torch.bfloat16) -> torch.Tensor: """ Unpacks uint8 values into FP4. Each uint8 contains two FP4 values (low nibble first). The 4-bit indices are mapped to FP4 values using kE2M1ToFloat. @@ -73,22 +70,20 @@ def unpack_fp4_from_uint8( @torch.compiler.disable() def _unpack_fp4_from_uint8_cpu( - a: torch.Tensor, m: int, n: int, dtype: Optional[torch.dtype] = torch.bfloat16 + a: torch.Tensor, m: int, n: int, dtype: torch.dtype | None = torch.bfloat16 ) -> torch.Tensor: return _unpack_fp4_from_uint8(a, m, n, dtype) # @torch.compile(fullgraph=True, dynamic=True) def _unpack_fp4_from_uint8_cuda( - a: torch.Tensor, m: int, n: int, dtype: Optional[torch.dtype] = torch.bfloat16 + a: torch.Tensor, m: int, n: int, dtype: torch.dtype | None = torch.bfloat16 ) -> torch.Tensor: return _unpack_fp4_from_uint8(a, m, n, dtype) # reference: : https://github.com/vllm-project/vllm/pull/16362 -def _unpack_fp4_from_uint8( - a: torch.Tensor, m: int, n: int, dtype: Optional[torch.dtype] = torch.bfloat16 -) -> torch.Tensor: +def _unpack_fp4_from_uint8(a: torch.Tensor, m: int, n: int, dtype: torch.dtype | None = torch.bfloat16) -> torch.Tensor: """ Unpacks uint8 values into fp4. Each uint8 consists of two fp4 values (i.e. first four bits correspond to one fp4 value, last four correspond to a diff --git a/auto_round/experimental/qmodules/fp8_static.py b/auto_round/experimental/qmodules/fp8_static.py index ca013fae2a..03fc1be07c 100644 --- a/auto_round/experimental/qmodules/fp8_static.py +++ b/auto_round/experimental/qmodules/fp8_static.py @@ -13,8 +13,6 @@ # limitations under the License. -from typing import Optional, Union - import torch from auto_round.experimental.qmodules.base import QModuleBase @@ -39,10 +37,10 @@ def __init__( self, in_features, out_features, - weight: Optional[torch.Tensor] = None, - weight_scale: Optional[torch.Tensor] = None, - bias: Union[torch.Tensor, bool, None] = None, - input_scale: Optional[torch.Tensor] = None, + weight: torch.Tensor | None = None, + weight_scale: torch.Tensor | None = None, + bias: torch.Tensor | bool | None = None, + input_scale: torch.Tensor | None = None, dtype=torch.bfloat16, ): super().__init__() diff --git a/auto_round/experimental/qmodules/mx.py b/auto_round/experimental/qmodules/mx.py index dea099760f..287b919372 100644 --- a/auto_round/experimental/qmodules/mx.py +++ b/auto_round/experimental/qmodules/mx.py @@ -13,8 +13,6 @@ # limitations under the License. -from typing import Optional, Union - import torch from auto_round.data_type.utils import get_quant_func @@ -57,9 +55,9 @@ def __init__( in_features, out_features, config: QuantizationScheme, - weight: Optional[torch.Tensor] = None, - weight_scale: Optional[torch.Tensor] = None, - bias: Union[torch.Tensor, bool, None] = None, + weight: torch.Tensor | None = None, + weight_scale: torch.Tensor | None = None, + bias: torch.Tensor | bool | None = None, dtype=torch.bfloat16, ): super().__init__() @@ -109,7 +107,7 @@ def __init__( ), ) - def initialize_weights(self, weight: Optional[torch.Tensor]) -> torch.Tensor: + def initialize_weights(self, weight: torch.Tensor | None) -> torch.Tensor: """ Initialize weights. This method should be overridden by subclasses. """ @@ -169,7 +167,7 @@ def forward(self, input: torch.Tensor) -> torch.Tensor: return out @classmethod - def from_original(cls, config: Optional[QuantizationScheme], original_layer: torch.nn.Linear): + def from_original(cls, config: QuantizationScheme | None, original_layer: torch.nn.Linear): """ Create an `MXQuantLinear` layer from an original linear layer. """ @@ -193,7 +191,7 @@ def __init__(self, *args, **kwargs): self.weight_name = "weight_packed" super().__init__(*args, **kwargs) - def initialize_weights(self, weight: Optional[torch.Tensor]) -> torch.Tensor: + def initialize_weights(self, weight: torch.Tensor | None) -> torch.Tensor: weight_dtype = torch.uint8 weight_in_features = self.in_features // 2 return torch.zeros((self.out_features, weight_in_features), dtype=weight_dtype) if weight is None else weight @@ -219,7 +217,7 @@ def __init__(self, *args, **kwargs): self.weight_name = "weight_packed" super().__init__(*args, **kwargs) - def initialize_weights(self, weight: Optional[torch.Tensor]) -> torch.Tensor: + def initialize_weights(self, weight: torch.Tensor | None) -> torch.Tensor: weight_dtype = torch.uint8 weight_in_features = self.in_features // 2 return torch.zeros((self.out_features, weight_in_features), dtype=weight_dtype) if weight is None else weight @@ -236,7 +234,7 @@ def unpack_data(self, packed_data: torch.Tensor) -> torch.Tensor: return unpacked_data @classmethod - def from_original(cls, config: Optional[QuantizationScheme], original_layer: torch.nn.Linear): + def from_original(cls, config: QuantizationScheme | None, original_layer: torch.nn.Linear): """ Create an `MXQuantLinear` layer from an original linear layer. """ @@ -260,7 +258,7 @@ def __init__(self, *args, **kwargs): self.weight_name = "weight" super().__init__(*args, **kwargs) - def initialize_weights(self, weight: Optional[torch.Tensor]) -> torch.Tensor: + def initialize_weights(self, weight: torch.Tensor | None) -> torch.Tensor: weight_dtype = torch.float8_e4m3fn weight_in_features = self.in_features return torch.zeros((self.out_features, weight_in_features), dtype=weight_dtype) if weight is None else weight diff --git a/auto_round/experimental/qmodules/mxint4_utils.py b/auto_round/experimental/qmodules/mxint4_utils.py index b772f54c1a..c8290307c5 100644 --- a/auto_round/experimental/qmodules/mxint4_utils.py +++ b/auto_round/experimental/qmodules/mxint4_utils.py @@ -27,7 +27,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -from typing import Optional import torch @@ -45,9 +44,7 @@ def get_e0m4_tensor(device): return _DEVICE_E0M4_TENSORS[device_str] -def unpack_int4_from_uint8( - a: torch.Tensor, m: int, n: int, dtype: Optional[torch.dtype] = torch.bfloat16 -) -> torch.Tensor: +def unpack_int4_from_uint8(a: torch.Tensor, m: int, n: int, dtype: torch.dtype | None = torch.bfloat16) -> torch.Tensor: """ Unpacks uint8 values into int4. Each uint8 contains two int4 values (low nibble first). The 4-bit indices are mapped to int4 values using kE0M4ToFloat. @@ -60,20 +57,20 @@ def unpack_int4_from_uint8( @torch.compiler.disable() def _unpack_int4_from_uint8_cpu( - a: torch.Tensor, m: int, n: int, dtype: Optional[torch.dtype] = torch.bfloat16 + a: torch.Tensor, m: int, n: int, dtype: torch.dtype | None = torch.bfloat16 ) -> torch.Tensor: return _unpack_int4_from_uint8(a, m, n, dtype) # @torch.compile(fullgraph=True, dynamic=True) def _unpack_int4_from_uint8_cuda( - a: torch.Tensor, m: int, n: int, dtype: Optional[torch.dtype] = torch.bfloat16 + a: torch.Tensor, m: int, n: int, dtype: torch.dtype | None = torch.bfloat16 ) -> torch.Tensor: return _unpack_int4_from_uint8(a, m, n, dtype) def _unpack_int4_from_uint8( - a: torch.Tensor, m: int, n: int, dtype: Optional[torch.dtype] = torch.bfloat16 + a: torch.Tensor, m: int, n: int, dtype: torch.dtype | None = torch.bfloat16 ) -> torch.Tensor: """ Unpacks uint8 values into int4. Each uint8 consists of two int4 values diff --git a/auto_round/experimental/qmodules/nvfp4.py b/auto_round/experimental/qmodules/nvfp4.py index dafeb25679..e9dc81e789 100644 --- a/auto_round/experimental/qmodules/nvfp4.py +++ b/auto_round/experimental/qmodules/nvfp4.py @@ -13,8 +13,6 @@ # limitations under the License. -from typing import Optional, Union - import torch from auto_round.data_type.nvfp import get_reciprocal, ref_nvfp4_quant @@ -59,9 +57,9 @@ def __init__( in_features: int, out_features: int, config: QuantizationScheme, - weight: Optional[torch.Tensor] = None, - weight_scale: Optional[torch.Tensor] = None, - bias: Union[torch.Tensor, bool, None] = None, + weight: torch.Tensor | None = None, + weight_scale: torch.Tensor | None = None, + bias: torch.Tensor | bool | None = None, dtype=torch.bfloat16, ): super().__init__() @@ -145,7 +143,7 @@ def load_state_dict(self, state_dict, strict=True, assign=False): self._convert_global_scale_to_float32(state_dict, "input_global_scale") return super().load_state_dict(state_dict, strict, assign) - def initialize_weights(self, weight: Optional[torch.Tensor]) -> torch.Tensor: + def initialize_weights(self, weight: torch.Tensor | None) -> torch.Tensor: """ Initialize weights. """ @@ -199,7 +197,7 @@ def forward(self, input: torch.Tensor) -> torch.Tensor: return out @classmethod - def from_original(cls, config: Optional[QuantizationScheme], original_layer: torch.nn.Linear): + def from_original(cls, config: QuantizationScheme | None, original_layer: torch.nn.Linear): """ Create an `NVFPQuantLinear` layer from an original linear layer. """ diff --git a/auto_round/experimental/qmodules/nvfp4_e5m3.py b/auto_round/experimental/qmodules/nvfp4_e5m3.py index c8b115e7e2..42f2b55283 100644 --- a/auto_round/experimental/qmodules/nvfp4_e5m3.py +++ b/auto_round/experimental/qmodules/nvfp4_e5m3.py @@ -13,7 +13,6 @@ # limitations under the License. import os -from typing import Optional, Union import torch @@ -33,7 +32,7 @@ _CACHE_WEIGHT_ENV = "AR_NVFP4_E5M3_CACHE_HP_WEIGHT" -def _resolve_cache_weight(cache_weight: Optional[bool], default: bool) -> bool: +def _resolve_cache_weight(cache_weight: bool | None, default: bool) -> bool: if cache_weight is not None: return cache_weight value = os.getenv(_CACHE_WEIGHT_ENV) @@ -53,11 +52,11 @@ def __init__( in_features: int, out_features: int, config: QuantizationScheme, - weight: Optional[torch.Tensor] = None, - weight_scale: Optional[torch.Tensor] = None, - bias: Union[torch.Tensor, bool, None] = None, + weight: torch.Tensor | None = None, + weight_scale: torch.Tensor | None = None, + bias: torch.Tensor | bool | None = None, dtype=torch.bfloat16, - cache_weight: Optional[bool] = None, + cache_weight: bool | None = None, ): super().__init__() assert dtype in self.SUPPORTED_COMPUTE_DTYPE diff --git a/auto_round/experimental/utils.py b/auto_round/experimental/utils.py index 2391d1798f..1f443bd533 100644 --- a/auto_round/experimental/utils.py +++ b/auto_round/experimental/utils.py @@ -107,7 +107,7 @@ def update_parameter_data(module: torch.nn.Module, new_val: torch.Tensor, name: module.register_parameter(name, torch.nn.Parameter(new_val)) else: logger.warning_once( - "Parameter %s not found in module %s, creating new parameter." % (name, module.__class__.__name__) + f"Parameter {name} not found in module {module.__class__.__name__}, creating new parameter." ) module.register_parameter(name, torch.nn.Parameter(new_val)) diff --git a/auto_round/export/export_to_autogptq/export.py b/auto_round/export/export_to_autogptq/export.py index 90be0c0d90..2199e2d09f 100644 --- a/auto_round/export/export_to_autogptq/export.py +++ b/auto_round/export/export_to_autogptq/export.py @@ -16,8 +16,9 @@ import inspect import json import os +from collections.abc import Callable from dataclasses import fields -from typing import Any, Callable, Dict, Union +from typing import Any # MIT License # @@ -83,7 +84,7 @@ from auto_round.export.export_to_autoround.utils import check_neq_config -def convert_to_autogptq_dynamic(regex_config: Dict[str, Dict[str, Any]]) -> Dict[str, Dict[str, Any]]: +def convert_to_autogptq_dynamic(regex_config: dict[str, dict[str, Any]]) -> dict[str, dict[str, Any]]: """ Convert AutoRound-style regex_config into AutoGPTQ-style QuantizerConfig.dynamic. @@ -101,7 +102,7 @@ def convert_to_autogptq_dynamic(regex_config: Dict[str, Dict[str, Any]]) -> Dict elif bits < 16: converted[f"+:{regex}"] = {"bits": bits} for key in GPTQ_REQUIRED_CONFIG_KEYS: # only save keys gptq supported - converted[f"+:{regex}"][key] = regex_config[name][key] + converted[f"+:{regex}"][key] = cfg[key] else: # skip quantization converted[f"-:{regex}"] = {} @@ -190,12 +191,12 @@ def pack_layer(name, model, backend, device=None): def save_quantized_as_autogptq( output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", + device: str | torch.device = "cpu", backend: str = "auto_gptq:exllamav2", - serialization_dict: dict = None, + serialization_dict: dict | None = None, **kwargs, ) -> torch.nn.Module: """Export the model to autogptq format to easily leverage cuda kernel.""" diff --git a/auto_round/export/export_to_autoround/export.py b/auto_round/export/export_to_autoround/export.py index fd02d98d9f..8dbdbff16c 100644 --- a/auto_round/export/export_to_autoround/export.py +++ b/auto_round/export/export_to_autoround/export.py @@ -17,9 +17,9 @@ import inspect import json import os +from collections.abc import Callable from dataclasses import fields from enum import Enum -from typing import Callable, Union import torch import torch.nn as nn @@ -242,12 +242,12 @@ def pack_layer(layer_name, model, backend, device=None): def save_quantized_as_autoround( output_dir: str, model: torch.nn.Module, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace=True, backend="auto_round:exllamav2", - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, **kwargs, ): """ @@ -279,7 +279,7 @@ def save_quantized_as_autoround( ): backend = backend.replace("auto_round", "auto_round:auto_gptq") - safe_serialization = True if "safe_serialization" not in kwargs.keys() else kwargs["safe_serialization"] + safe_serialization = kwargs.get("safe_serialization", True) if not inplace: model = copy.deepcopy(model.to("cpu")) diff --git a/auto_round/export/export_to_autoround/export_to_fp8.py b/auto_round/export/export_to_autoround/export_to_fp8.py index c0465dfcd3..861b8d18f2 100644 --- a/auto_round/export/export_to_autoround/export_to_fp8.py +++ b/auto_round/export/export_to_autoround/export_to_fp8.py @@ -15,8 +15,8 @@ import copy import json import os +from collections.abc import Callable from dataclasses import fields -from typing import Callable, Union import torch import transformers @@ -202,16 +202,16 @@ def pack_layer(layer_name, model, data_type, device=None, unsqueeze=False): def save_quantized_as_autoround( output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - backend: str = None, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + backend: str | None = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, quant_method: str = "auto-round", **kwargs, ): - safe_serialization = True if "safe_serialization" not in kwargs.keys() else kwargs["safe_serialization"] + safe_serialization = kwargs.get("safe_serialization", True) if not inplace: model = copy.deepcopy(model.to("cpu")) quantization_config = serialization_dict diff --git a/auto_round/export/export_to_autoround/export_to_nvfp_mx.py b/auto_round/export/export_to_autoround/export_to_nvfp_mx.py index e9be7d6346..ef87ffab34 100644 --- a/auto_round/export/export_to_autoround/export_to_nvfp_mx.py +++ b/auto_round/export/export_to_autoround/export_to_nvfp_mx.py @@ -16,8 +16,8 @@ import inspect import json import os +from collections.abc import Callable from dataclasses import fields -from typing import Callable, Union import torch import torch.nn as nn @@ -84,7 +84,7 @@ def pack_layer(name, model, backend, device=None): from auto_round.data_type.nvfp import calculate_gparam input_global_scale = calculate_gparam(layer.act_max, layer.group_size, "cpu") - setattr(layer, "input_global_scale", input_global_scale) + layer.input_global_scale = input_global_scale delattr(layer, "act_max") if type(layer) == nn.Linear: @@ -137,12 +137,12 @@ def pack_layer(name, model, backend, device=None): def save_quantized_as_fp( output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", + device: str | torch.device = "cpu", backend: str = "autoround:exllamav2", - serialization_dict: dict = None, + serialization_dict: dict | None = None, **kwargs, ) -> torch.nn.Module: """ @@ -170,7 +170,7 @@ def save_quantized_as_fp( data_type = serialization_dict.get("data_type", None) act_bits = serialization_dict.get("act_bits", None) act_data_type = serialization_dict.get("act_data_type", None) - safe_serialization = True if "safe_serialization" not in kwargs.keys() else kwargs["safe_serialization"] + safe_serialization = kwargs.get("safe_serialization", True) if not inplace: model = copy.deepcopy(model.to("cpu")) quantization_config = serialization_dict @@ -205,7 +205,7 @@ def save_quantized_as_fp( from auto_round.data_type.nvfp import calculate_gparam input_global_scale = calculate_gparam(layer.act_max, layer.group_size, model.device) - setattr(layer, "input_global_scale", input_global_scale) + layer.input_global_scale = input_global_scale delattr(layer, "act_max") # update fused input_global_scale from auto_round.data_type.utils import update_fused_layer_global_scales diff --git a/auto_round/export/export_to_autoround/qlinear_fp.py b/auto_round/export/export_to_autoround/qlinear_fp.py index 039c5e0b3d..d9082e1e46 100644 --- a/auto_round/export/export_to_autoround/qlinear_fp.py +++ b/auto_round/export/export_to_autoround/qlinear_fp.py @@ -190,7 +190,6 @@ def pack(self, linear, scales, zeros=None, g_idx=None, global_scale=None, input_ if input_global_scale is not None: # TODO: the shape of `input_global_scale` is [] in some cases — need to investigate why. self.input_global_scale = input_global_scale.to(torch.float32).to(device).reshape([1]) - return def pack_fp4_to_uint8(scaled_tensor: torch.Tensor): diff --git a/auto_round/export/export_to_autoround/qlinear_triton_act.py b/auto_round/export/export_to_autoround/qlinear_triton_act.py index 7d5f9dee77..57deb924dc 100644 --- a/auto_round/export/export_to_autoround/qlinear_triton_act.py +++ b/auto_round/export/export_to_autoround/qlinear_triton_act.py @@ -171,7 +171,7 @@ def pack(self, linear, scales, zeros, act_scales, w_bf16_to_fp8_scale, g_idx=Non zeros -= 1 shape = scales_t.shape value = 0 - for j in range(0, (32 // self.bits)): + for j in range(32 // self.bits): value |= zeros << (self.bits * j) qzeros = np.ones((shape[0], shape[1] // 32 * self.bits), dtype=np.uint32) * value qzeros = qzeros.astype(np.int32) diff --git a/auto_round/export/export_to_autoround/utils.py b/auto_round/export/export_to_autoround/utils.py index 5e9e081d24..526ecdc366 100644 --- a/auto_round/export/export_to_autoround/utils.py +++ b/auto_round/export/export_to_autoround/utils.py @@ -13,12 +13,11 @@ # limitations under the License. from dataclasses import fields -from typing import List from auto_round.schemes import QuantizationScheme -def check_neq_config(config: dict, **expected) -> List[str]: +def check_neq_config(config: dict, **expected) -> list[str]: """ Compare a config dict against expected values. Ensures all required keys are present in both config and expected. diff --git a/auto_round/export/export_to_awq/export.py b/auto_round/export/export_to_awq/export.py index fc91b3817c..306ed4c879 100644 --- a/auto_round/export/export_to_awq/export.py +++ b/auto_round/export/export_to_awq/export.py @@ -23,7 +23,7 @@ import copy import json import os -from typing import Callable, Union +from collections.abc import Callable import torch import torch.nn as nn @@ -61,7 +61,7 @@ def _collect_modules_to_not_convert( model: torch.nn.Module, layer_config: dict, regex_config: dict, - to_quant_block_names: list = None, + to_quant_block_names: list | None = None, ) -> list: """Collect all module names that should not be converted (not quantized). @@ -148,11 +148,11 @@ def pack_layer(name, model, backend, device=None): def save_quantized_as_autoawq( output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, **kwargs, ) -> torch.nn.Module: """Export the model to autogptq format to easily leverage cuda kernel.""" diff --git a/auto_round/export/export_to_awq/utils.py b/auto_round/export/export_to_awq/utils.py index 6629875298..ddf59ac0d1 100644 --- a/auto_round/export/export_to_awq/utils.py +++ b/auto_round/export/export_to_awq/utils.py @@ -356,10 +356,4 @@ def forward(self, x): return out.reshape(out_shape) def extra_repr(self) -> str: - return "in_features={}, out_features={}, bias={}, w_bit={}, group_size={}".format( - self.in_features, - self.out_features, - self.bias is not None, - self.w_bit, - self.group_size, - ) + return f"in_features={self.in_features}, out_features={self.out_features}, bias={self.bias is not None}, w_bit={self.w_bit}, group_size={self.group_size}" diff --git a/auto_round/export/export_to_gguf/convert.py b/auto_round/export/export_to_gguf/convert.py index c922be9a58..60bdbf3a3d 100644 --- a/auto_round/export/export_to_gguf/convert.py +++ b/auto_round/export/export_to_gguf/convert.py @@ -37,9 +37,10 @@ import json import os import re +from collections.abc import Iterator from functools import partial from itertools import chain -from typing import TYPE_CHECKING, Any, Iterator +from typing import TYPE_CHECKING import numpy as np import torch @@ -101,7 +102,7 @@ def wrapper_model_instance( def _need_low_cpu_mem(low_cpu_mem_usage): - if not low_cpu_mem_usage: + if not low_cpu_mem_usage: # noqa: SIM103 return False # process = psutil.Process(os.getpid()) @@ -130,7 +131,7 @@ def get_moe_name(cls, name, new_name): tensor_type = cls.tensor_map.get_type(new_name_tmp).name experts_name = name_tmp.split(".")[-1] - for k, v in type_mapping.items(): + for v in type_mapping.values(): if experts_name in v: idx = v.index(experts_name) name = name.replace(experts_name, type_mapping[tensor_type][idx]) @@ -534,11 +535,7 @@ def quantize_expert(arr, source_name): def get_qtype_by_layer_config(layer_config, name, data_qtype, *, explicit_only=False): name = name[: -len(".weight")] - if name not in layer_config and name.endswith("embed_tokens"): - embedding_names = [key for key in layer_config if key.endswith("embed_tokens")] - if len(embedding_names) == 1: - name = embedding_names[0] - elif name == "token_embd": + if (name not in layer_config and name.endswith("embed_tokens")) or name == "token_embd": embedding_names = [key for key in layer_config if key.endswith("embed_tokens")] if len(embedding_names) == 1: name = embedding_names[0] @@ -856,7 +853,7 @@ def prepare_tensors(cls): modify_name = _special_name_handle(cls, checkpoint_name) restored_outputs_completed = False - for new_name, data_torch in cls.modify_tensors(data_torch, modify_name, bid): + for new_name, data_torch in cls.modify_tensors(data_torch, modify_name, bid): # noqa: B020 restored_outputs_completed = True if _gguf_writer_has_tensor(cls.gguf_writer, new_name): logger.debug("%s already added to gguf_writer, skip", new_name) @@ -1012,10 +1009,9 @@ def prepare_tensors(cls): gguf.GGMLQuantizationType.Q2_K, gguf.GGMLQuantizationType.Q3_K, gguf.GGMLQuantizationType.Q4_K, + gguf.GGMLQuantizationType.Q5_K, ]: data_qtype = gguf.GGMLQuantizationType.Q5_0 - elif data_qtype == gguf.GGMLQuantizationType.Q5_K: - data_qtype = gguf.GGMLQuantizationType.Q5_0 elif data_qtype == gguf.GGMLQuantizationType.Q6_K: data_qtype = gguf.GGMLQuantizationType.Q8_0 diff --git a/auto_round/export/export_to_gguf/hf_checkpoint_restorer.py b/auto_round/export/export_to_gguf/hf_checkpoint_restorer.py index 6121ef887f..4b1922c36c 100644 --- a/auto_round/export/export_to_gguf/hf_checkpoint_restorer.py +++ b/auto_round/export/export_to_gguf/hf_checkpoint_restorer.py @@ -15,9 +15,9 @@ from __future__ import annotations from collections import defaultdict +from collections.abc import Callable, Iterator from copy import deepcopy from dataclasses import dataclass -from typing import Callable, Iterator import torch diff --git a/auto_round/export/export_to_gguf/llama_cpp_conversion.py b/auto_round/export/export_to_gguf/llama_cpp_conversion.py index e057dbbe71..99b3050621 100644 --- a/auto_round/export/export_to_gguf/llama_cpp_conversion.py +++ b/auto_round/export/export_to_gguf/llama_cpp_conversion.py @@ -48,15 +48,15 @@ class ConversionContext: @property def ModelBase(self): - return getattr(self.module, "ModelBase") + return self.module.ModelBase @property def ModelType(self): - return getattr(self.module, "ModelType") + return self.module.ModelType @property def get_model_architecture(self): - return getattr(self.module, "get_model_architecture") + return self.module.get_model_architecture def model_type(self, model_type: AutoRoundModelType | Any): if int(model_type) == int(AutoRoundModelType.MMPROJ): diff --git a/auto_round/export/export_to_gguf/special_handle.py b/auto_round/export/export_to_gguf/special_handle.py index 513db51bbe..1b0d7cb184 100644 --- a/auto_round/export/export_to_gguf/special_handle.py +++ b/auto_round/export/export_to_gguf/special_handle.py @@ -13,8 +13,8 @@ # limitations under the License. import json import os +from collections.abc import Iterable from functools import partial -from typing import Iterable import torch from safetensors import safe_open @@ -55,8 +55,9 @@ def granite_moe_modify_tensors(cls, original_modify_tensors, data_torch, name, b ( projection for projection in _GRANITE_AGGREGATED_EXPERT_TENSORS - if name.endswith(f"block_sparse_moe.experts.{projection}") - or name.endswith(f"block_sparse_moe.experts.{projection}.weight") + if name.endswith( + (f"block_sparse_moe.experts.{projection}", f"block_sparse_moe.experts.{projection}.weight") + ) ), None, ) @@ -142,7 +143,7 @@ def repack(name, data_torch, blocks0, blocks1): return blocks0, blocks1 for name, data_torch in cls.get_tensors(): - if GPTOSS_RELOAD and (name.endswith("mlp.experts.down_proj") or name.endswith("mlp.experts.gate_up_proj")): + if GPTOSS_RELOAD and name.endswith(("mlp.experts.down_proj", "mlp.experts.gate_up_proj")): block_name = name + "_blocks" block_data_torch = get_tensor_from_file(cls.model.name_or_path, block_name) blocks0, blocks1 = repack(block_name, block_data_torch, blocks0, blocks1) diff --git a/auto_round/export/export_to_gguf/sync_llama_cpp_conversion.py b/auto_round/export/export_to_gguf/sync_llama_cpp_conversion.py old mode 100644 new mode 100755 diff --git a/auto_round/export/export_to_llmcompressor/config.py b/auto_round/export/export_to_llmcompressor/config.py index 9c66c5ffa9..99eb1196df 100644 --- a/auto_round/export/export_to_llmcompressor/config.py +++ b/auto_round/export/export_to_llmcompressor/config.py @@ -53,17 +53,16 @@ def check_compressed_tensors_supported(raise_error: bool = False): # please refer to https://github.com/vllm-project/llm-compressor/blob/ # 29f4d5644b48e9c8ebb7e36d5be9f7c92747ceb7/src/llmcompressor/modifiers/quantization/quantization/mixin.py#L168 -def initialize_quantization(scheme, targets=["Linear"], config_groups=None, kv_cache_scheme=None, ignore=["lm_head"]): +def initialize_quantization(scheme, targets=None, config_groups=None, kv_cache_scheme=None, ignore=None): """ Attach quantization schemes to modules in the model and initialize the quantization config """ # apply scheme and status to model - scheme = scheme - targets = targets - config_groups = config_groups - kv_cache_scheme = kv_cache_scheme - ignore = ignore + if targets is None: + targets = ["Linear"] + if ignore is None: + ignore = ["lm_head"] using_mxfp4_for_mxfp8 = False check_compressed_tensors_supported(raise_error=True) if scheme is not None and config_groups is not None: diff --git a/auto_round/export/export_to_llmcompressor/export.py b/auto_round/export/export_to_llmcompressor/export.py index eb0d623982..12148253fc 100644 --- a/auto_round/export/export_to_llmcompressor/export.py +++ b/auto_round/export/export_to_llmcompressor/export.py @@ -15,8 +15,7 @@ import copy import math import os -from collections.abc import Mapping -from typing import Callable, Union +from collections.abc import Callable, Mapping import torch @@ -234,8 +233,8 @@ def pack_layer(name, model, device=None): weight_device = layer.weight.device scheme = construct_ct_scheme(layer) - setattr(layer, "quantization_scheme", scheme) - setattr(layer, "weight_scale", torch.nn.Parameter(layer.scale.to(weight_device))) + layer.quantization_scheme = scheme + layer.weight_scale = torch.nn.Parameter(layer.scale.to(weight_device)) # AutoRound zp is the UNSIGNED level convention (q in [0, 2^b-1], # W = (q - zp) * s); compressed-tensors packs the SIGNED convention (q in # [-2^(b-1), 2^(b-1)-1], zp added BEFORE the clamp into that range - see @@ -304,7 +303,7 @@ def pack_layer(name, model, device=None): ) zp = zp.to(torch.int8) - setattr(layer, "weight_zero_point", torch.nn.Parameter(zp.to(weight_device), requires_grad=False)) + layer.weight_zero_point = torch.nn.Parameter(zp.to(weight_device), requires_grad=False) delattr(layer, "scale") _compress_and_set_format(layer, scheme, device) @@ -314,11 +313,11 @@ def pack_layer(name, model, device=None): def save_quantized_as_llmcompressor( output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, **kwargs, ) -> torch.nn.Module: """ diff --git a/auto_round/export/export_to_llmcompressor/export_to_fp.py b/auto_round/export/export_to_llmcompressor/export_to_fp.py index 9e24a38e10..344b8f212a 100644 --- a/auto_round/export/export_to_llmcompressor/export_to_fp.py +++ b/auto_round/export/export_to_llmcompressor/export_to_fp.py @@ -17,7 +17,7 @@ import json import os import sys -from typing import Callable, Union +from collections.abc import Callable import torch import torch.nn as nn @@ -94,7 +94,7 @@ def pack_layer(name, model, device=None): from auto_round.data_type.nvfp import calculate_gparam input_global_scale = calculate_gparam(layer.act_max, layer.group_size) # , model.device - setattr(layer, "input_global_scale", input_global_scale) + layer.input_global_scale = input_global_scale delattr(layer, "act_max") # QuantLinear = get_fp_qlinear(backend, bits, group_size, sym) @@ -303,12 +303,12 @@ def _resolve_kv_cache_scheme( def save_quantized_as_fp( output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - backend: str = None, - serialization_dict: dict = None, + device: str | torch.device = "cpu", + backend: str | None = None, + serialization_dict: dict | None = None, **kwargs, ) -> torch.nn.Module: """ @@ -336,7 +336,7 @@ def save_quantized_as_fp( data_type = serialization_dict.get("data_type", None) act_bits = serialization_dict.get("act_bits", None) act_data_type = serialization_dict.get("act_data_type", None) - safe_serialization = True if "safe_serialization" not in kwargs.keys() else kwargs["safe_serialization"] + safe_serialization = kwargs.get("safe_serialization", True) if not inplace: model = copy.deepcopy(model.to("cpu")) processor = kwargs.get("processor", None) @@ -360,7 +360,7 @@ def save_quantized_as_fp( from auto_round.data_type.nvfp import calculate_gparam input_global_scale = calculate_gparam(layer.act_max, layer.group_size, model.device) - setattr(layer, "input_global_scale", input_global_scale) + layer.input_global_scale = input_global_scale delattr(layer, "act_max") # update fused input_global_scale from auto_round.data_type.utils import update_fused_layer_global_scales @@ -460,7 +460,7 @@ def save_quantized_as_fp( attention_config = _get_attention_config(model, static_attention_granularity) else: attention_config = None - setattr(quantization_config, "format", format) + quantization_config.format = format quantization_config = quantization_config.to_dict() quantization_config["provider"] = "auto-round" _configure_gaudi2_fp8_dtype(quantization_config) diff --git a/auto_round/export/export_to_llmcompressor/export_to_static_fp.py b/auto_round/export/export_to_llmcompressor/export_to_static_fp.py index 4741112f59..364d533fef 100644 --- a/auto_round/export/export_to_llmcompressor/export_to_static_fp.py +++ b/auto_round/export/export_to_llmcompressor/export_to_static_fp.py @@ -16,7 +16,7 @@ import json import os import sys -from typing import Callable, Union +from collections.abc import Callable import torch import transformers @@ -41,7 +41,7 @@ ) -def pack_layer(layer_name: str, model: torch.nn.Module, data_type: str, device: str = None) -> None: +def pack_layer(layer_name: str, model: torch.nn.Module, data_type: str, device: str | None = None) -> None: """ Packs a model layer for quantization based on its type and configuration. @@ -179,21 +179,19 @@ def _configure_gaudi2_fp8_dtype(quantization_config: dict) -> None: if is_gaudi2(): quantization_config["fp8_dtype_flavor"] = _GAUDI2_FP8_DTYPE_FLAVOR logger.warning_once( - ( - "Running on Intel Gaudi2 hardware." - f" Setting FP8 dtype flavor to {_GAUDI2_FP8_DTYPE_FLAVOR} for compatibility." - ) + "Running on Intel Gaudi2 hardware." + f" Setting FP8 dtype flavor to {_GAUDI2_FP8_DTYPE_FLAVOR} for compatibility." ) def save_quantized_as_static_fp( output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, **kwargs, ) -> torch.nn.Module: """ @@ -217,7 +215,7 @@ def save_quantized_as_static_fp( Raises: ValueError: If the backend is not supported. """ - safe_serialization = True if "safe_serialization" not in kwargs.keys() else kwargs["safe_serialization"] + safe_serialization = kwargs.get("safe_serialization", True) if not inplace: model = copy.deepcopy(model.to("cpu")) diff --git a/auto_round/export/export_to_llmcompressor/utils.py b/auto_round/export/export_to_llmcompressor/utils.py index 304a720bdf..a6020c8f0a 100644 --- a/auto_round/export/export_to_llmcompressor/utils.py +++ b/auto_round/export/export_to_llmcompressor/utils.py @@ -12,12 +12,11 @@ # See the License for the specific language governing permissions and # limitations under the License. -from typing import Dict, List from auto_round.utils import matches_any_regex, to_standard_regex -def generate_ignore_regex_list(regex_config: Dict[str, Dict], layer_config: Dict[str, Dict]) -> List[str]: +def generate_ignore_regex_list(regex_config: dict[str, dict], layer_config: dict[str, dict]) -> list[str]: """ Generate ignore regex list for llm_compressor based on regex_config and layer_config. @@ -34,7 +33,7 @@ def generate_ignore_regex_list(regex_config: Dict[str, Dict], layer_config: Dict List[str]: List of regex patterns to ignore during quantization. """ prefix = "re:" - ignore_regex: List[str] = [] + ignore_regex: list[str] = [] # Step 1: Add regex_config keys with bits >= 16 for key, cfg in regex_config.items(): diff --git a/auto_round/export/export_to_mlx/export.py b/auto_round/export/export_to_mlx/export.py index 2681ca8185..119b00bb8d 100644 --- a/auto_round/export/export_to_mlx/export.py +++ b/auto_round/export/export_to_mlx/export.py @@ -30,7 +30,7 @@ import copy import json import os -from typing import Callable, Union +from collections.abc import Callable import torch import torch.nn as nn @@ -600,11 +600,11 @@ def pack_layer(name, model, device=None, **kwargs): def save_quantized_as_mlx( output_dir: str, model: nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, **kwargs, ) -> nn.Module: """Save quantized model in MLX-compatible format. diff --git a/auto_round/export/formats/__init__.py b/auto_round/export/formats/__init__.py index d401af8c86..48038c8539 100644 --- a/auto_round/export/formats/__init__.py +++ b/auto_round/export/formats/__init__.py @@ -29,12 +29,12 @@ from auto_round.export.formats.resolver import resolve_formats __all__ = [ - "BackendDataType", - "AutoRoundFormat", "AutoAWQFormat", "AutoGPTQFormat", - "FakeFormat", + "AutoRoundFormat", + "BackendDataType", "FP8Format", + "FakeFormat", "GGUFFormat", "LLMCompressorFormat", "MLXFormat", diff --git a/auto_round/export/formats/backends/__init__.py b/auto_round/export/formats/backends/__init__.py index 3f77b50cbc..e74dbf3e19 100644 --- a/auto_round/export/formats/backends/__init__.py +++ b/auto_round/export/formats/backends/__init__.py @@ -26,8 +26,8 @@ "AutoAWQFormat", "AutoGPTQFormat", "AutoRoundFormat", - "FakeFormat", "FP8Format", + "FakeFormat", "GGUFFormat", "LLMCompressorFormat", "MLXFormat", diff --git a/auto_round/export/formats/backends/auto_awq.py b/auto_round/export/formats/backends/auto_awq.py index 00905cb5b1..f7f8c43889 100644 --- a/auto_round/export/formats/backends/auto_awq.py +++ b/auto_round/export/formats/backends/auto_awq.py @@ -13,7 +13,7 @@ # limitations under the License. import re -from typing import Callable, Union +from collections.abc import Callable import torch import transformers @@ -113,11 +113,11 @@ def save_quantized( self, output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, **kwargs, ) -> torch.nn.Module: backend = self.get_backend_name() diff --git a/auto_round/export/formats/backends/auto_gptq.py b/auto_round/export/formats/backends/auto_gptq.py index 03a14a33d3..805b8a157a 100644 --- a/auto_round/export/formats/backends/auto_gptq.py +++ b/auto_round/export/formats/backends/auto_gptq.py @@ -13,7 +13,7 @@ # limitations under the License. import re -from typing import Callable, Union +from collections.abc import Callable import torch @@ -91,11 +91,11 @@ def save_quantized( self, output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, **kwargs, ) -> torch.nn.Module: backend = self.get_backend_name() diff --git a/auto_round/export/formats/backends/autoround.py b/auto_round/export/formats/backends/autoround.py index 84134f6dfc..867aff2e55 100644 --- a/auto_round/export/formats/backends/autoround.py +++ b/auto_round/export/formats/backends/autoround.py @@ -13,7 +13,8 @@ # limitations under the License. import re -from typing import Any, Callable, Union +from collections.abc import Callable +from typing import Any import torch @@ -71,9 +72,12 @@ def __init__(self, format: str, scheme: QuantizationScheme, ctx: Any): ) if enable_awq: self.backend = AutoAWQFormat("auto_round:auto_awq", scheme, ctx) - elif scheme.is_nv_fp() or scheme.is_mx_fp() or scheme.data_type == BackendDataType.NVFP4_E5M3.value: - self.backend = AutoRoundFormat(scheme.data_type, scheme, ctx) - elif scheme.is_mx_int() and scheme.bits == 4: # only add mx_int4 now + elif ( + scheme.is_nv_fp() + or scheme.is_mx_fp() + or scheme.data_type == BackendDataType.NVFP4_E5M3.value + or (scheme.is_mx_int() and scheme.bits == 4) # only add mx_int4 now + ): self.backend = AutoRoundFormat(scheme.data_type, scheme, ctx) elif scheme.is_act_static(): # static wfp8afp8 self.backend = AutoRoundFormat(BackendDataType.FP8_STATIC.value, scheme, ctx) @@ -146,13 +150,10 @@ def pack_layer(self, layer_name, model, device=None, **kwargs): f"auto_round:{BackendDataType.MX_FP.value}", f"auto_round:{BackendDataType.MX_FP_RCEIL.value}", f"auto_round:{BackendDataType.NV_FP4_WITH_STATIC_GS.value}", + f"auto_round:{BackendDataType.MX_INT.value}", ]: from auto_round.export.export_to_autoround.export_to_nvfp_mx import pack_layer - pack_func = pack_layer - elif self.output_format in [f"auto_round:{BackendDataType.MX_INT.value}"]: - from auto_round.export.export_to_autoround.export_to_nvfp_mx import pack_layer - pack_func = pack_layer elif self.output_format in [ f"auto_round:{BackendDataType.FP8.value}", @@ -172,11 +173,11 @@ def save_quantized( self, output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, **kwargs, ) -> torch.nn.Module: if self.backend is not None: diff --git a/auto_round/export/formats/backends/fake.py b/auto_round/export/formats/backends/fake.py index 2c87053dd1..b5e34caaa8 100644 --- a/auto_round/export/formats/backends/fake.py +++ b/auto_round/export/formats/backends/fake.py @@ -16,7 +16,8 @@ import glob import json import os -from typing import Any, Callable, Union +from collections.abc import Callable +from typing import Any import torch @@ -127,11 +128,11 @@ def save_quantized( self, output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, **kwargs, ): has_fake_act_quant = False diff --git a/auto_round/export/formats/backends/fp8.py b/auto_round/export/formats/backends/fp8.py index 431e458d0b..390956c281 100644 --- a/auto_round/export/formats/backends/fp8.py +++ b/auto_round/export/formats/backends/fp8.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -from typing import Callable, Union +from collections.abc import Callable import torch @@ -58,11 +58,11 @@ def save_quantized( self, output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, **kwargs, ) -> torch.nn.Module: from auto_round.export.export_to_autoround.export_to_fp8 import save_quantized_as_autoround diff --git a/auto_round/export/formats/backends/gguf.py b/auto_round/export/formats/backends/gguf.py index cccb867196..8a215930c1 100644 --- a/auto_round/export/formats/backends/gguf.py +++ b/auto_round/export/formats/backends/gguf.py @@ -15,8 +15,9 @@ import copy import os import re +from collections.abc import Callable from dataclasses import fields -from typing import Any, Callable, Union +from typing import Any import torch import transformers @@ -176,11 +177,11 @@ def save_quantized( self, output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, **kwargs, ) -> torch.nn.Module: from auto_round.export.export_to_gguf.export import save_quantized_as_gguf @@ -203,7 +204,7 @@ def gguf_args_check( scheme: QuantizationScheme, model, platform: str, - formats: Union[str, list[str]] = None, + formats: str | list[str] | None = None, model_type=ModelType.TEXT, ) -> QuantizationScheme: import argparse @@ -285,9 +286,9 @@ def immediate_pack( name: str, model: torch.nn.Module, device: torch.device, - output_dir: str = None, + output_dir: str | None = None, mllm: bool = False, - layer_config: dict = None, + layer_config: dict | None = None, tokenizer=None, processor=None, image_processor=None, @@ -376,9 +377,7 @@ def _search_gguf_type(gguf_type): def gguf_type_fallback(gguf_type: str) -> str: gguf_type = gguf_type.lower() - if gguf_type in ("gguf:q2_k", "gguf:q3_k", "gguf:q4_k"): - gguf_type = "gguf:q5_0" - elif gguf_type == "gguf:q5_k": + if gguf_type in ("gguf:q2_k", "gguf:q3_k", "gguf:q4_k", "gguf:q5_k"): gguf_type = "gguf:q5_0" elif gguf_type == "gguf:q6_k": gguf_type = "gguf:q8_0" @@ -473,9 +472,7 @@ def _get_digital_in_layer_name(layer_name): def _gguf_type_fallback(gguf_type: str) -> str: gguf_type = gguf_type.lower() - if gguf_type in ("gguf:q2_k", "gguf:q3_k", "gguf:q4_k"): - gguf_type = "gguf:q5_0" - elif gguf_type == "gguf:q5_k": + if gguf_type in ("gguf:q2_k", "gguf:q3_k", "gguf:q4_k", "gguf:q5_k"): gguf_type = "gguf:q5_0" elif gguf_type == "gguf:q6_k": gguf_type = "gguf:q8_0" @@ -749,7 +746,7 @@ def _set_config(config, target_config): base_target_bits = int(inner_gguf_format[6]) def _resolve_gguf_name(layer_name): - if model_type != ModelType.TEXT and any([key in layer_name for key in MM_MODULE_KEYS]): + if model_type != ModelType.TEXT and any(key in layer_name for key in MM_MODULE_KEYS): gguf_layer_name = tensor_map_vision.get_name(layer_name) if gguf_layer_name is None: for key in MM_MODULE_KEYS: diff --git a/auto_round/export/formats/backends/llm_compressor.py b/auto_round/export/formats/backends/llm_compressor.py index 02493a0b97..749e41078c 100644 --- a/auto_round/export/formats/backends/llm_compressor.py +++ b/auto_round/export/formats/backends/llm_compressor.py @@ -13,7 +13,8 @@ # limitations under the License. import re -from typing import Any, Callable, Union +from collections.abc import Callable +from typing import Any import torch @@ -157,9 +158,11 @@ def check_and_reset_format( if scheme.act_bits <= 8 and (not scheme.is_act_standard_fp() or scheme.act_dynamic): if scheme.act_data_type == "nvfp4_v2": return None, scheme, layer_config, quant_block_list - if (scheme.is_act_nv_fp() and "static_gs" in scheme.act_data_type) or scheme.is_act_mx_fp(): - return None, scheme, layer_config, quant_block_list - elif scheme.is_dynamic_afp8() and scheme.is_block_wfp8(): + if ( + (scheme.is_act_nv_fp() and "static_gs" in scheme.act_data_type) + or scheme.is_act_mx_fp() + or (scheme.is_dynamic_afp8() and scheme.is_block_wfp8()) + ): return None, scheme, layer_config, quant_block_list else: bits, group_size, sym, act_bits = 8, -1, True, 8 @@ -189,15 +192,11 @@ def pack_layer(self, layer_name, model, device=None, **kwargs): from auto_round.export.export_to_llmcompressor.export_to_static_fp import pack_layer return pack_layer(layer_name, model, self.get_backend_name(), device=device) - elif re.search(f"{BackendDataType.INT8.value}", self.output_format): - from auto_round.export.export_to_llmcompressor.export import pack_layer - - return pack_layer(layer_name, model, device=device) - elif re.search(f"{BackendDataType.FP8_BLOCK.value}", self.output_format): - from auto_round.export.export_to_llmcompressor.export import pack_layer - - return pack_layer(layer_name, model, device=device) - elif re.search(f"{BackendDataType.WINT_A16.value}", self.output_format): + elif ( + re.search(f"{BackendDataType.INT8.value}", self.output_format) + or re.search(f"{BackendDataType.FP8_BLOCK.value}", self.output_format) + or re.search(f"{BackendDataType.WINT_A16.value}", self.output_format) + ): from auto_round.export.export_to_llmcompressor.export import pack_layer return pack_layer(layer_name, model, device=device) @@ -209,11 +208,11 @@ def save_quantized( self, output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, **kwargs, ) -> torch.nn.Module: backend = self.get_backend_name() diff --git a/auto_round/export/formats/backends/mlx.py b/auto_round/export/formats/backends/mlx.py index 37717a2970..28ff2e9098 100644 --- a/auto_round/export/formats/backends/mlx.py +++ b/auto_round/export/formats/backends/mlx.py @@ -13,7 +13,7 @@ # limitations under the License. import re -from typing import Callable, Union +from collections.abc import Callable import torch @@ -55,11 +55,11 @@ def save_quantized( self, output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, **kwargs, ) -> torch.nn.Module: from auto_round.export.export_to_mlx.export import save_quantized_as_mlx diff --git a/auto_round/export/formats/backends/svdquant_nunchaku.py b/auto_round/export/formats/backends/svdquant_nunchaku.py index d8c36c424a..a5b3275d8a 100644 --- a/auto_round/export/formats/backends/svdquant_nunchaku.py +++ b/auto_round/export/formats/backends/svdquant_nunchaku.py @@ -15,7 +15,8 @@ from __future__ import annotations import os -from typing import Any, Callable, Union +from collections.abc import Callable +from typing import Any import torch @@ -132,11 +133,11 @@ def save_quantized( self, output_dir: str, model: torch.nn.Module = None, - tokenizer: Callable = None, - layer_config: dict = None, + tokenizer: Callable | None = None, + layer_config: dict | None = None, inplace: bool = True, - device: Union[str, torch.device] = "cpu", - serialization_dict: dict = None, + device: str | torch.device = "cpu", + serialization_dict: dict | None = None, *, config=None, residual_provider=None, diff --git a/auto_round/export/formats/base.py b/auto_round/export/formats/base.py index 618e3e0fa9..320e8293cc 100644 --- a/auto_round/export/formats/base.py +++ b/auto_round/export/formats/base.py @@ -18,9 +18,10 @@ import os import re from abc import ABC, abstractmethod +from collections.abc import Callable from dataclasses import asdict from enum import Enum -from typing import Any, Callable, Optional, Union +from typing import Any import torch import transformers @@ -162,7 +163,7 @@ def is_supported_immediate_saving(self) -> bool: return self.backend.is_supported_immediate_saving() if self.backend is not None else True @classmethod - def is_support_scheme(cls: OutputFormat, scheme: Union[str, QuantizationScheme]) -> bool: + def is_support_scheme(cls: OutputFormat, scheme: str | QuantizationScheme) -> bool: if isinstance(scheme, str) and scheme.upper() in cls.support_schemes: return True if isinstance(scheme, QuantizationScheme): @@ -177,7 +178,7 @@ def check_and_reset_format( self, scheme: QuantizationScheme, ctx: Any, - ) -> tuple[Optional[str], QuantizationScheme, dict, list]: + ) -> tuple[str | None, QuantizationScheme, dict, list]: layer_config, quant_block_list = ctx.layer_config, ctx.quant_block_list if self.backend is not None: new_format, scheme, layer_config, quant_block_list = self.backend.check_and_reset_format(scheme, ctx) diff --git a/auto_round/export/formats/resolver.py b/auto_round/export/formats/resolver.py index 73c6739c13..46c94a73f2 100644 --- a/auto_round/export/formats/resolver.py +++ b/auto_round/export/formats/resolver.py @@ -14,8 +14,8 @@ from __future__ import annotations +from collections.abc import Iterable from types import SimpleNamespace -from typing import Iterable import torch @@ -94,15 +94,15 @@ def _precise_gguf_name(formats: list[OutputFormat]) -> str | None: def resolve_formats( scheme: ResolvedScheme, *, - format: str = None, - layer_config: dict = None, + format: str | None = None, + layer_config: dict | None = None, scale_dtype=None, quant_block_list=None, mllm: bool = False, iters: int = 0, enable_alg_ext: bool = False, quant_nontext_module: bool = False, - platform: str = None, + platform: str | None = None, is_auto_scheme: bool = False, model=None, ) -> FormatResolution: diff --git a/auto_round/export/svdquant_adapters/__init__.py b/auto_round/export/svdquant_adapters/__init__.py index 1404640b92..00782535db 100644 --- a/auto_round/export/svdquant_adapters/__init__.py +++ b/auto_round/export/svdquant_adapters/__init__.py @@ -89,11 +89,11 @@ def resolve_svdquant_model_adapter( __all__ = [ - "FLUX_TOP_LEVEL_TENSOR_KEYS", "FLUX_SVDQUANT_TARGET_MODULES", + "FLUX_TOP_LEVEL_TENSOR_KEYS", "SDXL_SVDQUANT_TARGET_MODULES", - "SDXLSVDQuantNunchakuAdapter", "FluxSVDQuantNunchakuAdapter", + "SDXLSVDQuantNunchakuAdapter", "detect_svdquant_model_adapter", "flux_onefile_tensor_count", "resolve_svdquant_model_adapter", diff --git a/auto_round/export/svdquant_nunchaku.py b/auto_round/export/svdquant_nunchaku.py index 7dfa313766..c2d695ab49 100644 --- a/auto_round/export/svdquant_nunchaku.py +++ b/auto_round/export/svdquant_nunchaku.py @@ -16,8 +16,9 @@ import json import os +from collections.abc import Iterable, Mapping from dataclasses import dataclass, field -from typing import Iterable, Mapping, Protocol +from typing import Protocol import torch diff --git a/auto_round/export/utils.py b/auto_round/export/utils.py index 121f89e729..f76d308091 100644 --- a/auto_round/export/utils.py +++ b/auto_round/export/utils.py @@ -92,7 +92,7 @@ def _state_dict_has_meta_tensor(model: nn.Module) -> bool: return False -def is_immediate_saving_mode(model: nn.Module, serialization_dict: dict = None) -> bool: +def is_immediate_saving_mode(model: nn.Module, serialization_dict: dict | None = None) -> bool: """Determine if the model was saved via ShardWriter (immediate saving mode). Resolution order: @@ -106,9 +106,7 @@ def is_immediate_saving_mode(model: nn.Module, serialization_dict: dict = None) return True if unsupported_meta_device(model): return True - if _state_dict_has_meta_tensor(model): - return True - return False + return _state_dict_has_meta_tensor(model) def is_local_pipeline_model_dir(model_dir: str) -> bool: @@ -507,8 +505,8 @@ def filter_quantization_config(quantization_config): default_dict["lr"] = 1.0 / iters if iters > 0 else 5e-3 default_dict["minmax_lr"] = default_dict["lr"] - for key in default_dict: - if key in quantization_config and default_dict[key] == quantization_config[key]: + for key, default_value in default_dict.items(): + if key in quantization_config and default_value == quantization_config[key]: quantization_config.pop(key) for k in list(quantization_config.keys()): if quantization_config[k] is None: @@ -536,9 +534,7 @@ def filter_quantization_config(quantization_config): if callable(key): quantization_config.pop(key) elif isinstance(quantization_config[key], (list, tuple)): - if any([callable(item) for item in quantization_config[key]]): - quantization_config.pop(key) - elif len(quantization_config[key]) == 0: + if any(callable(item) for item in quantization_config[key]) or len(quantization_config[key]) == 0: quantization_config.pop(key) if key in clean_list and key in quantization_config: quantization_config.pop(key) diff --git a/auto_round/formats.py b/auto_round/formats.py index e28ed025c5..c4866ade45 100644 --- a/auto_round/formats.py +++ b/auto_round/formats.py @@ -72,16 +72,16 @@ def get_formats(format: str, ar): __all__ = [ "AutoAWQFormat", - "AutoRoundExportFormat", "AutoGPTQFormat", + "AutoRoundExportFormat", "AutoRoundFormat", "BackendDataType", - "FakeFormat", "FP8Format", + "FakeFormat", "GGUFFormat", "LLMCompressorFormat", "MLXFormat", "OutputFormat", - "resolve_formats", "get_formats", + "resolve_formats", ] diff --git a/auto_round/inference/backend.py b/auto_round/inference/backend.py index d45549a75f..b252360074 100644 --- a/auto_round/inference/backend.py +++ b/auto_round/inference/backend.py @@ -15,7 +15,7 @@ import platform from dataclasses import dataclass, field from importlib import import_module -from typing import TYPE_CHECKING, Any, Optional +from typing import TYPE_CHECKING, Any import torch from packaging.version import Version @@ -29,6 +29,8 @@ BackendInfos = {} +import sys + import cpuinfo if TYPE_CHECKING: @@ -91,18 +93,18 @@ class BackendInfo: packing_format: list[str] bits: list[int] compute_dtype: list[str] = None - data_type: Optional[list[str]] = None - group_size: Optional[list[int]] = None - act_bits: Optional[list[int]] = None - act_group_size: Optional[list[int]] = None - act_sym: Optional[list[bool]] = None - act_data_type: Optional[list[str]] = None - act_dynamic: Optional[list[bool]] = None + data_type: list[str] | None = None + group_size: list[int] | None = None + act_bits: list[int] | None = None + act_group_size: list[int] | None = None + act_sym: list[bool] | None = None + act_data_type: list[str] | None = None + act_dynamic: list[bool] | None = None priority: int = 0 ##higher is better checkers: list[Any] = field(default_factory=list) - alias: Optional[list[str]] = None - requirements: Optional[list[str]] = None - systems: Optional[list[str]] = None + alias: list[str] | None = None + requirements: list[str] | None = None + systems: list[str] | None = None BACKEND_ACT_ATTRS = [ @@ -211,8 +213,8 @@ def fp8_static_scheme_checker( in_feature: int, out_feature: int, config: QuantizationScheme, - in_feature_multiplier: Optional[int] = None, - out_feature_multiplier: Optional[int] = None, + in_feature_multiplier: int | None = None, + out_feature_multiplier: int | None = None, ): from auto_round.schemes import FP8_STATIC @@ -1154,7 +1156,7 @@ def get_autogptq_infer_linear(backend, bits=4, group_size=128, sym=False): return QuantLinear -def find_backend(backend: str, orig_backend: str = None): +def find_backend(backend: str, orig_backend: str | None = None): """ Finds the matching backend key based on the target backend name or its aliases. @@ -1196,7 +1198,7 @@ def get_all_compatible_backend( # Find compatible backends compatible_backends = [ key - for key in BackendInfos.keys() + for key in BackendInfos if check_compatible(key, device, config, packing_format, in_features, out_features, check_requirements=False) ] @@ -1240,8 +1242,8 @@ def get_layer_backend( if backend == "auto": backends = BackendInfos.keys() else: - for key in BackendInfos.keys(): - if backend == key or (BackendInfos[key].alias and backend in BackendInfos[key].alias): + for key, info in BackendInfos.items(): + if backend == key or (info.alias and backend in info.alias): backends.append(key) # Find and store other compatible backends @@ -1281,8 +1283,7 @@ def get_highest_priority_backend( ) -> str | None: current_system = platform.system().lower() supported_backends = [] - for key in BackendInfos.keys(): - backend = BackendInfos[key] + for key, backend in BackendInfos.items(): # Filter by operating system (e.g. MLX is Darwin-only; ark CPU # backends are non-Darwin only). if backend.systems is not None: @@ -1384,10 +1385,10 @@ def build_pip_commands(gptq_req, other_reqs): for msg in install_instructions: log(msg) if logger_level == "error" and len(pip_cmds) == 0: - exit(-1) + sys.exit(-1) joined_cmds = " and ".join(f"`{cmd}`" for cmd in pip_cmds) if joined_cmds: log(joined_cmds) if logger_level == "error": - exit(-1) + sys.exit(-1) diff --git a/auto_round/inference/convert_model.py b/auto_round/inference/convert_model.py index 7ebf1718b0..eec7967c97 100644 --- a/auto_round/inference/convert_model.py +++ b/auto_round/inference/convert_model.py @@ -11,10 +11,10 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. +import itertools import os import re from types import SimpleNamespace -from typing import Union import torch import torch.nn as nn @@ -70,7 +70,7 @@ def skip_not_convert_modules(model, quantization_config, layer_names, layer_conf modules_to_not_convert = _get_modules_to_not_convert(model, modules_to_not_convert) if modules_to_not_convert: for layer_name in layer_names: - if any([re.search(re.compile(n), layer_name) for n in modules_to_not_convert]): + if any(re.search(re.compile(n), layer_name) for n in modules_to_not_convert): layer_configs[layer_name] = {"bits": 16, "act_bits": 16} return layer_configs @@ -101,9 +101,9 @@ def get_keys_to_not_convert(model): tied_params = find_tied_parameters(tied_model) # For compatibility with Accelerate < 0.18 if isinstance(tied_params, dict): - tied_keys = sum(list(tied_params.values()), []) + list(tied_params.keys()) + tied_keys = list(itertools.chain.from_iterable(tied_params.values())) + list(tied_params.keys()) else: - tied_keys = sum(tied_params, []) + tied_keys = list(itertools.chain.from_iterable(tied_params)) has_tied_params = len(tied_keys) > 0 # If there is not tied weights, we want to keep the lm_head(output_embedding) in full precision @@ -369,7 +369,7 @@ def get_layer_config(model, quantization_config): quantization_config.modules_in_block_to_quantize ) # Flatten the list for layer_name in layer_names: - if not any([re.search(re.compile(n), layer_name) is not None for n in modules_in_block_to_quantize]): + if not any(re.search(re.compile(n), layer_name) is not None for n in modules_in_block_to_quantize): extra_config[layer_name] = {"bits": 16} # Default to 16-bit for unquantized layers # Expand GPTQ 'dynamic' config (regex-based) @@ -408,7 +408,7 @@ def get_layer_config(model, quantization_config): return layer_configs -def get_device(obj: Union[torch.Tensor, nn.Module]) -> torch.device: +def get_device(obj: torch.Tensor | nn.Module) -> torch.device: if isinstance(obj, torch.Tensor): return obj.device return next(obj.parameters()).device @@ -601,7 +601,7 @@ def _create_quant_layer(layer, layer_backend, config, in_features, out_features, ) -def infer_target_device(device_map: Union[dict, int, str, None] = None) -> str: +def infer_target_device(device_map: dict | int | str | None = None) -> str: """Infers the target device from a device_map. Args: diff --git a/auto_round/logger.py b/auto_round/logger.py index e439d2cadc..c67e885281 100644 --- a/auto_round/logger.py +++ b/auto_round/logger.py @@ -14,15 +14,16 @@ import logging import warnings -from functools import lru_cache, wraps -from typing import TYPE_CHECKING, Any, Callable, Dict, List, Mapping, Optional, TypeVar +from collections.abc import Callable, Mapping +from functools import cache, wraps +from typing import TypeVar import auto_round.envs as envs T = TypeVar("T", bound="Callable") # used by `deprecated` -@lru_cache(maxsize=None) +@cache def warning_once(self, msg, *args): """ Log a warning message only once per unique message/arguments combination. @@ -34,7 +35,7 @@ def warning_once(self, msg, *args): logger.warning(msg, *args, stacklevel=2) -@lru_cache(maxsize=None) +@cache def info_once(self, msg, *args): """ Log an info message only once per unique message/arguments combination. @@ -115,7 +116,7 @@ def format(self, record): logger.addHandler(fh) -def deprecated(future_name: Optional[str] = None, message: Optional[str] = None) -> Callable[[T], T]: +def deprecated(future_name: str | None = None, message: str | None = None) -> Callable[[T], T]: """ Decorator to mark functions as deprecated diff --git a/auto_round/modeling/finegrained_fp8_patch.py b/auto_round/modeling/finegrained_fp8_patch.py index 06d85835bc..8b0cfb978d 100644 --- a/auto_round/modeling/finegrained_fp8_patch.py +++ b/auto_round/modeling/finegrained_fp8_patch.py @@ -133,7 +133,7 @@ def __init__(self, hf_quantizer): def convert(self, input_dict: torch.Tensor, **kwargs) -> dict[str, torch.Tensor]: # Unpack single key/value (value may be wrapped in a list) - target_keys, value = tuple(input_dict.items())[0] + target_keys, value = next(iter(input_dict.items())) value = value[0] # Resolve block size (support dict-like or attr-like quant_config) @@ -152,10 +152,8 @@ def convert(self, input_dict: torch.Tensor, **kwargs) -> dict[str, torch.Tensor] # Enforce exact tiling like your original if rows % block_m != 0 or cols % block_n != 0: raise ValueError( - ( - f"Matrix dimensions ({rows}, {cols}) must be divisible by block sizes" - f" ({block_m}, {block_n}). for {target_keys}" - ) + f"Matrix dimensions ({rows}, {cols}) must be divisible by block sizes" + f" ({block_m}, {block_n}). for {target_keys}" ) # Leading dims can be empty (2D) or include num_experts/... (3D+) diff --git a/auto_round/modeling/finegrained_fp8_patch_v4.py b/auto_round/modeling/finegrained_fp8_patch_v4.py index 11c2f6882f..7b80edce1d 100644 --- a/auto_round/modeling/finegrained_fp8_patch_v4.py +++ b/auto_round/modeling/finegrained_fp8_patch_v4.py @@ -12,7 +12,6 @@ # See the License for the specific language governing permissions and # limitations under the License. # Copied from https://github.com/huggingface/transformers/blob/v4.57.3/src/transformers/integrations/finegrained_fp8.py -from typing import Optional from transformers.utils import is_accelerate_available, is_torch_available, logging @@ -48,7 +47,7 @@ def __init__( out_features: int, bias: bool = False, dtype=None, - block_size: Optional[tuple[int, int]] = None, + block_size: tuple[int, int] | None = None, device=None, activation_scheme="dynamic", ): diff --git a/auto_round/modeling/fp8_quant.py b/auto_round/modeling/fp8_quant.py index 3f4a250579..75dbc78417 100644 --- a/auto_round/modeling/fp8_quant.py +++ b/auto_round/modeling/fp8_quant.py @@ -173,10 +173,8 @@ def apply_fp8_expert_replacement_patch(): auto_round_logger.debug("Applied FP8 expert replacement patch to transformers.") OriginalFineGrainedFP8HfQuantizer.validate_environment = oot_validate_environment auto_round_logger.debug( - ( - "Patched FineGrainedFP8HfQuantizer.validate_environment to bypass device " - "capability check for loading FP8 models on unsupported GPUs." - ) + "Patched FineGrainedFP8HfQuantizer.validate_environment to bypass device " + "capability check for loading FP8 models on unsupported GPUs." ) except ImportError as e: auto_round_logger.warning(f"Could not apply FP8 expert replacement patch as {e}.") diff --git a/auto_round/modeling/fused_moe/__init__.py b/auto_round/modeling/fused_moe/__init__.py index d29e61a91e..6e508e9101 100644 --- a/auto_round/modeling/fused_moe/__init__.py +++ b/auto_round/modeling/fused_moe/__init__.py @@ -31,7 +31,7 @@ resolve_experts_implementation, ) -__all__ = [ +__all__ = [ # noqa: RUF022 "ReplacementModuleBase", "apply_replacements", "materialize_model_", diff --git a/auto_round/modeling/fused_moe/deepseek_v2.py b/auto_round/modeling/fused_moe/deepseek_v2.py index c7d386763b..c65e7b66df 100644 --- a/auto_round/modeling/fused_moe/deepseek_v2.py +++ b/auto_round/modeling/fused_moe/deepseek_v2.py @@ -14,8 +14,8 @@ import inspect import warnings +from collections.abc import Callable from functools import partial -from typing import Callable, Optional import torch from transformers.modeling_rope_utils import dynamic_rope_update @@ -67,12 +67,12 @@ def rotary_emb_forward(module, x, position_ids): def attn_forward( module, hidden_states: torch.Tensor, - attention_mask: Optional[torch.Tensor] = None, - cache_position: Optional[torch.LongTensor] = None, - position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None, - position_ids: Optional[torch.Tensor] = None, + attention_mask: torch.Tensor | None = None, + cache_position: torch.LongTensor | None = None, + position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None, + position_ids: torch.Tensor | None = None, **kwargs, -) -> tuple[torch.Tensor, Optional[torch.Tensor], Optional[tuple[torch.Tensor]]]: +) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: if "padding_mask" in kwargs: warnings.warn( "Passing `padding_mask` is deprecated and will be removed in v4.37. " diff --git a/auto_round/modeling/fused_moe/grouped_experts.py b/auto_round/modeling/fused_moe/grouped_experts.py index c9a755316f..d2858ac099 100644 --- a/auto_round/modeling/fused_moe/grouped_experts.py +++ b/auto_round/modeling/fused_moe/grouped_experts.py @@ -221,9 +221,7 @@ def _projection_is_supported(layer: nn.Module) -> bool: # forward hook on each expert Linear, and the grouped path -- which multiplies the # weights directly and never calls ``Linear.forward`` -- would silently skip them, # leaving every expert without ``act_max`` (breaking static-act export, e.g. NVFP4). - if layer._forward_pre_hooks or layer._forward_hooks: - return False - return True + return not (layer._forward_pre_hooks or layer._forward_hooks) if not _is_wrapper_linear(layer): return False orig = layer.orig_layer @@ -231,9 +229,7 @@ def _projection_is_supported(layer: nn.Module) -> bool: return False # Conv1D / LinearAllreduce need their own forward if orig._forward_pre_hooks or orig._forward_hooks: return False # e.g. online Hadamard rotation must run per layer - if getattr(layer, "enable_act_quant", False) and not _act_quant_is_row_independent(layer): - return False - return True + return not (getattr(layer, "enable_act_quant", False) and not _act_quant_is_row_independent(layer)) def _compute_device(layer: nn.Module) -> torch.device: diff --git a/auto_round/modeling/fused_moe/moe_experts_interface.py b/auto_round/modeling/fused_moe/moe_experts_interface.py index cf0954b1b6..dc9596854c 100644 --- a/auto_round/modeling/fused_moe/moe_experts_interface.py +++ b/auto_round/modeling/fused_moe/moe_experts_interface.py @@ -122,8 +122,6 @@ class _ExpertContainer(nn.Module): which matches the standard checkpoint format without any hooks. """ - pass - def _install_compact_expert_repr(module: nn.Module) -> None: """Install compact __repr__ on the module's class. diff --git a/auto_round/modeling/fused_moe/qwen3_vl_moe.py b/auto_round/modeling/fused_moe/qwen3_vl_moe.py index eb76cfc341..2df48b7210 100644 --- a/auto_round/modeling/fused_moe/qwen3_vl_moe.py +++ b/auto_round/modeling/fused_moe/qwen3_vl_moe.py @@ -50,7 +50,7 @@ def __init__( calibrate_all_experts: bool = False, ): super().__init__(original) - text_config: "Qwen3VLMoeTextConfig" = config.get_text_config() + text_config: Qwen3VLMoeTextConfig = config.get_text_config() self.hidden_size = text_config.hidden_size self.num_experts = text_config.num_experts diff --git a/auto_round/modeling/fused_moe/replace_modules.py b/auto_round/modeling/fused_moe/replace_modules.py index f40867bb83..5048e8ab83 100644 --- a/auto_round/modeling/fused_moe/replace_modules.py +++ b/auto_round/modeling/fused_moe/replace_modules.py @@ -15,7 +15,6 @@ from abc import ABC, abstractmethod from contextlib import contextmanager from dataclasses import dataclass -from typing import Dict, Type import torch from tqdm import tqdm @@ -124,7 +123,7 @@ def is_custom_model(model: torch.nn.Module) -> bool: def _find_first_moe_block(model: torch.nn.Module) -> tuple[str, torch.nn.Module] | tuple[None, None]: """Return ``(name, module)`` of the first experts-like module, or ``(None, None)``.""" for name, module in model.named_modules(): - if name.endswith(".experts") or name.endswith(".moe"): + if name.endswith((".experts", ".moe")): return name, module return None, None @@ -226,7 +225,7 @@ class ReplacementModuleBase(ABC, torch.nn.Module): """ # Registry: module class name -> replacement module class - _replacement_registry: Dict[str, Type["ReplacementModuleBase"]] = {} + _replacement_registry: dict[str, type["ReplacementModuleBase"]] = {} supports_gguf_fused_moe: bool = False def __init_subclass__(cls, **kwargs): @@ -263,7 +262,7 @@ def __init__(self, original: torch.nn.Module): self._materialized = False @classmethod - def get_replacement_class(cls, module_class_name: str) -> Type["ReplacementModuleBase"]: + def get_replacement_class(cls, module_class_name: str) -> type["ReplacementModuleBase"]: """Get replacement class for a given module class name.""" return cls._replacement_registry.get(module_class_name) @@ -292,7 +291,6 @@ def get_registered_modules(cls) -> list: @abstractmethod def original_module_class(cls) -> str: """Return the class name of the module this replaces.""" - pass @classmethod @abstractmethod @@ -302,7 +300,6 @@ def from_original( config, ) -> "ReplacementModuleBase": """Create replacement module from original module.""" - pass def materialize_weights(self): """Materialize weights if needed.""" @@ -316,7 +313,6 @@ def _materialize_weights(self) -> None: Subclasses should override this method to implement weight materialization logic. """ - pass def release_original_module(self) -> None: """Release reference to the original module to free memory.""" @@ -467,7 +463,7 @@ class ModuleReplacementTracker: def __new__(cls): if cls._instance is None: - cls._instance = super(ModuleReplacementTracker, cls).__new__(cls) + cls._instance = super().__new__(cls) return cls._instance def __init__(self): @@ -476,9 +472,9 @@ def __init__(self): return # Map from replacement module id to original module - self._replacement_to_original: Dict[int, torch.nn.Module] = {} + self._replacement_to_original: dict[int, torch.nn.Module] = {} # Map from module name to ReplacedModuleInfo - self._name_to_info: Dict[str, ReplacedModuleInfo] = {} + self._name_to_info: dict[str, ReplacedModuleInfo] = {} ModuleReplacementTracker._initialized = True diff --git a/auto_round/modeling/fused_moe/step3_5_moe.py b/auto_round/modeling/fused_moe/step3_5_moe.py old mode 100755 new mode 100644 diff --git a/auto_round/modeling/unfused_moe/__init__.py b/auto_round/modeling/unfused_moe/__init__.py index 15012aea9c..5f0b0747b6 100644 --- a/auto_round/modeling/unfused_moe/__init__.py +++ b/auto_round/modeling/unfused_moe/__init__.py @@ -147,7 +147,7 @@ def pre_check_config(model_name: str | torch.nn.Module, trust_remote_code: bool if isinstance(model_name, str): try: config = AutoConfig.from_pretrained(model_name, trust_remote_code=trust_remote_code) - except (OSError, EnvironmentError, ValueError): + except (OSError, ValueError): return False elif isinstance(model_name, torch.nn.Module): config = getattr(model_name, "config", None) @@ -190,7 +190,7 @@ def apply_model_monkey_patches(model_name: str, trust_remote_code: bool = True) return False # patch blocks config = AutoConfig.from_pretrained(model_name, trust_remote_code=trust_remote_code) - model_type = getattr(config, "model_type") + model_type = config.model_type cfg = MODEL_CONFIG[model_type] for orig_path, custom_path in cfg.get("block_patch", []): @@ -226,7 +226,7 @@ def apply_modeling_patch(model: torch.nn.Module) -> bool: res = pre_check_config(model) if not res: return False - model_type = getattr(model.config, "model_type") + model_type = model.config.model_type cfg = MODEL_CONFIG[model_type] # patch blocks for orig_path, custom_path in cfg.get("block_patch", []): diff --git a/auto_round/modeling/unfused_moe/glm_moe_light.py b/auto_round/modeling/unfused_moe/glm_moe_light.py index c16b52b568..b031ad7dd9 100644 --- a/auto_round/modeling/unfused_moe/glm_moe_light.py +++ b/auto_round/modeling/unfused_moe/glm_moe_light.py @@ -50,7 +50,6 @@ def experts_forward( top_k_index: torch.Tensor, top_k_weights: torch.Tensor, ) -> torch.Tensor: - """ """ return sequential_moe_forward(hidden_states, top_k_index, top_k_weights, self.experts, self.num_experts) def route_tokens_to_experts(self, router_logits): diff --git a/auto_round/modeling/unfused_moe/qwen3_moe.py b/auto_round/modeling/unfused_moe/qwen3_moe.py index 2c39210140..d2c9161397 100644 --- a/auto_round/modeling/unfused_moe/qwen3_moe.py +++ b/auto_round/modeling/unfused_moe/qwen3_moe.py @@ -39,7 +39,6 @@ def __init__(self, config): ) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - """ """ batch_size, sequence_length, hidden_dim = hidden_states.shape hidden_states = hidden_states.view(-1, hidden_dim) # router_logits: (batch * sequence_length, n_experts) diff --git a/auto_round/modeling/unfused_moe/qwen3_next.py b/auto_round/modeling/unfused_moe/qwen3_next.py index 3441d66da9..ad476e697e 100644 --- a/auto_round/modeling/unfused_moe/qwen3_next.py +++ b/auto_round/modeling/unfused_moe/qwen3_next.py @@ -42,7 +42,6 @@ def __init__(self, config): self.shared_expert_gate = torch.nn.Linear(config.hidden_size, 1, bias=False) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - """ """ batch_size, sequence_length, hidden_dim = hidden_states.shape hidden_states = hidden_states.view(-1, hidden_dim) # router_logits: (batch * sequence_length, n_experts) diff --git a/auto_round/scheme_entry.py b/auto_round/scheme_entry.py index 3326454dba..9ec229439d 100644 --- a/auto_round/scheme_entry.py +++ b/auto_round/scheme_entry.py @@ -25,7 +25,7 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Union +from typing import TYPE_CHECKING from auto_round.logger import logger from auto_round.schemes import parse_scheme @@ -118,7 +118,7 @@ def eager_validate_scheme(config, scheme=None, format=None) -> None: temp_config.check_config() # raises ValueError / NotImplementedError if invalid -def is_weight_scheme(scheme: Union[str, dict, object]) -> bool: +def is_weight_scheme(scheme: str | dict | object) -> bool: if isinstance(scheme, str): return scheme.upper().startswith("W") if isinstance(scheme, dict): @@ -134,7 +134,7 @@ def is_weight_scheme(scheme: Union[str, dict, object]) -> bool: return False -def is_gguf_k_target(value: Union[str, "AutoScheme", object]) -> bool: +def is_gguf_k_target(value: str | AutoScheme | object) -> bool: from auto_round.auto_scheme.gen_auto_scheme import AutoScheme if isinstance(value, str): diff --git a/auto_round/schemes.py b/auto_round/schemes.py index 736c62cac5..3b6de11c01 100644 --- a/auto_round/schemes.py +++ b/auto_round/schemes.py @@ -15,7 +15,7 @@ from copy import deepcopy from dataclasses import asdict, dataclass, fields from enum import Enum -from typing import TYPE_CHECKING, Any, Optional, Union +from typing import TYPE_CHECKING, Any, Union import torch @@ -144,15 +144,15 @@ from auto_round.utils import SUPPORTED_DTYPES, contain_any_mm_keys, infer_bits_by_data_type __all__ = [ - "QuantizationScheme", - "BackendDataType", - "GGUF_SCHEME_FACTS", "GGUF_PRESET_ALIASES", - "is_standard_fp", + "GGUF_SCHEME_FACTS", + "BackendDataType", + "QuantizationScheme", + "get_gguf_scheme", "is_mx_fp", - "is_nv_fp", "is_mx_int", - "get_gguf_scheme", + "is_nv_fp", + "is_standard_fp", "preset_name_to_scheme", ] @@ -200,14 +200,14 @@ class QuantizationScheme: group_size: int = 128 sym: bool = True data_type: str = "int" - act_bits: Optional[int] = None - act_group_size: Optional[int] = None - act_sym: Optional[bool] = None - act_data_type: Optional[str] = None - act_dynamic: Optional[bool] = None - super_bits: Optional[int] = None - super_group_size: Optional[int] = None - rotation_config: Optional[dict] = None + act_bits: int | None = None + act_group_size: int | None = None + act_sym: bool | None = None + act_data_type: str | None = None + act_dynamic: bool | None = None + super_bits: int | None = None + super_group_size: int | None = None + rotation_config: dict | None = None @classmethod def empty(cls): @@ -375,7 +375,7 @@ def preset_name_to_scheme(name: str) -> QuantizationScheme: return scheme_args -def scheme_to_preset_name(scheme: Union[str, QuantizationScheme]) -> str: +def scheme_to_preset_name(scheme: str | QuantizationScheme) -> str: """Get preset scheme name from a QuantizationScheme instance.""" if isinstance(scheme, str): name = scheme.upper() @@ -423,8 +423,8 @@ def _reconcile_bits_and_dtype(config: dict, prefix: str = ""): def _override_scheme_with_user_specify( - scheme: Union[str, dict, QuantizationScheme], user_scheme_overrides: dict[str, Any], return_str=True -) -> Union[str, QuantizationScheme]: + scheme: str | dict | QuantizationScheme, user_scheme_overrides: dict[str, Any], return_str=True +) -> str | QuantizationScheme: """ Updates a base quantization scheme with user-provided overrides. Handles GGUF formatting and synchronizes weight/activation parameters. @@ -518,8 +518,8 @@ def format_allows_w8_asym(format: str | None) -> bool: def parse_scheme( scheme: Union[str, dict, QuantizationScheme, "AutoScheme"], user_scheme_overrides: dict[str, Any], - format: str = None, -) -> tuple[Union[str, QuantizationScheme], bool, dict[str, Any]]: + format: str | None = None, +) -> tuple[str | QuantizationScheme, bool, dict[str, Any]]: """ Parses the final scheme. """ @@ -978,14 +978,12 @@ def _handle_special_schemes( layer_config[n] = {"bits": 4, "data_type": "int"} elif n != lm_head_name and mllm: layer_config[n] = {"bits": 16} - elif n != lm_head_name: - layer_config[n] = {"bits": 8, "data_type": "int"} - elif n == lm_head_name and quant_lm_head: + elif n != lm_head_name or quant_lm_head: layer_config[n] = {"bits": 8, "data_type": "int"} return layer_config -def get_gguf_scheme(scheme: Union[str, QuantizationScheme]) -> str: +def get_gguf_scheme(scheme: str | QuantizationScheme) -> str: if scheme is None: return "" if isinstance(scheme, str) and scheme.upper().startswith("GGUF"): diff --git a/auto_round/special_model_handler.py b/auto_round/special_model_handler.py index dd9202c60a..41c66099f5 100644 --- a/auto_round/special_model_handler.py +++ b/auto_round/special_model_handler.py @@ -14,9 +14,10 @@ import importlib import re import sys +from collections.abc import Callable from dataclasses import dataclass, field from types import SimpleNamespace -from typing import Any, Callable +from typing import Any import torch @@ -394,7 +395,7 @@ def _handle_special_model(model): def update_module( - model, formats: list[OutputFormat] = None, trust_remote_code: bool = True, cleanup_original: bool = True + model, formats: list[OutputFormat] | None = None, trust_remote_code: bool = True, cleanup_original: bool = True ): gguf_export = formats is not None and any(format_.is_gguf() for format_ in formats) model = apply_replacements(model, gguf_export=gguf_export) diff --git a/auto_round/utils/bit_packing.py b/auto_round/utils/bit_packing.py index f444cc4d12..a4dcf1a7ca 100644 --- a/auto_round/utils/bit_packing.py +++ b/auto_round/utils/bit_packing.py @@ -48,20 +48,18 @@ there is no ``32 % bits == 0`` constraint. """ -from typing import Optional - import torch __all__ = [ "AWQ_PACK_ORDER", "AWQ_REVERSE_ORDER", "SUPPORTED_PACKING_BITS", - "awq_reverse_reorder", "awq_reorder", + "awq_reverse_reorder", "pack_bitstream", - "unpack_bitstream", "packed_dim_size", "requires_generic_bit_packing", + "unpack_bitstream", ] # Weight bit-widths that ``pack_bitstream`` / ``unpack_bitstream`` handle. @@ -240,7 +238,7 @@ def pack_scalar_zero(value: int, bits: int, num_words: int, shape, device=None) return row.reshape(*(1,) * (len(shape) - 1), num_words).expand(*shape).contiguous() -def infer_packed_bits(num_values: int, num_words: int) -> Optional[int]: +def infer_packed_bits(num_values: int, num_words: int) -> int | None: """Best-effort recovery of ``bits`` from packed/unpacked dimension sizes.""" if num_values <= 0 or num_words <= 0: return None diff --git a/auto_round/utils/common.py b/auto_round/utils/common.py index 6f7dfbb3bf..390c2cdae9 100644 --- a/auto_round/utils/common.py +++ b/auto_round/utils/common.py @@ -126,7 +126,7 @@ def download_audiocaps_csv(): logger.info(f"AudioCaps dataset cached at: {cache_file}") except requests.RequestException as e: raise RuntimeError(f"Failed to download AudioCaps from {url}: {e}") from e - except IOError as e: + except OSError as e: raise RuntimeError(f"Failed to write AudioCaps cache to {cache_file}: {e}") from e return cache_file @@ -146,7 +146,7 @@ def torch_version_at_least(version_string): TORCH_VERSION_AT_LEAST_2_4 = torch_version_at_least("2.4.0") -class LazyImport(object): +class LazyImport: """Lazy import python module till use.""" def __init__(self, module_name): @@ -407,7 +407,7 @@ def monkey_patch_transformers(): if parsed_version >= version.parse("5.0.0"): from transformers.initialization import no_init_weights - setattr(transformers.modeling_utils, "no_init_weights", no_init_weights) + transformers.modeling_utils.no_init_weights = no_init_weights if parsed_version >= version.parse("5.2.0"): # transformers 5.2.0 added Transpose.convert() which calls get_parameter() on # quantized buffer tensors (weight_packed, weight_scale), causing AttributeError. @@ -728,11 +728,11 @@ def __init__(self): self._support_list = self._support_format + self._gguf_format def __contains__(self, key): - return True if key in self._support_list else False + return key in self._support_list def __str__(self): # Return "(%s)" % ', '.join(self._support_format + ("gguf:q*_0", "gguf:q*_1", "gguf:q*_k_s")) - return "(%s)" % ", ".join(self._support_list) + return f"({', '.join(self._support_list)})" def __getitem__(self, key): return self._support_list[key] diff --git a/auto_round/utils/device.py b/auto_round/utils/device.py index fda2e7cf02..6ba56d0904 100644 --- a/auto_round/utils/device.py +++ b/auto_round/utils/device.py @@ -19,10 +19,11 @@ import shutil import sys import tempfile +from collections.abc import Callable from contextlib import ContextDecorator, contextmanager from functools import lru_cache from threading import Lock -from typing import Any, Callable, Optional, Union +from typing import Any import cpuinfo import psutil @@ -80,7 +81,7 @@ def _use_hpu_compile_mode(): return TORCH_VERSION_AT_LEAST_2_4 and not is_hpu_lazy_mode() -def _bump_dynamo_cache_limit(min_size: Optional[int] = None): +def _bump_dynamo_cache_limit(min_size: int | None = None): """Raise torch._dynamo cache/recompile limits. The same quant function (e.g. ``quant_tensor_sym``) is reused across @@ -109,9 +110,7 @@ def _bump_dynamo_cache_limit(min_size: Optional[int] = None): pass -def compile_func( - fun: Union[torch.nn.Module, Callable], device: Union[str, torch.device, int] -) -> Union[torch.nn.Module, Callable]: +def compile_func(fun: torch.nn.Module | Callable, device: str | torch.device | int) -> torch.nn.Module | Callable: """Compile a function on the specified device. The shared dynamo cache-limit bump lives in :func:`_bump_dynamo_cache_limit`; @@ -162,11 +161,9 @@ def is_tbb_available(): # pragma: no cover return False if not _is_tbb_configured(): logger.warning_once( - ( - "TBB is installed but not configured correctly. \n" - "Please add the TBB library path to `LD_LIBRARY_PATH`, " - "for example: `export LD_LIBRARY_PATH=$LD_LIBRARY_PATH:/usr/local/lib/`." - ) + "TBB is installed but not configured correctly. \n" + "Please add the TBB library path to `LD_LIBRARY_PATH`, " + "for example: `export LD_LIBRARY_PATH=$LD_LIBRARY_PATH:/usr/local/lib/`." ) return False return True @@ -180,9 +177,7 @@ def can_pack_with_numba(): # pragma: no cover if not is_numba_available(): logger.warning_once("Numba is not installed, please install it with `pip install numba`.") return False - if not is_tbb_available(): - return False - return True + return is_tbb_available() ## check hpex @@ -416,7 +411,7 @@ def __exit__(self, exc_type, exc, exc_tb): return False -class CpuInfo(object): +class CpuInfo: """Get CPU Info.""" def __init__(self): @@ -596,7 +591,7 @@ def set_tuning_device_for_layer(model, name: str, device: str) -> None: def set_non_auto_device_map( - model: torch.nn.Module, device_map: Union[str, int, dict], quant_layer_names: Union[None, list, tuple] = None + model: torch.nn.Module, device_map: str | int | dict, quant_layer_names: None | list | tuple = None ) -> None: if not device_map or device_map == "auto" or isinstance(device_map, int): return @@ -1174,7 +1169,7 @@ def partition_dict_numbers(number_dict, n): # - Assign each item to the group with the current smallest sum # Complexity: O(m log n) which scales well for large m (layers) groups_sums = [0.0] * n - groups = [dict() for _ in range(n)] + groups = [{} for _ in range(n)] # Sort items descending by size items_sorted = sorted(items, key=lambda kv: kv[1], reverse=True) @@ -1242,7 +1237,7 @@ def set_avg_auto_device_map(model: torch.nn.Module, device_map): for device in device_list: if device.startswith("hpu") and len(device_list) > 1: logger.warning_once("Auto-scheme does not support multiple HPUs.") - if device.startswith("cpu") or device.startswith("hpu"): + if device.startswith(("cpu", "hpu")): continue gpu_devices.append(device) num_devices = len(gpu_devices) @@ -1260,11 +1255,9 @@ def set_avg_auto_device_map(model: torch.nn.Module, device_map): params_dict[n] = in_features * out_features res_list = partition_dict_numbers(params_dict, num_devices) - device_index = 0 - for res in res_list: + for device_index, res in enumerate(res_list): for key in res.keys(): set_tuning_device_for_layer(block_module, key, gpu_devices[device_index]) - device_index += 1 if __name__ == "__main__": @@ -1292,7 +1285,7 @@ def set_avg_auto_device_map(model: torch.nn.Module, device_map): print(f"Group {i + 1}: {group}, Sum: {sum(group.values())}") -def parse_available_devices(device_map: Union[str, torch.device, int, dict, None] = None) -> list: +def parse_available_devices(device_map: str | torch.device | int | dict | None = None) -> list: """ Parse the device map and return a list of all available devices. @@ -1404,7 +1397,7 @@ def parse_available_devices(device_map: Union[str, torch.device, int, dict, None raise TypeError(f"Unsupported device_map type: {type(device_map)}") -@lru_cache(maxsize=None) +@functools.cache def is_gaudi2(): try: import habana_frameworks.torch.utils.experimental as htexp @@ -1602,7 +1595,7 @@ def wrapper(*args, **kwargs): # This function is designed for Auto Scheme and Diffusion Pipeline, # which requires dispatching the whole model on all available devices. def dispatch_model_by_all_available_devices( - model: torch.nn.Module, device_map: Union[str, int, dict, None] + model: torch.nn.Module, device_map: str | int | dict | None ) -> torch.nn.Module: # Important Notice: This dispatch does not follow dict device_map, just extract all available devices and use them device_type = get_major_device() diff --git a/auto_round/utils/device_manager.py b/auto_round/utils/device_manager.py index f2972d7337..c06c3807d3 100644 --- a/auto_round/utils/device_manager.py +++ b/auto_round/utils/device_manager.py @@ -49,7 +49,6 @@ import gc import re import sys -from typing import Optional, Union import torch @@ -57,23 +56,23 @@ __all__ = [ "ARDevice", + "ClearMemory", "DeviceManager", + "clear_memory", + "default_enable_torch_compile", + "detect_device_count", "device_manager", - "normalize_default_device_map", "get_ar_device", - "default_enable_torch_compile", + "get_available_device_types", "get_current_device_manager", "get_current_device_type", - "is_device_available", - "get_available_device_types", - "get_major_device", - "detect_device_count", "get_device_and_parallelism", + "get_device_memory", + "get_major_device", "get_packing_device", "is_auto_device_mapping", - "get_device_memory", - "ClearMemory", - "clear_memory", + "is_device_available", + "normalize_default_device_map", ] @@ -89,7 +88,7 @@ _PREFERRED_ORDER = ("cuda", "xpu", "hpu") # add mps later -def normalize_default_device_map(device_map: Union[None, str, int, torch.device, dict]): +def normalize_default_device_map(device_map: None | str | int | torch.device | dict): """Normalize default device selection across entry points. On Apple Silicon, the default ``0`` / ``"0"`` / ``None`` / ``"auto"`` @@ -106,7 +105,7 @@ def normalize_default_device_map(device_map: Union[None, str, int, torch.device, return device_map -def _torch_accelerator_type() -> Optional[str]: +def _torch_accelerator_type() -> str | None: """Return the canonical accelerator type reported by ``torch.accelerator``. A PyTorch build exposes at most one accelerator backend; this returns its @@ -167,7 +166,7 @@ def _hpu_available() -> bool: return False -def _normalize_device_type(device: Union[None, str, int, torch.device]) -> Optional[str]: +def _normalize_device_type(device: None | str | int | torch.device) -> str | None: """Reduce any device spec to a bare backend type string (``"cuda"`` ...).""" if device is None: return get_current_device_type() @@ -240,7 +239,7 @@ def get_available_device_types() -> list[str]: class _DeviceIndexContext: """Fallback for ``torch.accelerator.device_index`` on older PyTorch/backends.""" - def __init__(self, device: "ARDevice", index: int): + def __init__(self, device: ARDevice, index: int): self._device = device self._index = index self._prev = None @@ -280,7 +279,7 @@ class ARDevice: #: PyTorch backend that lacks a dedicated subclass (e.g. a fresh ``npu``). device_type: str = "" - _registry: dict[str, type["ARDevice"]] = {} + _registry: dict[str, type[ARDevice]] = {} def __init_subclass__(cls, **kwargs): super().__init_subclass__(**kwargs) @@ -289,7 +288,7 @@ def __init_subclass__(cls, **kwargs): ARDevice._registry[dtype] = cls @classmethod - def create(cls, device_type: str) -> "ARDevice": + def create(cls, device_type: str) -> ARDevice: """Instantiate the most specific :class:`Device` for ``device_type``.""" subclass = cls._registry.get(device_type) if subclass is not None: @@ -297,7 +296,7 @@ def create(cls, device_type: str) -> "ARDevice": return ARDevice(device_type) @staticmethod - def get_device_module(device: Union[None, str, int, torch.device] = None): + def get_device_module(device: None | str | int | torch.device = None): """Return the backend runtime module for ``device`` (e.g. ``torch.cuda``). This is a thin, version-tolerant wrapper around ``torch.get_device_module`` @@ -320,7 +319,7 @@ def get_device_module(device: Union[None, str, int, torch.device] = None): pass return getattr(torch, device_type, None) - def __init__(self, device_type: Optional[str] = None): + def __init__(self, device_type: str | None = None): self.type = device_type or self.device_type # Prefer the unified ``torch.accelerator`` API for runtime ops when this @@ -355,14 +354,14 @@ def current_device(self) -> int: pass return 0 - def set_device(self, index: Union[int, str, torch.device]) -> None: + def set_device(self, index: int | str | torch.device) -> None: if self._module is None: return ok, _ = _module_call(self._module, ("set_device_index", "set_device_idx", "set_device"), index) if ok: return - def device(self, index: Union[int, str, torch.device, None] = None) -> torch.device: + def device(self, index: int | str | torch.device | None = None) -> torch.device: """Build a ``torch.device`` for this backend / card ``index``.""" if index is None: return torch.device(self.type) @@ -377,7 +376,7 @@ def device(self, index: Union[int, str, torch.device, None] = None) -> torch.dev # return [self.device(i) for i in range(self.device_count())] # -- runtime ------------------------------------------------------------ - def synchronize(self, index: Union[int, None] = None) -> None: + def synchronize(self, index: int | None = None) -> None: if self._module is None: return fn = getattr(self._module, "synchronize", None) @@ -510,7 +509,7 @@ class HpuARDevice(ARDevice): device_type = "hpu" @staticmethod - def get_device_module(device: Union[None, str, int, torch.device] = None): + def get_device_module(device: None | str | int | torch.device = None): """Return the backend runtime module for ``device`` (e.g. ``torch.cuda``). This is a thin, version-tolerant wrapper around ``torch.get_device_module`` @@ -535,7 +534,7 @@ def get_device_module(device: Union[None, str, int, torch.device] = None): except Exception: # pragma: no cover return None - def set_device(self, index: Union[int, str, torch.device]) -> None: + def set_device(self, index: int | str | torch.device) -> None: if self._module is None: return fn = getattr(self._module, "set_device", None) @@ -581,13 +580,13 @@ class MpsARDevice(ARDevice): device_type = "mps" - def __init__(self, device_type: Optional[str] = None): + def __init__(self, device_type: str | None = None): # Always use torch.mps directly, never torch.accelerator. self.type = "mps" self._module = getattr(torch, "mps", None) @staticmethod - def get_device_module(device: Union[None, str, int, torch.device] = None): + def get_device_module(device: None | str | int | torch.device = None): """Return the backend runtime module for ``device`` (e.g. ``torch.cuda``). This is a thin, version-tolerant wrapper around ``torch.get_device_module`` @@ -611,10 +610,10 @@ def is_available(self) -> bool: def current_device(self) -> int: return 0 - def set_device(self, index: Union[int, str, torch.device]) -> None: + def set_device(self, index: int | str | torch.device) -> None: return None - def device(self, index: Union[int, str, torch.device, None] = None) -> torch.device: + def device(self, index: int | str | torch.device | None = None) -> torch.device: """Build a ``torch.device`` for this backend / card ``index``.""" return torch.device("mps") @@ -678,7 +677,7 @@ class CpuARDevice(ARDevice): device_type = "cpu" @staticmethod - def get_device_module(device: Union[None, str, int, torch.device] = None): + def get_device_module(device: None | str | int | torch.device = None): return None # -- discovery ---------------------------------------------------------- @@ -691,20 +690,20 @@ def device_count(self) -> int: # A single logical device from torch's view. def current_device(self) -> int: return 0 - def set_device(self, index: Union[int, str, torch.device]) -> None: # no-op + def set_device(self, index: int | str | torch.device) -> None: # no-op return None - def device(self, index: Union[int, str, torch.device, None] = None) -> torch.device: + def device(self, index: int | str | torch.device | None = None) -> torch.device: return torch.device("cpu") # -- runtime ------------------------------------------------------------ - def synchronize(self, index: Union[int, None] = None) -> None: # no-op + def synchronize(self, index: int | None = None) -> None: # no-op return None def empty_cache(self) -> None: # no-op: CPU has no caching allocator. return gc.collect() - def get_device_capability(self, index: Union[int, None] = None): + def get_device_capability(self, index: int | None = None): return None def device_index(self, index: int): # nothing to switch on CPU. @@ -775,26 +774,26 @@ class DeviceManager: shared instance, so the active device / device_list is always single-sourced. """ - _instance: Optional["DeviceManager"] = None + _instance: DeviceManager | None = None def __new__(cls, *args, **kwargs): if cls._instance is None: cls._instance = super().__new__(cls) return cls._instance - def __init__(self, device_map: Union[None, str, torch.device, int, dict] = None): + def __init__(self, device_map: None | str | torch.device | int | dict = None): # Initialise backing state once; later constructions reuse the singleton. if not getattr(self, "_initialized", False): self._cache: dict[str, ARDevice] = {} self._device_map = None - self._device_list: Optional[list] = None - self._major_device: Optional[str] = None + self._device_list: list | None = None + self._major_device: str | None = None self._initialized = True if device_map is not None: self.configure(device_map) # -- device_map configuration ------------------------------------------ - def configure(self, device_map: Union[None, str, torch.device, int, dict] = 0) -> "DeviceManager": + def configure(self, device_map: None | str | torch.device | int | dict = 0) -> DeviceManager: """Resolve a ``device_map`` into a concrete device list and major device. Centralises the device-map parsing the compressors used to perform by @@ -834,7 +833,7 @@ def device(self) -> str: return self._major_device @device.setter - def device(self, value: Union[str, torch.device]) -> None: + def device(self, value: str | torch.device) -> None: """Override the major device (e.g. an OOM fallback to ``"cpu"``).""" self._major_device = str(value) if isinstance(value, torch.device) else value @@ -852,7 +851,7 @@ def register(self, device_cls: type[ARDevice]) -> None: self._cache.pop(dtype, None) # -- lookup ------------------------------------------------------------- - def get_ar_device(self, device_type: Union[None, str, int, torch.device] = None) -> ARDevice: + def get_ar_device(self, device_type: None | str | int | torch.device = None) -> ARDevice: """Return the cached :class:`Device` for ``device_type`` (default: current).""" normalized = _normalize_device_type(device_type) or "cpu" device = self._cache.get(normalized) @@ -889,13 +888,13 @@ def all_devices(self) -> list[torch.device]: device_manager = DeviceManager() -def get_ar_device(device_type: Union[None, str, int, torch.device] = None) -> ARDevice: +def get_ar_device(device_type: None | str | int | torch.device = None) -> ARDevice: """Return the cached :class:`Device` handle for a specific backend type.""" return device_manager.get_ar_device(device_type) def default_enable_torch_compile( - device: Union[None, str, int, torch.device] = None, platform_name: str | None = None + device: None | str | int | torch.device = None, platform_name: str | None = None ) -> bool: """Return the safe torch.compile default for a backend.""" return (platform_name or sys.platform) != "win32" @@ -914,7 +913,7 @@ def detect_device_count() -> int: return get_current_device_manager().device_count() -def get_device_and_parallelism(device: Union[str, torch.device, int, dict]) -> tuple[str, bool]: +def get_device_and_parallelism(device: str | torch.device | int | dict) -> tuple[str, bool]: """Resolve a device spec into ``(device, parallelism)``. The multi-card *parallelism* policy itself is kept as a standalone function @@ -961,7 +960,7 @@ def get_device_and_parallelism(device: Union[str, torch.device, int, dict]) -> t return device, parallelism -def get_packing_device(device: Union[str, torch.device, None] = "auto") -> torch.device: +def get_packing_device(device: str | torch.device | None = "auto") -> torch.device: """Selects the packing device. - ``"auto"``: choose best available (active accelerator > CPU). @@ -987,12 +986,10 @@ def get_packing_device(device: Union[str, torch.device, None] = "auto") -> torch raise TypeError(f"Unsupported device type: {type(device)} ({device})") -def is_auto_device_mapping(device_map: Union[str, int, dict, None]) -> bool: +def is_auto_device_mapping(device_map: str | int | dict | None) -> bool: if device_map is None or isinstance(device_map, int): return False - elif device_map == "auto": - return True - elif isinstance(device_map, str) and "," in device_map: + elif device_map == "auto" or (isinstance(device_map, str) and "," in device_map): return True elif isinstance(device_map, dict): return False @@ -1000,7 +997,7 @@ def is_auto_device_mapping(device_map: Union[str, int, dict, None]) -> bool: return False -def get_major_device(device_map: Union[None, str, torch.device, int, dict] = None) -> str: +def get_major_device(device_map: None | str | torch.device | int | dict = None) -> str: if device_map is None or isinstance(device_map, (str, torch.device, int)): """Detects the appropriate computation device. @@ -1044,8 +1041,6 @@ def is_valid_digit(s): if device == "tp": # pragma: no cover # should not specify card, e.g., cuda:0 device = get_current_device_type() or "cpu" - else: - device = device return device if isinstance(device_map, dict) and device_map: @@ -1091,8 +1086,8 @@ def get_device_memory(i: int = 0) -> int: def _clear_memory_for_cpu_and_cuda( - tensor: Union[torch.Tensor, list, None] = None, - device_list: Union[tuple, list, str, torch.device, None] = None, + tensor: torch.Tensor | list | None = None, + device_list: tuple | list | str | torch.device | None = None, ): # ------------------------ # Clear CPU-side references @@ -1152,7 +1147,7 @@ def _clear_memory_for_cpu_and_cuda( class ClearMemory: - def __init__(self, device_list: Union[list, tuple, None] = None): + def __init__(self, device_list: list | tuple | None = None): self._device_list = device_list @property @@ -1167,8 +1162,8 @@ def device_list(self, value): def __call__( self, - tensor: Union[torch.Tensor, None, list] = None, - device_list: Union[list, tuple, None] = None, + tensor: torch.Tensor | None | list = None, + device_list: list | tuple | None = None, ): # Lazy imports: these symbols live in utils/device.py. from auto_round.utils.device import _force_trim_malloc, is_hpex_available, memory_monitor diff --git a/auto_round/utils/disk_stream_util.py b/auto_round/utils/disk_stream_util.py index 1ea98c06ec..c8f9bb53a3 100644 --- a/auto_round/utils/disk_stream_util.py +++ b/auto_round/utils/disk_stream_util.py @@ -13,9 +13,8 @@ import json import logging import re -from functools import lru_cache +from functools import cache, lru_cache from pathlib import Path -from typing import Dict import torch import torch.nn as nn @@ -48,7 +47,7 @@ def __init__(self, checkpoint_dir: str): # The shard names come from the checkpoint's own index, i.e. from the # artifact being loaded -- validate them against the directory before # anything downstream gets a chance to join and open them. - self.weight_map: Dict[str, str] = validate_weight_map( + self.weight_map: dict[str, str] = validate_weight_map( json.load(f)["weight_map"], self.checkpoint_dir, index_path=index_path ) else: @@ -70,14 +69,14 @@ def tensor_shape(self, name: str) -> tuple[int, ...]: with safe_open(str(shard_path), framework="pt") as f: return tuple(f.get_slice(name).get_shape()) - def read_tensors(self, names: list[str], device: str = "cpu") -> Dict[str, torch.Tensor]: + def read_tensors(self, names: list[str], device: str = "cpu") -> dict[str, torch.Tensor]: """Read several tensors, grouped by shard file so each shard is opened and closed (unmapped) once regardless of how many tensors are pulled from it.""" - by_shard: Dict[str, list[str]] = {} + by_shard: dict[str, list[str]] = {} for name in names: by_shard.setdefault(self.weight_map[name], []).append(name) - result: Dict[str, torch.Tensor] = {} + result: dict[str, torch.Tensor] = {} for shard_name, shard_tensor_names in by_shard.items(): shard_path = resolve_within_directory(self.checkpoint_dir, shard_name) with safe_open(str(shard_path), framework="pt") as f: @@ -94,7 +93,7 @@ def tensor_names_with_prefix(self, prefix: str) -> list[str]: @lru_cache(maxsize=8) -def get_safetensors_index(checkpoint_dir: str) -> "SafetensorsIndex": +def get_safetensors_index(checkpoint_dir: str) -> SafetensorsIndex: """Shared ``SafetensorsIndex`` per checkpoint directory. Only the (cheap) name->shard map is cached; no ``safe_open`` handle is kept, so this @@ -152,7 +151,7 @@ def checkpoint_has_native_fused_moe_experts(checkpoint_dir: str) -> bool: # whose names already match the model are untouched. -@lru_cache(maxsize=None) +@cache def _reverse_renamings_for(model_type): """Invert the checkpoint-conversion WeightRenaming entries for one family. @@ -188,7 +187,7 @@ def _reverse_renamings_for(model_type): return tuple(reversed_transforms) -@lru_cache(maxsize=None) +@cache def _model_types_for_dir(checkpoint_dir: str): """Every model_type in the checkpoint's config, including nested sub-configs. @@ -360,7 +359,7 @@ def _slice_fused_expert(fused, proj, attr, is_gate_up, target_shape, transposed_ ) -@lru_cache(maxsize=None) +@cache def _expert_projection_renames_for(model_type): """Map a fused projection to the checkpoint-side per-expert projection names. @@ -407,7 +406,7 @@ def _expert_projection_renames_for(model_type): return tuple(renames.items()) -@lru_cache(maxsize=None) +@cache def _concat_converters_for(model_type): """Model-side params assembled by concatenating several checkpoint tensors. @@ -499,7 +498,7 @@ def _dot_natural_key(name: str): return parts -@lru_cache(maxsize=None) +@cache def _wildcard_concat_converters_for(model_type): """Wildcard shard-concat converters registered for one family. @@ -546,7 +545,7 @@ def _wildcard_concat_converters_for(model_type): return tuple(converters) -@lru_cache(maxsize=None) +@cache def _wildcard_split_converters_for(model_type): """Save-side inverse of :func:`_wildcard_concat_converters_for`. @@ -610,7 +609,7 @@ def _resolve_num_shards(config, num_shards_attribute): return None -def split_merged_concat_tensor(config, full_name: str, tensor: "torch.Tensor"): +def split_merged_concat_tensor(config, full_name: str, tensor: torch.Tensor): """Split a merged model-side concat parameter back into its checkpoint shards. Inverse of the load-time :func:`_assemble_sharded_tensor`. Qwen3-Next "Flash" @@ -1259,12 +1258,12 @@ class stream_block_forward: is for a plain inference-only forward pass (e.g. eval loss), not tuning. """ - def __init__(self, model: nn.Module, index: SafetensorsIndex, device: str, block_names: list[str] = None): + def __init__(self, model: nn.Module, index: SafetensorsIndex, device: str, block_names: list[str] | None = None): self.model = model self.index = index self.device = device self.block_names = block_names if block_names is not None else _default_block_names(model) - self._originals: Dict[str, "callable"] = {} + self._originals: dict[str, callable] = {} def __enter__(self): for block_name in self.block_names: diff --git a/auto_round/utils/distributed.py b/auto_round/utils/distributed.py index d523e0e65b..6a7936ce92 100644 --- a/auto_round/utils/distributed.py +++ b/auto_round/utils/distributed.py @@ -13,14 +13,14 @@ # limitations under the License. import os -from functools import lru_cache +from functools import cache import torch from auto_round.logger import logger -@lru_cache(maxsize=None) +@cache def is_distributed(): import torch.distributed as dist diff --git a/auto_round/utils/missing_tensors.py b/auto_round/utils/missing_tensors.py index 82a9125033..b894a5effa 100644 --- a/auto_round/utils/missing_tensors.py +++ b/auto_round/utils/missing_tensors.py @@ -536,9 +536,9 @@ def _woq_quantize_missing_tensors(target_dir: str, missing_tensors_dict: dict) - # Pre-compile all valid regex patterns once to avoid repeated re.compile() calls # for every tensor lookup (O(N×M) → O(M) compile + O(N×M) match). _compiled_patterns: list = [] - for pattern in extra_config: + for pattern, pattern_cfg in extra_config.items(): try: - _compiled_patterns.append((_re.compile(pattern), pattern, extra_config[pattern])) + _compiled_patterns.append((_re.compile(pattern), pattern, pattern_cfg)) except _re.error as exc: logger.warning( "Invalid regex key in extra_config ignored during pre-compilation: %r (%s)", @@ -630,7 +630,7 @@ def _is_eligible(k: str) -> bool: if isinstance(block_name_to_quantize, list) else [b.strip() for b in block_name_to_quantize.split(",") if b.strip()] ) - if not any(k.startswith(b + ".") or k.startswith(b + "[") for b in blocks): + if not any(k.startswith((b + ".", b + "[")) for b in blocks): return False layer_name = k[: -len(".weight")] layer_cfg = _resolve_layer_cfg(layer_name) diff --git a/auto_round/utils/model.py b/auto_round/utils/model.py index ae982495ff..7be3aae9cb 100644 --- a/auto_round/utils/model.py +++ b/auto_round/utils/model.py @@ -16,9 +16,11 @@ import json import os import re +import sys from collections import UserDict +from collections.abc import Callable from pathlib import Path -from typing import TYPE_CHECKING, Any, Callable, Optional, Union +from typing import TYPE_CHECKING, Any import psutil import torch @@ -191,7 +193,7 @@ def check_diffusers_installed(): # pragma: no cover return True except ImportError: logger.error("Please install diffusers via 'pip install diffusers'" " to run diffusion model") - exit(-1) + sys.exit(-1) def check_start_with_block_name(name: str, block_name_to_quantize: list): @@ -211,7 +213,7 @@ def check_start_with_block_name(name: str, block_name_to_quantize: list): return False -def download_or_get_path(repo_id: str, platform: str = None) -> str: +def download_or_get_path(repo_id: str, platform: str | None = None) -> str: from auto_round import envs if platform is None: @@ -226,7 +228,7 @@ def download_or_get_path(repo_id: str, platform: str = None) -> str: return download_hf_model(repo_id) -def download_modelscope_model(repo_id: str, local_dir: str = None, cache_dir: str = None): +def download_modelscope_model(repo_id: str, local_dir: str | None = None, cache_dir: str | None = None): from modelscope.utils.file_utils import get_modelscope_cache_dir # pylint: disable=E0401 system_cache = cache_dir if cache_dir is not None else get_modelscope_cache_dir() @@ -441,7 +443,7 @@ def llm_load_model( pretrained_model_name_or_path: str, platform: str = "hf", trust_remote_code: bool = True, - model_dtype: str = None, + model_dtype: str | None = None, device: str = "cpu", **kwargs, ): @@ -568,7 +570,7 @@ def llm_load_model( return model, tokenizer -def _find_pipeline_model_subfolder(model_dir_or_repo: str, file_list: list = None) -> tuple: +def _find_pipeline_model_subfolder(model_dir_or_repo: str, file_list: list | None = None) -> tuple: """Find model/processor subfolders from a pipeline's model_index.json. Works for both local directories and remote HF repos. @@ -647,7 +649,7 @@ def mllm_load_model( torch_dtype: str = "auto", use_auto_mapping: bool = True, trust_remote_code: bool = True, - model_dtype: str = None, + model_dtype: str | None = None, **kwargs, ): from auto_round.special_model_handler import MISTRAL_3_2_MODELS @@ -881,7 +883,7 @@ def mllm_load_model( else: raise - if any([name in model.name_or_path for name in MISTRAL_3_2_MODELS]): + if any(name in model.name_or_path for name in MISTRAL_3_2_MODELS): from mistral_common.tokens.tokenizers.mistral import MistralTokenizer # pylint: disable=E0401 if os.path.isdir(pretrained_model_name_or_path): @@ -895,7 +897,7 @@ def mllm_load_model( tokenizer = AutoTokenizer.from_pretrained( pretrained_model_name_or_path, trust_remote_code=trust_remote_code, - fix_mistral_regex=True if model_type in FIX_MISTRAL_REGEX_MODEL_TYPE_LIST else False, + fix_mistral_regex=model_type in FIX_MISTRAL_REGEX_MODEL_TYPE_LIST, **processor_load_kwargs, ) processor = AutoProcessor.from_pretrained( @@ -959,12 +961,12 @@ def _stable_audio_pipeline_fn( def diffusion_load_model( pretrained_model_name_or_path: str, platform: str = "hf", - device: Union[str, torch.device] = "cpu", - torch_dtype: Union[str, torch.dtype] = "auto", + device: str | torch.device = "cpu", + torch_dtype: str | torch.dtype = "auto", use_auto_mapping: bool = False, trust_remote_code: bool = True, - model_dtype: str = None, - default_torch_dtype: Union[str, torch.dtype] = "auto", + model_dtype: str | None = None, + default_torch_dtype: str | torch.dtype = "auto", **kwargs, ): from functools import partial @@ -1085,8 +1087,8 @@ def diffusion_load_model( _attach_diffusion_pipeline_fn(pipe) # meta model uses model.config.save_pretrained for config saving - setattr(model.config, "save_pretrained", partial(config_save_pretrained, model.config, "config.json", model=model)) - setattr(pipe.config, "save_pretrained", partial(config_save_pretrained, pipe.config, "model_index.json")) + model.config.save_pretrained = partial(config_save_pretrained, model.config, "config.json", model=model) + pipe.config.save_pretrained = partial(config_save_pretrained, pipe.config, "model_index.json") def model_save_pretrained(model, save_directory, **kwargs): super(model.__class__, model).save_pretrained(save_directory, **kwargs) @@ -1096,7 +1098,7 @@ def model_save_pretrained(model, save_directory, **kwargs): writer.write(json.dumps(dict(model.config), indent=2, sort_keys=True) + "\n") # non-meta model uses model.save_pretrained for model and config saving - setattr(model, "save_pretrained", partial(model_save_pretrained, model)) + model.save_pretrained = partial(model_save_pretrained, model) for comp_name in pipe.components: comp = getattr(pipe, comp_name, None) @@ -1107,21 +1109,19 @@ def model_save_pretrained(model, save_directory, **kwargs): and isinstance(comp, torch.nn.Module) ): comp._autoround_checkpoint_subfolder = comp_name - setattr( - comp.config, "save_pretrained", partial(config_save_pretrained, comp.config, "config.json", model=comp) - ) - setattr(comp, "save_pretrained", partial(model_save_pretrained, comp)) + comp.config.save_pretrained = partial(config_save_pretrained, comp.config, "config.json", model=comp) + comp.save_pretrained = partial(model_save_pretrained, comp) return pipe, model.to(device) def load_model( - pretrained_model_name_or_path: Union[str, torch.nn.Module], + pretrained_model_name_or_path: str | torch.nn.Module, platform: str = "hf", - model_dtype: str = None, + model_dtype: str | None = None, trust_remote_code: bool = True, device: str = "cpu", - use_auto_mapping: bool = None, + use_auto_mapping: bool | None = None, use_model_replacements: bool = False, **kwargs, ) -> tuple: @@ -1252,7 +1252,7 @@ def is_pure_text_model(model): _LLM_ONLY_MODEL_TYPES = {"bagel"} -def get_model_name_or_path(model_or_path: Union[str, torch.nn.Module]) -> Optional[str]: +def get_model_name_or_path(model_or_path: str | torch.nn.Module) -> str | None: if isinstance(model_or_path, str): return model_or_path return getattr(model_or_path, "_name_or_path", None) or getattr(model_or_path, "name_or_path", None) @@ -1276,14 +1276,14 @@ def get_model_name_or_path(model_or_path: Union[str, torch.nn.Module]) -> Option _CODE_MODEL_TASKS = {"code-generation", "software-engineering", "text-to-code"} -def _match_code_model_name(value) -> Optional[str]: +def _match_code_model_name(value) -> str | None: if not isinstance(value, str) or not value: return None value = re.split(r"[/\\]", value.rstrip("/\\"))[-1] value = re.sub(r"(?<=[a-z0-9])(?=[A-Z])", " ", value) token_matches = _CODE_MODEL_TOKENS.intersection(re.findall(r"[a-z]+", value.lower())) if token_matches: - return sorted(token_matches)[0] + return min(token_matches) components = re.findall(r"[a-z0-9]+", value.lower()) for family in sorted(_CODE_MODEL_FAMILIES): if any(component == family or re.fullmatch(rf"{family}\d+", component) for component in components): @@ -1329,7 +1329,7 @@ def _get_code_model_match(model_or_path, config=None): return None -def is_code_model(model_or_path: Union[str, torch.nn.Module], config=None) -> bool: +def is_code_model(model_or_path: str | torch.nn.Module, config=None) -> bool: """Return whether a pure-text model is explicitly specialized for code.""" match = _get_code_model_match(model_or_path, config) if match is None: @@ -1338,7 +1338,7 @@ def is_code_model(model_or_path: Union[str, torch.nn.Module], config=None) -> bo return True -def is_mllm_model(model_or_path: Union[str, torch.nn.Module], platform: str = None): +def is_mllm_model(model_or_path: str | torch.nn.Module, platform: str | None = None): from auto_round.utils.common import MM_KEYS model_path = get_model_name_or_path(model_or_path) @@ -1365,29 +1365,27 @@ def is_mllm_model(model_or_path: Union[str, torch.nn.Module], platform: str = No # Only try to download if the path looks like a HF repo id (not a local filesystem path). # Skip download for absolute paths or relative paths that contain current/parent dir markers. # model_path is None for a model or pipeline built in-process, which has no name or path - _is_local_path = isinstance(model_path, str) and ( - os.path.isabs(model_path) or model_path.startswith("./") or model_path.startswith("../") - ) + _is_local_path = isinstance(model_path, str) and (os.path.isabs(model_path) or model_path.startswith(("./", "../"))) if model_path and not os.path.isdir(model_path) and not _is_local_path: model_path = download_or_get_path(model_path, platform=platform) result = False if isinstance(model_path, str): - if os.path.exists(os.path.join(model_path, "preprocessor_config.json")): - result = True - elif os.path.exists(os.path.join(model_path, "processor_config.json")): + if os.path.exists(os.path.join(model_path, "preprocessor_config.json")) or os.path.exists( + os.path.join(model_path, "processor_config.json") + ): result = True elif os.path.exists(os.path.join(model_path, "config.json")): with open(os.path.join(model_path, "config.json")) as f: config = json.load(f) for key in config.keys(): - if any([k in key for k in MM_KEYS]): + if any(k in key for k in MM_KEYS): result = True break if not result and isinstance(model_or_path, torch.nn.Module): for name, module in model_or_path.named_modules(): - if any([k in name for k in MM_KEYS]): + if any(k in name for k in MM_KEYS): result = True break @@ -1398,7 +1396,7 @@ def is_mllm_model(model_or_path: Union[str, torch.nn.Module], platform: str = No return result -def is_gguf_model(model_path: Union[str, torch.nn.Module]) -> bool: +def is_gguf_model(model_path: str | torch.nn.Module) -> bool: is_gguf_file = False if isinstance(model_path, str): if os.path.isfile(model_path) and model_path.endswith(".gguf"): @@ -1415,7 +1413,7 @@ def is_gguf_model(model_path: Union[str, torch.nn.Module]) -> bool: MODULAR_PIPELINE_INDEX_NAME = "modular_model_index.json" -def _find_pipeline_index_file(model_dir_or_repo: str) -> Optional[str]: +def _find_pipeline_index_file(model_dir_or_repo: str) -> str | None: """Return the pipeline index file of a diffusers directory or repo, if it has one. Standard pipelines ship ``model_index.json``, Modular Diffusers pipelines ship @@ -1453,7 +1451,7 @@ def _get_modular_pipeline_class(): return None -def is_diffusion_model(model_or_path: Union[str, object], trust_remote_code: bool = True) -> bool: +def is_diffusion_model(model_or_path: str | object, trust_remote_code: bool = True) -> bool: from auto_round.utils.common import LazyImport # Then check if model_index.json exists for diffusion pipeline, @@ -1747,8 +1745,8 @@ def get_gguf_architecture(dir_model, model_type=ModelType.TEXT): def get_layer_names_in_block( model: torch.nn.Module, supported_types=(torch.nn.Linear, transformers.pytorch_utils.Conv1D), - quant_block_list: list = None, - class_names: tuple = None, + quant_block_list: list | None = None, + class_names: tuple | None = None, ) -> list[str]: """Retrieves the names of layers within each block of the model. @@ -1781,7 +1779,7 @@ def set_nested_attr(module, attr_name: str, value): attrs = attr_name.split(".") for attr in attrs[:-1]: if not hasattr(module, attr): - return None # No need to set act_max for fp layers + return # No need to set act_max for fp layers module = getattr(module, attr) setattr(module, attrs[-1], value) @@ -1885,7 +1883,7 @@ def _to_model_dtype(model, model_dtype): model = cast_model_dtype(model, torch.float32) except Exception: logger.error("please use more device to fit the device or just use one device") - exit() + sys.exit() return model @@ -1983,10 +1981,7 @@ def unsupported_meta_device(model): if param.device.type == "meta" or target_device.type == "meta": return True if target_device.type == "meta": - if hasattr(model, "path"): - return False - else: - return True + return not hasattr(model, "path") return False @@ -2017,11 +2012,11 @@ def to_device(input, device=torch.device("cpu")): return None if isinstance(input, torch.Tensor): return input.to(device) - if isinstance(input, dict) or isinstance(input, UserDict): + if isinstance(input, (dict, UserDict)): for inp in input.keys(): input[inp] = to_device(input[inp], device) - elif isinstance(input, list) or isinstance(input, tuple): + elif isinstance(input, (list, tuple)): if len(input) == 0: return input input_res = [] @@ -2147,9 +2142,7 @@ def is_moe_model_via_config(config) -> bool: config_str = str(config).lower() except Exception: config_str = str(config.to_dict()).lower() if hasattr(config, "to_dict") else "" - if "moe" in config_str or "expert" in config_str: - return True - return False + return "moe" in config_str or "expert" in config_str def to_dtype(input, dtype=torch.float32): @@ -2166,11 +2159,11 @@ def to_dtype(input, dtype=torch.float32): return None if isinstance(input, torch.Tensor): return input.to(dtype) - if isinstance(input, dict) or isinstance(input, UserDict): + if isinstance(input, (dict, UserDict)): for inp in input.keys(): input[inp] = to_dtype(input[inp], dtype) - elif isinstance(input, list) or isinstance(input, tuple): + elif isinstance(input, (list, tuple)): if len(input) == 0: return input input_res = [] @@ -2308,7 +2301,7 @@ def set_amax_for_all_moe_layers(model: torch.nn.Module, layer_name=None, attr_na ) except AttributeError as e: # Provide more helpful debugging information - expert_types = list(set(type(expert).__name__ for expert in sub_module.experts)) + expert_types = list({type(expert).__name__ for expert in sub_module.experts}) raise AttributeError( f"Failed to access attribute '{linear_name}' on experts. " f"MoE module type: {type(sub_module).__name__}, " @@ -2544,7 +2537,7 @@ def find_matching_blocks(model, all_blocks, to_quant_block_names): if not to_quant_block_names: return all_blocks to_quant_block_list = to_quant_block_names - if isinstance(to_quant_block_names, list) or isinstance(to_quant_block_names, tuple): + if isinstance(to_quant_block_names, (list, tuple)): return to_quant_block_names if isinstance(to_quant_block_names, str): to_quant_block_list = [name.strip() for name in to_quant_block_names.split(",")] @@ -2575,18 +2568,12 @@ def is_separate_lm_head(model: torch.nn.Module) -> bool: if "model.safetensors.index.json" in os.listdir(dir_path): with open(os.path.join(dir_path, "model.safetensors.index.json")) as f: index_mapping = json.load(f) - if lm_head_name in index_mapping["weight_map"]: - return True - else: - return False + return lm_head_name in index_mapping["weight_map"] else: from safetensors import safe_open f = safe_open(os.path.join(dir_path, "model.safetensors"), framework="pt") - if lm_head_name in f.keys(): - return True - else: - return False + return lm_head_name in f.keys() def is_separate_tensor(model: torch.nn.Module, tensor_name: str) -> bool: @@ -2599,18 +2586,12 @@ def is_separate_tensor(model: torch.nn.Module, tensor_name: str) -> bool: if "model.safetensors.index.json" in os.listdir(dir_path): with open(os.path.join(dir_path, "model.safetensors.index.json")) as f: index_mapping = json.load(f) - if tensor_name in index_mapping["weight_map"]: - return True - else: - return False + return tensor_name in index_mapping["weight_map"] else: from safetensors import safe_open f = safe_open(os.path.join(dir_path, "model.safetensors"), framework="pt") - if tensor_name in f.keys(): - return True - else: - return False + return tensor_name in f.keys() def handle_generation_config(model: torch.nn.Module): @@ -2728,10 +2709,12 @@ def rename_weights_files(path: str, prefix="diffusion_pytorch_model"): # rename index.json idx = os.path.join(path, "model.safetensors.index.json") if os.path.exists(idx): - d = json.load(open(idx)) + with open(idx) as f: + d = json.load(f) d["weight_map"] = {k: v.replace("model-", prefix + "-") for k, v in d["weight_map"].items()} new_idx = os.path.join(path, f"{prefix}.safetensors.index.json") - json.dump(d, open(new_idx, "w"), indent=2) + with open(new_idx, "w") as f: + json.dump(d, f, indent=2) os.remove(idx) @@ -2888,7 +2871,7 @@ def __init__(self, embedding: "torch.nn.Embedding", devices: list) -> None: def weight(self): # Exposed so callers that read ``embedding.weight.device`` keep working (the forward # re-routes ids to the correct shard regardless of which device they arrive on). - return getattr(self, "shard_0") + return self.shard_0 def forward(self, input_ids: "torch.Tensor") -> "torch.Tensor": flat = input_ids.reshape(-1) diff --git a/auto_round/utils/model_free_utils.py b/auto_round/utils/model_free_utils.py index 8c03a6df04..69731a074f 100644 --- a/auto_round/utils/model_free_utils.py +++ b/auto_round/utils/model_free_utils.py @@ -26,9 +26,10 @@ import re import shutil import warnings +from collections.abc import Callable from dataclasses import fields from functools import lru_cache -from typing import Any, Callable, Optional, Union +from typing import Any import torch @@ -282,7 +283,7 @@ def quantize_weight_rtn( bits: int, group_size: int, sym: bool = True, - device: Optional[torch.device] = None, + device: torch.device | None = None, disable_opt_rtn: bool = True, *, packing: str, @@ -458,13 +459,13 @@ class _PatternMatcher: """Precompile ignore and layer-config patterns for shard processing.""" __slots__ = ( - "_ignore_re", - "_skip_re", - "_layer_config", - "_default_scheme", "_compiled_lc", + "_default_scheme", "_ignore_cache", + "_ignore_re", + "_layer_config", "_scheme_cache", + "_skip_re", ) def __init__( @@ -964,7 +965,7 @@ def _dequantize_with_device_fallback( return on_cpu() -def _normalize_scheme(scheme: Union[str, QuantizationScheme]) -> QuantizationScheme: +def _normalize_scheme(scheme: str | QuantizationScheme) -> QuantizationScheme: """Convert *scheme* to a :class:`QuantizationScheme` instance. Raises ``ValueError`` for unknown preset names and ``TypeError`` for @@ -1046,7 +1047,7 @@ def _fused_expert_layer_name(tensor_name: str) -> str: def _quantize_moe_fused_expert_weight( tensor_name: str, tensor: torch.Tensor, - matcher: "_PatternMatcher", + matcher: _PatternMatcher, device: str = "cpu", disable_opt_rtn: bool = False, ) -> tuple[str, dict[str, torch.Tensor], str | None, str | None]: @@ -1339,7 +1340,7 @@ def _declared_int_packing(sym: bool) -> str: def _quantize_single_tensor( tensor_name: str, tensor: torch.Tensor, - matcher: "_PatternMatcher", + matcher: _PatternMatcher, device: str = "cpu", quantize_func: Callable = quantize_weight_rtn, disable_opt_rtn: bool = False, @@ -1908,7 +1909,7 @@ def _dequant_mxfp_tensors( def _handle_mxfp_source_tensors( raw_tensors: dict[str, torch.Tensor], - matcher: "_PatternMatcher", + matcher: _PatternMatcher, source_state: dict[str, int] | None = None, device: str = "cpu", shard_name: str | None = None, @@ -2050,13 +2051,13 @@ def _dequant_fp8_tensors( def _process_shard( shard_path: str, - default_scheme: dict = None, - layer_config: dict = None, - ignore_patterns: list[str] = None, + default_scheme: dict | None = None, + layer_config: dict | None = None, + ignore_patterns: list[str] | None = None, device: str = "cpu", *, shard_name: str | None = None, - matcher: "_PatternMatcher | None" = None, + matcher: _PatternMatcher | None = None, fp8_block_size: list | None = None, model_type: str | None = None, source_quantization_config: dict | None = None, @@ -2153,9 +2154,7 @@ def _process_shard( # so the saved model exports them in full precision. preserved_prefixes: set[str] = set() for tname in raw_tensors: - if ( - tname.endswith(".weight") or tname.endswith(".weight_packed") or tname.endswith(".qweight") - ) and matcher.should_skip(tname): + if tname.endswith((".weight", ".weight_packed", ".qweight")) and matcher.should_skip(tname): preserved_prefixes.add(tname.rsplit(".", 1)[0]) preserved_tensors: dict[str, torch.Tensor] = {} @@ -2459,7 +2458,7 @@ def _is_weight_shard(fname: str) -> bool: """ if fname.endswith(".index.json"): return False - return fname.endswith(".safetensors") or fname.endswith(".bin") + return fname.endswith((".safetensors", ".bin")) # Keep old name as an alias for backward compatibility. @@ -2800,7 +2799,7 @@ def _build_mxfp_autoround_quantization_config( quantized_layers: list[str], ignored_layers: list[str], layer_config: dict | None = None, - block_name_to_quantize: Optional[str] = None, + block_name_to_quantize: str | None = None, ) -> dict: """Build an auto-round style quantization_config for MXFP4 / MXFP8. @@ -3071,7 +3070,7 @@ def _derive_dominant_int_scheme( layer_config=layer_config, default_scheme=fallback, ) - counter: "Counter[tuple]" = Counter() + counter: Counter[tuple] = Counter() for layer in quantized_layers: scheme = temp_matcher.resolve_scheme(f"{layer}.weight") if scheme is None: @@ -3109,7 +3108,7 @@ def _build_quantization_config( ignore_patterns: list[str], quantized_layers: list[str], ignored_layers: list[str], - block_name_to_quantize: Optional[str] = None, + block_name_to_quantize: str | None = None, format: str = "auto_round", ) -> dict: """Build a quantization_config dict compatible with auto-round format.""" @@ -3268,8 +3267,8 @@ def _build_quantization_config( def _apply_scheme_overrides( - scheme: Union[str, QuantizationScheme], - scheme_overrides: Optional[dict] = None, + scheme: str | QuantizationScheme, + scheme_overrides: dict | None = None, ) -> QuantizationScheme: """Return the effective scheme after applying non-None overrides.""" scheme_obj = copy.deepcopy(_normalize_scheme(scheme)) @@ -3285,7 +3284,7 @@ def _apply_scheme_overrides( def _validate_supported_scheme( scheme_obj: QuantizationScheme, - scheme_input: Union[str, QuantizationScheme], + scheme_input: str | QuantizationScheme, ) -> None: """Raise ``ValueError`` if *scheme_obj* is not supported by model-free. @@ -3378,8 +3377,8 @@ def _validate_supported_scheme( def is_model_free_supported_scheme( - scheme: Union[str, QuantizationScheme], - scheme_overrides: Optional[dict] = None, + scheme: str | QuantizationScheme, + scheme_overrides: dict | None = None, ) -> bool: """Return True if *scheme* can be quantized via model-free mode. @@ -3477,7 +3476,7 @@ def _validate_auto_scheme_options(auto_scheme: Any) -> str: def _convert_auto_scheme_layer_config( generated: dict[str, dict], - preferred_base_scheme: Union[str, QuantizationScheme, None] = None, + preferred_base_scheme: str | QuantizationScheme | None = None, ) -> tuple[QuantizationScheme, dict[str, dict], list[str]]: """Convert an AutoScheme-generated ``layer_config`` into model-free inputs. @@ -3496,7 +3495,7 @@ def _convert_auto_scheme_layer_config( scheme_keys = {f.name for f in fields(QuantizationScheme)} per_layer: dict[str, dict] = {} fp16_layers: list[str] = [] - counter: "Counter[tuple]" = Counter() + counter: Counter[tuple] = Counter() for name, cfg in generated.items(): if not isinstance(cfg, dict): @@ -3524,9 +3523,9 @@ def _convert_auto_scheme_layer_config( # the quantization kernels ("mxfp8" / "MXFP4" → "mx_fp"). if data_type_raw: dt_lower = data_type_raw.lower() - if dt_lower.startswith("mxfp") or dt_lower.startswith("mx_fp"): + if dt_lower.startswith(("mxfp", "mx_fp")): clean["data_type"] = "mx_fp" - elif dt_lower.startswith("nvfp") or dt_lower.startswith("nv_fp"): + elif dt_lower.startswith(("nvfp", "nv_fp")): clean["data_type"] = "nv_fp" if bits >= 16: diff --git a/auto_round/utils/offload.py b/auto_round/utils/offload.py index 35d1f9ae2c..8aecc7d11d 100644 --- a/auto_round/utils/offload.py +++ b/auto_round/utils/offload.py @@ -51,7 +51,7 @@ import tempfile from collections import defaultdict from functools import partial -from typing import Any, Optional, Union +from typing import Any import torch @@ -190,7 +190,7 @@ def _clear_module_weights( # ===================================================================== -def _resolve_model_dir(model_dir: str, revision: Optional[str] = None) -> str: +def _resolve_model_dir(model_dir: str, revision: str | None = None) -> str: """Resolve a model name/path to a local directory containing weight files.""" if os.path.isdir(model_dir): return model_dir @@ -371,11 +371,11 @@ def __init__( self, enabled: bool = True, mode: str = "offload", - model_dir: Optional[str] = None, + model_dir: str | None = None, offload_dir_prefix: str = "ar_offload", cache_numel: bool = False, retain_saved_entries: bool = False, - model_revision: Optional[str] = None, + model_revision: str | None = None, ): from auto_round import envs @@ -391,7 +391,7 @@ def __init__( self.retain_saved_entries = retain_saved_entries # Disk state (offload mode) - self._tempdir: Optional[str] = None + self._tempdir: str | None = None self._saved: dict[str, dict] = {} # name -> {"save_path": str} # Cached weight map for clean mode (avoids repeated disk I/O) @@ -399,12 +399,12 @@ def __init__( # Hook state (for add_offload_hooks/remove_offload_hooks transparent offloading) self._hook_handles: list = [] - self._model_ref: Optional[torch.nn.Module] = None + self._model_ref: torch.nn.Module | None = None self._module_names: list[str] = [] - self._last_loaded: Optional[str] = None + self._last_loaded: str | None = None # Ensure-style state (for wrapping loops) - self._current_loaded: Optional[str] = None + self._current_loaded: str | None = None # ------------------------------------------------------------------ # Context manager @@ -426,7 +426,7 @@ def __del__(self): def __call__( self, model: torch.nn.Module, - names: Union[str, list[str], list[list[str]]], + names: str | list[str] | list[list[str]], *, skip_if_saved: bool = False, overwrite: bool = False, @@ -450,7 +450,7 @@ def __call__( def offload( self, model: torch.nn.Module, - names: Union[str, list[str], list[list[str]]], + names: str | list[str] | list[list[str]], *, skip_if_saved: bool = False, overwrite: bool = False, @@ -508,7 +508,7 @@ def offload( logger.info(f"offload done, freed {total_gb:.2f} GB") return total_gb - def _check_disk_space(self, model: torch.nn.Module, names: Union[str, list[str], list[list[str]]]) -> bool: + def _check_disk_space(self, model: torch.nn.Module, names: str | list[str] | list[list[str]]) -> bool: """Check whether there is enough disk space to offload the given modules. Args: @@ -576,7 +576,7 @@ def _offload( self._save_to_disk(name, module) self._clear(module, block_name=name) - def reload(self, model: torch.nn.Module, names: Union[str, list[str], None] = None) -> None: + def reload(self, model: torch.nn.Module, names: str | list[str] | None = None) -> None: """Reload previously offloaded module(s). For ``"offload"`` mode: loads from the temp directory, then @@ -692,7 +692,7 @@ def add_offload_hooks(self, model: torch.nn.Module, names: list[str]) -> None: clear_memory() logger.info("module weights cleared") - def remove_offload_hooks(self, model: torch.nn.Module, names: Optional[list[str]] = None) -> None: + def remove_offload_hooks(self, model: torch.nn.Module, names: list[str] | None = None) -> None: """Remove hooks and reload all managed modules. Args: @@ -1007,7 +1007,7 @@ def _needs_loading(module: torch.nn.Module) -> bool: return False @staticmethod - def _flatten_names(names: Union[list[str], list[list[str]]]) -> list[str]: + def _flatten_names(names: list[str] | list[list[str]]) -> list[str]: """Flatten a potentially nested list of names.""" flat = [] for item in names: diff --git a/auto_round/utils/path_safety.py b/auto_round/utils/path_safety.py index c337aaf01b..f2220f3a47 100644 --- a/auto_round/utils/path_safety.py +++ b/auto_round/utils/path_safety.py @@ -38,12 +38,12 @@ import os from pathlib import Path, PureWindowsPath -from typing import Any, Dict +from typing import Any __all__ = [ "UnsafeCheckpointPathError", - "sanitize_shard_name", "resolve_within_directory", + "sanitize_shard_name", "validate_weight_map", ] @@ -129,7 +129,7 @@ def validate_weight_map( base_dir: str | os.PathLike, *, index_path: str | os.PathLike | None = None, -) -> Dict[str, str]: +) -> dict[str, str]: """Validate every shard reference in an artifact-provided ``weight_map``. Returns a new dict with the same keys and the same relative shard names, so it @@ -144,7 +144,7 @@ def validate_weight_map( f"{label}: 'weight_map' must be a JSON object mapping tensor names to shards, " f"got {type(weight_map).__name__}" ) - validated: Dict[str, str] = {} + validated: dict[str, str] = {} for tensor_name, shard_name in weight_map.items(): resolve_within_directory(base_dir, shard_name, origin=f"{label}[{tensor_name}]") validated[tensor_name] = shard_name diff --git a/auto_round/utils/resume.py b/auto_round/utils/resume.py index 9299cbb6c0..79872bb039 100644 --- a/auto_round/utils/resume.py +++ b/auto_round/utils/resume.py @@ -28,7 +28,6 @@ import os import tempfile from pathlib import Path -from typing import Optional import torch @@ -172,7 +171,7 @@ def clear(self) -> None: def compute_run_signature( - model_dir: Optional[str], + model_dir: str | None, scheme_desc: str, dataset_desc: str, nsamples: int, diff --git a/auto_round/utils/weight_handler.py b/auto_round/utils/weight_handler.py old mode 100755 new mode 100644 index abfca6ba41..c14aad463d --- a/auto_round/utils/weight_handler.py +++ b/auto_round/utils/weight_handler.py @@ -59,10 +59,10 @@ def convert_layer(self, layer, dtype, device, to_cpu): ... import os from abc import ABC, abstractmethod +from collections.abc import Callable from contextlib import ContextDecorator from dataclasses import fields from enum import Enum, auto -from typing import Callable, Dict, Optional, Set, Type import psutil import torch @@ -179,7 +179,6 @@ def detect_layer(self, module: torch.nn.Module) -> bool: Returns: True if the module is of this weight type, False otherwise. """ - pass def attach_weight_shape(self, module: torch.nn.Module): """Optional helper to attach weight shape information to the module for detection.""" @@ -209,12 +208,11 @@ def convert_layer( Returns: A new high-precision layer with dequantized weights. """ - pass # --- Handler Registry --- -_WEIGHT_TYPE_HANDLERS: Dict[ModuleWeightType, WeightTypeHandler] = {} +_WEIGHT_TYPE_HANDLERS: dict[ModuleWeightType, WeightTypeHandler] = {} def register_weight_type_handler(weight_type: ModuleWeightType): @@ -232,7 +230,7 @@ class MXFP4Handler(WeightTypeHandler): ... """ - def decorator(handler_cls: Type[WeightTypeHandler]): + def decorator(handler_cls: type[WeightTypeHandler]): if not issubclass(handler_cls, WeightTypeHandler): raise TypeError(f"Handler {handler_cls.__name__} must be a subclass of WeightTypeHandler") _WEIGHT_TYPE_HANDLERS[weight_type] = handler_cls() @@ -241,7 +239,7 @@ def decorator(handler_cls: Type[WeightTypeHandler]): return decorator -def get_handler(weight_type: ModuleWeightType) -> Optional[WeightTypeHandler]: +def get_handler(weight_type: ModuleWeightType) -> WeightTypeHandler | None: """Get the registered handler for a weight type. Args: @@ -253,7 +251,7 @@ def get_handler(weight_type: ModuleWeightType) -> Optional[WeightTypeHandler]: return _WEIGHT_TYPE_HANDLERS.get(weight_type) -def get_all_handlers() -> Dict[ModuleWeightType, WeightTypeHandler]: +def get_all_handlers() -> dict[ModuleWeightType, WeightTypeHandler]: """Get all registered weight type handlers. Returns: @@ -265,7 +263,7 @@ def get_all_handlers() -> Dict[ModuleWeightType, WeightTypeHandler]: # ============================================================================ # Section 2: PUBLIC API - Detection and Conversion Functions # ============================================================================ -def detect_weight_type(module: torch.nn.Module) -> Optional[ModuleWeightType]: +def detect_weight_type(module: torch.nn.Module) -> ModuleWeightType | None: """Detect the weight type of a module or model. First checks if the module itself has a quantized_weight_type attribute. @@ -290,7 +288,7 @@ def detect_weight_type(module: torch.nn.Module) -> Optional[ModuleWeightType]: # --- Model Marking Functions --- -def check_and_mark_quantized_module(model: torch.nn.Module) -> Set[ModuleWeightType]: +def check_and_mark_quantized_module(model: torch.nn.Module) -> set[ModuleWeightType]: """Check if model contains quantized layers and mark them accordingly. This function scans the model (including the model itself) for quantized layers using @@ -303,7 +301,7 @@ def check_and_mark_quantized_module(model: torch.nn.Module) -> Set[ModuleWeightT Returns: A set of detected ModuleWeightType values. Empty set if no quantized layers found. """ - detected_types: Set[ModuleWeightType] = set() + detected_types: set[ModuleWeightType] = set() for weight_type, handler in _WEIGHT_TYPE_HANDLERS.items(): # Check model itself first if handler.detect_layer(model): @@ -332,7 +330,7 @@ def check_and_mark_quantized_module(model: torch.nn.Module) -> Set[ModuleWeightT return detected_types -def is_quantized_input_module(model: torch.nn.Module) -> Optional[ModuleWeightType]: +def is_quantized_input_module(model: torch.nn.Module) -> ModuleWeightType | None: """Check if a model has quantized input weights and return the weight type. This traverses all submodules to check for the `quantized_weight_type` attribute @@ -485,7 +483,7 @@ def _pad_block_fp8_weight_naive( @with_thread_limits() def _dequant_fp8_linear_weight( - weight: torch.Tensor, weight_scale: torch.Tensor, block_size: list = None, data_type: str = None + weight: torch.Tensor, weight_scale: torch.Tensor, block_size: list | None = None, data_type: str | None = None ) -> torch.Tensor: """Core dequantization logic for block-wise FP8 weights.""" dtype = torch.bfloat16 diff --git a/auto_round/wrapper.py b/auto_round/wrapper.py index 66897972e4..ac59c2916b 100644 --- a/auto_round/wrapper.py +++ b/auto_round/wrapper.py @@ -95,7 +95,7 @@ def __init__( enable_norm_bias_tuning (bool): Whether to enable normalization and tuning for the bias term. device (str): The computation device, such as 'cpu' or 'cuda'. """ - super(WrapperLinear, self).__init__() + super().__init__() self.orig_layer = orig_layer self.orig_layer.iters = kwargs.pop("iters", 200) self.disable_opt_rtn = disable_opt_rtn @@ -112,7 +112,7 @@ def __init__( from auto_round.data_type.nvfp import calculate_gparam weight_global_scale = calculate_gparam(self.orig_layer.weight, self.orig_layer.group_size) - setattr(self, "weight_global_scale", weight_global_scale) + self.weight_global_scale = weight_global_scale self.weight_global_scale = self.weight_global_scale.to(self.orig_layer.weight.device) if hasattr(self.orig_layer, "scale_dtype") and self.orig_layer.scale_dtype == torch.float32: self.q_scale_thresh = 1e-8 @@ -332,7 +332,7 @@ def _qdq_weight(self, value, min_scale, max_scale): weight_q = weight_q.t() return weight_q, scale, zp - def _qdq_act(self, x, act_min_scale=torch.tensor(1.0), act_max_scale=torch.tensor(1.0), act_max=None): + def _qdq_act(self, x, act_min_scale=None, act_max_scale=None, act_max=None): """Quantizes and dequantizes activations. Args: @@ -343,6 +343,10 @@ def _qdq_act(self, x, act_min_scale=torch.tensor(1.0), act_max_scale=torch.tenso Returns: tuple: Quantized activation, scale, and zero point. """ + if act_min_scale is None: + act_min_scale = torch.tensor(1.0) + if act_max_scale is None: + act_max_scale = torch.tensor(1.0) act_max_scale.data.clamp_(0, 1.0) act_min_scale.data.clamp_(0, 1.0) env_act_scale = envs.AR_ACT_SCALE # fixed activation ratio,prioritize to use this one if set @@ -608,7 +612,7 @@ def forward(self, x): class WrapperWALayer(torch.nn.Module): def __init__(self, orig_layer, enable_torch_compile=True, device="cpu"): - super(WrapperWALayer, self).__init__() + super().__init__() self.orig_layer = orig_layer self.enable_torch_compile = enable_torch_compile self.device = device @@ -692,7 +696,7 @@ class WrapperLayerNorm(torch.nn.Module): """ def __init__(self, orig_layer, bit=4, group_size=-1, device="cpu"): - super(WrapperLayerNorm, self).__init__() + super().__init__() self.orig_layer = orig_layer self.bits = bit self.group_size = group_size @@ -743,7 +747,7 @@ class WrapperLlamaNorm(torch.nn.Module): """ def __init__(self, orig_layer, bit=4, group_size=-1, device="cpu"): - super(WrapperLlamaNorm, self).__init__() + super().__init__() self.orig_layer = orig_layer self.bits = bit self.group_size = group_size @@ -803,7 +807,7 @@ class WrapperMultiblock(torch.nn.Module): """ def __init__(self, module_list): - super(WrapperMultiblock, self).__init__() + super().__init__() self.layers = torch.nn.ModuleList(module_list) def forward(self, x, *args, **kwargs): @@ -811,7 +815,7 @@ def forward(self, x, *args, **kwargs): for idx, decoder_layer in enumerate(self.layers): layer_outputs = decoder_layer(hidden_states, *args, **kwargs) hidden_states = layer_outputs - if isinstance(hidden_states, tuple) or isinstance(hidden_states, list): + if isinstance(hidden_states, (tuple, list)): hidden_states = layer_outputs[0] return hidden_states @@ -857,7 +861,7 @@ def wrapper_block( elif enable_norm_bias_tuning: if "norm" in m.__class__.__name__.lower(): - if m.__class__.__name__ in NORM_MAPPING.keys(): + if m.__class__.__name__ in NORM_MAPPING: wrapper_layer_class = NORM_MAPPING[m.__class__.__name__] new_m = wrapper_layer_class(m, device=device) set_module(block, n, new_m) diff --git a/auto_round_extension/ark/auto_round_kernel/__init__.py b/auto_round_extension/ark/auto_round_kernel/__init__.py index be679a144a..7d484963d4 100644 --- a/auto_round_extension/ark/auto_round_kernel/__init__.py +++ b/auto_round_extension/ark/auto_round_kernel/__init__.py @@ -455,7 +455,7 @@ def get_lib(A: torch.Tensor): # A: mxk, B: nxk, bias: n or [1, n] -def matmul_sycl_tla(A: torch.Tensor, B: torch.Tensor, bias: Optional[torch.Tensor] = None): +def matmul_sycl_tla(A: torch.Tensor, B: torch.Tensor, bias: torch.Tensor | None = None): if A.device.type != "xpu" or B.device.type != "xpu": raise NotImplementedError("matmul_sycl_tla is only supported on XPU") if A.ndim != 2 or B.ndim != 2: @@ -1724,7 +1724,7 @@ def forward( num_heads_kv: int | None = None, *, is_causal: bool = False, - scale: Optional[float] = None, + scale: float | None = None, tensor_layout: str = "HND", ) -> torch.Tensor: del num_heads_kv @@ -1757,7 +1757,7 @@ def ark_cpu_packed_kv_descriptor( def ark_cpu_packed_kv_alloc_from_descriptor( descriptor, *, - dtype: Optional[torch.dtype] = None, + dtype: torch.dtype | None = None, device: str = "cpu", ) -> tuple[torch.Tensor, torch.Tensor]: if cpu_lib is None or not hasattr(cpu_lib, "ark_cpu_packed_kv_elems_desc"): @@ -1798,10 +1798,10 @@ def ark_cpu_packed_kv_alloc( def ark_cpu_packed_kv_info( - batch: Optional[int] = None, - num_heads_kv: Optional[int] = None, - capacity: Optional[int] = None, - head_dim: Optional[int] = None, + batch: int | None = None, + num_heads_kv: int | None = None, + capacity: int | None = None, + head_dim: int | None = None, *, dtype: torch.dtype = torch.float16, descriptor=None, @@ -2027,7 +2027,7 @@ def ark_cpu_bestla_sdpa_packed( num_heads_kv: int, *, is_causal: bool = False, - scale: Optional[float] = None, + scale: float | None = None, tensor_layout: str = "HND", ) -> torch.Tensor: """Internal/experimental BestLA mixed-precision SDPA over a packed K/V cache. @@ -2081,7 +2081,7 @@ def ark_cpu_bestla_sdpa_packed_from_descriptor( seq_len_kv: int, *, is_causal: bool = False, - scale: Optional[float] = None, + scale: float | None = None, tensor_layout: str = "HND", ) -> torch.Tensor: """Descriptor-based internal/experimental packed BestLA SDPA forward.""" @@ -2122,7 +2122,7 @@ def sageattn( v: torch.Tensor, tensor_layout: str = "HND", is_causal: bool = False, - sm_scale: Optional[float] = None, + sm_scale: float | None = None, return_lse: bool = False, kernel: str = "v1_pvhalf", **kwargs, @@ -2452,8 +2452,8 @@ def moe_gemm_decode( weights: torch.Tensor, num_tokens_per_expert: torch.Tensor, *, - scales: Optional[torch.Tensor] = None, - zeros: Optional[torch.Tensor] = None, + scales: torch.Tensor | None = None, + zeros: torch.Tensor | None = None, weight_bits: int = 4, group_size: int = 128, asym: bool = False, @@ -2605,8 +2605,8 @@ def _validate_moe_quant_args( weights: torch.Tensor, num_tokens_per_expert: torch.Tensor, *, - scales: Optional[torch.Tensor], - zeros: Optional[torch.Tensor], + scales: torch.Tensor | None, + zeros: torch.Tensor | None, weight_bits: int, group_size: int, asym: bool, @@ -2763,17 +2763,17 @@ class MoeSymmetricGemm: weights_ptr: int scales_ptr: int zeros_ptr: int = 0 - decode_expert_id_per_token: Optional[torch.Tensor] = None + decode_expert_id_per_token: torch.Tensor | None = None @classmethod def prepare( cls, weights: torch.Tensor, - scales: Optional[torch.Tensor], + scales: torch.Tensor | None, *, weight_bits: int = 4, group_size: int = 128, - activation_dtype: Optional[torch.dtype] = None, + activation_dtype: torch.dtype | None = None, max_decode_tokens: int = 0, ) -> "MoeSymmetricGemm": """Prepare the symmetric quantized MoE GEMM fast path. @@ -2876,8 +2876,8 @@ def decode( activations: torch.Tensor, num_tokens_per_expert: torch.Tensor, *, - outputs: Optional[torch.Tensor] = None, - expert_id_per_token: Optional[torch.Tensor] = None, + outputs: torch.Tensor | None = None, + expert_id_per_token: torch.Tensor | None = None, ) -> torch.Tensor: total_tokens = int(activations.shape[0]) if outputs is None: @@ -2915,8 +2915,8 @@ def prefill( activations: torch.Tensor, num_tokens_per_expert: torch.Tensor, *, - outputs: Optional[torch.Tensor] = None, - workspace: Optional[torch.Tensor] = None, + outputs: torch.Tensor | None = None, + workspace: torch.Tensor | None = None, ) -> torch.Tensor: total_tokens = int(activations.shape[0]) if outputs is None: @@ -2953,10 +2953,10 @@ def moe( num_tokens_per_expert: torch.Tensor, *, phase: str = "auto", - decode_threshold: Optional[int] = None, - outputs: Optional[torch.Tensor] = None, - expert_id_per_token: Optional[torch.Tensor] = None, - workspace: Optional[torch.Tensor] = None, + decode_threshold: int | None = None, + outputs: torch.Tensor | None = None, + expert_id_per_token: torch.Tensor | None = None, + workspace: torch.Tensor | None = None, ) -> torch.Tensor: if phase not in _MOE_VALID_PHASES: raise ValueError(f"phase must be one of {_MOE_VALID_PHASES}, got {phase!r}") @@ -2985,10 +2985,10 @@ def apply( num_tokens_per_expert: torch.Tensor, *, phase: str = "auto", - decode_threshold: Optional[int] = None, - outputs: Optional[torch.Tensor] = None, - expert_id_per_token: Optional[torch.Tensor] = None, - workspace: Optional[torch.Tensor] = None, + decode_threshold: int | None = None, + outputs: torch.Tensor | None = None, + expert_id_per_token: torch.Tensor | None = None, + workspace: torch.Tensor | None = None, ) -> torch.Tensor: return self.moe( activations, @@ -3006,7 +3006,7 @@ def moe_gemm( weights: torch.Tensor, num_tokens_per_expert: torch.Tensor, *, - scales: Optional[torch.Tensor] = None, + scales: torch.Tensor | None = None, ) -> torch.Tensor: """MOE GEMM (Mixture of Experts Grouped GEMM). @@ -3258,12 +3258,12 @@ def moe_gemm_prefill( weights: torch.Tensor, num_tokens_per_expert: torch.Tensor, *, - scales: Optional[torch.Tensor] = None, - zeros: Optional[torch.Tensor] = None, + scales: torch.Tensor | None = None, + zeros: torch.Tensor | None = None, weight_bits: int = 4, group_size: int = 128, asym: bool = False, - scale_scheme: Optional[str] = None, + scale_scheme: str | None = None, ) -> torch.Tensor: """MoE Grouped GEMM optimized for the prefill phase, supporting all weight encodings of ``moe_gemm_decode`` (FP16/BF16, INT8 sym/asym, INT4 sym/asym, @@ -3583,13 +3583,13 @@ def moe( weights: torch.Tensor, num_tokens_per_expert: torch.Tensor, *, - scales: Optional[torch.Tensor] = None, - zeros: Optional[torch.Tensor] = None, + scales: torch.Tensor | None = None, + zeros: torch.Tensor | None = None, weight_bits: int = 4, group_size: int = 128, asym: bool = False, phase: str = "auto", - decode_threshold: Optional[int] = None, + decode_threshold: int | None = None, ) -> torch.Tensor: """Unified MoE GEMM entry point that dispatches to decode or prefill. diff --git a/auto_round_extension/ark/auto_round_kernel/qlinear.py b/auto_round_extension/ark/auto_round_kernel/qlinear.py index 749d2db1c2..82beecbfe9 100644 --- a/auto_round_extension/ark/auto_round_kernel/qlinear.py +++ b/auto_round_extension/ark/auto_round_kernel/qlinear.py @@ -89,7 +89,7 @@ def convert_dtype_torch2str(dtype): elif isinstance(dtype, str) and dtype in ["int8", "fp32", "fp16", "bf16"]: return dtype else: - assert False, "Unsupported pytorch dtype {} to str dtype".format(dtype) + assert False, f"Unsupported pytorch dtype {dtype} to str dtype" class QuantLinear(nn.Module): @@ -164,12 +164,9 @@ def __init__( self.bias = None def extra_repr(self) -> str: - return "in_features={}, out_features={}, bias={}, w_bit={}, group_size={}".format( - self.infeatures, - self.outfeatures, - self.bias is not None, - self.bits, - self.group_size, + return ( + f"in_features={self.infeatures}, out_features={self.outfeatures}, bias={self.bias is not None}, " + f"w_bit={self.bits}, group_size={self.group_size}" ) def post_init(self): @@ -563,14 +560,10 @@ def forward(self, x: torch.Tensor): return outputs.to(raw_input_dtype).view(out_shape) def extra_repr(self) -> str: - return "in_features={}, out_features={}, bias={}, bits={}, " "group_size={}, data_type={}".format( - self.infeatures, - self.outfeatures, - self.bias is not None, - self.bits, - self.group_size, - self.data_type, + return ( + f"in_features={self.infeatures}, out_features={self.outfeatures}, bias={self.bias is not None}, bits={self.bits}, " + f"group_size={self.group_size}, data_type={self.data_type}" ) -__all__ = ["QuantLinear", "QuantLinearGPTQ", "QuantLinearAWQ", "QuantLinearFP8"] +__all__ = ["QuantLinear", "QuantLinearAWQ", "QuantLinearFP8", "QuantLinearGPTQ"] diff --git a/auto_round_extension/ark/auto_round_kernel/sparge_preprocess_triton.py b/auto_round_extension/ark/auto_round_kernel/sparge_preprocess_triton.py index 2fe1091f80..9672bd7acc 100644 --- a/auto_round_extension/ark/auto_round_kernel/sparge_preprocess_triton.py +++ b/auto_round_extension/ark/auto_round_kernel/sparge_preprocess_triton.py @@ -4,7 +4,8 @@ from __future__ import annotations import logging -from typing import Any, Callable +from collections.abc import Callable +from typing import Any import torch import triton diff --git a/auto_round_extension/ark/auto_round_kernel/sparse_attention.py b/auto_round_extension/ark/auto_round_kernel/sparse_attention.py index e45febf887..9e86faae3a 100644 --- a/auto_round_extension/ark/auto_round_kernel/sparse_attention.py +++ b/auto_round_extension/ark/auto_round_kernel/sparse_attention.py @@ -6,7 +6,7 @@ import os import warnings from dataclasses import dataclass -from typing import Any, Optional +from typing import Any import torch @@ -417,7 +417,7 @@ def sageattn( v: torch.Tensor, tensor_layout: str = "HND", is_causal: bool = False, - sm_scale: Optional[float] = None, + sm_scale: float | None = None, return_lse: bool = False, kernel: str = "v1_pvhalf", **kwargs, @@ -499,7 +499,7 @@ def _normalize_sparse_mask( def _normalize_per_head_hparam( - value: float | int | torch.Tensor, + value: float | torch.Tensor, num_heads: int, device: torch.device, name: str, @@ -825,7 +825,7 @@ def _prefix_protection_requested() -> bool: def _get_explicit_protected_prefix( - ctx: "_SpargePreprocessContext", + ctx: _SpargePreprocessContext, ) -> tuple[int, int]: protected_tokens = _get_protected_kv_tokens() protected_sparse_blocks = _get_protected_kv_blocks() @@ -847,7 +847,7 @@ def _get_explicit_protected_prefix( def _get_prefix_protection_blocks( - ctx: "_SpargePreprocessContext", + ctx: _SpargePreprocessContext, *, raw_block_count: int, tile_block_count: int, diff --git a/auto_round_extension/ark/auto_round_kernel/utils.py b/auto_round_extension/ark/auto_round_kernel/utils.py index 17c4642606..c88738bf00 100644 --- a/auto_round_extension/ark/auto_round_kernel/utils.py +++ b/auto_round_extension/ark/auto_round_kernel/utils.py @@ -15,14 +15,14 @@ import logging import re import subprocess -from functools import lru_cache +from functools import cache import torch logger = logging.getLogger(__name__) -@lru_cache(maxsize=None) +@cache def is_oneapi_ge_2026() -> bool: try: output = subprocess.check_output( @@ -41,7 +41,7 @@ def is_oneapi_ge_2026() -> bool: B70_IDENTIFIERS = (B70_DEVICE_ID, "b70") -@lru_cache(maxsize=None) +@cache def is_b70(device: int = 0) -> bool: try: pro = torch.xpu.get_device_properties(device) diff --git a/auto_round_extension/ark/auto_round_kernel/xpu_loader.py b/auto_round_extension/ark/auto_round_kernel/xpu_loader.py index 6a2ec2cfe8..1eea91286b 100644 --- a/auto_round_extension/ark/auto_round_kernel/xpu_loader.py +++ b/auto_round_extension/ark/auto_round_kernel/xpu_loader.py @@ -15,8 +15,8 @@ import importlib.util import sys import sysconfig +from collections.abc import Iterable from pathlib import Path -from typing import Iterable _DEFAULT_MODULE_NAME = "auto_round_kernel._local.auto_round_kernel_xpu" diff --git a/auto_round_extension/ark/benchmarks/bench_mxfp4_hadamard.py b/auto_round_extension/ark/benchmarks/bench_mxfp4_hadamard.py old mode 100644 new mode 100755 index 04af7f551a..b5852133cd --- a/auto_round_extension/ark/benchmarks/bench_mxfp4_hadamard.py +++ b/auto_round_extension/ark/benchmarks/bench_mxfp4_hadamard.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # Copyright (C) 2026 Intel Corporation # SPDX-License-Identifier: Apache-2.0 diff --git a/auto_round_extension/ark/benchmarks/bench_sparse_topk.py b/auto_round_extension/ark/benchmarks/bench_sparse_topk.py old mode 100644 new mode 100755 index 9efe91e070..e77ed2294b --- a/auto_round_extension/ark/benchmarks/bench_sparse_topk.py +++ b/auto_round_extension/ark/benchmarks/bench_sparse_topk.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (C) 2026 Intel Corporation # # SPDX-License-Identifier: Apache-2.0 diff --git a/auto_round_extension/ark/examples/run_flux.py b/auto_round_extension/ark/examples/run_flux.py index 57fcf902d6..c73cddae8e 100644 --- a/auto_round_extension/ark/examples/run_flux.py +++ b/auto_round_extension/ark/examples/run_flux.py @@ -420,23 +420,9 @@ def prepare_block_benchmark_state(): img_ids = img_ids[0] image_rotary_emb = transformer.pos_embed(torch.cat((text_ids, img_ids), dim=0)) - with torch.no_grad(): - with transformer.cache_context("cond"): - if benchmark_block_kind == "single": - for block in joint_blocks: - encoder_hidden_states, hidden_states = block( - hidden_states=hidden_states, - encoder_hidden_states=encoder_hidden_states, - temb=temb, - image_rotary_emb=image_rotary_emb, - joint_attention_kwargs=joint_attention_kwargs, - ) - preceding_blocks = ( - joint_blocks[:benchmark_block_index] - if benchmark_block_kind == "joint" - else single_blocks[:benchmark_block_index] - ) - for block in preceding_blocks: + with torch.no_grad(), transformer.cache_context("cond"): + if benchmark_block_kind == "single": + for block in joint_blocks: encoder_hidden_states, hidden_states = block( hidden_states=hidden_states, encoder_hidden_states=encoder_hidden_states, @@ -444,6 +430,19 @@ def prepare_block_benchmark_state(): image_rotary_emb=image_rotary_emb, joint_attention_kwargs=joint_attention_kwargs, ) + preceding_blocks = ( + joint_blocks[:benchmark_block_index] + if benchmark_block_kind == "joint" + else single_blocks[:benchmark_block_index] + ) + for block in preceding_blocks: + encoder_hidden_states, hidden_states = block( + hidden_states=hidden_states, + encoder_hidden_states=encoder_hidden_states, + temb=temb, + image_rotary_emb=image_rotary_emb, + joint_attention_kwargs=joint_attention_kwargs, + ) return { "block": target_block, diff --git a/auto_round_extension/ark/setup.py b/auto_round_extension/ark/setup.py index 13c030f47d..6857ca8c8e 100644 --- a/auto_round_extension/ark/setup.py +++ b/auto_round_extension/ark/setup.py @@ -71,9 +71,7 @@ def detect_oneapi_version(): Returns a string like '2025.3' or None if detection fails. """ try: - result = subprocess.run( - ["icx", "--version"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True, check=True - ) + result = subprocess.run(["icx", "--version"], capture_output=True, text=True, check=True) match = re.search(r"Compiler\s+(\d{4}\.\d+)", result.stdout) if match: return match.group(1) @@ -354,7 +352,7 @@ def run(self): version=get_build_version(), description="Auto Round Kernel binary package", author_email="yu.luo@intel.com", - long_description=open("README.md", "r", encoding="utf-8").read(), + long_description=Path("README.md").read_text(encoding="utf-8"), long_description_content_type="text/markdown", keywords="quantization,auto-around,LLM,kernel", license="Apache 2.0", diff --git a/auto_round_extension/ark/test/bench_ark_cpu_sdpa.py b/auto_round_extension/ark/test/bench_ark_cpu_sdpa.py old mode 100644 new mode 100755 index c0c0fd612b..1c68a59a86 --- a/auto_round_extension/ark/test/bench_ark_cpu_sdpa.py +++ b/auto_round_extension/ark/test/bench_ark_cpu_sdpa.py @@ -201,8 +201,10 @@ def torch_call(): ark_ms, torch_ms = _measure_pair(ark_call, torch_call, warmup, runs) return ( torch_ms / ark_ms, - f"{shape.name:<22}{route.name:<18}{shape.label:<24}{resolved_name:<12}" - f"{ark_ms:>10.3f}{torch_ms:>14.3f}{torch_ms / ark_ms:>10.2f}x", + ( + f"{shape.name:<22}{route.name:<18}{shape.label:<24}{resolved_name:<12}" + f"{ark_ms:>10.3f}{torch_ms:>14.3f}{torch_ms / ark_ms:>10.2f}x" + ), ) diff --git a/auto_round_extension/ark/test/conftest.py b/auto_round_extension/ark/test/conftest.py index 66b2387755..dcb8b5a870 100644 --- a/auto_round_extension/ark/test/conftest.py +++ b/auto_round_extension/ark/test/conftest.py @@ -1,5 +1,3 @@ -#!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2026 Intel Corporation # diff --git a/auto_round_extension/ark/test/test_flash_attn.py b/auto_round_extension/ark/test/test_flash_attn.py old mode 100644 new mode 100755 index 27aad54260..07f0633b68 --- a/auto_round_extension/ark/test/test_flash_attn.py +++ b/auto_round_extension/ark/test/test_flash_attn.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2026 Intel Corporation # @@ -55,7 +54,7 @@ def reference_attention(Q, K, V, scale, is_causal=True): scale=scale, attn_mask=None, is_causal=is_causal, - enable_gqa=True if K.shape[1] != Q.shape[1] else False, + enable_gqa=K.shape[1] != Q.shape[1], ) return ref diff --git a/auto_round_extension/ark/test/test_matmul.py b/auto_round_extension/ark/test/test_matmul.py old mode 100644 new mode 100755 index 86f06a8ee1..bcad260b55 --- a/auto_round_extension/ark/test/test_matmul.py +++ b/auto_round_extension/ark/test/test_matmul.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2023 Intel Corporation # diff --git a/auto_round_extension/ark/test/test_moe.py b/auto_round_extension/ark/test/test_moe.py old mode 100644 new mode 100755 index e4d7c39e7f..3b1603d4d6 --- a/auto_round_extension/ark/test/test_moe.py +++ b/auto_round_extension/ark/test/test_moe.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2026 Intel Corporation # diff --git a/auto_round_extension/ark/test/test_moe_decode_perf.py b/auto_round_extension/ark/test/test_moe_decode_perf.py old mode 100644 new mode 100755 index ce32f765f1..6554abf83d --- a/auto_round_extension/ark/test/test_moe_decode_perf.py +++ b/auto_round_extension/ark/test/test_moe_decode_perf.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2026 Intel Corporation # @@ -115,16 +114,13 @@ def _decode_skip_reason() -> str: # Surface diagnostics on collection so the user always sees why the suite # would skip, without having to add extra flags. +_xpu_lib_state = "loaded" if ark.xpu_lib is not None else "None" +_has_decode = hasattr(ark.xpu_lib, "moe_gemm_decode") if ark.xpu_lib is not None else False print( - "[moe-decode-perf] xpu_available=%s xpu_lib=%s has_moe_gemm_decode=%s" - % ( - _xpu_available(), - "loaded" if ark.xpu_lib is not None else "None", - hasattr(ark.xpu_lib, "moe_gemm_decode") if ark.xpu_lib is not None else False, - ) + f"[moe-decode-perf] xpu_available={_xpu_available()} xpu_lib={_xpu_lib_state} has_moe_gemm_decode={_has_decode}" ) if _DECODE_SKIP: - print("[moe-decode-perf] suite will SKIP. reason: %s" % _DECODE_SKIP) + print(f"[moe-decode-perf] suite will SKIP. reason: {_DECODE_SKIP}") # --------------------------------------------------------------------------- diff --git a/auto_round_extension/ark/test/test_moe_prefill_accuracy.py b/auto_round_extension/ark/test/test_moe_prefill_accuracy.py old mode 100644 new mode 100755 index e9f7ca859f..c6a46c6a05 --- a/auto_round_extension/ark/test/test_moe_prefill_accuracy.py +++ b/auto_round_extension/ark/test/test_moe_prefill_accuracy.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2026 Intel Corporation # diff --git a/auto_round_extension/ark/test/test_moe_prefill_perf.py b/auto_round_extension/ark/test/test_moe_prefill_perf.py old mode 100644 new mode 100755 index 5a61d816ce..b9e1ba0dbd --- a/auto_round_extension/ark/test/test_moe_prefill_perf.py +++ b/auto_round_extension/ark/test/test_moe_prefill_perf.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2026 Intel Corporation # @@ -150,16 +149,11 @@ def _quantized_prefill_skip_reason() -> str: _QUANT_PREFILL_SKIP = _quantized_prefill_skip_reason() # Surface diagnostics on collection -print( - "[moe-prefill-perf] xpu_available=%s xpu_lib=%s has_moe_gemm=%s" - % ( - _xpu_available(), - "loaded" if ark.xpu_lib is not None else "None", - hasattr(ark.xpu_lib, "moe_gemm") if ark.xpu_lib is not None else False, - ) -) +_xpu_lib_state = "loaded" if ark.xpu_lib is not None else "None" +_has_prefill = hasattr(ark.xpu_lib, "moe_gemm") if ark.xpu_lib is not None else False +print(f"[moe-prefill-perf] xpu_available={_xpu_available()} xpu_lib={_xpu_lib_state} has_moe_gemm={_has_prefill}") if _PREFILL_SKIP: - print("[moe-prefill-perf] suite will SKIP. reason: %s" % _PREFILL_SKIP) + print(f"[moe-prefill-perf] suite will SKIP. reason: {_PREFILL_SKIP}") # --------------------------------------------------------------------------- diff --git a/auto_round_extension/ark/test/test_moe_unified.py b/auto_round_extension/ark/test/test_moe_unified.py old mode 100644 new mode 100755 index ad29f11589..1e969222cb --- a/auto_round_extension/ark/test/test_moe_unified.py +++ b/auto_round_extension/ark/test/test_moe_unified.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2026 Intel Corporation # diff --git a/auto_round_extension/ark/test/test_mxfp4_hadamard.py b/auto_round_extension/ark/test/test_mxfp4_hadamard.py old mode 100644 new mode 100755 index 9c19f35ee4..c94843cb3a --- a/auto_round_extension/ark/test/test_mxfp4_hadamard.py +++ b/auto_round_extension/ark/test/test_mxfp4_hadamard.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2026 Intel Corporation # @@ -131,7 +130,7 @@ def test_e8m0_matches_floor_log2_contract(self): y = hadamard_transform_reference(x.reshape(-1, HADAMARD_DIM), get_hadamard_matrix(HADAMARD_DIM)) amax = y.abs().amax(dim=-1) for group, group_amax in enumerate(amax.tolist()): - expected = min(max(int(math.floor(math.log2(group_amax))) - 2 + 127, 0), 254) + expected = min(max(math.floor(math.log2(group_amax)) - 2 + 127, 0), 254) assert scale.reshape(-1)[group].item() == expected # amax / scale must land in [4, 8): the E8M0 exponent is standard. ratio = group_amax / (2.0 ** (expected - 127)) diff --git a/auto_round_extension/ark/test/test_packq.py b/auto_round_extension/ark/test/test_packq.py old mode 100644 new mode 100755 index 7f73fc2095..4b1eac62f5 --- a/auto_round_extension/ark/test/test_packq.py +++ b/auto_round_extension/ark/test/test_packq.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2023 Intel Corporation # diff --git a/auto_round_extension/ark/test/test_sagev1_varlen.py b/auto_round_extension/ark/test/test_sagev1_varlen.py old mode 100644 new mode 100755 index f3cf1493d2..09378724d7 --- a/auto_round_extension/ark/test/test_sagev1_varlen.py +++ b/auto_round_extension/ark/test/test_sagev1_varlen.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2026 Intel Corporation # @@ -348,8 +347,8 @@ def benchmark_sagev1_varlen_case( head_dim, dtype, device, - min_seq_q=1 if total_q > batch else 1, - min_seq_kv=1 if total_kv > batch else 1, + min_seq_q=1, + min_seq_kv=1, ) scale = 1.0 / math.sqrt(head_dim) diff --git a/auto_round_extension/ark/test/test_sdpa_varlen.py b/auto_round_extension/ark/test/test_sdpa_varlen.py old mode 100644 new mode 100755 index ec16c2d13c..fcafff5f7d --- a/auto_round_extension/ark/test/test_sdpa_varlen.py +++ b/auto_round_extension/ark/test/test_sdpa_varlen.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2026 Intel Corporation # @@ -359,8 +358,8 @@ def benchmark_sdpa_varlen_case( head_dim, dtype, device, - min_seq_q=1 if total_q > batch else 1, - min_seq_kv=1 if total_kv > batch else 1, + min_seq_q=1, + min_seq_kv=1, ) scale = 1.0 / math.sqrt(head_dim) diff --git a/auto_round_extension/ark/test/test_weightonly.py b/auto_round_extension/ark/test/test_weightonly.py old mode 100644 new mode 100755 index 10bbb17e22..e7dd0863de --- a/auto_round_extension/ark/test/test_weightonly.py +++ b/auto_round_extension/ark/test/test_weightonly.py @@ -1,5 +1,4 @@ #!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2023 Intel Corporation # @@ -342,9 +341,8 @@ def _woq_call(x): graph = torch.xpu.CUDAGraph() try: - with torch.xpu.stream(capture_stream): - with torch.xpu.graph(graph): - graph_output = _woq_call(graph_input) + with torch.xpu.stream(capture_stream), torch.xpu.graph(graph): + graph_output = _woq_call(graph_input) except RuntimeError as err: pytest.skip(f"XPU graph capture is unavailable in this environment: {err}") diff --git a/auto_round_extension/ark/test/ut_utils.py b/auto_round_extension/ark/test/ut_utils.py index c1a93d2479..df8d00eb55 100644 --- a/auto_round_extension/ark/test/ut_utils.py +++ b/auto_round_extension/ark/test/ut_utils.py @@ -1,5 +1,3 @@ -#!/usr/bin/env python -# -*- coding: utf-8 -*- # # Copyright (c) 2023 Intel Corporation # diff --git a/auto_round_extension/ark/test/validate_non_int8_cpu_sdpa.py b/auto_round_extension/ark/test/validate_non_int8_cpu_sdpa.py old mode 100644 new mode 100755 index 3ec7081908..0c1a1827fc --- a/auto_round_extension/ark/test/validate_non_int8_cpu_sdpa.py +++ b/auto_round_extension/ark/test/validate_non_int8_cpu_sdpa.py @@ -312,7 +312,7 @@ def main(): real_cmd = cmd[cmd.index(item) :] break - result = subprocess.run(real_cmd, cwd=repo_root, env=env) + result = subprocess.run(real_cmd, cwd=repo_root, env=env, check=False) if result.returncode != 0: print(f" FAILED (exit {result.returncode})\n") overall = False diff --git a/auto_round_extension/ark/tools/analyze_sycl_tla_templates.py b/auto_round_extension/ark/tools/analyze_sycl_tla_templates.py old mode 100644 new mode 100755 diff --git a/auto_round_extension/ark/tools/lm_eval_with_ark_sdpa.py b/auto_round_extension/ark/tools/lm_eval_with_ark_sdpa.py old mode 100644 new mode 100755 diff --git a/auto_round_extension/ark/tools/measure_sycl_tla_compile_memory.py b/auto_round_extension/ark/tools/measure_sycl_tla_compile_memory.py old mode 100644 new mode 100755 index c2d2a75b81..70cbe74607 --- a/auto_round_extension/ark/tools/measure_sycl_tla_compile_memory.py +++ b/auto_round_extension/ark/tools/measure_sycl_tla_compile_memory.py @@ -41,8 +41,7 @@ def measure_command(command, time_binary): completed = subprocess.run( [time_binary, "-v", *command["argv"]], cwd=command["directory"], - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, + capture_output=True, text=True, check=False, ) diff --git a/auto_round_extension/cuda/cute_nvfp4_e5m3.py b/auto_round_extension/cuda/cute_nvfp4_e5m3.py index 4be4769e25..4aa087d0b7 100644 --- a/auto_round_extension/cuda/cute_nvfp4_e5m3.py +++ b/auto_round_extension/cuda/cute_nvfp4_e5m3.py @@ -14,9 +14,8 @@ """Optional CuTe DSL dispatch for NVFP4 E5M3 activation QDQ.""" -from functools import lru_cache +from functools import cache, lru_cache from importlib.util import find_spec -from typing import Optional import torch @@ -167,7 +166,7 @@ def launch_weight_dq(weight_packed: cute.Tensor, weight_scale: cute.Tensor, outp return launch_weight_dq -@lru_cache(maxsize=None) +@cache def _get_compiled_qdq_kernel(device_index: int, dtype: torch.dtype): import cutlass.cute as cute from cutlass.cute.runtime import from_dlpack @@ -181,7 +180,7 @@ def _get_compiled_qdq_kernel(device_index: int, dtype: torch.dtype): ) -@lru_cache(maxsize=None) +@cache def _get_compiled_weight_dq_kernel(device_index: int, dtype: torch.dtype): import cutlass.cute as cute from cutlass.cute.runtime import from_dlpack @@ -197,7 +196,7 @@ def _get_compiled_weight_dq_kernel(device_index: int, dtype: torch.dtype): ) -def try_cute_nvfp4_v2_qdq(activation: torch.Tensor, group_size: int) -> Optional[torch.Tensor]: +def try_cute_nvfp4_v2_qdq(activation: torch.Tensor, group_size: int) -> torch.Tensor | None: """Run a CuTe DSL group-size-16 FP4 QDQ kernel when eligible.""" if not can_use_cute_nvfp4_v2_qdq(activation, group_size): return None @@ -222,7 +221,7 @@ def try_cute_nvfp4_v2_qdq(activation: torch.Tensor, group_size: int) -> Optional def try_cute_nvfp4_e5m3_weight_dq( weight_packed: torch.Tensor, weight_scale: torch.Tensor, dtype: torch.dtype -) -> Optional[torch.Tensor]: +) -> torch.Tensor | None: """Dequantize packed FP4 E5M3 weights with CuTe.""" if ( not is_cute_dsl_available() @@ -261,13 +260,13 @@ def try_cute_nvfp4_e5m3_linear( activation: torch.Tensor, weight_packed: torch.Tensor, weight_scale: torch.Tensor, - bias: Optional[torch.Tensor], -) -> Optional[torch.Tensor]: + bias: torch.Tensor | None, +) -> torch.Tensor | None: """Reserved second-stage dispatch point for fused QDQ, unpack, and GEMM. Returning ``None`` keeps the reference Linear path active until the packed-weight mainloop has been validated against the existing E5M3 checkpoint format. """ del activation, weight_packed, weight_scale, bias - fused_output: Optional[torch.Tensor] = None + fused_output: torch.Tensor | None = None return fused_output diff --git a/auto_round_extension/cuda/gptqmodel_marlin.py b/auto_round_extension/cuda/gptqmodel_marlin.py index a804572f52..7966c9bea1 100644 --- a/auto_round_extension/cuda/gptqmodel_marlin.py +++ b/auto_round_extension/cuda/gptqmodel_marlin.py @@ -17,7 +17,7 @@ # Adapted from vllm # at https://github.com/vllm-project/vllm/blob/main/vllm/model_executor/layers/quantization/gptq_marlin.py -from typing import Any, Dict, List, Optional, Tuple +from typing import Any import numpy as np import torch @@ -72,7 +72,7 @@ def get_marlin_layer(): ##use an ugly wrapper to import gptqmodel on demand def set_weight_attrs( weight: torch.Tensor, - weight_attrs: Optional[Dict[str, Any]], + weight_attrs: dict[str, Any] | None, ): """Set attributes on a weight tensor. @@ -109,7 +109,7 @@ def marlin_make_workspace_new(device: torch.device, max_blocks_per_sm: int = 1) sms = torch.cuda.get_device_properties(device).multi_processor_count return torch.zeros(sms * max_blocks_per_sm, dtype=torch.int, device=device, requires_grad=False) - def marlin_sort_g_idx(g_idx: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]: + def marlin_sort_g_idx(g_idx: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: g_idx_sort_indices = torch.argsort(g_idx).to(torch.int) return g_idx[g_idx_sort_indices], g_idx_sort_indices @@ -137,10 +137,10 @@ def marlin_permute_scales(s: torch.Tensor, size_k: int, size_n: int, group_size: return s def get_scale_perms(): - scale_perm: List[int] = [] + scale_perm: list[int] = [] for i in range(8): scale_perm.extend([i + 8 * j for j in range(8)]) - scale_perm_single: List[int] = [] + scale_perm_single: list[int] = [] for i in range(4): scale_perm_single.extend([2 * i + j for j in [0, 1, 8, 9, 16, 17, 24, 25]]) return scale_perm, scale_perm_single @@ -316,7 +316,7 @@ def __init__( ) # toggle fp32 mode depending on MARLIN or MARLIN_FP16 backend - self.fp32 = True if self.backend in [BACKEND.MARLIN, BACKEND.AUTO] else False + self.fp32 = self.backend in [BACKEND.MARLIN, BACKEND.AUTO] if not self.fp32: logger.warning_once( @@ -501,7 +501,7 @@ def post_init(self): super().post_init() - def list_buffers(self) -> List: + def list_buffers(self) -> list: buf = super().list_buffers() if hasattr(self, "workspace") and self.workspace is not None: buf.append(self.workspace) diff --git a/auto_round_extension/humming/qlinear_humming.py b/auto_round_extension/humming/qlinear_humming.py index 421920f8e4..ec74743b59 100644 --- a/auto_round_extension/humming/qlinear_humming.py +++ b/auto_round_extension/humming/qlinear_humming.py @@ -226,9 +226,9 @@ class QuantLinearAWQ(QuantLinear): __all__ = [ + "SUPPORTED_BITS", "QuantLinear", "QuantLinearAWQ", "QuantLinearGPTQ", - "SUPPORTED_BITS", "is_humming_available", ] diff --git a/auto_round_extension/mlx/__init__.py b/auto_round_extension/mlx/__init__.py index 4a8119fcec..a1b7da04ac 100644 --- a/auto_round_extension/mlx/__init__.py +++ b/auto_round_extension/mlx/__init__.py @@ -14,4 +14,4 @@ from auto_round_extension.mlx.qlinear_mlx import QuantLinearMLX, MLX_AVAILABLE -__all__ = ["QuantLinearMLX", "MLX_AVAILABLE"] +__all__ = ["MLX_AVAILABLE", "QuantLinearMLX"] diff --git a/auto_round_extension/torch/qlinear_torch.py b/auto_round_extension/torch/qlinear_torch.py index 4660577004..cc364327cd 100644 --- a/auto_round_extension/torch/qlinear_torch.py +++ b/auto_round_extension/torch/qlinear_torch.py @@ -234,7 +234,7 @@ def pack_248_bits(self, linear, scales, zeros, g_idx=None, device=None): else: shape = scales_t.shape value = 0 - for j in range(0, (32 // self.bits)): + for j in range(32 // self.bits): value |= zeros << (self.bits * j) qzeros = torch.ones((shape[0], shape[1] // 32 * self.bits), dtype=torch.int32) * value self.qzeros = qzeros.cpu() diff --git a/auto_round_extension/torch/qlinear_torch_zp.py b/auto_round_extension/torch/qlinear_torch_zp.py index d6fef7beb8..117997ce7e 100644 --- a/auto_round_extension/torch/qlinear_torch_zp.py +++ b/auto_round_extension/torch/qlinear_torch_zp.py @@ -165,7 +165,7 @@ def pack_248_bits(self, linear, scales, zeros, g_idx=None, device=None): zeros = int(min(max(zeros - 1, 0), self.maxq)) shape = scales_t.shape value = 0 - for j in range(0, (32 // self.bits)): + for j in range(32 // self.bits): value |= zeros << (self.bits * j) qzeros = torch.ones((shape[0], shape[1] // 32 * self.bits), dtype=torch.int32) * value self.qzeros = qzeros.cpu() diff --git a/auto_round_extension/triton/neuqi_sweep.py b/auto_round_extension/triton/neuqi_sweep.py index 251d5557be..98d639631c 100644 --- a/auto_round_extension/triton/neuqi_sweep.py +++ b/auto_round_extension/triton/neuqi_sweep.py @@ -124,7 +124,7 @@ def _neuqi_sweep_kernel( qwt = tl.load(qw_ptr + c_off[:, None] * G + g_off[None, :], mask=cm[:, None] & gm[None, :], other=0.0) else: qwt = tl.zeros((BC, GP), dtype=tl.float32) + 1.0 - for k in range(0, K): + for k in range(K): sc = tl.load(scales_ptr + c_off * K + k, mask=cm, other=1.0) # [BC] x = d / sc[:, None] r = _rint_f32(x) @@ -144,7 +144,7 @@ def _neuqi_sweep_kernel( best = tl.where(better, acc, best) bz = tl.where(better, z, bz) else: - for z in range(0, NZ): + for z in range(NZ): q = tl.minimum(tl.maximum(r + z, 0.0), MAXQ) err = sc[:, None] * (q - z) - d loss2 = err * err * qwt @@ -190,7 +190,7 @@ def _sym_search_kernel( best = tl.zeros((BC,), dtype=tl.float32) + float("inf") bk = tl.zeros((BC,), dtype=tl.int32) bmir = tl.zeros((BC,), dtype=tl.int32) - for k in range(0, K): + for k in range(K): sc = tl.load(scales_ptr + c_off * K + k, mask=cm, other=1.0) # [BC] x = d / sc[:, None] fl = tl.floor(x) @@ -264,7 +264,7 @@ def _neuqi_shared_kernel( best = tl.zeros((BC,), dtype=tl.float32) + float("inf") bk = tl.zeros((BC,), dtype=tl.int32) bz = tl.zeros((BC,), dtype=tl.int32) - for k in range(0, K): + for k in range(K): f = tl.load(fracs_ptr + k) invf = tl.load(invf_ptr + k) x = d * invf @@ -285,7 +285,7 @@ def _neuqi_shared_kernel( l_best = tl.where(better, acc, l_best) l_bz = tl.where(better, z, l_bz) else: - for z in range(0, MAXQ + 1): + for z in range(MAXQ + 1): q = tl.minimum(tl.maximum(r + z, 0.0), MAXQ) err = f * (q - z) - d loss2 = err * err * qwt @@ -469,7 +469,7 @@ def _sym_shared_kernel( best = tl.zeros((BC,), dtype=tl.float32) + float("inf") bk = tl.zeros((BC,), dtype=tl.int32) bmir = tl.zeros((BC,), dtype=tl.int32) - for k in range(0, K): + for k in range(K): f = tl.load(fracs_ptr + k) invf = tl.load(invf_ptr + k) x = d * invf diff --git a/auto_round_extension/triton/qlinear_tritonv2.py b/auto_round_extension/triton/qlinear_tritonv2.py index 2fcb87f927..28798ee2ce 100644 --- a/auto_round_extension/triton/qlinear_tritonv2.py +++ b/auto_round_extension/triton/qlinear_tritonv2.py @@ -13,6 +13,7 @@ # limitations under the License. import math +import sys from logging import getLogger import numpy as np @@ -30,7 +31,7 @@ except ImportError as e: if torch.xpu.is_available(): logger.error("please make sure your triton version is same with `pytorch-triton-xpu` library ") - exit(-1) + sys.exit(-1) triton_import_exception = e def error_raiser_triton(*args, **kwargs): @@ -154,7 +155,7 @@ def pack(self, linear, scales, zeros, g_idx=None, device=None): else: shape = scales_t.shape value = 0 - for j in range(0, (32 // self.bits)): + for j in range(32 // self.bits): value |= zeros << (self.bits * j) qzeros = np.ones((shape[0], shape[1] // 32 * self.bits), dtype=np.uint32) * value qzeros = qzeros.astype(np.int32) @@ -210,7 +211,7 @@ def warmup(cls, model, transpose=False, seqlen=2048): logger.info(f"Found {len(kn_values)} unique KN Linear values.") logger.info("Warming up autotune cache ...") with torch.no_grad(): - for m in tqdm(range(0, math.ceil(math.log2(seqlen)) + 1)): + for m in tqdm(range(math.ceil(math.log2(seqlen)) + 1)): m = 2**m for (k, n), ( qweight, diff --git a/auto_round_extension/triton/qlinear_tritonv2_zp.py b/auto_round_extension/triton/qlinear_tritonv2_zp.py index 58ff9de0a7..e10bcb21fd 100644 --- a/auto_round_extension/triton/qlinear_tritonv2_zp.py +++ b/auto_round_extension/triton/qlinear_tritonv2_zp.py @@ -13,6 +13,7 @@ # limitations under the License. import math +import sys from logging import getLogger import torch @@ -27,7 +28,7 @@ except ImportError as e: if torch.xpu.is_available(): logger.error("please make sure your triton version is same with `pytorch-triton-xpu` library ") - exit(-1) + sys.exit(-1) triton_import_exception = e def error_raiser_triton(*args, **kwargs): @@ -221,7 +222,7 @@ def warmup(cls, model, transpose=False, seqlen=2048): logger.info(f"Found {len(kn_values)} unique KN Linear values.") logger.info("Warming up autotune cache ...") with torch.no_grad(): - for m in tqdm(range(0, math.ceil(math.log2(seqlen)) + 1)): + for m in tqdm(range(math.ceil(math.log2(seqlen)) + 1)): m = 2**m for (k, n), ( qweight, diff --git a/auto_round_extension/triton/triton_utils/custom_autotune.py b/auto_round_extension/triton/triton_utils/custom_autotune.py index 5b5b5b14d8..1ac6fe36ba 100644 --- a/auto_round_extension/triton/triton_utils/custom_autotune.py +++ b/auto_round_extension/triton/triton_utils/custom_autotune.py @@ -36,7 +36,6 @@ import builtins import math import time -from typing import Dict import triton @@ -54,7 +53,7 @@ def __init__( configs, key, reset_to_zero, - prune_configs_by: Dict = None, + prune_configs_by: dict | None = None, nearest_power_of_two: bool = False, ): if not configs: @@ -207,9 +206,9 @@ def matmul248_kernel_config_pruner(configs, nargs): """ The main purpose of this function is to shrink BLOCK_SIZE_* when the corresponding dimension is smaller. """ - m = max(2 ** int(math.ceil(math.log2(nargs["M"]))), 16) - n = max(2 ** int(math.ceil(math.log2(nargs["N"]))), 16) - k = max(2 ** int(math.ceil(math.log2(nargs["K"]))), 16) + m = max(2 ** math.ceil(math.log2(nargs["M"])), 16) + n = max(2 ** math.ceil(math.log2(nargs["N"])), 16) + k = max(2 ** math.ceil(math.log2(nargs["K"])), 16) used = set() for config in configs: diff --git a/auto_round_extension/triton/triton_utils/kernels.py b/auto_round_extension/triton/triton_utils/kernels.py index efe809fb46..7ab222509b 100644 --- a/auto_round_extension/triton/triton_utils/kernels.py +++ b/auto_round_extension/triton/triton_utils/kernels.py @@ -186,7 +186,7 @@ def quant_matmul_248_kernel( zeros_shifter = (offs_bn % infearure_per_bits) * bits accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) - for k in range(0, num_pid_k): + for k in range(num_pid_k): g_idx = tl.load(g_ptrs) # Fetch scales and zeros; these are per-outfeature and thus reused in the inner loop @@ -348,7 +348,7 @@ def transpose_quant_matmul_248_kernel( zeros_shifter = (offs_n % infearure_per_bits) * bits accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_K), dtype=tl.float32) - for k in range(0, num_pid_n): + for k in range(num_pid_n): # Fetch scales and zeros; these are per-outfeature and thus reused in the inner loop scales = tl.load(scales_ptrs) # (BLOCK_SIZE_K, BLOCK_SIZE_N,) zeros = tl.load(zeros_ptrs) # (BLOCK_SIZE_K, BLOCK_SIZE_N,) diff --git a/auto_round_extension/triton/triton_utils_zp/custom_autotune.py b/auto_round_extension/triton/triton_utils_zp/custom_autotune.py index 5b5b5b14d8..1ac6fe36ba 100644 --- a/auto_round_extension/triton/triton_utils_zp/custom_autotune.py +++ b/auto_round_extension/triton/triton_utils_zp/custom_autotune.py @@ -36,7 +36,6 @@ import builtins import math import time -from typing import Dict import triton @@ -54,7 +53,7 @@ def __init__( configs, key, reset_to_zero, - prune_configs_by: Dict = None, + prune_configs_by: dict | None = None, nearest_power_of_two: bool = False, ): if not configs: @@ -207,9 +206,9 @@ def matmul248_kernel_config_pruner(configs, nargs): """ The main purpose of this function is to shrink BLOCK_SIZE_* when the corresponding dimension is smaller. """ - m = max(2 ** int(math.ceil(math.log2(nargs["M"]))), 16) - n = max(2 ** int(math.ceil(math.log2(nargs["N"]))), 16) - k = max(2 ** int(math.ceil(math.log2(nargs["K"]))), 16) + m = max(2 ** math.ceil(math.log2(nargs["M"])), 16) + n = max(2 ** math.ceil(math.log2(nargs["N"])), 16) + k = max(2 ** math.ceil(math.log2(nargs["K"])), 16) used = set() for config in configs: diff --git a/auto_round_extension/triton/triton_utils_zp/kernels.py b/auto_round_extension/triton/triton_utils_zp/kernels.py index b8715519b9..f79e386d7b 100644 --- a/auto_round_extension/triton/triton_utils_zp/kernels.py +++ b/auto_round_extension/triton/triton_utils_zp/kernels.py @@ -185,7 +185,7 @@ def quant_matmul_248_kernel( zeros_shifter = (offs_bn % infearure_per_bits) * bits accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) - for k in range(0, num_pid_k): + for k in range(num_pid_k): g_idx = tl.load(g_ptrs) # Fetch scales and zeros; these are per-outfeature and thus reused in the inner loop @@ -346,7 +346,7 @@ def transpose_quant_matmul_248_kernel( zeros_shifter = (offs_n % infearure_per_bits) * bits accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_K), dtype=tl.float32) - for k in range(0, num_pid_n): + for k in range(num_pid_n): # Fetch scales and zeros; these are per-outfeature and thus reused in the inner loop scales = tl.load(scales_ptrs) # (BLOCK_SIZE_K, BLOCK_SIZE_N,) zeros = tl.load(zeros_ptrs) # (BLOCK_SIZE_K, BLOCK_SIZE_N,) diff --git a/pyproject.toml b/pyproject.toml index f95721f733..b13c77f97e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -74,11 +74,10 @@ indent-width = 4 target-version = "py311" [tool.ruff.lint] -# Enable Pyflakes (`F`) and a subset of the pycodestyle (`E`) codes by default. -# Unlike Flake8, Ruff doesn't enable pycodestyle warnings (`W`) or -# McCabe complexity (`C901`) by default. -select = ["E4", "E7", "E9", "F", "NPY", "FURB"] +# Using default rules: https://docs.astral.sh/ruff/default-rules/ ignore = [ + "B023", # Function definition does not bind loop variable (closures here are invoked within the same iteration) + "BLE001", # Do not catch blind exception: `Exception` "E402", # Module level import not at top of file "E501", # Line too long (121 > 120 characters) "E721", # Do not compare types, use isinstance() @@ -89,6 +88,16 @@ ignore = [ "F403", # from {name} import * used; unable to detect undefined names "F841", # Local variable is assigned to but never used{name} "FURB171", # Membership test against single-item container + "I001", # Import block is un-sorted or un-formatted (handled by isort hook) + "PLR0402", # Use `from torch import nn` in lieu of alias (`import torch.nn as nn` is the PyTorch idiom) + "RUF012", # Mutable class attributes should be annotated with `typing.ClassVar` + "RUF059", # Unpacked variable is never used + "RUF100", # Unused `noqa` directive (keeps noqa for globally ignored rules) + "S110", # `try`-`except`-`pass` detected + "S112", # `try`-`except`-`continue` detected + "SIM102", # Use a single `if` statement instead of nested `if` statements + "SIM118", # Use `key in dict` instead of `key in dict.keys()` (safetensors handles only expose `.keys()`) + "TRY004", # Prefer `TypeError` exception for invalid type (changing exception types breaks callers) ] # Allow fix for all enabled rules (when `--fix`) is provided. @@ -98,6 +107,12 @@ unfixable = [] # Allow unused variables when underscore-prefixed. dummy-variable-rgx = "^(_+|(_+[a-zA-Z0-9_]*[a-zA-Z0-9]+?))$" +[tool.ruff.lint.flake8-comprehensions] +allow-dict-calls-with-keyword-arguments = true + +[tool.ruff.lint.flake8-bugbear] +extend-immutable-calls = ["torch.device"] + [tool.ruff.format] # Like Black, use double quotes for strings. quote-style = "double" diff --git a/setup.py b/setup.py index 6131dcb4bc..7de42ca2bb 100644 --- a/setup.py +++ b/setup.py @@ -1,9 +1,9 @@ +import builtins import os import re import subprocess import sys from functools import lru_cache -from io import open from setuptools import find_packages, setup @@ -11,10 +11,10 @@ os.environ["CXX"] = "g++" try: filepath = "./auto_round/version.py" - with open(filepath) as version_file: + with builtins.open(filepath) as version_file: (__version__,) = re.findall('__version__ = "(.*)"', version_file.read()) except Exception as error: - assert False, "Error: Could not open '%s' due %s\n" % (filepath, error) + assert False, f"Error: Could not open '{filepath}' due {error}\n" # All BUILD_* flags are initially set to `False` and # will be updated to `True` if the corresponding environment check passes. @@ -40,7 +40,7 @@ def get_build_version(): # it avoids the sdist/wheel version mismatch that occurs because git is not # available inside the extracted sdist directory. if os.path.exists("PKG-INFO"): - with open("PKG-INFO", encoding="utf-8") as f: + with builtins.open("PKG-INFO", encoding="utf-8") as f: for line in f: if line.startswith("Version:"): return line.split(":", 1)[1].strip() @@ -103,7 +103,7 @@ def is_cpu_env(): def fetch_requirements(path): requirements = [] - with open(path, "r") as fd: + with builtins.open(path, "r") as fd: requirements = [r.strip() for r in fd] return requirements @@ -162,13 +162,16 @@ def fetch_requirements(path): install_requires = INSTALL_CFG.get("install_requires", []) extras_require = INSTALL_CFG.get("extras_require", {}) + with builtins.open("README.md", "r", encoding="utf-8") as f: + long_description = f.read() + setup( name=package_name, author="Intel AIPT Team", version=get_build_version(), author_email="wenhua.cheng@intel.com, weiwei1.zhang@intel.com, heng.guo@intel.com", description="Repository of AutoRound: Advanced Weight-Only Quantization Algorithm for LLMs", - long_description=open("README.md", "r", encoding="utf-8").read(), + long_description=long_description, long_description_content_type="text/markdown", keywords="quantization,auto-around,LLM,SignRound", license="Apache 2.0", diff --git a/test/conftest.py b/test/conftest.py index 006079a3bc..7b85a985db 100644 --- a/test/conftest.py +++ b/test/conftest.py @@ -1,6 +1,6 @@ import os import sys -from typing import Mapping +from collections.abc import Mapping import pytest diff --git a/test/e2e/test_cpu/conftest.py b/test/e2e/test_cpu/conftest.py index a603c55d5e..de985a2523 100644 --- a/test/e2e/test_cpu/conftest.py +++ b/test/e2e/test_cpu/conftest.py @@ -32,7 +32,6 @@ import sys import time from dataclasses import asdict, dataclass, field -from typing import List, Optional import pytest @@ -144,7 +143,7 @@ class ModelCase: # Default CPU matrix: real, small (≤1.5B) LLMs that finish quantize+eval # in a few minutes on a 32 GiB host. These are the models most likely # to actually run on CPU in production. -DEFAULT_MODEL_CASES: List[ModelCase] = [ +DEFAULT_MODEL_CASES: list[ModelCase] = [ # Qwen family - small + well supported across all formats. ModelCase("Qwen/Qwen3-0.6B", 4, 128, True, "auto_round", min_ram_gib=8, eval_limit=80), ModelCase("Qwen/Qwen3-0.6B", 4, 128, True, "auto_gptq", min_ram_gib=8, eval_limit=80), @@ -166,7 +165,7 @@ class ModelCase: ] # Heavier cases - 1.5B-2B; need ~24 GiB free RAM and longer wall-clock. -LARGE_MODEL_CASES: List[ModelCase] = [ +LARGE_MODEL_CASES: list[ModelCase] = [ ModelCase("Qwen/Qwen2.5-1.5B-Instruct", 4, 128, True, "auto_round", min_ram_gib=14, eval_limit=80), ModelCase("Qwen/Qwen2.5-1.5B-Instruct", 4, 128, True, "auto_gptq", min_ram_gib=14, eval_limit=80), ModelCase("Qwen/Qwen2.5-1.5B-Instruct", 4, 128, True, "auto_awq", min_ram_gib=14, eval_limit=80), @@ -175,7 +174,7 @@ class ModelCase: ] -def _resolve_mem_override(request) -> Optional[int]: +def _resolve_mem_override(request) -> int | None: return request.config.getoption("--e2e-cpu-mem-gib") if request else None @@ -185,7 +184,7 @@ def _resolve_mem_override(request) -> Optional[int]: @pytest.fixture(scope="session") -def model_matrix(request) -> List[ModelCase]: +def model_matrix(request) -> list[ModelCase]: preset = os.environ.get("E2E_CPU_PRESET", "default") if preset == "default": return DEFAULT_MODEL_CASES @@ -261,9 +260,9 @@ class EvalResult: bits: int group_size: int sym: bool - task: Optional[str] = None - metric: Optional[str] = None - value: Optional[float] = None + task: str | None = None + metric: str | None = None + value: float | None = None extra: dict = field(default_factory=dict) wall_time_s: float = 0.0 @@ -295,8 +294,8 @@ def quantize_and_save( iters: int = 200, nsamples: int = 128, seqlen: int = 2048, - extra_kwargs: Optional[dict] = None, - scheme: Optional[str] = None, + extra_kwargs: dict | None = None, + scheme: str | None = None, ): """Quantize a model with the Python API and save it. @@ -342,7 +341,7 @@ def run_lm_eval( limit: int = 100, batch_size: str = "auto", model_type: str = "hf", - extra_model_args: Optional[dict] = None, + extra_model_args: dict | None = None, ): """Run ``lm-eval`` over a saved checkpoint. @@ -366,7 +365,7 @@ def run_lm_eval( ) -def extract_metric(results: dict, task: str, metric: str = "acc,none") -> Optional[float]: +def extract_metric(results: dict, task: str, metric: str = "acc,none") -> float | None: """Pull a single metric out of the lm-eval results dict (may be missing).""" try: return float(results["results"][task][metric]) @@ -379,7 +378,7 @@ def extract_metric(results: dict, task: str, metric: str = "acc,none") -> Option # --------------------------------------------------------------------------- -def run_cli(argv: List[str], env: Optional[dict] = None, timeout: int = 60 * 60) -> int: +def run_cli(argv: list[str], env: dict | None = None, timeout: int = 60 * 60) -> int: """Spawn ``python -m auto_round `` and return the exit code. Used by the CLI e2e tests; intentionally goes through the actual diff --git a/test/e2e/test_cpu/test_bf16_vs_quant_quality.py b/test/e2e/test_cpu/test_bf16_vs_quant_quality.py index 1826caefc1..1a47ee7307 100644 --- a/test/e2e/test_cpu/test_bf16_vs_quant_quality.py +++ b/test/e2e/test_cpu/test_bf16_vs_quant_quality.py @@ -39,7 +39,6 @@ quantize_and_save, record, ) -from typing import Dict, Optional import pytest @@ -75,7 +74,7 @@ def _case_id(model_id: str, scheme: str) -> str: # --------------------------------------------------------------------------- -def _evaluate_bf16(model_id: str, tasks: str, limit: int) -> Dict[str, Optional[float]]: +def _evaluate_bf16(model_id: str, tasks: str, limit: int) -> dict[str, float | None]: """Run ``lm-eval`` on the bf16 model and return a {task: acc} dict.""" from auto_round.eval.evaluation import simple_evaluate_user_model from auto_round.utils import llm_load_model diff --git a/test/e2e/test_cpu/test_diffusion_quantize_e2e.py b/test/e2e/test_cpu/test_diffusion_quantize_e2e.py index 126cf3252c..464a60c765 100644 --- a/test/e2e/test_cpu/test_diffusion_quantize_e2e.py +++ b/test/e2e/test_cpu/test_diffusion_quantize_e2e.py @@ -36,7 +36,6 @@ EvalResult, record, ) -from typing import Optional import pytest import torch diff --git a/test/e2e/test_cpu/test_gguf_conversion_e2e.py b/test/e2e/test_cpu/test_gguf_conversion_e2e.py index 91ce205ac0..bc1afd1fa2 100644 --- a/test/e2e/test_cpu/test_gguf_conversion_e2e.py +++ b/test/e2e/test_cpu/test_gguf_conversion_e2e.py @@ -43,7 +43,6 @@ assert_non_garbage_output, record, ) -from typing import List, Optional import pytest @@ -88,7 +87,7 @@ def _case_id(gguf_type: str) -> str: def _find_gguf(save_dir: str) -> str: - matches: List[str] = [] + matches: list[str] = [] for root, _, files in os.walk(save_dir): for name in files: if name.endswith(".gguf"): diff --git a/test/e2e/test_cpu/test_gguf_cpu_inference.py b/test/e2e/test_cpu/test_gguf_cpu_inference.py index 5440dc36e4..bbb992c9ab 100644 --- a/test/e2e/test_cpu/test_gguf_cpu_inference.py +++ b/test/e2e/test_cpu/test_gguf_cpu_inference.py @@ -42,7 +42,6 @@ quantize_and_save, record, ) -from typing import List import pytest @@ -99,7 +98,7 @@ def _build_llamacpp(gguf_path: str, n_ctx: int = 512, n_threads: int = 0): def _find_gguf(save_dir: str) -> str: """Locate the .gguf file produced by ``auto-round --format gguf:*``.""" - matches: List[str] = [] + matches: list[str] = [] for root, _, files in os.walk(save_dir): for name in files: if name.endswith(".gguf"): diff --git a/test/e2e/test_cpu/test_moe_e2e.py b/test/e2e/test_cpu/test_moe_e2e.py index 9b6e7493a8..1a8417012d 100644 --- a/test/e2e/test_cpu/test_moe_e2e.py +++ b/test/e2e/test_cpu/test_moe_e2e.py @@ -39,7 +39,6 @@ quantize_and_save, record, ) -from typing import List import pytest import torch diff --git a/test/e2e/test_cpu/test_save_load_roundtrip.py b/test/e2e/test_cpu/test_save_load_roundtrip.py index dca6742d88..dd5dfff8b5 100644 --- a/test/e2e/test_cpu/test_save_load_roundtrip.py +++ b/test/e2e/test_cpu/test_save_load_roundtrip.py @@ -46,7 +46,6 @@ assert_non_garbage_output, record, ) -from typing import List, Optional import pytest import torch @@ -117,7 +116,7 @@ def _reload_and_generate_gguf(saved_dir: str) -> str: """Reload a GGUF checkpoint via llama.cpp and run a short generation.""" from llama_cpp import Llama - matches: List[str] = [] + matches: list[str] = [] for root, _, files in os.walk(saved_dir): for name in files: if name.endswith(".gguf"): diff --git a/test/e2e/test_cuda/conftest.py b/test/e2e/test_cuda/conftest.py index 7adf16464c..a6d9fb926d 100644 --- a/test/e2e/test_cuda/conftest.py +++ b/test/e2e/test_cuda/conftest.py @@ -35,7 +35,6 @@ import sys import time from dataclasses import dataclass, field -from typing import List, Optional import pytest import torch @@ -126,7 +125,7 @@ class ModelCase: # every user runs, (b) the W2A16 low-memory path, (c) the activation # quant path and (d) the GPTQ/AWQ back-compat paths. All models are # small enough to fit on a single 24 GiB GPU at W4A16 with offloading. -DEFAULT_MODEL_CASES: List[ModelCase] = [ +DEFAULT_MODEL_CASES: list[ModelCase] = [ ModelCase("Qwen/Qwen3-1.7B", 4, 128, True, "auto_round", min_gpu_gib=10), ModelCase("Qwen/Qwen3-1.7B", 4, 128, True, "auto_gptq", min_gpu_gib=10), ModelCase("Qwen/Qwen3-1.7B", 4, 128, True, "auto_awq", min_gpu_gib=10), @@ -136,7 +135,7 @@ class ModelCase: # A more demanding matrix for nightly runs (8 GiB / 7B-class). These # require ~16 GiB free at fp16 master + W4A16 weights. -LARGE_MODEL_CASES: List[ModelCase] = [ +LARGE_MODEL_CASES: list[ModelCase] = [ ModelCase("Qwen/Qwen2.5-7B-Instruct", 4, 128, True, "auto_round", min_gpu_gib=18), ModelCase("Qwen/Qwen2.5-7B-Instruct", 4, 128, True, "auto_gptq", min_gpu_gib=18), ModelCase("meta-llama/Llama-3.2-3B-Instruct", 4, 128, True, "auto_round", min_gpu_gib=12), @@ -159,7 +158,7 @@ def pytest_addoption(parser): @pytest.fixture(scope="session") -def model_matrix(request) -> List[ModelCase]: +def model_matrix(request) -> list[ModelCase]: preset = request.config.getoption("--e2e-model-preset") if preset == "default": return DEFAULT_MODEL_CASES @@ -211,7 +210,7 @@ def quantize_and_save( iters: int = 200, nsamples: int = 128, seqlen: int = 2048, - extra_kwargs: Optional[dict] = None, + extra_kwargs: dict | None = None, ): """Run the full AutoRound pipeline and return the saved checkpoint dir. @@ -264,14 +263,14 @@ class BenchResult: # End-to-end output tokens per second (prompt processing + decoding). output_tokens_per_s: float # Decoding-only tokens per second (excludes prompt eval), if measurable. - gen_tokens_per_s: Optional[float] + gen_tokens_per_s: float | None # Time-to-first-token seconds (mean over the batch, if available). - ttft_s: Optional[float] + ttft_s: float | None # Generated text for the first prompt; useful for sanity checks. sample_output: str -def _standard_prompts() -> List[str]: +def _standard_prompts() -> list[str]: """A fixed prompt list so different runs are comparable.""" return [ "The capital of France is", @@ -285,7 +284,7 @@ def _standard_prompts() -> List[str]: ] -def make_bench_prompts(tokenizer, num_prompts: int, target_input_tokens: int = 64) -> List[str]: +def make_bench_prompts(tokenizer, num_prompts: int, target_input_tokens: int = 64) -> list[str]: """Pad each base prompt with lorem-style text to ~target_input_tokens. The result is a list of prompts whose prompt-eval cost is similar @@ -296,7 +295,7 @@ def make_bench_prompts(tokenizer, num_prompts: int, target_input_tokens: int = 6 " Lorem ipsum dolor sit amet, consectetur adipiscing elit. " "Sed do eiusmod tempor incididunt ut labore et dolore magna aliqua. " ) - out: List[str] = [] + out: list[str] = [] i = 0 while len(out) < num_prompts: prompt = base[i % len(base)] + pad * 4 diff --git a/test/e2e/test_cuda/test_sglang_throughput.py b/test/e2e/test_cuda/test_sglang_throughput.py index 3ed58a3aca..0d031b9348 100644 --- a/test/e2e/test_cuda/test_sglang_throughput.py +++ b/test/e2e/test_cuda/test_sglang_throughput.py @@ -41,7 +41,6 @@ make_bench_prompts, quantize_and_save, ) -from typing import List import pytest import torch diff --git a/test/e2e/test_cuda/test_vllm_throughput.py b/test/e2e/test_cuda/test_vllm_throughput.py index 03d03d0111..6f7cb11d7a 100644 --- a/test/e2e/test_cuda/test_vllm_throughput.py +++ b/test/e2e/test_cuda/test_vllm_throughput.py @@ -51,7 +51,6 @@ make_bench_prompts, quantize_and_save, ) -from typing import List import pytest import torch diff --git a/test/helpers.py b/test/helpers.py index a05b0a9aa1..68fc6cf97e 100644 --- a/test/helpers.py +++ b/test/helpers.py @@ -381,14 +381,14 @@ def _get_module(cls_name, mod_name, folder_name): config = json.load(f) _reduce_config_layers(config, num_layers, num_experts) _apply_config_overrides(config, config_overrides) - return getattr(getattr(diffusers_module, mod_name), "from_config")(config) + return getattr(diffusers_module, mod_name).from_config(config) else: config = transformers.AutoConfig.from_pretrained( os.path.join(local_dir, folder_name, "config.json") ) _reduce_config_layers(config, num_layers, num_experts) _apply_config_overrides(config, config_overrides) - return getattr(getattr(transformers_module, mod_name), "_from_config")(config) + return getattr(transformers_module, mod_name)._from_config(config) with open(os.path.join(local_dir, "model_index.json"), "r", encoding="utf-8") as f: model_index = json.load(f) @@ -452,7 +452,7 @@ def slice_layers(module): return sliced kwargs["dtype"] = "auto" if "auto" not in kwargs else kwargs["dtype"] - kwargs["trust_remote_code"] = True if "trust_remote_code" not in kwargs else kwargs["trust_remote_code"] + kwargs["trust_remote_code"] = kwargs.get("trust_remote_code", True) if is_mllm: model, processor, tokenizer, image_processor = mllm_load_model(model_name_or_path, **kwargs) if hasattr(model.config, "vision_config"): diff --git a/test/integration/test_cpu/test_inc_integration.py b/test/integration/test_cpu/test_inc_integration.py index b5d71b6ff1..257d68aad7 100644 --- a/test/integration/test_cpu/test_inc_integration.py +++ b/test/integration/test_cpu/test_inc_integration.py @@ -31,7 +31,7 @@ @torch.no_grad() def run_fn(model, dataloader): for data in dataloader: - if isinstance(data, tuple) or isinstance(data, list): + if isinstance(data, (tuple, list)): model(*data) elif isinstance(data, dict): model(**data) diff --git a/test/unit/common/algorithms/transforms/hadamard/test_dispatcher.py b/test/unit/common/algorithms/transforms/hadamard/test_dispatcher.py index bd490df565..8abf418c41 100644 --- a/test/unit/common/algorithms/transforms/hadamard/test_dispatcher.py +++ b/test/unit/common/algorithms/transforms/hadamard/test_dispatcher.py @@ -289,16 +289,18 @@ def test_transform_with_unsupported_hadamard_type_raises(self): fuse_online_to_weight=None, ) model = nn.Linear(8, 8) - with patch( - "auto_round.algorithms.transforms.hadamard.dispatcher.resolve_hadamard_backend", - return_value="transform", - ): - with patch( + with ( + patch( + "auto_round.algorithms.transforms.hadamard.dispatcher.resolve_hadamard_backend", + return_value="transform", + ), + patch( "auto_round.algorithms.transforms.hadamard.dispatcher._to_config", return_value=fake_cfg, - ): - with pytest.raises(ValueError, match="only supports hadamard or random_hadamard"): - apply_hadamard_rotation(model, fake_cfg, "mx_fp", compute_device="cpu") + ), + pytest.raises(ValueError, match="only supports hadamard or random_hadamard"), + ): + apply_hadamard_rotation(model, fake_cfg, "mx_fp", compute_device="cpu") def test_rotation_config_stored_on_model(self): """After apply, ``_rotation_config`` is set on the model (inplace path).""" diff --git a/test/unit/common/calibration/test_diffusion.py b/test/unit/common/calibration/test_diffusion.py index 3a6de3177e..b578e2828d 100644 --- a/test/unit/common/calibration/test_diffusion.py +++ b/test/unit/common/calibration/test_diffusion.py @@ -157,10 +157,13 @@ def test_calib_string_dataset_reloads_dataloader(self, calibrator): calibrator.pipe = FakePipeline(fn=lambda *args, **kwargs: None) calibrator._requires_calibration_image = lambda: False - with patch( - "auto_round.compressors.diffusion.dataset.get_diffusion_dataloader", - return_value=(new_dataloader, 2), - ), patch("auto_round.calibration.diffusion.tqdm", FakeTqdm): + with ( + patch( + "auto_round.compressors.diffusion.dataset.get_diffusion_dataloader", + return_value=(new_dataloader, 2), + ), + patch("auto_round.calibration.diffusion.tqdm", FakeTqdm), + ): calibrator.calib(nsamples=2, bs=1) assert calibrator.dataloader is new_dataloader @@ -203,12 +206,14 @@ def test_calib_exits_on_multi_device_offload(self, calibrator): ) calibrator.dataset = "mock" - with patch( - "auto_round.compressors.diffusion.dataset.get_diffusion_dataloader", - return_value=([], 2), + with ( + patch( + "auto_round.compressors.diffusion.dataset.get_diffusion_dataloader", + return_value=([], 2), + ), + pytest.raises(SystemExit), ): - with pytest.raises(SystemExit): - calibrator.calib(nsamples=1, bs=1) + calibrator.calib(nsamples=1, bs=1) def test_calib_moves_pipeline_to_target_device(self, calibrator): seen = [] @@ -224,9 +229,12 @@ def fake_to(device): calibrator.pipe.to = fake_to - with patch("auto_round.calibration.diffusion.tqdm", FakeTqdm), patch( - "auto_round.calibration.diffusion.device_manager", - SimpleNamespace(device="cuda:0"), + with ( + patch("auto_round.calibration.diffusion.tqdm", FakeTqdm), + patch( + "auto_round.calibration.diffusion.device_manager", + SimpleNamespace(device="cuda:0"), + ), ): calibrator.calib(nsamples=2, bs=1) @@ -294,9 +302,11 @@ def failing_pipe(*args, **kwargs): calibrator.pipe = FakePipeline(fn=failing_pipe) calibrator._requires_calibration_image = lambda: False - with patch("auto_round.calibration.diffusion.tqdm", FakeTqdm): - with pytest.raises(NotImplementedError, match="unsupported op"): - calibrator.calib(nsamples=1, bs=1) + with ( + patch("auto_round.calibration.diffusion.tqdm", FakeTqdm), + pytest.raises(NotImplementedError, match="unsupported op"), + ): + calibrator.calib(nsamples=1, bs=1) def test_calib_other_exceptions_propagate(self, calibrator): def failing_pipe(*args, **kwargs): @@ -306,9 +316,8 @@ def failing_pipe(*args, **kwargs): calibrator.pipe = FakePipeline(fn=failing_pipe) calibrator._requires_calibration_image = lambda: False - with patch("auto_round.calibration.diffusion.tqdm", FakeTqdm): - with pytest.raises(RuntimeError, match="unexpected"): - calibrator.calib(nsamples=1, bs=1) + with patch("auto_round.calibration.diffusion.tqdm", FakeTqdm), pytest.raises(RuntimeError, match="unexpected"): + calibrator.calib(nsamples=1, bs=1) def test_calib_single_sample_stops_early(self, calibrator): seen = [] @@ -330,9 +339,8 @@ def test_calib_zero_samples_exits(self, calibrator): calibrator.pipe = FakePipeline(fn=lambda *args, **kwargs: None) calibrator._requires_calibration_image = lambda: False - with patch("auto_round.calibration.diffusion.tqdm", FakeTqdm): - with pytest.raises(SystemExit): - calibrator.calib(nsamples=1, bs=1) + with patch("auto_round.calibration.diffusion.tqdm", FakeTqdm), pytest.raises(SystemExit): + calibrator.calib(nsamples=1, bs=1) def test_calib_insufficient_samples_warns_and_truncates(self, calibrator): def fake_pipe(prompts, **kwargs): @@ -357,6 +365,8 @@ def fake_pipe(prompts, **kwargs): calibrator.pipe = FakePipeline(fn=fake_pipe) calibrator._requires_calibration_image = lambda: False - with patch("auto_round.calibration.diffusion.tqdm", FakeTqdm): - with pytest.raises(ValueError, match="valid sample count is less than batch_size"): - calibrator.calib(nsamples=3, bs=2) + with ( + patch("auto_round.calibration.diffusion.tqdm", FakeTqdm), + pytest.raises(ValueError, match="valid sample count is less than batch_size"), + ): + calibrator.calib(nsamples=3, bs=2) diff --git a/test/unit/common/compressors/test_compressors_init.py b/test/unit/common/compressors/test_compressors_init.py index 672c4095b0..08fc468655 100644 --- a/test/unit/common/compressors/test_compressors_init.py +++ b/test/unit/common/compressors/test_compressors_init.py @@ -17,7 +17,7 @@ class TestCompressorsLazyImports: def test_auto_round_lazy_import(self): with pytest.raises(AttributeError, match="has no attribute"): - getattr(compressors, "AutoRound") + _ = compressors.AutoRound def test_base_compressor_lazy_import(self): BaseCompressor = compressors.BaseCompressor @@ -41,7 +41,7 @@ def test_model_free_compressor_lazy_import(self): def test_unknown_attribute_raises(self): with pytest.raises(AttributeError, match="has no attribute"): - getattr(compressors, "UnknownClass123") + _ = compressors.UnknownClass123 def test_all_contains_expected(self): assert "BaseOrchestrator" in compressors.__all__ diff --git a/test/unit/common/eval/test_evaluation.py b/test/unit/common/eval/test_evaluation.py index 028345c067..d9a9da0813 100644 --- a/test/unit/common/eval/test_evaluation.py +++ b/test/unit/common/eval/test_evaluation.py @@ -242,20 +242,24 @@ def test_falls_back_to_dispatch_block_wise(self): def test_raises_when_meta_device(self): m = nn.Linear(4, 4) m.dtype = torch.bfloat16 - with patch("auto_round.eval.evaluation._normalize_model_eval_dtype", return_value=m): - with patch("auto_round.eval.evaluation.dispatch_model_block_wise") as mock_dispatch: - result = prepare_model_for_eval(m, "cpu", "auto") - assert result is m - mock_dispatch.assert_called_once() + with ( + patch("auto_round.eval.evaluation._normalize_model_eval_dtype", return_value=m), + patch("auto_round.eval.evaluation.dispatch_model_block_wise") as mock_dispatch, + ): + result = prepare_model_for_eval(m, "cpu", "auto") + assert result is m + mock_dispatch.assert_called_once() def test_multi_device_dispatch(self): m = nn.Linear(4, 4) m.hf_device_map = {"linear": "cpu", "linear2": "cpu"} - with patch("auto_round.eval.evaluation._normalize_model_eval_dtype", return_value=m): - with patch("accelerate.big_modeling.dispatch_model") as mock_dispatch: - result = prepare_model_for_eval(m, "cpu", "auto") - assert result is m - mock_dispatch.assert_called_once() + with ( + patch("auto_round.eval.evaluation._normalize_model_eval_dtype", return_value=m), + patch("accelerate.big_modeling.dispatch_model") as mock_dispatch, + ): + result = prepare_model_for_eval(m, "cpu", "auto") + assert result is m + mock_dispatch.assert_called_once() class TestSimpleEvaluate: @@ -274,16 +278,18 @@ class TestSimpleEvaluateUserModel: def test_creates_hflm(self): mock_hflm = MagicMock() - with patch.dict( - "sys.modules", - { - "lm_eval": MagicMock(), - "lm_eval.models": MagicMock(), - "lm_eval.models.huggingface": MagicMock(HFLM=mock_hflm), - }, + with ( + patch.dict( + "sys.modules", + { + "lm_eval": MagicMock(), + "lm_eval.models": MagicMock(), + "lm_eval.models.huggingface": MagicMock(HFLM=mock_hflm), + }, + ), + patch("lm_eval.simple_evaluate", return_value={"results": {}}) as mock_eval, ): - with patch("lm_eval.simple_evaluate", return_value={"results": {}}) as mock_eval: - model = MagicMock() - tokenizer = MagicMock() - result = simple_evaluate_user_model(model, tokenizer, batch_size=4) - assert "results" in result or mock_hflm.called + model = MagicMock() + tokenizer = MagicMock() + result = simple_evaluate_user_model(model, tokenizer, batch_size=4) + assert "results" in result or mock_hflm.called diff --git a/test/unit/common/export/test_gguf_conversion.py b/test/unit/common/export/test_gguf_conversion.py index aa3de40162..85de13ec0b 100644 --- a/test/unit/common/export/test_gguf_conversion.py +++ b/test/unit/common/export/test_gguf_conversion.py @@ -1,6 +1,3 @@ -#!/usr/bin/env python3 -# -*- coding: utf-8 -*- - """Comprehensive unit tests for GGUF conversion modules with low coverage. Tests conversion modules that have 0% coverage in the test suite, covering @@ -343,9 +340,11 @@ def test_prepare_tensors_unprocessed_experts_error(self): obj.tensor_map.mapping = {"tensor": ("KEY", "tensor_name")} # Mock super().prepare_tensors() to skip the base class work - with patch("auto_round.export.export_to_gguf.conversion.base.ModelBase.prepare_tensors"): - with pytest.raises(ValueError, match="Unprocessed experts"): - obj.prepare_tensors() + with ( + patch("auto_round.export.export_to_gguf.conversion.base.ModelBase.prepare_tensors"), + pytest.raises(ValueError, match="Unprocessed experts"), + ): + obj.prepare_tensors() # ============================================================================== @@ -813,9 +812,11 @@ def test_olmoe_prepare_tensors_unprocessed_experts_error(self): obj._experts = [{"unprocessed.tensor": None}] obj.tensor_map.mapping = {"tensor": ("KEY", "tensor_name")} - with patch("auto_round.export.export_to_gguf.conversion.base.ModelBase.prepare_tensors"): - with pytest.raises(ValueError, match="Unprocessed experts"): - obj.prepare_tensors() + with ( + patch("auto_round.export.export_to_gguf.conversion.base.ModelBase.prepare_tensors"), + pytest.raises(ValueError, match="Unprocessed experts"), + ): + obj.prepare_tensors() # ============================================================================== @@ -1387,9 +1388,11 @@ def test_smallthinker_prepare_tensors_unprocessed_experts_error(self): obj._experts = [{"unprocessed.tensor": None}] obj.tensor_map.mapping = {"tensor": ("KEY", "tensor_name")} - with patch("auto_round.export.export_to_gguf.conversion.base.ModelBase.prepare_tensors"): - with pytest.raises(ValueError, match="Unprocessed experts"): - obj.prepare_tensors() + with ( + patch("auto_round.export.export_to_gguf.conversion.base.ModelBase.prepare_tensors"), + pytest.raises(ValueError, match="Unprocessed experts"), + ): + obj.prepare_tensors() def test_smallthinker_set_gguf_parameters_expert_gating(self): """Test expert gating function is set correctly.""" @@ -3044,9 +3047,11 @@ def test_ministral3_asserts_yarn_rope_type(self): }, ) obj.rope_parameters = {"rope_type": "linear", "mscale_all_dim": 1.0, "llama_4_scaling_beta": 0.5} - with patch.object(cls.__mro__[1], "set_gguf_parameters", lambda self: None): - with pytest.raises(AssertionError, match="rope_type must be 'yarn'"): - obj.set_gguf_parameters() + with ( + patch.object(cls.__mro__[1], "set_gguf_parameters", lambda self: None), + pytest.raises(AssertionError, match="rope_type must be 'yarn'"), + ): + obj.set_gguf_parameters() # ============================================================================== @@ -5686,9 +5691,11 @@ def test_bert_set_gguf_parameters(self): obj = _make_mock_model(BertModel) obj.cls_out_labels = None - with patch.object(BertModel.__mro__[1], "set_gguf_parameters", lambda self: None): - with patch.object(obj, "_try_set_pooling_type"): - obj.set_gguf_parameters() + with ( + patch.object(BertModel.__mro__[1], "set_gguf_parameters", lambda self: None), + patch.object(obj, "_try_set_pooling_type"), + ): + obj.set_gguf_parameters() obj.gguf_writer.add_causal_attention.assert_called_once_with(False) def test_bert_set_gguf_parameters_with_classifier_labels(self): @@ -5697,9 +5704,11 @@ def test_bert_set_gguf_parameters_with_classifier_labels(self): obj = _make_mock_model(BertModel) obj.cls_out_labels = {"0": "NEGATIVE", "1": "POSITIVE"} - with patch.object(BertModel.__mro__[1], "set_gguf_parameters", lambda self: None): - with patch.object(obj, "_try_set_pooling_type"): - obj.set_gguf_parameters() + with ( + patch.object(BertModel.__mro__[1], "set_gguf_parameters", lambda self: None), + patch.object(obj, "_try_set_pooling_type"), + ): + obj.set_gguf_parameters() obj.gguf_writer.add_classifier_output_labels.assert_called_once_with(["NEGATIVE", "POSITIVE"]) def test_bert_filter_tensors_strips_bert_prefix(self): diff --git a/test/unit/common/export/test_gguf_dtype_helpers.py b/test/unit/common/export/test_gguf_dtype_helpers.py index 71368b8c3f..f8606cb240 100644 --- a/test/unit/common/export/test_gguf_dtype_helpers.py +++ b/test/unit/common/export/test_gguf_dtype_helpers.py @@ -195,7 +195,7 @@ def test_first_eighth(self): from auto_round.export.export_to_gguf.gguf_dtype import _use_more_bits # 8 layers: first 8/8=1 layer uses more bits - for i in range(0, 1): + for i in range(1): assert _use_more_bits(i, 8) is True def test_last_eighth(self): diff --git a/test/unit/common/export/test_mlx_export.py b/test/unit/common/export/test_mlx_export.py index d9cef58859..cb41d2a19f 100644 --- a/test/unit/common/export/test_mlx_export.py +++ b/test/unit/common/export/test_mlx_export.py @@ -552,7 +552,8 @@ def test_saves_config_json(self, tmp_path): ) cfg_path = os.path.join(output_dir, "config.json") assert os.path.exists(cfg_path) - cfg = json.load(open(cfg_path)) + with open(cfg_path) as f: + cfg = json.load(f) assert "quantization" in cfg def test_autoround_format_flag(self, tmp_path): diff --git a/test/unit/common/export/test_qlinear_fp_helpers.py b/test/unit/common/export/test_qlinear_fp_helpers.py index 050c565350..3effc88810 100644 --- a/test/unit/common/export/test_qlinear_fp_helpers.py +++ b/test/unit/common/export/test_qlinear_fp_helpers.py @@ -67,7 +67,7 @@ def test_construction_4bit_nv(self): ) assert layer.weight_global_scale.shape == (1,) # act_bits > 8 -> input_global_scale NOT registered - assert not hasattr(layer, "input_global_scale") or layer.input_global_scale is None or True + assert not hasattr(layer, "input_global_scale") or layer.input_global_scale is None def test_construction_4bit_nv_act_global(self): from auto_round.export.export_to_autoround.qlinear_fp import QuantLinear diff --git a/test/unit/common/modeling/test_fp8_quant.py b/test/unit/common/modeling/test_fp8_quant.py index 167dc6eafd..52e3ba33f9 100644 --- a/test/unit/common/modeling/test_fp8_quant.py +++ b/test/unit/common/modeling/test_fp8_quant.py @@ -94,15 +94,17 @@ def __init__(self): model = _NoLinear() config = _QuantConfigStub(dequantize=False) - with patch("transformers.integrations.finegrained_fp8.FP8Linear"): - with patch( + with ( + patch("transformers.integrations.finegrained_fp8.FP8Linear"), + patch( "transformers.integrations.finegrained_fp8.should_convert_module", return_value=True, - ): - with patch("transformers.integrations.finegrained_fp8.logger") as mock_logger: - result = oot_replace_with_fp8_linear(model, quantization_config=config) - assert result is model - assert mock_logger.warning.called + ), + patch("transformers.integrations.finegrained_fp8.logger") as mock_logger, + ): + result = oot_replace_with_fp8_linear(model, quantization_config=config) + assert result is model + assert mock_logger.warning.called def test_replaces_linear_modules(self): """All ``nn.Linear`` children should be replaced with the FP8 class.""" @@ -113,21 +115,23 @@ def test_replaces_linear_modules(self): def _make_module(*args, **kwargs): return nn.Linear(8, 8) - with patch( - "transformers.integrations.finegrained_fp8.FP8Linear", - side_effect=_make_module, - ): - with patch( + with ( + patch( + "transformers.integrations.finegrained_fp8.FP8Linear", + side_effect=_make_module, + ), + patch( "transformers.integrations.finegrained_fp8.should_convert_module", return_value=True, - ): - with patch( - "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5_4_0", - return_value=False, - ): - result = oot_replace_with_fp8_linear(model, quantization_config=config) - # FP8Linear was called for each nn.Linear child. - assert result is model + ), + patch( + "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5_4_0", + return_value=False, + ), + ): + result = oot_replace_with_fp8_linear(model, quantization_config=config) + # FP8Linear was called for each nn.Linear child. + assert result is model def test_with_modules_to_not_convert(self): """Names listed in ``modules_to_not_convert`` are skipped.""" @@ -135,19 +139,21 @@ def test_with_modules_to_not_convert(self): model = _TinyModel(with_bias=True) config = _QuantConfigStub(dequantize=False) - with patch("transformers.integrations.finegrained_fp8.FP8Linear") as mock_fp8: - with patch( + with ( + patch("transformers.integrations.finegrained_fp8.FP8Linear") as mock_fp8, + patch( "transformers.integrations.finegrained_fp8.should_convert_module", return_value=False, - ): - oot_replace_with_fp8_linear( - model, - modules_to_not_convert=["fc1"], - quantization_config=config, - ) - # No replacement calls should have happened - # because should_convert_module returned False everywhere. - assert not mock_fp8.called + ), + ): + oot_replace_with_fp8_linear( + model, + modules_to_not_convert=["fc1"], + quantization_config=config, + ) + # No replacement calls should have happened + # because should_convert_module returned False everywhere. + assert not mock_fp8.called def test_pre_quantized(self): """The ``pre_quantized=True`` path passes ``dtype=None`` instead of @@ -162,23 +168,25 @@ def _capture(*args, **kwargs): captured_kwargs.append(kwargs) return nn.Linear(8, 8) - with patch("transformers.integrations.finegrained_fp8.FP8Linear", side_effect=_capture): - with patch( + with ( + patch("transformers.integrations.finegrained_fp8.FP8Linear", side_effect=_capture), + patch( "transformers.integrations.finegrained_fp8.should_convert_module", return_value=True, - ): - with patch( - "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5_4_0", - return_value=True, - ): - oot_replace_with_fp8_linear( - model, - quantization_config=config, - pre_quantized=True, - ) - # Every captured call must include ``dtype=None``. - for kw in captured_kwargs: - assert kw.get("dtype") is None + ), + patch( + "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5_4_0", + return_value=True, + ), + ): + oot_replace_with_fp8_linear( + model, + quantization_config=config, + pre_quantized=True, + ) + # Every captured call must include ``dtype=None``. + for kw in captured_kwargs: + assert kw.get("dtype") is None def test_bias_kwarg_name_pre_v5_4(self): """On transformers < 5.4, the bias flag is passed as ``bias``.""" @@ -191,23 +199,25 @@ def _capture(*args, **kwargs): captured_kwargs.append(kwargs) return nn.Linear(8, 8) - with patch("transformers.integrations.finegrained_fp8.FP8Linear", side_effect=_capture): - with patch( + with ( + patch("transformers.integrations.finegrained_fp8.FP8Linear", side_effect=_capture), + patch( "transformers.integrations.finegrained_fp8.should_convert_module", return_value=True, - ): - with patch( - "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5_4_0", - return_value=False, - ): - oot_replace_with_fp8_linear(model, quantization_config=config) - # At least one replacement happened. - assert len(captured_kwargs) >= 1 - for kw in captured_kwargs: - # On pre-5.4, ``bias`` is the kwarg (not ``has_bias``). - assert "bias" in kw - assert "has_bias" not in kw - assert kw["bias"] is True + ), + patch( + "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5_4_0", + return_value=False, + ), + ): + oot_replace_with_fp8_linear(model, quantization_config=config) + # At least one replacement happened. + assert len(captured_kwargs) >= 1 + for kw in captured_kwargs: + # On pre-5.4, ``bias`` is the kwarg (not ``has_bias``). + assert "bias" in kw + assert "has_bias" not in kw + assert kw["bias"] is True def test_bias_kwarg_name_v5_4_plus(self): """On transformers >= 5.4, the bias flag is passed as ``has_bias``.""" @@ -220,21 +230,23 @@ def _capture(*args, **kwargs): captured_kwargs.append(kwargs) return nn.Linear(8, 8) - with patch("transformers.integrations.finegrained_fp8.FP8Linear", side_effect=_capture): - with patch( + with ( + patch("transformers.integrations.finegrained_fp8.FP8Linear", side_effect=_capture), + patch( "transformers.integrations.finegrained_fp8.should_convert_module", return_value=True, - ): - with patch( - "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5_4_0", - return_value=True, - ): - oot_replace_with_fp8_linear(model, quantization_config=config) - assert len(captured_kwargs) >= 1 - for kw in captured_kwargs: - assert "has_bias" in kw - assert "bias" not in kw - assert kw["has_bias"] is True + ), + patch( + "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5_4_0", + return_value=True, + ), + ): + oot_replace_with_fp8_linear(model, quantization_config=config) + assert len(captured_kwargs) >= 1 + for kw in captured_kwargs: + assert "has_bias" in kw + assert "bias" not in kw + assert kw["has_bias"] is True def test_no_bias_flag_passed_correctly(self): """When the linear module has ``bias=False``, the OOT function must @@ -249,19 +261,21 @@ def _capture(*args, **kwargs): captured_kwargs.append(kwargs) return nn.Linear(8, 8) - with patch("transformers.integrations.finegrained_fp8.FP8Linear", side_effect=_capture): - with patch( + with ( + patch("transformers.integrations.finegrained_fp8.FP8Linear", side_effect=_capture), + patch( "transformers.integrations.finegrained_fp8.should_convert_module", return_value=True, - ): - with patch( - "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5_4_0", - return_value=False, - ): - oot_replace_with_fp8_linear(model, quantization_config=config) - assert len(captured_kwargs) >= 1 - for kw in captured_kwargs: - assert kw.get("bias") is False + ), + patch( + "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5_4_0", + return_value=False, + ), + ): + oot_replace_with_fp8_linear(model, quantization_config=config) + assert len(captured_kwargs) >= 1 + for kw in captured_kwargs: + assert kw.get("bias") is False def test_returns_self(self): """The function returns the (mutated) model object.""" @@ -272,20 +286,22 @@ def test_returns_self(self): def _make_module(*args, **kwargs): return nn.Linear(8, 8) - with patch( - "transformers.integrations.finegrained_fp8.FP8Linear", - side_effect=_make_module, - ): - with patch( + with ( + patch( + "transformers.integrations.finegrained_fp8.FP8Linear", + side_effect=_make_module, + ), + patch( "transformers.integrations.finegrained_fp8.should_convert_module", return_value=True, - ): - with patch( - "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5_4_0", - return_value=True, - ): - result = oot_replace_with_fp8_linear(model, quantization_config=config) - assert result is model + ), + patch( + "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5_4_0", + return_value=True, + ), + ): + result = oot_replace_with_fp8_linear(model, quantization_config=config) + assert result is model # --------------------------------------------------------------------------- @@ -333,23 +349,27 @@ class TestApplyFp8ExpertReplacementPatch: def test_no_cuda_does_nothing(self): """On a non-CUDA host the function must be a no-op without raising.""" - with patch("torch.cuda.is_available", return_value=False): - with patch( + with ( + patch("torch.cuda.is_available", return_value=False), + patch( "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5", return_value=True, - ): - # Should not raise. - assert apply_fp8_expert_replacement_patch() is None + ), + ): + # Should not raise. + assert apply_fp8_expert_replacement_patch() is None def test_old_transformers_does_nothing(self): """With transformers < 5 the function must be a no-op.""" - with patch("torch.cuda.is_available", return_value=True): - with patch( + with ( + patch("torch.cuda.is_available", return_value=True), + patch( "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5", return_value=False, - ): - assert apply_fp8_expert_replacement_patch() is None + ), + ): + assert apply_fp8_expert_replacement_patch() is None def test_import_error_is_swallowed(self): """If the local import of ``transformers.integrations.finegrained_fp8`` @@ -365,14 +385,16 @@ def fake_import(name, globals=None, locals=None, fromlist=(), level=0): raise ImportError("boom") return real_import(name, globals, locals, fromlist, level) - with patch("torch.cuda.is_available", return_value=True): - with patch( + with ( + patch("torch.cuda.is_available", return_value=True), + patch( "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5", return_value=True, - ): - with patch("builtins.__import__", side_effect=fake_import): - # Should not raise despite the ImportError. - assert apply_fp8_expert_replacement_patch() is None + ), + patch("builtins.__import__", side_effect=fake_import), + ): + # Should not raise despite the ImportError. + assert apply_fp8_expert_replacement_patch() is None def test_replaces_upstream_replace_with_fp8_linear(self): """When transformers >= 5 and CUDA is available, the upstream @@ -385,13 +407,15 @@ def test_replaces_upstream_replace_with_fp8_linear(self): original = upstream.replace_with_fp8_linear try: - with patch("torch.cuda.is_available", return_value=True): - with patch( + with ( + patch("torch.cuda.is_available", return_value=True), + patch( "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5", return_value=True, - ): - apply_fp8_expert_replacement_patch() - assert upstream.replace_with_fp8_linear is fp8q.oot_replace_with_fp8_linear + ), + ): + apply_fp8_expert_replacement_patch() + assert upstream.replace_with_fp8_linear is fp8q.oot_replace_with_fp8_linear finally: upstream.replace_with_fp8_linear = original @@ -408,13 +432,15 @@ def test_patches_validate_environment(self): original = FineGrainedFP8HfQuantizer.validate_environment try: - with patch("torch.cuda.is_available", return_value=True): - with patch( + with ( + patch("torch.cuda.is_available", return_value=True), + patch( "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5", return_value=True, - ): - apply_fp8_expert_replacement_patch() - assert FineGrainedFP8HfQuantizer.validate_environment is fp8q.oot_validate_environment + ), + ): + apply_fp8_expert_replacement_patch() + assert FineGrainedFP8HfQuantizer.validate_environment is fp8q.oot_validate_environment finally: FineGrainedFP8HfQuantizer.validate_environment = original @@ -463,9 +489,11 @@ class TestPatchBehaviorMatrix: def test_all_combinations_no_raise(self, cuda_available, transformers_v5): """Every combination of gating conditions must not raise.""" - with patch("torch.cuda.is_available", return_value=cuda_available): - with patch( + with ( + patch("torch.cuda.is_available", return_value=cuda_available), + patch( "auto_round.modeling.fp8_quant.is_transformers_version_greater_or_equal_5", return_value=transformers_v5, - ): - assert apply_fp8_expert_replacement_patch() is None + ), + ): + assert apply_fp8_expert_replacement_patch() is None diff --git a/test/unit/common/models/test_bagel.py b/test/unit/common/models/test_bagel.py index 137005a27a..d8b9a9048b 100644 --- a/test/unit/common/models/test_bagel.py +++ b/test/unit/common/models/test_bagel.py @@ -454,9 +454,7 @@ def test_autoconfig_valueerror_caught_in_base_compressor(self): # The branch diff adds ValueError to the except clause around line 294 # Verify the pattern appears (line numbers may shift) - assert ( - "except (OSError, EnvironmentError, ValueError)" in content - ), "BaseCompressor should catch ValueError alongside OSError/EnvironmentError" + assert "except (OSError, ValueError)" in content, "BaseCompressor should catch ValueError alongside OSError" def test_autoconfig_valueerror_caught_in_model_context(self): """ModelContext: ValueError should be caught in AutoConfig.from_pretrained.""" @@ -465,9 +463,7 @@ def test_autoconfig_valueerror_caught_in_model_context(self): content = f.read() # The branch diff adds ValueError to the except clause around line 146 - assert ( - "except (OSError, EnvironmentError, ValueError)" in content - ), "ModelContext should catch ValueError alongside OSError/EnvironmentError" + assert "except (OSError, ValueError)" in content, "ModelContext should catch ValueError alongside OSError" # ================= Test: mllm_load_model for bagel ================= diff --git a/test/unit/common/models/test_block_names.py b/test/unit/common/models/test_block_names.py index e264614572..771ba34ffa 100644 --- a/test/unit/common/models/test_block_names.py +++ b/test/unit/common/models/test_block_names.py @@ -5,7 +5,7 @@ # ================= simple multimodal model ================= class TextEncoder(nn.Module): def __init__(self, input_size, hidden_size): - super(TextEncoder, self).__init__() + super().__init__() self.fc = nn.Linear(input_size, hidden_size) def forward(self, x): @@ -14,7 +14,7 @@ def forward(self, x): class VisionEncoder(nn.Module): def __init__(self, input_size, hidden_size): - super(VisionEncoder, self).__init__() + super().__init__() self.fc = nn.Linear(input_size, hidden_size) def forward(self, x): @@ -31,7 +31,7 @@ class VisionEncoderModuleList(nn.ModuleList): class SimpleMultimodalModel(nn.Module): def __init__(self, text_input_size, image_input_size, hidden_size, num_text_encoders, num_image_encoders): - super(SimpleMultimodalModel, self).__init__() + super().__init__() self.text_encoders = TextEncoderModuleList( [TextEncoder(text_input_size, hidden_size) for _ in range(num_text_encoders)] ) @@ -56,7 +56,7 @@ def forward(self, text_input, image_input): # ================= simple MoE model ================= class Expert(nn.Module): def __init__(self, input_size, hidden_size): - super(Expert, self).__init__() + super().__init__() self.fc = nn.Linear(input_size, hidden_size) def forward(self, x): @@ -73,7 +73,7 @@ class ExpertModuleList(nn.ModuleList): class NestedMoEModel(nn.Module): def __init__(self, input_size, hidden_size, num_groups, experts_per_group): - super(NestedMoEModel, self).__init__() + super().__init__() self.expert_groups = ExpertModuleList( [ ExpertGroupModuleList([Expert(input_size, hidden_size) for _ in range(experts_per_group)]) diff --git a/test/unit/common/models/test_cosmos3.py b/test/unit/common/models/test_cosmos3.py index 8f0713cbf1..d73862ad54 100644 --- a/test/unit/common/models/test_cosmos3.py +++ b/test/unit/common/models/test_cosmos3.py @@ -38,9 +38,8 @@ def test_forward_mode_sets_and_clears_state(self): def test_forward_mode_cleanup_on_exception(self): from auto_round.special_model_handler import _cosmos3_forward_mode, _cosmos3_forward_state - with pytest.raises(RuntimeError): - with _cosmos3_forward_mode(): - raise RuntimeError("simulated error") + with pytest.raises(RuntimeError), _cosmos3_forward_mode(): + raise RuntimeError("simulated error") assert not getattr(_cosmos3_forward_state, "active", False) def test_forward_mode_nested_context(self): @@ -591,12 +590,14 @@ def test_non_cosmos3_pipeline_skips_cosmos3_loader(self): tmpdir = self._write_model_index("DDPMScheduler") try: - with patch("auto_round.special_model_handler.load_cosmos3_diffusion") as mock_cosmos: - with patch("auto_round.utils.common.LazyImport", return_value=MagicMock()): - try: - diffusion_load_model(tmpdir) - except Exception: - pass - mock_cosmos.assert_not_called() + with ( + patch("auto_round.special_model_handler.load_cosmos3_diffusion") as mock_cosmos, + patch("auto_round.utils.common.LazyImport", return_value=MagicMock()), + ): + try: + diffusion_load_model(tmpdir) + except Exception: + pass + mock_cosmos.assert_not_called() finally: shutil.rmtree(tmpdir) diff --git a/test/unit/common/models/test_moe_model.py b/test/unit/common/models/test_moe_model.py index 91fcd71788..e2d0b87a79 100644 --- a/test/unit/common/models/test_moe_model.py +++ b/test/unit/common/models/test_moe_model.py @@ -13,7 +13,7 @@ def quantize_model(model, output_dir, scheme, iters=0, ignore_layers="self_attn,router,lm_head,mlp.gate"): """Helper function to quantize the model with the given scheme.""" - disable_opt_rtn = True if iters == 0 else False + disable_opt_rtn = iters == 0 autoround = AutoRound( model, scheme=scheme, diff --git a/test/unit/common/models/test_unfused_moe_init.py b/test/unit/common/models/test_unfused_moe_init.py index 8dfae1fcb4..03506fe5bb 100644 --- a/test/unit/common/models/test_unfused_moe_init.py +++ b/test/unit/common/models/test_unfused_moe_init.py @@ -106,19 +106,21 @@ def test_get_checkpoint_conversion_mapping_ar_passthrough(): transformers mapping function. We mock that to verify the call. """ sentinel = ["x", "y"] - with mock.patch( - "transformers.conversion_mapping.orig_get_checkpoint_conversion_mapping", - create=True, - return_value=sentinel, - new_callable=mock.MagicMock, - ) as fake: - # Some transformers versions don't expose ``orig_*`` yet; guard - # against that by also patching the public name. - with mock.patch( + # Some transformers versions don't expose ``orig_*`` yet; guard + # against that by also patching the public name. + with ( + mock.patch( + "transformers.conversion_mapping.orig_get_checkpoint_conversion_mapping", + create=True, + return_value=sentinel, + new_callable=mock.MagicMock, + ) as fake, + mock.patch( "transformers.conversion_mapping.get_checkpoint_conversion_mapping", side_effect=lambda mt: sentinel, - ): - result = get_checkpoint_conversion_mapping_ar("not_in_our_list") + ), + ): + result = get_checkpoint_conversion_mapping_ar("not_in_our_list") assert result == sentinel diff --git a/test/unit/common/models/test_vlm_ram_reduction.py b/test/unit/common/models/test_vlm_ram_reduction.py index 2e78813d95..bd841146b7 100644 --- a/test/unit/common/models/test_vlm_ram_reduction.py +++ b/test/unit/common/models/test_vlm_ram_reduction.py @@ -167,11 +167,10 @@ def __init__(self): pipe = FakePipe() pipe.components = {"transformer": pipe.transformer, "vae": pipe.vae} - with patch("torch.cuda.is_available", return_value=False): - with patch("torch.cuda.device_count", return_value=0): - # Single device path (falls back to pipe.to) - result = dispatch_model_by_all_available_devices(pipe, "cpu") - assert result is pipe + with patch("torch.cuda.is_available", return_value=False), patch("torch.cuda.device_count", return_value=0): + # Single device path (falls back to pipe.to) + result = dispatch_model_by_all_available_devices(pipe, "cpu") + assert result is pipe def test_multi_device_respects_non_main_memory(self): """With multiple devices, memory reservation must be computed for non-main components.""" diff --git a/test/unit/common/schemes/test_w8_asym_policy.py b/test/unit/common/schemes/test_w8_asym_policy.py index 732439deb6..f9ba87564f 100644 --- a/test/unit/common/schemes/test_w8_asym_policy.py +++ b/test/unit/common/schemes/test_w8_asym_policy.py @@ -82,7 +82,7 @@ def _scheme(self): return AutoScheme(options=["W8A16", "W4A16"], avg_bits=6.0) def _w8_option(self, scheme): - return [o for o in scheme.options if getattr(o, "bits", None) == 8][0] + return next(o for o in scheme.options if getattr(o, "bits", None) == 8) @pytest.mark.parametrize( "fmt,opt_in,expect_sym", diff --git a/test/unit/common/utils/test_common_pure_helpers.py b/test/unit/common/utils/test_common_pure_helpers.py index 2d712867aa..1744c3cd25 100644 --- a/test/unit/common/utils/test_common_pure_helpers.py +++ b/test/unit/common/utils/test_common_pure_helpers.py @@ -381,7 +381,7 @@ def test_nested_dict(self): def test_invalid_input_raises(self): from auto_round.utils.common import parse_layer_config_arg - with pytest.raises(Exception): + with pytest.raises(ValueError): parse_layer_config_arg("") def test_multi_key_unquoted_dict(self): diff --git a/test/unit/common/utils/test_device.py b/test/unit/common/utils/test_device.py index 836652807b..dc2962a819 100644 --- a/test/unit/common/utils/test_device.py +++ b/test/unit/common/utils/test_device.py @@ -81,8 +81,9 @@ def test_compile_mode_true(self): from auto_round.utils.device import _use_hpu_compile_mode # Mock both is_hpu_lazy_mode and TORCH_VERSION_AT_LEAST_2_4 (imported inside function) - with patch("auto_round.utils.device.is_hpu_lazy_mode", return_value=False), patch.dict( - "sys.modules", {"auto_round.utils.common": MagicMock(TORCH_VERSION_AT_LEAST_2_4=True)} + with ( + patch("auto_round.utils.device.is_hpu_lazy_mode", return_value=False), + patch.dict("sys.modules", {"auto_round.utils.common": MagicMock(TORCH_VERSION_AT_LEAST_2_4=True)}), ): result = _use_hpu_compile_mode() assert result is True @@ -98,8 +99,9 @@ def test_compile_mode_false_torch_old(self): """Test compile mode False when torch < 2.4.""" from auto_round.utils.device import _use_hpu_compile_mode - with patch("auto_round.utils.device.is_hpu_lazy_mode", return_value=False), patch.dict( - "sys.modules", {"auto_round.utils.common": MagicMock(TORCH_VERSION_AT_LEAST_2_4=False)} + with ( + patch("auto_round.utils.device.is_hpu_lazy_mode", return_value=False), + patch.dict("sys.modules", {"auto_round.utils.common": MagicMock(TORCH_VERSION_AT_LEAST_2_4=False)}), ): result = _use_hpu_compile_mode() assert result is False @@ -118,11 +120,13 @@ def test_bump_with_explicit_min_size(self): mock_config.accumulated_cache_size_limit = 8 mock_config.recompile_limit = 8 - with patch.dict("sys.modules", {"torch._dynamo.config": mock_config}): - with patch("torch._dynamo.config", mock_config): - _bump_dynamo_cache_limit(min_size=32) - # Function should attempt to set values >= 32 - # Best effort - it may or may not raise depending on imports + with ( + patch.dict("sys.modules", {"torch._dynamo.config": mock_config}), + patch("torch._dynamo.config", mock_config), + ): + _bump_dynamo_cache_limit(min_size=32) + # Function should attempt to set values >= 32 + # Best effort - it may or may not raise depending on imports def test_bump_without_value_uses_default(self): """Test _bump_dynamo_cache_limit without min_size uses env default.""" @@ -1125,17 +1129,19 @@ def test_multi_device_uses_accelerate(self): # Multi-device path: provide 2 "cpu" entries so ``len(devices) > 1``. # After the inner loop dedupes, ``device == "cpu"`` is used to index # the mocked max_memory dict. - with patch( - "auto_round.utils.device.parse_available_devices", - return_value=["cpu", "cpu"], - ), patch( - "auto_round.utils.device.get_max_memory", return_value={"cpu": 1024} - ), patch("auto_round.utils.device.get_balanced_memory", return_value={"cpu": 512}), patch( - "auto_round.utils.device.infer_auto_device_map", - return_value={"0": "cpu"}, - ) as mock_infer, patch( - "auto_round.utils.device.dispatch_model", return_value="MOCKED" - ) as mock_dispatch: + with ( + patch( + "auto_round.utils.device.parse_available_devices", + return_value=["cpu", "cpu"], + ), + patch("auto_round.utils.device.get_max_memory", return_value={"cpu": 1024}), + patch("auto_round.utils.device.get_balanced_memory", return_value={"cpu": 512}), + patch( + "auto_round.utils.device.infer_auto_device_map", + return_value={"0": "cpu"}, + ) as mock_infer, + patch("auto_round.utils.device.dispatch_model", return_value="MOCKED") as mock_dispatch, + ): result = dispatch_model_block_wise(model, device_map="cpu,cpu", max_mem_ratio=0.5) assert mock_infer.called assert mock_dispatch.called @@ -1164,9 +1170,11 @@ def test_auto_branch_uses_max_memory(self): from auto_round.utils.device import dispatch_model_by_all_available_devices model = MagicMock(spec=nn.Module) - with patch("auto_round.utils.device.get_balanced_memory", return_value={0: 1024}) as balanced, patch( - "auto_round.utils.device.infer_auto_device_map", return_value={"0": "cpu"} - ), patch("auto_round.utils.device.dispatch_model", return_value="AUTO_MODEL"): + with ( + patch("auto_round.utils.device.get_balanced_memory", return_value={0: 1024}) as balanced, + patch("auto_round.utils.device.infer_auto_device_map", return_value={"0": "cpu"}), + patch("auto_round.utils.device.dispatch_model", return_value="AUTO_MODEL"), + ): with patch( "auto_round.utils.device.parse_available_devices", return_value=["cpu"], @@ -1455,9 +1463,11 @@ def test_context_manager_runs(self): def test_context_manager_with_warning_level(self): from auto_round.utils.device import dump_memory_usage_ctx - with patch("auto_round.utils.device.logger.warning") as warn: - with dump_memory_usage_ctx(msg="warn-ctx", log_level="warning"): - pass + with ( + patch("auto_round.utils.device.logger.warning") as warn, + dump_memory_usage_ctx(msg="warn-ctx", log_level="warning"), + ): + pass assert warn.called def test_decorator_runs_function(self): @@ -1720,12 +1730,13 @@ def test_context_manager_restores_state(self): from auto_round.utils.device import fake_cuda_for_hpu original = MagicMock(return_value=True) - with patch("auto_round.utils.device.is_hpex_available", return_value=True), patch( - "torch.cuda.is_available", original + with ( + patch("auto_round.utils.device.is_hpex_available", return_value=True), + patch("torch.cuda.is_available", original), + fake_cuda_for_hpu(), ): - with fake_cuda_for_hpu(): - # Should be temporarily faked. - pass + # Should be temporarily faked. + pass # After exit, original is restored. # We can't strictly assert identity due to dynamic restoration, # but should at least have called __exit__ without raising. @@ -1738,13 +1749,15 @@ def test_with_existing_triton(self): from auto_round.utils.device import fake_triton_for_hpu # Create a fake triton module - with patch.dict( - sys.modules, - {"triton": MagicMock(), "triton.language": MagicMock()}, + with ( + patch.dict( + sys.modules, + {"triton": MagicMock(), "triton.language": MagicMock()}, + ), + patch("auto_round.utils.device.is_hpex_available", return_value=True), + fake_triton_for_hpu(), ): - with patch("auto_round.utils.device.is_hpex_available", return_value=True): - with fake_triton_for_hpu(): - pass + pass # =========================================================================== @@ -1853,9 +1866,10 @@ def test_counter_increments(self): from auto_round.utils import device as device_mod device_mod._malloc_trim_counter = 0 - with patch.dict(os.environ, {"AR_ENABLE_MALLOC_TRIM": "1"}, clear=False), patch( - "auto_round.utils.device.ctypes.CDLL" - ) as mock_cdll: + with ( + patch.dict(os.environ, {"AR_ENABLE_MALLOC_TRIM": "1"}, clear=False), + patch("auto_round.utils.device.ctypes.CDLL") as mock_cdll, + ): mock_libc = MagicMock() mock_cdll.return_value = mock_libc from auto_round.utils.device import _maybe_trim_malloc @@ -1870,11 +1884,14 @@ def test_counter_increments(self): def test_invalid_every_falls_back_to_default(self): from auto_round.utils.device import _maybe_trim_malloc - with patch.dict( - os.environ, - {"AR_ENABLE_MALLOC_TRIM": "1", "AR_MALLOC_TRIM_EVERY": "notanumber"}, - clear=False, - ), patch("auto_round.utils.device.ctypes.CDLL") as mock_cdll: + with ( + patch.dict( + os.environ, + {"AR_ENABLE_MALLOC_TRIM": "1", "AR_MALLOC_TRIM_EVERY": "notanumber"}, + clear=False, + ), + patch("auto_round.utils.device.ctypes.CDLL") as mock_cdll, + ): mock_libc = MagicMock() mock_cdll.return_value = mock_libc _maybe_trim_malloc() @@ -1885,11 +1902,14 @@ def test_negative_or_zero_every_normalised_to_one(self): from auto_round.utils import device as device_mod from auto_round.utils.device import _maybe_trim_malloc - with patch.dict( - os.environ, - {"AR_ENABLE_MALLOC_TRIM": "1", "AR_MALLOC_TRIM_EVERY": "0"}, - clear=False, - ), patch("auto_round.utils.device.ctypes.CDLL") as mock_cdll: + with ( + patch.dict( + os.environ, + {"AR_ENABLE_MALLOC_TRIM": "1", "AR_MALLOC_TRIM_EVERY": "0"}, + clear=False, + ), + patch("auto_round.utils.device.ctypes.CDLL") as mock_cdll, + ): mock_libc = MagicMock() mock_cdll.return_value = mock_libc diff --git a/test/unit/common/utils/test_device_manager.py b/test/unit/common/utils/test_device_manager.py index f6a22de257..f93445d24c 100644 --- a/test/unit/common/utils/test_device_manager.py +++ b/test/unit/common/utils/test_device_manager.py @@ -358,18 +358,22 @@ def test_returns_cpu_when_nothing_available(self): ) # Force every discovery path to report "nothing" - with patch.object( - __import__("auto_round.utils.device_manager", fromlist=["_hpu_available"]), - "_hpu_available", - return_value=False, - ), patch.object( - __import__("auto_round.utils.device_manager", fromlist=["_torch_accelerator_type"]), - "_torch_accelerator_type", - return_value=None, - ), patch.object( - __import__("auto_round.utils.device_manager", fromlist=["_PREFERRED_ORDER"]), - "_PREFERRED_ORDER", - (), + with ( + patch.object( + __import__("auto_round.utils.device_manager", fromlist=["_hpu_available"]), + "_hpu_available", + return_value=False, + ), + patch.object( + __import__("auto_round.utils.device_manager", fromlist=["_torch_accelerator_type"]), + "_torch_accelerator_type", + return_value=None, + ), + patch.object( + __import__("auto_round.utils.device_manager", fromlist=["_PREFERRED_ORDER"]), + "_PREFERRED_ORDER", + (), + ), ): # Need to clear the lru_cache get_current_device_type.cache_clear() @@ -415,8 +419,9 @@ def test_is_device_available_when_accelerator(self): def test_get_available_device_types_cpu_only(self): from auto_round.utils.device_manager import get_available_device_types - with patch("auto_round.utils.device_manager._hpu_available", return_value=False), patch( - "auto_round.utils.device_manager._torch_accelerator_type", return_value=None + with ( + patch("auto_round.utils.device_manager._hpu_available", return_value=False), + patch("auto_round.utils.device_manager._torch_accelerator_type", return_value=None), ): assert get_available_device_types() == [] @@ -766,12 +771,14 @@ class TestGetDeviceMemory: def test_cpu_raises_runtime_error(self): from auto_round.utils.device_manager import get_device_memory - with patch( - "auto_round.utils.device_manager.get_current_device_type", - return_value="cpu", + with ( + patch( + "auto_round.utils.device_manager.get_current_device_type", + return_value="cpu", + ), + pytest.raises(RuntimeError), ): - with pytest.raises(RuntimeError): - get_device_memory() + get_device_memory() # --------------------------------------------------------------------------- diff --git a/test/unit/common/utils/test_distributed.py b/test/unit/common/utils/test_distributed.py index a83c8951b2..519d094550 100644 --- a/test/unit/common/utils/test_distributed.py +++ b/test/unit/common/utils/test_distributed.py @@ -27,26 +27,32 @@ class TestIsDistributed: """Test is_distributed with mocked torch.distributed.""" def test_not_initialized(self): - with patch("auto_round.utils.distributed.is_distributed", return_value=False): - with patch("torch.distributed.is_initialized", return_value=False): - # Clear cache to ensure fresh evaluation - is_distributed.cache_clear() - result = is_distributed() - assert result is False + with ( + patch("auto_round.utils.distributed.is_distributed", return_value=False), + patch("torch.distributed.is_initialized", return_value=False), + ): + # Clear cache to ensure fresh evaluation + is_distributed.cache_clear() + result = is_distributed() + assert result is False def test_initialized_single_device(self): - with patch("torch.distributed.is_initialized", return_value=True): - with patch("torch.distributed.get_world_size", return_value=1): - is_distributed.cache_clear() - result = is_distributed() - assert result is False + with ( + patch("torch.distributed.is_initialized", return_value=True), + patch("torch.distributed.get_world_size", return_value=1), + ): + is_distributed.cache_clear() + result = is_distributed() + assert result is False def test_initialized_multi_device(self): - with patch("torch.distributed.is_initialized", return_value=True): - with patch("torch.distributed.get_world_size", return_value=4): - is_distributed.cache_clear() - result = is_distributed() - assert result is True + with ( + patch("torch.distributed.is_initialized", return_value=True), + patch("torch.distributed.get_world_size", return_value=4), + ): + is_distributed.cache_clear() + result = is_distributed() + assert result is True def test_dist_not_initialized(self): with patch("torch.distributed.is_initialized", return_value=False): @@ -55,18 +61,22 @@ def test_dist_not_initialized(self): is_distributed.cache_clear() def test_dist_initialized_single_world(self): - with patch("torch.distributed.is_initialized", return_value=True): - with patch("torch.distributed.get_world_size", return_value=1): - is_distributed.cache_clear() - assert is_distributed() is False - is_distributed.cache_clear() + with ( + patch("torch.distributed.is_initialized", return_value=True), + patch("torch.distributed.get_world_size", return_value=1), + ): + is_distributed.cache_clear() + assert is_distributed() is False + is_distributed.cache_clear() def test_dist_initialized_multi_world(self): - with patch("torch.distributed.is_initialized", return_value=True): - with patch("torch.distributed.get_world_size", return_value=2): - is_distributed.cache_clear() - assert is_distributed() is True - is_distributed.cache_clear() + with ( + patch("torch.distributed.is_initialized", return_value=True), + patch("torch.distributed.get_world_size", return_value=2), + ): + is_distributed.cache_clear() + assert is_distributed() is True + is_distributed.cache_clear() class TestNoopSync: @@ -164,13 +174,15 @@ def test_single_device_ddp(self): pytest.skip("CUDA not available") model = nn.Linear(4, 4).to("cpu") - with patch("torch.distributed.is_initialized", return_value=True): - with patch("torch.distributed.get_world_size", return_value=2): - with patch("torch.distributed.get_rank", return_value=0): - with patch("torch.nn.parallel.DistributedDataParallel"): - ar = SimpleNamespace() - block, sync_fn = setup_ddp_if_needed_(ar, model, [0]) - assert sync_fn is _noop_sync + with ( + patch("torch.distributed.is_initialized", return_value=True), + patch("torch.distributed.get_world_size", return_value=2), + patch("torch.distributed.get_rank", return_value=0), + patch("torch.nn.parallel.DistributedDataParallel"), + ): + ar = SimpleNamespace() + block, sync_fn = setup_ddp_if_needed_(ar, model, [0]) + assert sync_fn is _noop_sync def test_single_device_returns_noop_when_distributed(self): """Test the multi-GPU case which doesn't need to move to GPU device.""" @@ -201,12 +213,14 @@ def test_multi_device_manual_reduce(self): model = nn.Linear(4, 4) - with patch("torch.distributed.is_initialized", return_value=True): - with patch("torch.distributed.get_world_size", return_value=4): - with patch("torch.distributed.get_rank", return_value=0): - ar = SimpleNamespace() - block, sync_fn = setup_ddp_if_needed_(ar, model, [0, 1]) - # Should not be the noop - assert sync_fn is not _noop_sync - # Calling it should not raise - sync_fn() + with ( + patch("torch.distributed.is_initialized", return_value=True), + patch("torch.distributed.get_world_size", return_value=4), + patch("torch.distributed.get_rank", return_value=0), + ): + ar = SimpleNamespace() + block, sync_fn = setup_ddp_if_needed_(ar, model, [0, 1]) + # Should not be the noop + assert sync_fn is not _noop_sync + # Calling it should not raise + sync_fn() diff --git a/test/unit/common/utils/test_generation.py b/test/unit/common/utils/test_generation.py index b5635face7..f7bf507ced 100644 --- a/test/unit/common/utils/test_generation.py +++ b/test/unit/common/utils/test_generation.py @@ -63,7 +63,7 @@ def test_autoround_sym(self, dataloader): for bits in [4]: model = AutoModelForCausalLM.from_pretrained(self.model_name, torch_dtype="auto", trust_remote_code=True) tokenizer = AutoTokenizer.from_pretrained(self.model_name, trust_remote_code=True) - bits, group_size, sym = bits, 128, True + group_size, sym = 128, True autoround = AutoRound( model, tokenizer, diff --git a/test/unit/common/utils/test_missing_tensors.py b/test/unit/common/utils/test_missing_tensors.py index 2867450a42..e727b4b8e6 100644 --- a/test/unit/common/utils/test_missing_tensors.py +++ b/test/unit/common/utils/test_missing_tensors.py @@ -134,8 +134,8 @@ def test_2d_and_non_expert_pass_through(self): } result = split_fused_expert_tensors(tensors) assert set(result.keys()) == set(tensors.keys()) - for k in tensors: - assert torch.equal(result[k], tensors[k]) + for k, v in tensors.items(): + assert torch.equal(result[k], v) def test_warns_on_3d_tensor_with_unsupported_parent(self, caplog, _autoround_log_propagate): tensors = { diff --git a/test/unit/common/utils/test_model_utils.py b/test/unit/common/utils/test_model_utils.py index c561de4afa..235712c1ff 100644 --- a/test/unit/common/utils/test_model_utils.py +++ b/test/unit/common/utils/test_model_utils.py @@ -964,7 +964,7 @@ class MockModule: set_attr(model, "inner.new_attr", "new_value") - assert getattr(model.inner, "new_attr") == "new_value" + assert model.inner.new_attr == "new_value" def test_set_attr_missing_parent(self): """Test set_attr does not raise when the parent path doesn't exist.""" diff --git a/test/unit/common/utils/test_weight_handler.py b/test/unit/common/utils/test_weight_handler.py index 418ee67d45..8a2a1eb44f 100644 --- a/test/unit/common/utils/test_weight_handler.py +++ b/test/unit/common/utils/test_weight_handler.py @@ -871,7 +871,6 @@ class CustomWeightType: # Actually, ModuleWeightType is an Enum, so we can't easily create a new one # Let's just verify that unknown combinations return None # The function should return None for any unregistered type - pass # ============================================================================== diff --git a/test/unit/envs.py b/test/unit/envs.py index 5b6b5730c8..dab67083da 100644 --- a/test/unit/envs.py +++ b/test/unit/envs.py @@ -14,8 +14,9 @@ import importlib.util import unittest +from collections.abc import Callable from functools import wraps -from typing import Callable, Literal +from typing import Literal import torch from transformers.utils.versions import require_version diff --git a/test/unit/test_cpu/core/test_autoround.py b/test/unit/test_cpu/core/test_autoround.py index d639be0859..f1dff0278e 100644 --- a/test/unit/test_cpu/core/test_autoround.py +++ b/test/unit/test_cpu/core/test_autoround.py @@ -238,7 +238,6 @@ def test_disable_minmax_tuning(self, dataloader): ) autoround.quantize() - # def test_signround(self, tiny_opt_model_path, dataloader): model_name = tiny_opt_model_path bits, group_size, sym = 4, -1, False diff --git a/test/unit/test_cpu/core/test_resume_integration.py b/test/unit/test_cpu/core/test_resume_integration.py index 6396a65aa5..2225754163 100644 --- a/test/unit/test_cpu/core/test_resume_integration.py +++ b/test/unit/test_cpu/core/test_resume_integration.py @@ -60,10 +60,12 @@ def crash_after_first_block(self, block_name, q_input, input_ids): if len(crashed_after) == 1: raise RuntimeError("simulated crash") - with mock.patch.object(ResumeState, "mark_block_done", crash_after_first_block): - with pytest.raises(RuntimeError, match="simulated crash"): - ar = AutoRound(model=tiny_opt_model_path, scheme="W4A16", iters=1, nsamples=1) - ar.quantize() + with ( + mock.patch.object(ResumeState, "mark_block_done", crash_after_first_block), + pytest.raises(RuntimeError, match="simulated crash"), + ): + ar = AutoRound(model=tiny_opt_model_path, scheme="W4A16", iters=1, nsamples=1) + ar.quantize() assert crashed_after == ["model.decoder.layers.0"] diff --git a/test/unit/test_cpu/export/test_llmc_format.py b/test/unit/test_cpu/export/test_llmc_format.py index 8078530471..e981551769 100644 --- a/test/unit/test_cpu/export/test_llmc_format.py +++ b/test/unit/test_cpu/export/test_llmc_format.py @@ -265,7 +265,8 @@ def test_llmcompressor_fp8(self, tmp_path): from safetensors import safe_open - config = json.load(open(os.path.join(quantized_model_path, "config.json"))) + with open(os.path.join(quantized_model_path, "config.json")) as f: + config = json.load(f) assert "group_0" in config["quantization_config"]["config_groups"] assert config["quantization_config"]["config_groups"]["group_0"]["input_activations"]["num_bits"] == 8 assert config["quantization_config"]["config_groups"]["group_0"]["weights"]["strategy"] == "channel" @@ -289,7 +290,8 @@ def test_autoround_llmcompressor_fp8(self, tmp_path): import json - config = json.load(open(os.path.join(quantized_model_path, "config.json"))) + with open(os.path.join(quantized_model_path, "config.json")) as f: + config = json.load(f) assert "group_0" in config["quantization_config"]["config_groups"] assert config["quantization_config"]["config_groups"]["group_0"]["input_activations"]["num_bits"] == 8 assert config["quantization_config"]["config_groups"]["group_0"]["weights"]["strategy"] == "tensor" diff --git a/test/unit/test_cpu/models/test_block_names.py b/test/unit/test_cpu/models/test_block_names.py index c97999593e..0876d65552 100644 --- a/test/unit/test_cpu/models/test_block_names.py +++ b/test/unit/test_cpu/models/test_block_names.py @@ -61,7 +61,7 @@ def test_mm_block_name(self, tiny_qwen_vl_model_path): model = Qwen2VLForConditionalGeneration.from_pretrained(model_name, trust_remote_code=True, device_map="auto") block_name = get_block_names(model, quant_vision=True) assert len(block_name) == 2 - assert all(["visual.merger.mlp" not in n for n in block_name]) + assert all("visual.merger.mlp" not in n for n in block_name) block_name = get_block_names(model, quant_vision=False) assert len(block_name) == 1 assert block_name == get_block_names(model) diff --git a/test/unit/test_cpu/quantization/test_model_free_parity.py b/test/unit/test_cpu/quantization/test_model_free_parity.py old mode 100755 new mode 100644 diff --git a/test/unit/test_cpu/quantization/test_neuqi.py b/test/unit/test_cpu/quantization/test_neuqi.py index b4e4b63c33..ab4769c4c5 100644 --- a/test/unit/test_cpu/quantization/test_neuqi.py +++ b/test/unit/test_cpu/quantization/test_neuqi.py @@ -795,10 +795,10 @@ def test_two_stage_core_uses_shared_launch(self, monkeypatch): def spy_launch(dn, q, fracs, invf, nm): spy_launch.called = True - return None # decline -> per-candidate fallback + # implicit None: decline -> per-candidate fallback spy_launch.called = False - monkeypatch.setattr(N, "_sym_coarse_pass_shared", lambda *a, **k: (None if not spy_launch.called else None)) + monkeypatch.setattr(N, "_sym_coarse_pass_shared", lambda *a, **k: None) # direct wiring check: the core calls _sym_coarse_pass_shared at all real = N._sym_coarse_pass_shared @@ -996,7 +996,7 @@ def test_search_wires_shared_coarse(self, monkeypatch): def wrapper(data_, qw_, s0_, coarse_, maxq_): wrapper.called = True - return None # decline -> batched fallback still correct + # implicit None: decline -> batched fallback still correct wrapper.called = False monkeypatch.setattr(N, "_zp_coarse_pass_shared", wrapper) @@ -1520,7 +1520,7 @@ def _quantizer(self, enable_neuqi=True, disable_opt_rtn=False, is_moe=True): cfg = OptimizedRTNConfig(bits=4, group_size=32, sym=False, enable_neuqi=enable_neuqi) cfg.disable_opt_rtn = disable_opt_rtn - cfg.orig_disable_opt_rtn = False if not disable_opt_rtn else True + cfg.orig_disable_opt_rtn = bool(disable_opt_rtn) q = OptimizedRTNQuantizer(cfg) if is_moe: # model_context is a read-only property backed by the run ctx; diff --git a/test/unit/test_cuda/advanced/test_multiple_card.py b/test/unit/test_cuda/advanced/test_multiple_card.py index e84d954707..1e5ed76aea 100644 --- a/test/unit/test_cuda/advanced/test_multiple_card.py +++ b/test/unit/test_cuda/advanced/test_multiple_card.py @@ -128,8 +128,8 @@ def test_device_map_for_triton(self): model_name = "OPEA/Qwen2.5-0.5B-Instruct-int4-sym-inc" device_map = {} - for i in range(0, 32): - key = f"model.layers.{str(i)}" + for i in range(32): + key = f"model.layers.{i!s}" device_map[key] = "cuda:0" device_map["model.layers.1"] = "cpu" device_map["model.layers.2"] = "cpu" diff --git a/test/unit/test_cuda/backends/test_marlin_backend.py b/test/unit/test_cuda/backends/test_marlin_backend.py index 5b868c5baf..bf142f12ed 100644 --- a/test/unit/test_cuda/backends/test_marlin_backend.py +++ b/test/unit/test_cuda/backends/test_marlin_backend.py @@ -67,7 +67,7 @@ def test_marlin_group_size(self, dataloader): print(f"{group_size}!!!!!!!!!!!!!!!!!") model = AutoModelForCausalLM.from_pretrained(self.model_name, torch_dtype="auto", trust_remote_code=True) tokenizer = AutoTokenizer.from_pretrained(self.model_name, trust_remote_code=True) - bits, group_size, sym = 4, group_size, True + bits, sym = 4, True autoround = AutoRound( model, tokenizer, @@ -97,7 +97,7 @@ def test_marlin_group_size(self, dataloader): print(f"{group_size}!!!!!!!!!!!!!!!!!") model = AutoModelForCausalLM.from_pretrained(self.model_name, torch_dtype="auto", trust_remote_code=True) tokenizer = AutoTokenizer.from_pretrained(self.model_name, trust_remote_code=True) - bits, group_size, sym = 4, group_size, True + bits, sym = 4, True autoround = AutoRound( model, tokenizer, diff --git a/test/unit/test_cuda/models/test_get_block_name.py b/test/unit/test_cuda/models/test_get_block_name.py index 633b797c26..b3190273f3 100644 --- a/test/unit/test_cuda/models/test_get_block_name.py +++ b/test/unit/test_cuda/models/test_get_block_name.py @@ -29,7 +29,7 @@ def setup_class(self): def teardown_class(self): shutil.rmtree("runs", ignore_errors=True) - def check_block_names(self, block_names, prefixs=[], n_layers=[]): + def check_block_names(self, block_names, prefixs=(), n_layers=()): assert len(block_names) == len(prefixs) == len(n_layers) for i, block_name in enumerate(block_names): prefix = prefixs[i] @@ -213,7 +213,7 @@ def test_flux(self): block_names = get_block_names(model) self.check_block_names(block_names, ["transformer_blocks", "single_transformer_blocks"], [19, 38]) - assert any(["context_embedder" not in n for n in block_names]) + assert any("context_embedder" not in n for n in block_names) block_names = get_block_names(model, quant_vision=True) self.check_block_names(block_names, ["transformer_blocks", "single_transformer_blocks"], [19, 38]) diff --git a/test/unit/test_cuda/models/test_mllm.py b/test/unit/test_cuda/models/test_mllm.py index c91b99325d..630667d938 100644 --- a/test/unit/test_cuda/models/test_mllm.py +++ b/test/unit/test_cuda/models/test_mllm.py @@ -137,8 +137,8 @@ def test_mm_block_name(self): model = MllamaForConditionalGeneration.from_pretrained(model_name, trust_remote_code=True, device_map="auto") block_name = get_block_names(model, quant_vision=True) assert len(block_name) == 3 - assert any(["vision_model.global_transformer.layers.0" not in n for n in block_name]) - assert any(["vision_model.transformer.layers.0" not in n for n in block_name]) + assert any("vision_model.global_transformer.layers.0" not in n for n in block_name) + assert any("vision_model.transformer.layers.0" not in n for n in block_name) block_name = get_block_names(model, quant_vision=False) assert len(block_name) == 1 assert get_block_names(model) == block_name diff --git a/test/unit/test_mlx/test_mlx_format.py b/test/unit/test_mlx/test_mlx_format.py index a302e3bd8b..51ccfe9e87 100644 --- a/test/unit/test_mlx/test_mlx_format.py +++ b/test/unit/test_mlx/test_mlx_format.py @@ -78,7 +78,7 @@ def _quantize_and_save( fmt: str, sym: bool = True, group_size: int = 128, - layer_config: dict = None, + layer_config: dict | None = None, ): """Run AutoRound RTN quantization on Qwen3-0.6B and export to ``fmt``.""" ar = AutoRound( diff --git a/test/unit/test_xpu/test_autoround.py b/test/unit/test_xpu/test_autoround.py index da32e35a63..ee4264cc6e 100644 --- a/test/unit/test_xpu/test_autoround.py +++ b/test/unit/test_xpu/test_autoround.py @@ -16,12 +16,10 @@ class TestAutoRoundXPU: @classmethod def setup_class(self): self.device = "xpu" - pass @classmethod def teardown_class(self): shutil.rmtree("runs", ignore_errors=True) - pass @pytest.fixture(autouse=True) def _save_dir(self, tmp_path):