/
Handdas
/
FirstOrderAlgorythms
Обзор
Документация
Войти
/
Handdas
/
FirstOrderAlgorythms
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
CI/CD
Аналитика
Безопасность
master
blackbox_optimizer/core/base.py
67 строк
2 KB
Your Name
initial commit
09 июн 2026, 10:13
09 июн 2026, 10:13
f88c506
Код
Авторство
О чём код?
"""Core abstractions for first-order gradient optimizers.""" from __future__ import annotations import logging from abc import ABC, abstractmethod from dataclasses import dataclass, field from typing import Optional, Protocol import numpy as np from .nfe_counter import NFECounter logger = logging.getLogger(__name__) class LossAndGradFn(Protocol): """Callable protocol compatible with WIND-like objective API.""" def __call__(self, params: np.ndarray) -> tuple[float, np.ndarray]: ... @dataclass class OptimizationTrace: """Metrics collected during optimization.""" iterations: int = 0 losses: list[float] = field(default_factory=list) gradient_norms: list[float] = field(default_factory=list) wall_times: list[float] = field(default_factory=list) class BaseOptimizer(ABC): """Unified OOP base for first-order optimizers with `step()` API.""" def __init__( self, learning_rate: float, *, weight_decay: float = 0.0, clip_grad_norm: Optional[float] = None, ) -> None: self.learning_rate = float(learning_rate) self.weight_decay = float(weight_decay) self.clip_grad_norm = clip_grad_norm self.nfe_counter = NFECounter() self.iteration = 0 self.logger = logging.getLogger(self.__class__.__name__) @abstractmethod def step(self, params: np.ndarray, grad: np.ndarray) -> np.ndarray: """Apply one optimization step and return updated params.""" def reset_state(self) -> None: """Reset optimizer state between runs.""" self.iteration = 0 self.nfe_counter.reset() def _regularize_grad(self, params: np.ndarray, grad: np.ndarray) -> np.ndarray: reg_grad: np.ndarray = grad.astype(float, copy=True) if self.weight_decay != 0.0: reg_grad = reg_grad + self.weight_decay * params if self.clip_grad_norm is not None: norm = float(np.linalg.norm(reg_grad)) if norm > self.clip_grad_norm and norm > 0: reg_grad *= self.clip_grad_norm / norm return reg_grad