Source code for arcade.classifieroptimization.base

"""Base optimizer class and registry."""

from __future__ import annotations

import importlib
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Any

_OPTIMIZER_REGISTRY: dict[str, type["BaseOptimizer"]] = {}


[docs] def register_optimizer(name: str): """Decorator to register an optimizer class by name.""" def decorator(cls: type["BaseOptimizer"]) -> type["BaseOptimizer"]: _OPTIMIZER_REGISTRY[name] = cls cls.name = name return cls return decorator
[docs] def get_optimizer(name: str) -> type["BaseOptimizer"]: """Look up an optimizer by registry name or dotted import path.""" if name in _OPTIMIZER_REGISTRY: return _OPTIMIZER_REGISTRY[name] if "." in name: module_path, cls_name = name.rsplit(".", 1) mod = importlib.import_module(module_path) cls = getattr(mod, cls_name) if isinstance(cls, type) and issubclass(cls, BaseOptimizer): return cls raise TypeError(f"{name} resolved to {cls!r}, not a BaseOptimizer subclass.") raise KeyError( f"Unknown optimizer {name!r}. Registered: {sorted(_OPTIMIZER_REGISTRY)}." )
[docs] def list_optimizers() -> list[str]: """Return the names of all registered optimizers.""" return sorted(_OPTIMIZER_REGISTRY)
[docs] @dataclass class OptimizationResult: """Result returned by an optimizer's ``optimize()`` method.""" model: Any original_params: int = 0 optimized_params: int = 0 accuracy_before: float = 0.0 accuracy_after: float = 0.0 pareto_points: list[Any] = field(default_factory=list) metadata: dict[str, Any] = field(default_factory=dict)
[docs] class BaseOptimizer(ABC): """Interface for model optimization (pruning, distillation, NAS, etc.). Subclass and decorate with ``@register_optimizer("name")``. Example:: @register_optimizer("my_opt") class MyOptimizer(BaseOptimizer): def optimize(self, model, train_data, val_data=None, **kwargs): ... return OptimizationResult(model=optimized) """ name: str = "base"
[docs] @abstractmethod def optimize( self, model: Any, train_data: tuple[Any, Any], val_data: tuple[Any, Any] | None = None, **kwargs: Any, ) -> OptimizationResult: """Apply optimization to a model. Args: model: The model to optimize. train_data: Tuple of (features, labels) for training. val_data: Optional tuple of (features, labels) for validation. **kwargs: Optimizer-specific options. Returns: An OptimizationResult with the optimized model and metrics. """ ...