Source code for arcade.stages.optimization

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