update prodigyplusschedulefree, support Flex with arg bypass_flux_guidance

https://github.com/kohya-ss/sd-scripts/pull/1893
This commit is contained in:
kijai
2025-01-31 21:52:02 +02:00
parent 5f254225c7
commit f6af45a169
7 changed files with 309 additions and 146 deletions
+6
View File
@@ -669,6 +669,9 @@ class FluxTrainer:
if not args.apply_t5_attn_mask:
t5_attn_mask = None
if args.bypass_flux_guidance:
flux_utils.bypass_flux_guidance(flux)
with accelerator.autocast():
# YiYi notes: divide it by 1000 for now because we scale it by 1000 in the transformer model (we should not keep it but I want to keep the inputs same for the model for testing)
model_pred = flux(
@@ -685,6 +688,9 @@ class FluxTrainer:
# unpack latents
model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width)
if args.bypass_flux_guidance:
flux_utils.restore_flux_guidance(flux)
# apply model prediction type
model_pred, weighting = flux_train_utils.apply_model_prediction_type(args, model_pred, noisy_model_input, sigmas)
+6
View File
@@ -357,6 +357,9 @@ class FluxNetworkTrainer(NetworkTrainer):
return model_pred
if args.bypass_flux_guidance:
flux_utils.bypass_flux_guidance(unet)
model_pred = call_dit(
img=packed_noisy_model_input,
img_ids=img_ids,
@@ -371,6 +374,9 @@ class FluxNetworkTrainer(NetworkTrainer):
# unpack latents
model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width)
if args.bypass_flux_guidance: #for flex
flux_utils.restore_flux_guidance(unet)
# apply model prediction type
model_pred, weighting = flux_train_utils.apply_model_prediction_type(args, model_pred, noisy_model_input, sigmas)
+5
View File
@@ -578,3 +578,8 @@ def add_flux_train_arguments(parser: argparse.ArgumentParser):
default=3.0,
help="Discrete flow shift for the Euler Discrete Scheduler, default is 3.0. / Euler Discrete Schedulerの離散フローシフト、デフォルトは3.0。",
)
parser.add_argument(
"--bypass_flux_guidance"
, action="store_true"
, help="bypass flux guidance module for Flex.1-Alpha Training"
)
+7
View File
@@ -21,6 +21,13 @@ MODEL_VERSION_FLUX_V1 = "flux1"
MODEL_NAME_DEV = "dev"
MODEL_NAME_SCHNELL = "schnell"
# bypass guidance
def bypass_flux_guidance(transformer):
transformer.params.guidance_embed = False
# restore the forward function
def restore_flux_guidance(transformer):
transformer.params.guidance_embed = True
def analyze_checkpoint_state(ckpt_path: str) -> Tuple[bool, bool, Tuple[int, int], List[str]]:
"""
+4 -2
View File
@@ -396,15 +396,16 @@ class OptimizerConfigProdigyPlusScheduleFree:
"split_groups": ("BOOLEAN",{"default": True, "tooltip": "Track individual adaptation values for each parameter group."}),
#"beta3": ("FLOAT",{"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.0001, "tooltip": " 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",{"default": 0, "min": 0.0, "max": 1.0, "step": 0.0001, "tooltip": "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)."}),
"use_bias_correction": ("BOOLEAN",{"default": False, "tooltip": "Turn on Adafactor-style bias correction, which scales beta2 directly."}),
"use_bias_correction": ("BOOLEAN",{"default": False, "tooltip": "Use the RAdam variant of schedule-free"}),
"min_snr_gamma": ("FLOAT",{"default": 5.0, "min": 0.0, "step": 0.01, "tooltip": "gamma for reducing the weight of high loss timesteps. Lower numbers have stronger effect. 5 is recommended by the paper"}),
"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"}),
"use_orthograd": ("BOOLEAN",{"default": False, "tooltip": "Experimental. Updates weights using the component of the gradient that is orthogonal to the current weight direction, as described in (https://arxiv.org/pdf/2501.04697). Can help prevent overfitting and improve generalisation."}),
"use_focus ": ("BOOLEAN",{"default": False, "tooltip": "Experimental. Modifies the update step to better handle noise at large step sizes. (https://arxiv.org/abs/2501.12243). This method is incompatible with factorisation, Muon and Adam-atan2."}),
"extra_optimizer_args": ("STRING",{"multiline": True, "default": "", "tooltip": "additional optimizer args"}),
},
}
@@ -497,6 +498,7 @@ class InitFluxLoRATraining:
dataset_toml = toml.dumps(json.loads(dataset_config))
parser = train_network_setup_parser()
flux_train_utils.add_flux_train_arguments(parser)
if additional_args is not None:
print(f"additional_args: {additional_args}")
args, _ = parser.parse_known_args(args=shlex.split(additional_args))
+156 -68
View File
@@ -10,16 +10,20 @@ class CoreOptimiser(torch.optim.Optimizer):
use_bias_correction=False,
d0=1e-6, d_coef=1.0,
prodigy_steps=0,
use_speed=False,
eps=1e-8,
split_groups=True,
split_groups_mean=True,
factored=True,
factored_fp32=True,
fused_back_pass=False,
use_stableadamw=True,
use_muon_pp=False,
use_cautious=False,
use_grams=False,
use_adopt=False,
use_orthograd=False,
use_focus=False,
stochastic_rounding=True):
if not 0.0 < d0:
@@ -50,6 +54,17 @@ class CoreOptimiser(torch.optim.Optimizer):
print(f"[{self.__class__.__name__}] 'use_grams' has been disabled (mutually exclusive with 'use_cautious').")
use_grams = False
if use_focus:
if factored:
print(f"[{self.__class__.__name__}] 'factored' has been disabled (incompatible with 'use_focus').")
factored = False
if use_muon_pp:
print(f"[{self.__class__.__name__}] 'use_muon_pp' has been disabled (incompatible with 'use_focus').")
use_muon_pp = False
if eps is None:
print(f"[{self.__class__.__name__}] Adam-atan2 ('eps=None') has been disabled (incompatible with 'use_focus').")
# We skip the Adam-atan2 branch entirely when FOCUS is enabled.
defaults = dict(lr=lr, betas=betas, beta3=beta3,
eps=eps,
weight_decay=weight_decay,
@@ -58,15 +73,19 @@ class CoreOptimiser(torch.optim.Optimizer):
k=1, train_mode=True,
weight_sum=0,
prodigy_steps=prodigy_steps,
use_speed=use_speed,
use_bias_correction=use_bias_correction,
d_numerator=0.0,
d_denom=0,
factored=factored,
factored_fp32=factored_fp32,
use_stableadamw=use_stableadamw,
use_muon_pp=use_muon_pp,
use_cautious=use_cautious,
use_grams=use_grams,
use_adopt=use_adopt,
use_orthograd=use_orthograd,
use_focus=use_focus,
stochastic_rounding=stochastic_rounding)
super().__init__(params, defaults)
@@ -113,10 +132,21 @@ class CoreOptimiser(torch.optim.Optimizer):
def get_sliced_tensor(self, tensor, slice_p=11):
return tensor.ravel()[::slice_p]
@torch.no_grad()
def check_running_values_for_group(self, p, group):
if not self.split_groups:
group = self.param_groups[0]
if group['running_d_numerator'].device != p.device:
group['running_d_numerator'] = group['running_d_numerator'].to(p.device)
if group['running_d_denom'].device != p.device:
group['running_d_denom'] = group['running_d_denom'].to(p.device)
@torch.no_grad()
def get_running_values_for_group(self, group):
if not self.split_groups:
group = self.param_groups[0]
return group['running_d_numerator'], group['running_d_denom']
@torch.no_grad()
@@ -135,10 +165,11 @@ class CoreOptimiser(torch.optim.Optimizer):
@torch.no_grad()
def newton_schulz_(self, G, steps=6, eps=1e-7):
# Inline reshaping step within the method itself.
X = G.view(G.size(0), -1)
G_shape = G.shape
G = G.view(G.size(0), -1)
a, b, c = (3.4445, -4.7750, 2.0315)
X = X.to(dtype=torch.bfloat16, copy=True)
X = G.to(dtype=torch.bfloat16, copy=True)
if G.size(0) > G.size(1):
X = X.T
@@ -153,10 +184,33 @@ class CoreOptimiser(torch.optim.Optimizer):
# 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))
G.copy_(X)
del X
return G
return G.view(G_shape)
# Implementation from: https://github.com/LucasPrietoAl/grokking-at-the-edge-of-numerical-stability/blob/main/orthograd.py
def orthograd(self, p, grad):
G_shape = grad.shape
w = p.view(-1)
g = grad.view(-1)
proj = torch.dot(w, g) / (torch.dot(w, w) + 1e-30)
g_orth = g.to(dtype=torch.float32, copy=True).sub_(w, alpha=proj)
g_orth_scaled = g_orth.mul_(g.norm(2) / (g_orth.norm(2) + 1e-30))
return g_orth_scaled.view(G_shape)
def orthograd_(self, p, grad):
G_shape = grad.shape
w = p.view(-1)
g = grad.view(-1)
proj = torch.dot(w, g) / (torch.dot(w, w) + 1e-30)
g_orth = g.sub_(w, alpha=proj)
g_orth_scaled = g_orth.mul_(g.norm(2) / (g_orth.norm(2) + 1e-30))
return g_orth_scaled.view(G_shape)
# Implementation by Nerogar. From: https://github.com/pytorch/pytorch/issues/120376#issuecomment-1974828905
def copy_stochastic_(self, target, source):
@@ -224,14 +278,18 @@ class CoreOptimiser(torch.optim.Optimizer):
if needs_init:
grad = p.grad
dtype = torch.bfloat16 if p.dtype == torch.float32 else p.dtype
dtype = torch.bfloat16 if grad.dtype == torch.float32 else grad.dtype
sliced_data = self.get_sliced_tensor(p)
if group['use_focus']:
state['exp_avg_sq'] = torch.zeros_like(grad, memory_format=torch.preserve_format).detach()
state['muon'] = False
else:
# NOTE: We don't initialise z/exp_avg here -- subclass needs to do that.
state['muon'] = group['use_muon_pp'] and len(grad.shape) >= 2
if state['muon']:
state["rms_sq"] = 0
state["rms_sq"] = None if group['use_speed'] else torch.tensor(0.0, dtype=dtype, device=p.device)
else:
factored_dims = self.factored_dims(
grad.shape,
@@ -240,20 +298,20 @@ class CoreOptimiser(torch.optim.Optimizer):
)
if factored_dims is not None:
# Store reduction variables so we don't have to recalculate each step.
dc, dr = factored_dims
row_shape = list(p.grad.shape)
row_shape = list(grad.shape)
row_shape[dr] = 1
col_shape = list(p.grad.shape)
col_shape = list(grad.shape)
col_shape[dc] = 1
reduce_dc = dc - 1 if dc > dr else dc
# Store reduction variables so we don't have to recalculate each step.
# 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(),
factored_dtype = torch.float32 if group['factored_fp32'] else grad.dtype
state["exp_avg_sq"] = [torch.zeros(row_shape, dtype=factored_dtype, device=p.device).detach(),
torch.zeros(col_shape, dtype=factored_dtype, device=p.device).detach(),
dr, dc, reduce_dc]
else:
state['exp_avg_sq'] = torch.zeros_like(p, memory_format=torch.preserve_format).detach()
state['exp_avg_sq'] = torch.zeros_like(grad, memory_format=torch.preserve_format).detach()
# If the initial weights are zero, don't bother storing them.
if p.any() > 0:
@@ -266,7 +324,7 @@ class CoreOptimiser(torch.optim.Optimizer):
return state, needs_init
@torch.no_grad()
def update_d_and_reset(self, group):
def update_d_stats_and_reset(self, group):
k = group['k']
prodigy_steps = group['prodigy_steps']
@@ -274,8 +332,6 @@ class CoreOptimiser(torch.optim.Optimizer):
return
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)
@@ -283,8 +339,6 @@ class CoreOptimiser(torch.optim.Optimizer):
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()
@@ -299,21 +353,38 @@ class CoreOptimiser(torch.optim.Optimizer):
else:
d_numerator += d_numerator_item
d_hat = math.atan2(d_coef * d_numerator, d_denom_item)
d = max(d, 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):
@torch.no_grad()
def calculate_d(self, group):
k = group['k']
prodigy_steps = group['prodigy_steps']
if prodigy_steps > 0 and k >= prodigy_steps:
return
d = group['d']
d_hat = math.atan2(group['d_coef'] * group['d_numerator'], group['d_denom'])
if group['use_speed']:
d_hat = max(d, d_hat)
d = min(d_hat, d ** 0.975)
else:
d = max(d, d_hat)
group['d_prev'] = group['d']
group['d'] = d
def on_start_step(self, p, group):
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)
# Check running values are on-device.
self.check_running_values_for_group(p, group)
def on_end_step(self):
self.parameters_to_process -= 1
@@ -321,25 +392,26 @@ class CoreOptimiser(torch.optim.Optimizer):
if self.parameters_to_process == 0:
# Update d for next optimiser step.
if self.split_groups:
i = 0
for group in self.param_groups:
for i, group in enumerate(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)
self.update_d_stats_and_reset(group)
for group in self.param_groups:
self.calculate_d(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.
first_group = self.param_groups[0]
self.update_d_and_reset(first_group)
self.update_d_stats_and_reset(first_group)
self.calculate_d(first_group)
i = 0
for group in self.param_groups:
for i, group in enumerate(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}.")
@@ -348,63 +420,69 @@ class CoreOptimiser(torch.optim.Optimizer):
group['d_denom'] = first_group['d_denom']
group['weight_sum'] = group.get('running_weight_sum', 0)
group['k'] += 1
i += 1
def get_dlr(self, group):
return (self.shared_d if self.split_groups and self.shared_d else group['d']) * group['lr']
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.
dlr = (self.shared_d if self.split_groups and self.shared_d else group['d']) * group['lr']
return dlr * group.get('rect', 1.0)
def update_prodigy(self, state, group, grad, data):
k = group['k']
prodigy_steps = group['prodigy_steps']
if prodigy_steps <= 0 or k < prodigy_steps:
beta3 = group['beta3']
d, d0 = group['d'], group['d0']
# Slow down, rather than speed up, as we approach the
# appropriate LR.
d_k = (d0 / d) * d
d_update = group['d'] ** 0.5
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_k * num_scale)
x0_dot = torch.dot(sliced_grad, x0_minus)
s.mul_(beta3).add_(sliced_grad, alpha=d_k)
if group['use_speed']:
d_update *= group['d0']
x0_dot, sliced_grad = x0_dot.sign(), sliced_grad.sign()
s = state['s']
s.mul_(beta3).add_(sliced_grad, alpha=d_update)
running_d_numerator.add_(x0_dot, alpha=d_update)
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 update_(self, num, denom, group):
def update_(self, num, denom, group, w):
if group['use_focus']:
# FOCUS: First Order Concentrated Updating Scheme: https://arxiv.org/pdf/2501.12243
gamma = 0.1
# Original form.
# update = torch.sign(num) + gamma * torch.sign(w - denom)
denom = denom.sub_(w).sign_().mul_(-gamma)
update = num.sign_().add_(denom)
else:
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)
# 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.atan2_(denom.mul_(b)).mul_(a)
update = num.atan2_(denom.mul_(b)).mul_(b)
else:
update = num.div_(denom.add_(eps))
return update, 1.0
return update
def get_denom(self, state):
def get_denom(self, state, group):
exp_avg_sq = state['exp_avg_sq']
# Adam EMA updates
@@ -415,53 +493,63 @@ class CoreOptimiser(torch.optim.Optimizer):
row_factor = row_var.div(row_col_mean).sqrt_()
col_factor = col_var.sqrt()
denom = row_factor * col_factor
elif group['use_focus']:
denom = exp_avg_sq.clone()
else:
denom = exp_avg_sq.sqrt()
return denom
def update_first_moment(self, state, group, grad):
def update_first_moment(self, state, group, grad, beta1):
exp_avg = state['exp_avg']
beta1, _ = group['betas']
d_k = group['d_prev'] / group['d']
return exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1)
return exp_avg.mul_(beta1 * d_k).add_(grad, weight=1 - beta1)
def update_second_moment(self, state, group, grad, beta2, return_denom=True, denom_before_update=False):
def update_second_moment(self, state, group, grad, beta2, w, return_denom=True, denom_before_update=False):
exp_avg_sq = state['exp_avg_sq']
d_k = group['d_prev'] / group['d']
denom = None
if return_denom and denom_before_update:
denom = self.get_denom(state)
denom = self.get_denom(state, group)
# Adam EMA updates
if group['use_focus']:
exp_avg_sq.mul_(beta2 * d_k * d_k).add_(w, alpha=1 - beta2)
else:
if isinstance(exp_avg_sq, list):
row_var, col_var, dr, dc, _ = exp_avg_sq
row_var.lerp_(
grad.norm(dim=dr, keepdim=True).square_().div_(grad.shape[dr]),
weight=1 - beta2
row_var.mul_(beta2 * d_k * d_k).add_(
grad.norm(dim=dr, keepdim=True).square_().mul_(1 / grad.shape[dr]),
alpha=1 - beta2
)
col_var.lerp_(
grad.norm(dim=dc, keepdim=True).square_().div_(grad.shape[dc]),
weight=1 - beta2
col_var.mul_(beta2 * d_k * d_k).add_(
grad.norm(dim=dc, keepdim=True).square_().mul_(1 / grad.shape[dc]),
alpha=1 - beta2
)
else:
exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2)
exp_avg_sq.mul_(beta2 * d_k * d_k).addcmul_(grad, grad, value=1 - beta2)
if return_denom and denom is None:
denom = self.get_denom(state)
denom = self.get_denom(state, group)
return denom
def get_rms(self, tensor, eps=1e-8):
return tensor.norm().div(tensor.numel() ** 0.5).clamp_min(eps)
return tensor.norm(2).div(tensor.numel() ** 0.5).clamp_min(eps)
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)))
# Prodigy works best with unscaled gradients during early steps.
if not group['use_speed'] and group['d'] <= group['d0']:
return 50
return 1
def try_hook_kohya_fbp(self):
self.kohya_original_patch_adafactor_fused = None
@@ -62,7 +62,11 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
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).
Use the RAdam variant of schedule-free (https://github.com/facebookresearch/schedule_free/blob/main/schedulefree/radam_schedulefree.py).
This combines bias correction with automatic warmup. Please note this will significantly dampen Prodigy's adaptive stepsize
calculations -- it can take up to 10 times longer to start adjusting the learning rate. This can be mitigated somewhat by enabling
SPEED (use_speed=True).
(default: False).
d0 (float):
Initial estimate for Prodigy. Also serves as the minimum learning rate.
(default: 1e-6).
@@ -74,6 +78,11 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
Freeze Prodigy stepsize adjustments after a certain optimiser step and releases all state memory required
by Prodigy.
(default: 0)
use_speed (boolean):
Highly experimental. Signed Prodigy with ExponEntial D. This decouples the adaptive stepsize calculations from
the magnitude of the weights and gradient. This can provide faster, more accurate LRs in some scenarios,
but may fail in situations where the optimal LR is very close to (or less than) d0.
(default: False):
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.
@@ -88,6 +97,12 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
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)
factored_fp32 (boolean):
Force the use of float32 for the factored second moment. Because the factorisation is an approximation, it can
be beneficial to use high precision to avoid stability issues. However, if you're training in lower precision
for short durations, setting this to False will slightly reduce memory usage.
Ignored if factored is False.
(default: True)
fused_back_pass (boolean):
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).
@@ -118,6 +133,16 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
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)
use_orthograd (boolean):
Experimental. Updates weights using the component of the gradient that is orthogonal to the current
weight direction, as described in "Grokking at the Edge of Numerical Stability" (https://arxiv.org/pdf/2501.04697).
Can help prevent overfitting and improve generalisation.
(default: False)
use_focus (boolean):
Experimental. Modifies the update step to better handle noise at large step sizes. From
"FOCUS: First-Order Concentrated Update Scheme" (https://arxiv.org/abs/2501.12243). This method is
incompatible with factorisation, Muon and Adam-atan2.
(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.
@@ -130,27 +155,32 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
use_bias_correction=False,
d0=1e-6, d_coef=1.0,
prodigy_steps=0,
use_speed=False,
eps=1e-8,
split_groups=True,
split_groups_mean=True,
factored=True,
factored_fp32=True,
fused_back_pass=False,
use_stableadamw=True,
use_muon_pp=False,
use_cautious=False,
use_grams=False,
use_adopt=False,
use_orthograd=False,
use_focus=False,
stochastic_rounding=True):
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,
d0=d0, d_coef=d_coef, prodigy_steps=prodigy_steps, use_speed=use_speed,
eps=eps, split_groups=split_groups,
split_groups_mean=split_groups_mean, factored=factored,
split_groups_mean=split_groups_mean, factored=factored, factored_fp32=factored_fp32,
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)
use_adopt=use_adopt, use_orthograd=use_orthograd, use_focus=use_focus,
stochastic_rounding=stochastic_rounding)
@torch.no_grad()
def eval(self):
@@ -188,9 +218,7 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
return state
@torch.no_grad()
def update_params(self, y, z, update, group):
dlr = self.get_dlr(group)
def update_params(self, y, z, update, group, dlr):
beta1, _ = group['betas']
decay = group['weight_decay']
@@ -208,21 +236,23 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
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
cautious, grams = group['use_cautious'], group['use_grams']
if cautious or grams:
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))
if 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
mask = update.mul_(u).sign_().clamp_min_(0)
mask.mul_(mask.numel() / mask.sum().add(1))
u.mul_(mask)
y.sub_(u)
del mask, u
elif group['use_grams']:
elif 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_()))
u.abs_().mul_(update.sign_())
y.sub_(u)
del u
else:
y.lerp_(end=z, weight=ckp1)
@@ -233,7 +263,7 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
@torch.no_grad()
def step_param(self, p, group):
self.on_start_step()
self.on_start_step(p, group)
if not group['train_mode']:
raise Exception("Not in train mode!")
@@ -241,8 +271,6 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
weight_sum = group['weight_sum']
if p.grad is not None:
grad = p.grad.to(dtype=torch.float32, copy=True)
use_adopt = group['use_adopt']
stochastic = group['stochastic_rounding']
_, beta2 = group['betas']
@@ -250,40 +278,61 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
state = self.initialise_state(p, group)
z_state = state['z']
y, z = (p.float(), z_state.float()) if stochastic else (p, z_state)
grad = self.orthograd(z_state, p.grad) if group['use_orthograd'] else p.grad.to(dtype=torch.float32, copy=True)
dlr = self.get_dlr(group)
if group['use_bias_correction']:
beta2_t = beta2 ** k
bias_correction2 = 1 - beta2_t
# maximum length of the approximated SMA
rho_inf = 2 / (1 - beta2) - 1
# compute the length of the approximated SMA
rho_t = rho_inf - 2 * k * beta2_t / bias_correction2
rect = (
((rho_t - 4) * (rho_t - 2) * rho_inf / ((rho_inf - 4) * (rho_inf - 2) * rho_t)) ** 0.5
if rho_t > 4.0
else 0.0
)
dlr *= rect
beta2 = 1 - (1 - beta2) / (1 - beta2_t)
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))
if group['use_speed']:
grad_rms = state['rms_sq']
if grad_rms is None:
grad_rms = state['rms_sq'] = 1 / self.get_rms(grad)
update = grad.mul_(grad_rms)
else:
d_k = group['d_prev'] / group['d']
rms_sq = state["rms_sq"].mul_(beta2 * d_k * d_k).add_(self.get_rms(grad).square(), alpha=1 - beta2)
update = grad.mul_(1 / rms_sq.sqrt().add(1e-8))
else:
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)
if use_adopt and group['k'] == 1:
self.update_second_moment(state, group, grad, 0, return_denom=False)
self.update_second_moment(state, group, grad, 0, y, return_denom=False)
else:
denom = self.update_second_moment(state, group, grad, beta2, denom_before_update=use_adopt)
update, num_scale = self.update_(grad, denom, group)
denom = self.update_second_moment(state, group, grad, beta2, y, denom_before_update=use_adopt)
if group['use_bias_correction'] and rho_t <= 4.0:
update = grad
else:
update = self.update_(grad, denom, group, y)
del denom
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)
rms = self.get_rms(update, 1).div(clip_threshold).clamp_min(1)
update.mul_(1 / rms)
z_state = state['z']
self.update_prodigy(state, group, p.grad, z_state, 1.0)
self.update_prodigy(state, group, p.grad, p)
y, z = (p.float(), z_state.float()) if stochastic else (p, z_state)
weight_sum = self.update_params(y, z, update, group)
weight_sum = self.update_params(y, z, update, group, dlr)
self.smart_copy(p, y, stochastic, True)
self.smart_copy(z_state, z, stochastic, True)