Update prodigy-plus-schedule-free

This commit is contained in:
kijai
2025-01-10 12:12:58 +02:00
parent a15cfb181a
commit 136697a655
5 changed files with 996 additions and 866 deletions
File diff suppressed because it is too large Load Diff
Binary file not shown.

Before

Width:  |  Height:  |  Size: 2.5 MiB

+1
View File
@@ -357,6 +357,7 @@ class OptimizerConfigProdigyPlusScheduleFree:
"use_stableadamw": ("BOOLEAN",{"default": True, "tooltip": "Scales parameter updates by the root-mean-square of the normalised gradient, in essence identical to Adafactor's gradient scaling. Set to False if the adaptive learning rate never improves."}),
"use_cautious" : ("BOOLEAN",{"default": False, "tooltip": "Experimental. Perform 'cautious' updates, as proposed in https://arxiv.org/pdf/2411.16085. Modifies the update to isolate and boost values that align with the current gradient."}),
"use_adopt": ("BOOLEAN",{"default": False, "tooltip": "Experimental. Performs a modified step where the second moment is updated after the parameter update, so as not to include the current gradient in the denominator. This is a partial implementation of ADOPT (https://arxiv.org/abs/2411.02853), as we don't have a first moment to use for the update."}),
"use_grams": ("BOOLEAN",{"default": False, "tooltip": "Perform 'grams' updates, as proposed in https://arxiv.org/abs/2412.17107. Modifies the update using sign operations that align with the current gradient. Note that we do not have access to a first moment, so this deviates from the paper (we apply the sign directly to the update). May have a limited effect."}),
"stochastic_rounding": ("BOOLEAN",{"default": True, "tooltip": "Use stochastic rounding for bfloat16 weights"}),
"extra_optimizer_args": ("STRING",{"multiline": True, "default": "", "tooltip": "additional optimizer args"}),
+239 -174
View File
@@ -1,23 +1,24 @@
import math
import torch
from statistics import mean, harmonic_mean, geometric_mean
from statistics import harmonic_mean
class CoreOptimiser(torch.optim.Optimizer):
def __init__(self, params, lr=1.0,
betas=(0.9, 0.99), beta3=None, beta4=0,
betas=(0.9, 0.99), beta3=None,
weight_decay=0.0,
weight_decay_by_lr=True,
use_bias_correction=False,
d0=1e-6, d_coef=1.0,
prodigy_steps=0,
warmup_steps=0,
eps=1e-8,
split_groups=True,
split_groups_mean="harmonic_mean",
split_groups_mean=True,
factored=True,
fused_back_pass=False,
use_stableadamw=True,
use_muon_pp=False,
use_cautious=False,
use_grams=False,
use_adopt=False,
stochastic_rounding=True):
@@ -33,29 +34,38 @@ class CoreOptimiser(torch.optim.Optimizer):
raise ValueError("Invalid beta parameter at index 1: {}".format(betas[1]))
if beta3 is not None and not 0.0 <= beta3 < 1.0:
raise ValueError("Invalid beta3 parameter: {}".format(beta3))
if beta4 is not None and not 0.0 <= beta4 < 1.0:
raise ValueError("Invalid beta4 parameter: {}".format(beta4))
if split_groups_mean not in {None, "mean", "harmonic_mean", "geometric_mean"}:
raise ValueError(f"Invalid value for split_groups_mean: '{split_groups_mean}'. Must be one of {None, 'mean', 'harmonic_mean', 'geometric_mean'}")
if use_adopt and use_muon_pp:
print(f"[{self.__class__.__name__}] Muon and ADOPT cannot be used at the same time. Muon has been disabled.")
use_muon_pp = False
self.try_hook_kohya_fbp()
defaults = dict(lr=lr, betas=betas, beta3=beta3, beta4=beta4,
if beta3 is None:
beta3 = betas[1] ** 0.5
if eps is None:
print(f"[{self.__class__.__name__}] 'eps' is None, Adam-atan2 enabled.")
if use_stableadamw:
print(f"[{self.__class__.__name__}] 'use_stableadamw' has been disabled (mutually exclusive with Adam-atan2).")
use_stableadamw = False
if use_cautious and use_grams:
print(f"[{self.__class__.__name__}] 'use_grams' has been disabled (mutually exclusive with 'use_cautious').")
use_grams = False
defaults = dict(lr=lr, betas=betas, beta3=beta3,
eps=eps,
weight_decay=weight_decay,
d=d0, d0=d0, d_coef=d_coef,
weight_decay_by_lr=weight_decay_by_lr,
d=d0, d_prev=d0, d0=d0, d_coef=d_coef,
k=1, train_mode=True,
weight_sum=0,
prodigy_steps=prodigy_steps,
warmup_steps=warmup_steps,
use_bias_correction=use_bias_correction,
d_numerator=0.0,
d_denom=0,
factored=factored,
use_stableadamw=use_stableadamw,
use_muon_pp=use_muon_pp,
use_cautious=use_cautious,
use_grams=use_grams,
use_adopt=use_adopt,
stochastic_rounding=stochastic_rounding)
@@ -70,7 +80,7 @@ class CoreOptimiser(torch.optim.Optimizer):
self.split_groups_mean = split_groups_mean
# Properties for fused backward pass.
self.groups_to_process = None
self.parameters_to_process = None
self.shared_d = None
self.fused_back_pass = fused_back_pass
@@ -110,42 +120,42 @@ class CoreOptimiser(torch.optim.Optimizer):
return group['running_d_numerator'], group['running_d_denom']
@torch.no_grad()
def get_d_mean(self, groups, mode):
if mode is None:
return None
elif mode == "harmonic_mean":
return harmonic_mean(group['d'] for group in groups)
elif mode == "geometric_mean":
return geometric_mean(group['d'] for group in groups)
elif mode == "mean":
return mean(group['d'] for group in groups)
raise ValueError(f"Invalid value for split_groups_mean: '{mode}'. Must be one of {None, 'mean', 'harmonic_mean', 'geometric_mean'}")
def get_d_mean(self):
if self.split_groups and self.split_groups_mean:
return harmonic_mean(group['d'] for group in self.param_groups)
return None
@torch.no_grad()
def get_d_max(self, group):
if self.split_groups:
return max(group['d'] for group in self.param_groups)
return group['d']
# From: https://github.com/KellerJordan/Muon/blob/master/muon.py
@torch.no_grad()
def newton_schulz_(self, G, steps=6, eps=1e-7):
# Inline reshaping step within the method itself.
original_shape = None
if len(G.shape) > 2:
original_shape = G.shape
G = G.view(G.size(0), -1)
X = G.view(G.size(0), -1)
a, b, c = (3.4445, -4.7750, 2.0315)
X = G.bfloat16()
X /= (X.norm() + eps) # ensure top singular value <= 1
X = X.to(dtype=torch.bfloat16, copy=True)
if G.size(0) > G.size(1):
X = X.T
X /= X.norm().add(eps) # ensure top singular value <= 1
for _ in range(steps):
A = X @ X.T
B = b * A + c * A @ A
X = a * X + B @ X
if G.size(0) > G.size(1):
X = X.T
if X is not G:
G.copy_(X)
del X
if original_shape is not None:
G = G.view(*original_shape)
# Gradient scaling adaptation from: https://github.com/leloykun/adaptive-muon
X = torch.einsum('ij,ij->', G.type_as(X), X).clamp(-1.0, 1.0) * X
G.copy_(X.view_as(G))
del X
return G
# Implementation by Nerogar. From: https://github.com/pytorch/pytorch/issues/120376#issuecomment-1974828905
@@ -167,6 +177,18 @@ class CoreOptimiser(torch.optim.Optimizer):
# copy the higher 16 bit into the target tensor
target.copy_(result.view(dtype=torch.float32))
def smart_copy(self, target, source, stochastic_rounding, smart_delete_source):
if target is source:
return
if stochastic_rounding and target.dtype == torch.bfloat16 and source.dtype == torch.float32:
self.copy_stochastic_(target, source)
else:
target.copy_(source)
if smart_delete_source:
del source
# Modified Adafactor factorisation implementation by Ross Wightman
# https://github.com/huggingface/pytorch-image-models/pull/2320
@torch.no_grad()
@@ -192,11 +214,11 @@ class CoreOptimiser(torch.optim.Optimizer):
return int(sorted_dims[-2][1]), int(sorted_dims[-1][1])
@torch.no_grad()
def initialise_state(self, p, factored, use_muon_pp):
def initialise_state(self, p, group):
raise Exception("Not implemented!")
@torch.no_grad()
def initialise_state_internal(self, p, factored, use_muon_pp):
def initialise_state_internal(self, p, group):
state = self.state[p]
needs_init = len(state) == 0
@@ -206,12 +228,14 @@ class CoreOptimiser(torch.optim.Optimizer):
sliced_data = self.get_sliced_tensor(p)
# NOTE: We don't initialise z/exp_avg here -- subclass needs to do that.
state['muon'] = use_muon_pp and len(grad.shape) >= 2 and grad.size(0) < 10000
state['muon'] = group['use_muon_pp'] and len(grad.shape) >= 2
if not state['muon']:
if state['muon']:
state["rms_sq"] = 0
else:
factored_dims = self.factored_dims(
grad.shape,
factored=factored,
factored=group['factored'],
min_dim_size_to_factor=32
)
@@ -226,19 +250,19 @@ class CoreOptimiser(torch.optim.Optimizer):
# Always store second moment low ranks in fp32 to avoid precision issues. Memory difference
# between bf16/fp16 and fp32 is negligible here.
state["exp_avg_sq"] = [torch.zeros(row_shape, dtype=torch.float32, device=p.device).detach(),
torch.zeros(col_shape, dtype=torch.float32, device=p.device).detach(),
dr, dc, reduce_dc]
torch.zeros(col_shape, dtype=torch.float32, device=p.device).detach(),
dr, dc, reduce_dc]
else:
state['exp_avg_sq'] = torch.zeros_like(p, memory_format=torch.preserve_format).detach()
# If the initial weights are zero, don't bother storing them.
if p.count_nonzero() > 0:
if p.any() > 0:
state['p0'] = sliced_data.to(dtype=dtype, memory_format=torch.preserve_format, copy=True).detach()
else:
state['p0'] = torch.tensor(0.0, dtype=dtype, device=p.device)
state['s'] = torch.zeros_like(sliced_data, memory_format=torch.preserve_format, dtype=dtype).detach()
return state, needs_init
@torch.no_grad()
@@ -249,141 +273,139 @@ class CoreOptimiser(torch.optim.Optimizer):
if prodigy_steps > 0 and k >= prodigy_steps:
return
beta1, beta2 = group['betas']
beta3, beta4 = group['beta3'], group['beta4']
if beta3 is None:
beta3 = beta2 ** 0.5
if beta4 is None:
beta4 = beta1 ** 0.5
d = group['d']
d0 = group['d0']
d, d0 = group['d'], group['d0']
d_prev = group['d_prev']
d_coef = group['d_coef']
beta3 = group['beta3']
running_d_numerator, running_d_denom = self.get_running_values_for_group(group)
d_numerator = group['d_numerator']
d_numerator *= beta3
d_prev = d
d_numerator_item = running_d_numerator.item()
d_denom_item = running_d_denom.item()
# Prevent the accumulation of negative values in the numerator in early training.
# We still allow negative updates once progress starts being made, as this is
# important for regulating the adaptive stepsize.
if d_numerator_item > 0 or d > d0:
d_numerator = max(0, d_numerator + d_numerator_item)
# Force Prodigy to be extremely confident before increasing the LR when gradient
# and weights drift.
if d_numerator_item < 0:
if d > d0:
# Prevent the accumulation of negative values in the numerator in early training.
# We still allow negative updates once progress starts being made, as this is
# important for regulating the adaptive stepsize.
d_numerator = min(d_numerator, d_numerator_item)
else:
d_numerator += d_numerator_item
d_hat = math.atan2(d_coef * d_numerator, d_denom_item)
d = max(d, d_hat)
if d_denom_item > 0:
d_hat = max(math.atan2(d_coef * d_numerator, d_denom_item), d)
d = d * beta4 + d_hat * (1 - beta4) if beta4 > 0 else d_hat
group['d'] = d
group['d_prev'] = d_prev
group['d_numerator'] = d_numerator
group['d_denom'] = d_denom_item
running_d_numerator.zero_()
running_d_denom.zero_()
def on_start_step(self, group):
if self.groups_to_process is None:
# Optimiser hasn't run yet, so initialise.
self.groups_to_process = {i: len(group['params']) for i, group in enumerate(self.param_groups)}
elif len(self.groups_to_process) == 0:
# Start of new optimiser run, so grab updated d.
self.groups_to_process = {i: len(group['params']) for i, group in enumerate(self.param_groups)}
def on_start_step(self):
if self.parameters_to_process is None or self.parameters_to_process == 0:
# Optimiser hasn't run yet (or is starting a new step), so initialise.
self.parameters_to_process = sum(len(group['params']) for group in self.param_groups)
def on_end_step(self):
self.parameters_to_process -= 1
if not self.split_groups:
# When groups aren't split, calculate d for the first group,
if self.parameters_to_process == 0:
# Update d for next optimiser step.
if self.split_groups:
i = 0
for group in self.param_groups:
if group['prodigy_steps'] > 0 and group['k'] == group['prodigy_steps']:
print(f"[{self.__class__.__name__}] Prodigy stepsize adaptation disabled after {group['k']} steps for param_group {i}.")
self.update_d_and_reset(group)
group['weight_sum'] = group.get('running_weight_sum', 0)
group['k'] += 1
i += 1
self.shared_d = self.get_d_mean()
else:
# When groups aren't split, calculate d for the first group (which collects stats for all groups in non-split mode),
# then copy to all other groups.
self.update_d_and_reset(group)
for g in self.param_groups:
g['d'] = group['d']
first_group = self.param_groups[0]
self.update_d_and_reset(first_group)
i = 0
for group in self.param_groups:
if group['prodigy_steps'] > 0 and group['k'] == group['prodigy_steps']:
print(f"[{self.__class__.__name__}] Prodigy stepsize adaptation disabled after {group['k']} steps for param_group {i}.")
self.shared_d = self.get_d_mean(self.param_groups, self.split_groups_mean) if self.split_groups else None
group['d'] = first_group['d']
group['d_numerator'] = first_group['d_numerator']
group['d_denom'] = first_group['d_denom']
group['weight_sum'] = group.get('running_weight_sum', 0)
group['k'] += 1
i += 1
def on_end_step(self, group):
group_index = self.param_groups.index(group)
# Decrement params processed so far.
self.groups_to_process[group_index] -= 1
# End of param loop for group, update calculations.
if self.groups_to_process[group_index] == 0:
k = group['k']
prodigy_steps = group['prodigy_steps']
if prodigy_steps > 0 and k == prodigy_steps:
print(f"[{self.__class__.__name__}] Prodigy stepsize adaptation disabled after {k} steps for param_group {group_index}.")
self.groups_to_process.pop(group_index)
if self.split_groups: # When groups are split, calculate per-group d.
self.update_d_and_reset(group)
group['k'] = k + 1
return True
return False
def get_dlr(self, group):
lr = group['lr']
k = group['k']
return (self.shared_d if self.split_groups and self.shared_d else group['d']) * group['lr']
warmup_steps = group['warmup_steps']
def update_prodigy(self, state, group, grad, data, num_scale):
# num_scale is used to compensate the numerator calculations when
# clipping/scaling is applied to the incoming update. If we don't
# do this, it will dampen Prodigy's 'd' predictions.
d = group['d']
dlr = (self.shared_d if self.split_groups and self.shared_d else d) * lr
# Apply warmup separate to the denom and numerator updates.
if k < warmup_steps:
dlr *= k / warmup_steps
return dlr
def update_prodigy(self, state, group, grad, data, dlr):
k = group['k']
prodigy_steps = group['prodigy_steps']
if prodigy_steps <= 0 or k < prodigy_steps:
d, d0 = group['d'], group['d0']
beta3 = group['beta3']
d, d0 = group['d'], group['d0']
if beta3 is None:
beta3 = group['betas'][1] ** 0.5
# Slow down, rather than speed up, as we approach the
# appropriate LR.
d_k = (d0 / d) * d
sliced_grad = self.get_sliced_tensor(grad)
sliced_data = self.get_sliced_tensor(data)
running_d_numerator, running_d_denom = self.get_running_values_for_group(group)
s = state['s']
x0_minus = state['p0'] - sliced_data
running_d_numerator.add_(torch.dot(sliced_grad, x0_minus), alpha=(d / d0) * dlr)
del x0_minus
s.mul_(beta3).add_(sliced_grad, alpha=(d / d0) * dlr)
running_d_numerator.add_(torch.dot(sliced_grad, x0_minus), alpha=d_k * num_scale)
s.mul_(beta3).add_(sliced_grad, alpha=d_k)
running_d_denom.add_(s.abs().sum())
del x0_minus
elif 's' in state: # Free the memory used by Prodigy, as we no longer need it.
del state['s']
del state['p0']
def get_update(self, num, denom, group):
d = group['d']
def update_(self, num, denom, group):
eps = group['eps']
if eps is None:
# Approximate scaling for a regular Adam-style update.
b = self.get_clip_threshold(group)
a = 1 / math.atan(1 / b)
if group['eps'] is None:
# Adam-atan2. Use atan2 rather than epsilon and division
# for parameter updates (https://arxiv.org/abs/2407.05872).
# Has the nice property of "clipping" the gradient as well.
update = num.mul_(d).atan2_(denom)
update = num.atan2_(denom.mul_(b)).mul_(a)
else:
# Assume eps as already been added.
update = num.div_(denom).mul_(d)
update = num.div_(denom.add_(eps))
return update
return update, 1.0
def get_denom(self, state, group):
def get_denom(self, state):
exp_avg_sq = state['exp_avg_sq']
eps = group['eps']
# Adam EMA updates
if isinstance(exp_avg_sq, list):
@@ -395,63 +417,104 @@ class CoreOptimiser(torch.optim.Optimizer):
denom = row_factor * col_factor
else:
denom = exp_avg_sq.sqrt()
if eps is not None:
denom.add_(group['d'] * eps)
return denom
def update_first_moment(self, exp_avg, group, grad):
d = group['d']
def update_first_moment(self, state, group, grad):
exp_avg = state['exp_avg']
beta1, _ = group['betas']
exp_avg.mul_(beta1).add_(grad, value=d * (1 - beta1))
return exp_avg
def update_second_moment(self, state, group, grad, beta2, return_denom=True):
d = group['d']
return exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1)
def update_second_moment(self, state, group, grad, beta2, return_denom=True, denom_before_update=False):
exp_avg_sq = state['exp_avg_sq']
# Adafactor / PaLM beta2 decay. Clip beta2 as per Scaling ViT paper.
if group['use_bias_correction']:
beta2 = min(1 - group['k'] ** -0.8, beta2)
denom = None
one_minus_beta2_d = d * d * (1 - beta2)
if return_denom and denom_before_update:
denom = self.get_denom(state)
# Adam EMA updates
if isinstance(exp_avg_sq, list):
row_var, col_var, dr, dc, _ = exp_avg_sq
row_var.mul_(beta2).add_(
grad.norm(dim=dr, keepdim=True).square_().div_(grad.shape[dr]),
alpha=one_minus_beta2_d)
col_var.mul_(beta2).add_(
grad.norm(dim=dc, keepdim=True).square_().div_(grad.shape[dc]),
alpha=one_minus_beta2_d)
row_var.lerp_(
grad.norm(dim=dr, keepdim=True).square_().div_(grad.shape[dr]),
weight=1 - beta2
)
col_var.lerp_(
grad.norm(dim=dc, keepdim=True).square_().div_(grad.shape[dc]),
weight=1 - beta2
)
else:
exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=one_minus_beta2_d)
exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)
return self.get_denom(state, group) if return_denom else None
def rms_(self, tensor, rms_min):
if rms_min is not None:
rms = tensor.norm().div(tensor.numel() ** 0.5).add(rms_min)
tensor.div_(rms)
return tensor
if return_denom and denom is None:
denom = self.get_denom(state)
# "Cautious Optimizer (C-Optim): Improving Training with One Line of Code"
# https://github.com/kyleliang919/c-optim
def cautious_(self, update, grad, reuse_grad):
if reuse_grad:
mask = grad.mul_(update) > 0
else:
mask = grad.mul(update) > 0
return denom
mask_scale = mask.numel() / mask.sum().add(1)
update.mul_(mask).mul_(mask_scale)
del mask
def get_rms(self, tensor, eps=1e-8):
return tensor.norm().div(tensor.numel() ** 0.5).clamp_min(eps)
return update
def rms_(self, tensor, eps):
return tensor.div_(self.get_rms(tensor, eps))
def get_clip_threshold(self, group):
return max(1, 8 * (0.99 ** (group['k'] - 1)))
def try_hook_kohya_fbp(self):
self.kohya_original_patch_adafactor_fused = None
try:
# Import and patching will fail if not Kohya.
import library.adafactor_fused
# Get the original method so we can restore it later.
self.kohya_original_patch_adafactor_fused = library.adafactor_fused.patch_adafactor_fused
# Define the override.
def prodigy_patch_adafactor_fused(optimizer):
unwrapped_optimiser = None
if hasattr(optimizer, "optimizer"):
# If the optimiser is wrapped, forward the calls to the actual optimiser.
def _step(self, *args, **kwargs):
return self.optimizer.step(*args, **kwargs)
def _step_param(self, *args, **kwargs):
return self.optimizer.step_param(*args, **kwargs)
optimizer.step = _step.__get__(optimizer)
optimizer.step_param = _step_param.__get__(optimizer)
unwrapped_optimiser = optimizer.optimizer
else:
unwrapped_optimiser = optimizer
print(f"[{self.__class__.__name__}] Kohya pipeline detected with fused backward pass. Gradient hook patch successful.")
library.adafactor_fused.patch_adafactor_fused = unwrapped_optimiser.kohya_original_patch_adafactor_fused # Restore the original method.
unwrapped_optimiser.fused_back_pass = True
unwrapped_optimiser.kohya_original_patch_adafactor_fused = None
# Patch the method.
library.adafactor_fused.patch_adafactor_fused = prodigy_patch_adafactor_fused
except:
pass
def try_unhook_kohya_fbp(self):
if self.kohya_original_patch_adafactor_fused is None:
return
try:
# Import and patching will fail if not Kohya.
import library.adafactor_fused
# User did not opt for fused backward pass, so remove our hook.
library.adafactor_fused.patch_adafactor_fused = self.kohya_original_patch_adafactor_fused
except:
pass
self.kohya_original_patch_adafactor_fused = None
@torch.no_grad()
def step_param(self, p, group):
@@ -463,6 +526,8 @@ class CoreOptimiser(torch.optim.Optimizer):
@torch.no_grad()
def step(self, closure=None):
self.try_unhook_kohya_fbp()
if self.fused_back_pass:
return
@@ -5,8 +5,7 @@ from .core_optimiser import CoreOptimiser
class ProdigyPlusScheduleFree(CoreOptimiser):
r"""
An optimiser based on Prodigy that includes schedule-free logic. Has additional improvements in the form of optional StableAdamW
gradient scaling and Adam-atan2 updates, per parameter group adaptation, lower memory utilisation, fused back pass support and
tweaks to mitigate uncontrolled LR growth.
gradient scaling and Adam-atan2 updates, per parameter group adaptation, lower memory utilisation and fused back pass support.
Based on code from:
https://github.com/facebookresearch/schedule_free
@@ -26,120 +25,132 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
ability for the optimiser to predict stepsizes. Gradient clipping/normalisation is already handled in the following configurations:
1) `use_stableadamw=True,eps=1e8` (or any reasonable positive epsilon)
2) `eps=None` (Adam-atan2, scale invariant and can mess with Prodigy's stepsize calculations in some scenarios)
2) `eps=None` (Adam-atan2, scale invariant. Will disable StableAdamW if enabled.)
A new parameter, `beta4`, allows `d` to be updated via a moving average, rather than being immediately updated. This can help
smooth out learning rate adjustments. Values of 0.9-0.99 are recommended if trying out the feature. If set to None, the
square root of `beta1` is used, while a setting of 0 (the default) disables the feature.
By default, `split_groups` is set to `True`, so each parameter group will have its own adaptation values. So if you're training
different networks together, they won't contaminate each other's learning rates. The disadvantage of this approach is that some
networks can take a long time to reach a good learning rate when trained alongside others (for example, SDXL's Unet).
It's recommended to use a higher `d0` (1e-5, 5e-5, 1e-4) so these networks don't get stuck at a low learning rate.
By default, `split_groups` and `split_groups_mean` are set to `True`, so each parameter group will have its own `d` values, however,
they will all use the harmonic mean for the dynamic learning rate. To make each group use its own dynamic LR, set `split_groups_mean` to False.
To use the reference Prodigy behaviour where all groups are combined, set `split_groups` to False.
For Prodigy's reference behaviour, which lumps all parameter groups together, set `split_groups` to `False`.
In some scenarios, it can be advantageous to freeze Prodigy's adaptive stepsize after a certain number of steps. This
can be controlled via the `prodigy_steps` settings.
can be controlled via the `prodigy_steps` settings. This will also free any Prodigy-specific memory used by the
optimiser (though with all the memory-related improvements, this should not be significant unless you're training
very large models).
Arguments:
params (iterable):
Iterable of parameters to optimize or dicts defining parameter groups.
lr (float):
Learning rate adjustment parameter. Increases or decreases the Prodigy learning rate.
(default: 1.0)
betas (Tuple[float, float], optional):
Coefficients used for computing running averages of gradient and its square
Coefficients used for computing running averages of gradient and its square.
(default: (0.9, 0.99))
eps (float):
Term added to the denominator outside of the root operation to improve numerical stability. If set to None,
Adam-atan2 is used instead. This removes the need for epsilon tuning, but may not work well in all situations.
(default: 1e-8).
beta3 (float):
Coefficient for computing the Prodigy stepsize using running averages.
If set to None, uses the value of square root of beta2 (default: None).
beta4 (float):
Coefficient for updating the learning rate from Prodigy's adaptive stepsize. Smooths out spikes in learning rate adjustments.
If set to None, beta1 is used instead. (default 0, which disables smoothing and uses original Prodigy behaviour).
Coefficient for computing the Prodigy stepsize using running averages. If set to None, uses the value of
square root of beta2
(default: None).
weight_decay (float):
Decoupled weight decay. Value is multiplied by the adaptive learning rate.
Decoupled weight decay. Use the weight_decay_by_lr setting to determine if decay should be multiplied by the
adaptive learning rate.
(default: 0).
weight_decay_by_lr (boolean):
If True, weight_decay is multiplied by the adaptive learning rate (as per the PyTorch implementation of AdamW).
If False, weight_decay will have a much stronger effect.
(default: True).
use_bias_correction (boolean):
Turn on Adafactor-style bias correction, which scales beta2 directly. (default False).
Turn on Adafactor-style bias correction, which scales beta2 directly. (default: False).
d0 (float):
Initial estimate for Prodigy (default 1e-6).
Initial estimate for Prodigy. Also serves as the minimum learning rate.
(default: 1e-6).
d_coef (float):
Coefficient in the expression for the estimate of d (default 1.0). Values such as 0.5 and 2.0 typically work as well.
Coefficient in the expression for the estimate of d. Values such as 0.5 and 2.0 typically work as well.
Changing this parameter is the preferred way to tune the method.
(default: 1.0)
prodigy_steps (int):
Freeze Prodigy stepsize adjustments after a certain optimiser step.
(default 0)
warmup_steps (int):
Enables a linear learning rate warmup (default 0). Use this over the warmup settings of your LR scheduler.
Freeze Prodigy stepsize adjustments after a certain optimiser step and releases all state memory required
by Prodigy.
(default: 0)
split_groups (boolean):
Track individual adaptation values for each parameter group. For example, if training
a text encoder beside a Unet. Note this can have a significant impact on training dynamics.
Set to False for original Prodigy behaviour, where all groups share the same values.
(default True)
split_groups_mean (str: None, "mean", "harmonic_mean", "geometric_mean"):
When split_groups is True, use specified mean of learning rates for all groups. This favours
(default: True)
split_groups_mean (boolean):
When split_groups is True, use the harmonic mean of learning rates for all groups. This favours
a more conservative LR. Calculation remains per-group. If split_groups is False, this value has no effect.
Set to None to have each group use its own learning rate calculation.
(default "harmonic_mean")
Set to False to have each group use its own learning rate.
(default: True)
factored (boolean):
Use factored approximation of the second moment, similar to Adafactor. Reduces memory usage. Disable
if training results in NaNs or the learning rate fails to grow.
(default True)
(default: True)
fused_back_pass (boolean):
Stops the optimiser from running the normal step method. Set to True if using fused backward pass.
(default False)
Stops the optimiser from running the normal step method. Set to True if using fused backward pass. Really only
needed for scripts and UIs that call the regular step method even when using fused backward pass (OneTrainer).
(default: False)
use_stableadamw (boolean):
Scales parameter updates by the root-mean-square of the normalised gradient, in essence identical to
Adafactor's gradient scaling. Set to False if the adaptive learning rate never improves.
(default True)
(default: True)
use_muon_pp (boolean):
Experimental. Perform orthogonalisation post-processing on 2D+ parameter updates ala Shampoo/SOAP/Muon.
Experimental. Perform orthogonalisation on the gradient before it is used for updates ala Shampoo/SOAP/Muon.
(https://github.com/KellerJordan/Muon/blob/master/muon.py). Not suitable for all training scenarios.
May not work well with small batch sizes or finetuning. (default False)
May not work well with small batch sizes or finetuning.
(default: False)
use_cautious (boolean):
Experimental. Perform "cautious" updates, as proposed in https://arxiv.org/pdf/2411.16085. Modifies
the update to isolate and boost values that align with the current gradient.
(default False)
the update to isolate and boost values that align with the current gradient. Note that we do not have
access to a first moment, so this deviates from the paper (we apply the mask directly to the update).
May have a limited effect.
(default: False)
use_grams (boolean):
Experimental. Perform "grams" updates, as proposed in https://arxiv.org/abs/2412.17107. Modifies
the update using sign operations that align with the current gradient. Note that we do not have
access to a first moment, so this deviates from the paper (we apply the sign directly to the update).
May have a limited effect.
(default: False)
use_adopt (boolean):
Experimental. Performs a modified step where the second moment is updated after the parameter update,
so as not to include the current gradient in the denominator. This is a partial implementation of ADOPT
(https://arxiv.org/abs/2411.02853), as we don't have a first moment to use for the update.
(default False)
(default: False)
stochastic_rounding (boolean):
Use stochastic rounding for bfloat16 weights (https://github.com/pytorch/pytorch/issues/120376). Brings
bfloat16 training performance close to that of float32.
(default True)
(default: True)
"""
def __init__(self, params, lr=1.0,
betas=(0.9, 0.99), beta3=None, beta4=0,
betas=(0.9, 0.99), beta3=None,
weight_decay=0.0,
weight_decay_by_lr=True,
use_bias_correction=False,
d0=1e-6, d_coef=1.0,
prodigy_steps=0,
warmup_steps=0,
eps=1e-8,
split_groups=True,
split_groups_mean="harmonic_mean",
split_groups_mean=True,
factored=True,
fused_back_pass=False,
use_stableadamw=True,
use_muon_pp=False,
use_cautious=False,
use_grams=False,
use_adopt=False,
stochastic_rounding=True):
super().__init__(params=params, lr=lr, betas=betas, beta3=beta3, beta4=beta4,
weight_decay=weight_decay, use_bias_correction=use_bias_correction,
d0=d0, d_coef=d_coef, prodigy_steps=prodigy_steps,
warmup_steps=warmup_steps, eps=eps, split_groups=split_groups,
split_groups_mean=split_groups_mean, factored=factored,
fused_back_pass=fused_back_pass, use_stableadamw=use_stableadamw,
use_muon_pp=use_muon_pp, use_cautious=use_cautious, use_adopt=use_adopt,
stochastic_rounding=stochastic_rounding)
super().__init__(params=params, lr=lr, betas=betas, beta3=beta3,
weight_decay=weight_decay, weight_decay_by_lr=weight_decay_by_lr,
use_bias_correction=use_bias_correction,
d0=d0, d_coef=d_coef, prodigy_steps=prodigy_steps,
eps=eps, split_groups=split_groups,
split_groups_mean=split_groups_mean, factored=factored,
fused_back_pass=fused_back_pass, use_stableadamw=use_stableadamw,
use_muon_pp=use_muon_pp, use_cautious=use_cautious, use_grams=use_grams,
use_adopt=use_adopt, stochastic_rounding=stochastic_rounding)
@torch.no_grad()
def eval(self):
@@ -168,96 +179,116 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
group['train_mode'] = True
@torch.no_grad()
def initialise_state(self, p, factored, use_muon_pp):
state, needs_init = self.initialise_state_internal(p, factored, use_muon_pp)
def initialise_state(self, p, group):
state, needs_init = self.initialise_state_internal(p, group)
if needs_init:
state['z'] = p.detach().clone(memory_format=torch.preserve_format)
state['z'] = p.detach().clone(memory_format=torch.preserve_format)
return state
@torch.no_grad()
def update_params(self, y, z, update, dlr, group):
# Weight decay.
weight_decay = group['weight_decay']
if weight_decay != 0:
update.add_(y, alpha=weight_decay)
@torch.no_grad()
def update_params(self, y, z, update, group):
dlr = self.get_dlr(group)
beta1, _ = group['betas']
decay = group['weight_decay']
weight = dlr ** 2
weight_sum = group['weight_sum'] + weight
ckp1 = weight / weight_sum if weight_sum else 0
y.lerp_(end=z, weight=ckp1)
y.add_(update, alpha=dlr * (group['betas'][0] * (1 - ckp1) - 1))
z.sub_(update, alpha=dlr)
xy_step = 1 - beta1 * (1 - ckp1)
if decay != 0:
# Weight decay at Y.
if group['weight_decay_by_lr']:
decay *= dlr
y.sub_(y, alpha=decay * xy_step)
z.sub_(y, alpha=decay)
if group['use_cautious']:
# "Cautious Optimizer (C-Optim): Improving Training with One Line of Code": https://github.com/kyleliang919/c-optim
# ScheduleFree implementation by nhamanasu: https://github.com/facebookresearch/schedule_free/pull/54
u = (y - z).mul_(ckp1).add_(update, alpha=dlr * xy_step)
z.sub_(update, alpha=dlr)
mask = (u * update > 0).to(update.dtype)
mask.mul_(mask.numel() / (mask.sum() + 1))
u.mul_(mask)
y.sub_(u)
del mask, u
elif group['use_grams']:
# "Grams: Gradient Descent with Adaptive Momentum Scaling": https://arxiv.org/abs/2412.17107
u = (y - z).mul_(ckp1).add_(update, alpha=dlr * xy_step)
z.sub_(update, alpha=dlr) # Update z now so we can do sign in-place.
y.sub_(u.abs_().mul_(update.sign_()))
del u
else:
y.lerp_(end=z, weight=ckp1)
y.sub_(update, alpha=dlr * xy_step)
z.sub_(update, alpha=dlr)
return weight_sum
@torch.no_grad()
def step_param(self, p, group):
self.on_start_step()
if not group['train_mode']:
raise Exception("Not in train mode!")
self.on_start_step(group)
weight_sum = group['weight_sum']
if p.grad is not None:
grad = p.grad
grad = p.grad.to(dtype=torch.float32, copy=True)
state = self.initialise_state(p, group['factored'], group['use_muon_pp'])
use_adopt = group['use_adopt']
use_adopt = group['use_adopt']
stochastic = group['stochastic_rounding']
_, beta2 = group['betas']
k = group['k']
if use_adopt and group['k'] == 1:
self.update_second_moment(state, group, grad.float(), 0, return_denom=False)
state = self.initialise_state(p, group)
update = None
if state['muon']:
grad = self.newton_schulz_(grad)
grad_rms = self.get_rms(grad).item() ** 2
rms_sq = (state["rms_sq"] * beta2) + (grad_rms * (1 - beta2))
state["rms_sq"] = rms_sq
update = grad.mul_(1.0 / ((rms_sq ** 0.5) + 1e-12))
else:
dlr = self.get_dlr(group)
rms_min = 1.0 if group['use_stableadamw'] else None
y, z = p, state['z']
if group['use_bias_correction']:
# Adafactor / PaLM beta2 decay. Clip beta2 as per Scaling ViT paper.
beta2 = min(beta2, 1 - k ** -0.8)
beta2 = (1 - beta2) / (1 - beta2 ** k)
self.update_prodigy(state, group, grad, z, dlr)
grad_mask = grad.clone() if group['use_cautious'] else None
if state['muon']:
# newton_schulz_ casts to bf16 internally, so do float cast afterwards.
update = self.newton_schulz_(grad).float()
rms_min = 1e-30
if use_adopt and group['k'] == 1:
self.update_second_moment(state, group, grad, 0, return_denom=False)
else:
grad = grad.float()
_, beta2 = group['betas']
if use_adopt:
denom = self.get_denom(state, group)
self.update_second_moment(state, group, grad, beta2, return_denom=False)
else:
denom = self.update_second_moment(state, group, grad, beta2)
update = self.get_update(grad, denom, group)
denom = self.update_second_moment(state, group, grad, beta2, denom_before_update=use_adopt)
update, num_scale = self.update_(grad, denom, group)
del denom
if group['eps'] is None:
rms_min = None
if update is not None:
if group['use_stableadamw']:
clip_threshold = self.get_clip_threshold(group)
num_scale = max(1, self.get_rms(update, 1.0).item() / clip_threshold)
update.mul_(1 / num_scale)
self.rms_(update, rms_min)
z_state = state['z']
self.update_prodigy(state, group, p.grad, z_state, 1.0)
if grad_mask is not None:
self.cautious_(update, grad_mask, reuse_grad=True)
y, z = (p.float(), z_state.float()) if stochastic else (p, z_state)
weight_sum = self.update_params(y, z, update, group)
if group['stochastic_rounding'] and y.dtype == z.dtype == torch.bfloat16:
y_fp32, z_fp32 = y.float(), z.float()
weight_sum = self.update_params(y_fp32, z_fp32, update, dlr, group)
self.copy_stochastic_(y, y_fp32)
self.copy_stochastic_(z, z_fp32)
del y_fp32, z_fp32
else:
weight_sum = self.update_params(y, z, update, dlr, group)
self.smart_copy(p, y, stochastic, True)
self.smart_copy(z_state, z, stochastic, True)
del update
if self.on_end_step(group):
group['weight_sum'] = weight_sum
group['running_weight_sum'] = weight_sum
self.on_end_step()