"""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.
"""
...