"""Post-training optimization stage."""
from __future__ import annotations
from typing import Any
from ..optimization import run_optimization_stage
from .base import BaseOptimizationStage
[docs]
class DefaultOptimizationStage(BaseOptimizationStage):
"""Apply pruning, distillation, NAS, and quantization."""
[docs]
def run(self, cfg: Any, classify_out: dict[str, Any]) -> dict[str, Any]:
data_out = classify_out.get("_data_out", {})
clean = {k: v for k, v in classify_out.items() if k != "_data_out"}
return run_optimization_stage(cfg, data_out, clean)