import torch from torch.optim.optimizer import Optimizer from transformers.optimization import Adafactor class BabkaSotona(Optimizer): """ Гибридный оптимизатор: Aurora (для 2D-весов начала модели) + Adafactor (для 1D-весов и хвоста модели). """ def __init__( self, params, lr=1e-5, aurora_weight_decay=0.01, adafactor_weight_decay=0.001, adafactor_lr_scale=10.0, # Делитель LR для Adafactor (раз во сколько он должен быть меньше) tail_fraction=0.3, # Доля тензоров с конца модели, которые целиком уходят в Adafactor aurora_mu=0.9, nesterov=False, pp_iterations=2, pp_beta=0.5, eps=1e-7, adafactor_kwargs=None ): self.adafactor_lr_scale = adafactor_lr_scale # Нормализуем входные параметры param_list = list(params) if len(param_list) > 0 and isinstance(param_list[0], dict): actual_params = [] for group in param_list: actual_params.extend(group['params']) else: actual_params = param_list actual_params = [p for p in actual_params if p.requires_grad] # Определяем границу "хвоста" total_tensors = len(actual_params) split_idx = int(total_tensors * (1.0 - tail_fraction)) aurora_params = [] adafactor_params = [] aurora_numel = 0 adafactor_numel = 0 for i, p in enumerate(actual_params): if i >= split_idx: # Зона хвоста: ВСЕ параметры идут в Adafactor adafactor_params.append(p) adafactor_numel += p.numel() else: # Основная зона: 2D идет в Aurora, остальное (1D) в Adafactor if p.ndim == 2: aurora_params.append(p) aurora_numel += p.numel() else: adafactor_params.append(p) adafactor_numel += p.numel() # --- Вывод дебаг-информации --- print(f"\n[{self.__class__.__name__} DEBUG INFO]") print(f" 🔹 Aurora Params (2D): {len(aurora_params):>4} tensors | {aurora_numel:>12,} parameters") print(f" 🔸 Adafactor Params : {len(adafactor_params):>4} tensors | {adafactor_numel:>12,} parameters") print(f" (Tail fraction: {tail_fraction:.0%} | Adafactor LR scale: /{adafactor_lr_scale})") print(f"--------------------------------------------------\n") defaults = dict(lr=lr) param_groups = [ { 'params': aurora_params, 'is_aurora': True, 'weight_decay': aurora_weight_decay, 'mu': aurora_mu, 'nesterov': nesterov, 'pp_iterations': pp_iterations, 'pp_beta': pp_beta, 'eps': eps }, { 'params': adafactor_params, 'is_aurora': False, 'weight_decay': adafactor_weight_decay } ] super().__init__(param_groups, defaults) if adafactor_kwargs is None: adafactor_kwargs = { "eps": (1e-30, 1e-3), "clip_threshold": 1.0, "decay_rate": -0.8, "beta1": None, "relative_step": False, "scale_parameter": False, "warmup_init": False } # Инициализируем Adafactor с уже уменьшенным LR scaled_lr = lr / self.adafactor_lr_scale adafactor_group = [{'params': adafactor_params, 'lr': scaled_lr, 'weight_decay': adafactor_weight_decay}] self.adafactor = Adafactor(adafactor_group, **adafactor_kwargs) # Связываем state_dict для корректного сохранения чекпоинтов self.adafactor.state = self.state def load_state_dict(self, state_dict): super().load_state_dict(state_dict) self.adafactor.state = self.state def _polar(self, G: torch.Tensor) -> torch.Tensor: """Полярная декомпозиция (Newton-Schulz). Здесь внутреннее транспонирование необходимо математически для правильного умножения X @ X.mT, но оно не влияет на логику внешнего апдейта.""" assert G.ndim >= 2 X = G.bfloat16() if G.size(-2) > G.size(-1): X = X.mT X = X / (X.norm(dim=(-2, -1), keepdim=True) + 1e-7) a, b, c = 2, -1.5, 0.5 for _ in range(12): A = X @ X.mT B = b * A + c * A @ A X = a * X + B @ X if G.size(-2) > G.size(-1): X = X.mT return X @torch.no_grad() def step(self, closure=None): loss = None if closure is not None: with torch.enable_grad(): loss = closure() # Синхронизация LR: применяем делитель для Adafactor ada_group_idx = 0 for group in self.param_groups: if not group.get('is_aurora', False): self.adafactor.param_groups[ada_group_idx]['lr'] = group['lr'] / self.adafactor_lr_scale self.adafactor.param_groups[ada_group_idx]['weight_decay'] = group['weight_decay'] ada_group_idx += 1 if len(self.adafactor.param_groups[0]['params']) > 0: self.adafactor.step() for group in self.param_groups: if group.get('is_aurora', False): eta = group['lr'] weight_decay = group['weight_decay'] mu = group['mu'] nesterov = group['nesterov'] pp_iterations = group['pp_iterations'] pp_beta = group['pp_beta'] eps = group['eps'] for p in group['params']: if p.grad is None: continue G = p.grad W = p.data state = self.state[p] if 'momentum' not in state: state['momentum'] = torch.zeros_like(W) momentum = state['momentum'] momentum.lerp_(G, 1 - mu) update = G.lerp_(momentum, mu) if nesterov else momentum.clone() m, n = update.size(-2), update.size(-1) if m == n: update = self._polar(update) else: # ВАЖНО: Убрано внешнее транспонирование (transposed = m < n). # Прямое применение диагонального прекондиционирования Aurora G32 = update.to(torch.float32) target_row_sq = n / m row_norm = G32.norm(dim=-1, keepdim=True).clamp_(min=eps) D = 1.0 / row_norm for k in range(pp_iterations): U = self._polar(D * G32) if k < pp_iterations - 1: row_sq = U.to(torch.float32).pow(2).sum(dim=-1, keepdim=True).clamp_(min=eps * eps) D = D * (target_row_sq / row_sq).pow(pp_beta) update = U update *= max(1, G.size(-2) / G.size(-1)) ** 0.5 if not update.isfinite().all(): raise RuntimeError("Aurora: non-finite update. Check gradients or conditioning.") W.mul_(1 - eta * weight_decay) W.add_(update, alpha=-eta) return loss