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: if not args.apply_t5_attn_mask:
t5_attn_mask = None t5_attn_mask = None
if args.bypass_flux_guidance:
flux_utils.bypass_flux_guidance(flux)
with accelerator.autocast(): 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) # 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( model_pred = flux(
@@ -685,6 +688,9 @@ class FluxTrainer:
# unpack latents # unpack latents
model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width) 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 # apply model prediction type
model_pred, weighting = flux_train_utils.apply_model_prediction_type(args, model_pred, noisy_model_input, sigmas) 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 return model_pred
if args.bypass_flux_guidance:
flux_utils.bypass_flux_guidance(unet)
model_pred = call_dit( model_pred = call_dit(
img=packed_noisy_model_input, img=packed_noisy_model_input,
img_ids=img_ids, img_ids=img_ids,
@@ -371,6 +374,9 @@ class FluxNetworkTrainer(NetworkTrainer):
# unpack latents # unpack latents
model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width) 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 # apply model prediction type
model_pred, weighting = flux_train_utils.apply_model_prediction_type(args, model_pred, noisy_model_input, sigmas) 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, default=3.0,
help="Discrete flow shift for the Euler Discrete Scheduler, default is 3.0. / Euler Discrete Schedulerの離散フローシフト、デフォルトは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_DEV = "dev"
MODEL_NAME_SCHNELL = "schnell" 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]]: 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."}), "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)."}), #"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)."}), #"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"}), "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_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_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_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."}), "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"}), "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"}), "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)) dataset_toml = toml.dumps(json.loads(dataset_config))
parser = train_network_setup_parser() parser = train_network_setup_parser()
flux_train_utils.add_flux_train_arguments(parser)
if additional_args is not None: if additional_args is not None:
print(f"additional_args: {additional_args}") print(f"additional_args: {additional_args}")
args, _ = parser.parse_known_args(args=shlex.split(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, use_bias_correction=False,
d0=1e-6, d_coef=1.0, d0=1e-6, d_coef=1.0,
prodigy_steps=0, prodigy_steps=0,
use_speed=False,
eps=1e-8, eps=1e-8,
split_groups=True, split_groups=True,
split_groups_mean=True, split_groups_mean=True,
factored=True, factored=True,
factored_fp32=True,
fused_back_pass=False, fused_back_pass=False,
use_stableadamw=True, use_stableadamw=True,
use_muon_pp=False, use_muon_pp=False,
use_cautious=False, use_cautious=False,
use_grams=False, use_grams=False,
use_adopt=False, use_adopt=False,
use_orthograd=False,
use_focus=False,
stochastic_rounding=True): stochastic_rounding=True):
if not 0.0 < d0: 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').") print(f"[{self.__class__.__name__}] 'use_grams' has been disabled (mutually exclusive with 'use_cautious').")
use_grams = False 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, defaults = dict(lr=lr, betas=betas, beta3=beta3,
eps=eps, eps=eps,
weight_decay=weight_decay, weight_decay=weight_decay,
@@ -58,15 +73,19 @@ class CoreOptimiser(torch.optim.Optimizer):
k=1, train_mode=True, k=1, train_mode=True,
weight_sum=0, weight_sum=0,
prodigy_steps=prodigy_steps, prodigy_steps=prodigy_steps,
use_speed=use_speed,
use_bias_correction=use_bias_correction, use_bias_correction=use_bias_correction,
d_numerator=0.0, d_numerator=0.0,
d_denom=0, d_denom=0,
factored=factored, factored=factored,
factored_fp32=factored_fp32,
use_stableadamw=use_stableadamw, use_stableadamw=use_stableadamw,
use_muon_pp=use_muon_pp, use_muon_pp=use_muon_pp,
use_cautious=use_cautious, use_cautious=use_cautious,
use_grams=use_grams, use_grams=use_grams,
use_adopt=use_adopt, use_adopt=use_adopt,
use_orthograd=use_orthograd,
use_focus=use_focus,
stochastic_rounding=stochastic_rounding) stochastic_rounding=stochastic_rounding)
super().__init__(params, defaults) super().__init__(params, defaults)
@@ -113,10 +132,21 @@ class CoreOptimiser(torch.optim.Optimizer):
def get_sliced_tensor(self, tensor, slice_p=11): def get_sliced_tensor(self, tensor, slice_p=11):
return tensor.ravel()[::slice_p] 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() @torch.no_grad()
def get_running_values_for_group(self, group): def get_running_values_for_group(self, group):
if not self.split_groups: if not self.split_groups:
group = self.param_groups[0] group = self.param_groups[0]
return group['running_d_numerator'], group['running_d_denom'] return group['running_d_numerator'], group['running_d_denom']
@torch.no_grad() @torch.no_grad()
@@ -135,10 +165,11 @@ class CoreOptimiser(torch.optim.Optimizer):
@torch.no_grad() @torch.no_grad()
def newton_schulz_(self, G, steps=6, eps=1e-7): def newton_schulz_(self, G, steps=6, eps=1e-7):
# Inline reshaping step within the method itself. # 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) 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): if G.size(0) > G.size(1):
X = X.T X = X.T
@@ -153,10 +184,33 @@ class CoreOptimiser(torch.optim.Optimizer):
# Gradient scaling adaptation from: https://github.com/leloykun/adaptive-muon # 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 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 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 # Implementation by Nerogar. From: https://github.com/pytorch/pytorch/issues/120376#issuecomment-1974828905
def copy_stochastic_(self, target, source): def copy_stochastic_(self, target, source):
@@ -224,14 +278,18 @@ class CoreOptimiser(torch.optim.Optimizer):
if needs_init: if needs_init:
grad = p.grad 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) 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. # 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 state['muon'] = group['use_muon_pp'] and len(grad.shape) >= 2
if state['muon']: 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: else:
factored_dims = self.factored_dims( factored_dims = self.factored_dims(
grad.shape, grad.shape,
@@ -240,20 +298,20 @@ class CoreOptimiser(torch.optim.Optimizer):
) )
if factored_dims is not None: if factored_dims is not None:
# Store reduction variables so we don't have to recalculate each step.
dc, dr = factored_dims dc, dr = factored_dims
row_shape = list(p.grad.shape) row_shape = list(grad.shape)
row_shape[dr] = 1 row_shape[dr] = 1
col_shape = list(p.grad.shape) col_shape = list(grad.shape)
col_shape[dc] = 1 col_shape[dc] = 1
reduce_dc = dc - 1 if dc > dr else dc 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 factored_dtype = torch.float32 if group['factored_fp32'] else grad.dtype
# between bf16/fp16 and fp32 is negligible here. state["exp_avg_sq"] = [torch.zeros(row_shape, dtype=factored_dtype, device=p.device).detach(),
state["exp_avg_sq"] = [torch.zeros(row_shape, dtype=torch.float32, device=p.device).detach(), torch.zeros(col_shape, dtype=factored_dtype, device=p.device).detach(),
torch.zeros(col_shape, dtype=torch.float32, device=p.device).detach(),
dr, dc, reduce_dc] dr, dc, reduce_dc]
else: 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 the initial weights are zero, don't bother storing them.
if p.any() > 0: if p.any() > 0:
@@ -266,7 +324,7 @@ class CoreOptimiser(torch.optim.Optimizer):
return state, needs_init return state, needs_init
@torch.no_grad() @torch.no_grad()
def update_d_and_reset(self, group): def update_d_stats_and_reset(self, group):
k = group['k'] k = group['k']
prodigy_steps = group['prodigy_steps'] prodigy_steps = group['prodigy_steps']
@@ -274,8 +332,6 @@ class CoreOptimiser(torch.optim.Optimizer):
return return
d, d0 = group['d'], group['d0'] d, d0 = group['d'], group['d0']
d_prev = group['d_prev']
d_coef = group['d_coef']
beta3 = group['beta3'] beta3 = group['beta3']
running_d_numerator, running_d_denom = self.get_running_values_for_group(group) 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 = group['d_numerator']
d_numerator *= beta3 d_numerator *= beta3
d_prev = d
d_numerator_item = running_d_numerator.item() d_numerator_item = running_d_numerator.item()
d_denom_item = running_d_denom.item() d_denom_item = running_d_denom.item()
@@ -299,21 +353,38 @@ class CoreOptimiser(torch.optim.Optimizer):
else: else:
d_numerator += d_numerator_item 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_numerator'] = d_numerator
group['d_denom'] = d_denom_item group['d_denom'] = d_denom_item
running_d_numerator.zero_() running_d_numerator.zero_()
running_d_denom.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: 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. # 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) 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): def on_end_step(self):
self.parameters_to_process -= 1 self.parameters_to_process -= 1
@@ -321,25 +392,26 @@ class CoreOptimiser(torch.optim.Optimizer):
if self.parameters_to_process == 0: if self.parameters_to_process == 0:
# Update d for next optimiser step. # Update d for next optimiser step.
if self.split_groups: if self.split_groups:
i = 0 for i, group in enumerate(self.param_groups):
for group in self.param_groups:
if group['prodigy_steps'] > 0 and group['k'] == group['prodigy_steps']: 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}.") 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['weight_sum'] = group.get('running_weight_sum', 0)
group['k'] += 1 group['k'] += 1
i += 1
self.shared_d = self.get_d_mean() self.shared_d = self.get_d_mean()
else: else:
# When groups aren't split, calculate d for the first group (which collects stats for all groups in non-split mode), # 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. # then copy to all other groups.
first_group = self.param_groups[0] 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 i, group in enumerate(self.param_groups):
for group in self.param_groups:
if group['prodigy_steps'] > 0 and group['k'] == group['prodigy_steps']: 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}.") 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['d_denom'] = first_group['d_denom']
group['weight_sum'] = group.get('running_weight_sum', 0) group['weight_sum'] = group.get('running_weight_sum', 0)
group['k'] += 1 group['k'] += 1
i += 1
def get_dlr(self, group): def get_dlr(self, group):
return (self.shared_d if self.split_groups and self.shared_d else group['d']) * group['lr'] 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, 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.
def update_prodigy(self, state, group, grad, data):
k = group['k'] k = group['k']
prodigy_steps = group['prodigy_steps'] prodigy_steps = group['prodigy_steps']
if prodigy_steps <= 0 or k < prodigy_steps: if prodigy_steps <= 0 or k < prodigy_steps:
beta3 = group['beta3'] beta3 = group['beta3']
d, d0 = group['d'], group['d0'] d_update = group['d'] ** 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_grad = self.get_sliced_tensor(grad)
sliced_data = self.get_sliced_tensor(data) sliced_data = self.get_sliced_tensor(data)
running_d_numerator, running_d_denom = self.get_running_values_for_group(group) running_d_numerator, running_d_denom = self.get_running_values_for_group(group)
s = state['s']
x0_minus = state['p0'] - sliced_data 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()) running_d_denom.add_(s.abs().sum())
del x0_minus del x0_minus
elif 's' in state: # Free the memory used by Prodigy, as we no longer need it. elif 's' in state: # Free the memory used by Prodigy, as we no longer need it.
del state['s'] del state['s']
del state['p0'] 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'] eps = group['eps']
if eps is None: if eps is None:
# Approximate scaling for a regular Adam-style update. # Approximate scaling for a regular Adam-style update.
b = self.get_clip_threshold(group) b = self.get_clip_threshold(group)
a = 1 / math.atan(1 / b)
# Adam-atan2. Use atan2 rather than epsilon and division # Adam-atan2. Use atan2 rather than epsilon and division
# for parameter updates (https://arxiv.org/abs/2407.05872). # for parameter updates (https://arxiv.org/abs/2407.05872).
# Has the nice property of "clipping" the gradient as well. # 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: else:
update = num.div_(denom.add_(eps)) 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'] exp_avg_sq = state['exp_avg_sq']
# Adam EMA updates # Adam EMA updates
@@ -415,53 +493,63 @@ class CoreOptimiser(torch.optim.Optimizer):
row_factor = row_var.div(row_col_mean).sqrt_() row_factor = row_var.div(row_col_mean).sqrt_()
col_factor = col_var.sqrt() col_factor = col_var.sqrt()
denom = row_factor * col_factor denom = row_factor * col_factor
elif group['use_focus']:
denom = exp_avg_sq.clone()
else: else:
denom = exp_avg_sq.sqrt() denom = exp_avg_sq.sqrt()
return denom return denom
def update_first_moment(self, state, group, grad): def update_first_moment(self, state, group, grad, beta1):
exp_avg = state['exp_avg'] 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'] exp_avg_sq = state['exp_avg_sq']
d_k = group['d_prev'] / group['d']
denom = None denom = None
if return_denom and denom_before_update: if return_denom and denom_before_update:
denom = self.get_denom(state) denom = self.get_denom(state, group)
# Adam EMA updates # 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): if isinstance(exp_avg_sq, list):
row_var, col_var, dr, dc, _ = exp_avg_sq row_var, col_var, dr, dc, _ = exp_avg_sq
row_var.lerp_( row_var.mul_(beta2 * d_k * d_k).add_(
grad.norm(dim=dr, keepdim=True).square_().div_(grad.shape[dr]), grad.norm(dim=dr, keepdim=True).square_().mul_(1 / grad.shape[dr]),
weight=1 - beta2 alpha=1 - beta2
) )
col_var.lerp_( col_var.mul_(beta2 * d_k * d_k).add_(
grad.norm(dim=dc, keepdim=True).square_().div_(grad.shape[dc]), grad.norm(dim=dc, keepdim=True).square_().mul_(1 / grad.shape[dc]),
weight=1 - beta2 alpha=1 - beta2
) )
else: 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: if return_denom and denom is None:
denom = self.get_denom(state) denom = self.get_denom(state, group)
return denom return denom
def get_rms(self, tensor, eps=1e-8): 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): def rms_(self, tensor, eps):
return tensor.div_(self.get_rms(tensor, eps)) return tensor.div_(self.get_rms(tensor, eps))
def get_clip_threshold(self, group): 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): def try_hook_kohya_fbp(self):
self.kohya_original_patch_adafactor_fused = None self.kohya_original_patch_adafactor_fused = None
@@ -62,7 +62,11 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
If False, weight_decay will have a much stronger effect. If False, weight_decay will have a much stronger effect.
(default: True). (default: True).
use_bias_correction (boolean): 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): d0 (float):
Initial estimate for Prodigy. Also serves as the minimum learning rate. Initial estimate for Prodigy. Also serves as the minimum learning rate.
(default: 1e-6). (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 Freeze Prodigy stepsize adjustments after a certain optimiser step and releases all state memory required
by Prodigy. by Prodigy.
(default: 0) (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): split_groups (boolean):
Track individual adaptation values for each parameter group. For example, if training 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. 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 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. if training results in NaNs or the learning rate fails to grow.
(default: True) (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): fused_back_pass (boolean):
Stops the optimiser from running the normal step method. Set to True if using fused backward pass. Really only 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). 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 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. (https://arxiv.org/abs/2411.02853), as we don't have a first moment to use for the update.
(default: False) (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): stochastic_rounding (boolean):
Use stochastic rounding for bfloat16 weights (https://github.com/pytorch/pytorch/issues/120376). Brings Use stochastic rounding for bfloat16 weights (https://github.com/pytorch/pytorch/issues/120376). Brings
bfloat16 training performance close to that of float32. bfloat16 training performance close to that of float32.
@@ -130,27 +155,32 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
use_bias_correction=False, use_bias_correction=False,
d0=1e-6, d_coef=1.0, d0=1e-6, d_coef=1.0,
prodigy_steps=0, prodigy_steps=0,
use_speed=False,
eps=1e-8, eps=1e-8,
split_groups=True, split_groups=True,
split_groups_mean=True, split_groups_mean=True,
factored=True, factored=True,
factored_fp32=True,
fused_back_pass=False, fused_back_pass=False,
use_stableadamw=True, use_stableadamw=True,
use_muon_pp=False, use_muon_pp=False,
use_cautious=False, use_cautious=False,
use_grams=False, use_grams=False,
use_adopt=False, use_adopt=False,
use_orthograd=False,
use_focus=False,
stochastic_rounding=True): stochastic_rounding=True):
super().__init__(params=params, lr=lr, betas=betas, beta3=beta3, super().__init__(params=params, lr=lr, betas=betas, beta3=beta3,
weight_decay=weight_decay, weight_decay_by_lr=weight_decay_by_lr, weight_decay=weight_decay, weight_decay_by_lr=weight_decay_by_lr,
use_bias_correction=use_bias_correction, 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, 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, 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_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() @torch.no_grad()
def eval(self): def eval(self):
@@ -188,9 +218,7 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
return state return state
@torch.no_grad() @torch.no_grad()
def update_params(self, y, z, update, group): def update_params(self, y, z, update, group, dlr):
dlr = self.get_dlr(group)
beta1, _ = group['betas'] beta1, _ = group['betas']
decay = group['weight_decay'] decay = group['weight_decay']
@@ -208,21 +236,23 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
y.sub_(y, alpha=decay * xy_step) y.sub_(y, alpha=decay * xy_step)
z.sub_(y, alpha=decay) z.sub_(y, alpha=decay)
if group['use_cautious']: cautious, grams = group['use_cautious'], group['use_grams']
# "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 if cautious or grams:
u = (y - z).mul_(ckp1).add_(update, alpha=dlr * xy_step) u = (y - z).mul_(ckp1).add_(update, alpha=dlr * xy_step)
z.sub_(update, alpha=dlr) 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) u.mul_(mask)
y.sub_(u) elif grams:
del mask, u
elif group['use_grams']:
# "Grams: Gradient Descent with Adaptive Momentum Scaling": https://arxiv.org/abs/2412.17107 # "Grams: Gradient Descent with Adaptive Momentum Scaling": https://arxiv.org/abs/2412.17107
u = (y - z).mul_(ckp1).add_(update, alpha=dlr * xy_step) u.abs_().mul_(update.sign_())
z.sub_(update, alpha=dlr) # Update z now so we can do sign in-place.
y.sub_(u.abs_().mul_(update.sign_())) y.sub_(u)
del u del u
else: else:
y.lerp_(end=z, weight=ckp1) y.lerp_(end=z, weight=ckp1)
@@ -233,7 +263,7 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
@torch.no_grad() @torch.no_grad()
def step_param(self, p, group): def step_param(self, p, group):
self.on_start_step() self.on_start_step(p, group)
if not group['train_mode']: if not group['train_mode']:
raise Exception("Not in train mode!") raise Exception("Not in train mode!")
@@ -241,8 +271,6 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
weight_sum = group['weight_sum'] weight_sum = group['weight_sum']
if p.grad is not None: if p.grad is not None:
grad = p.grad.to(dtype=torch.float32, copy=True)
use_adopt = group['use_adopt'] use_adopt = group['use_adopt']
stochastic = group['stochastic_rounding'] stochastic = group['stochastic_rounding']
_, beta2 = group['betas'] _, beta2 = group['betas']
@@ -250,40 +278,61 @@ class ProdigyPlusScheduleFree(CoreOptimiser):
state = self.initialise_state(p, group) 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 update = None
if state['muon']: if state['muon']:
grad = self.newton_schulz_(grad) grad = self.newton_schulz_(grad)
grad_rms = self.get_rms(grad).item() ** 2 if group['use_speed']:
grad_rms = state['rms_sq']
rms_sq = (state["rms_sq"] * beta2) + (grad_rms * (1 - beta2)) if grad_rms is None:
state["rms_sq"] = rms_sq grad_rms = state['rms_sq'] = 1 / self.get_rms(grad)
update = grad.mul_(grad_rms)
update = grad.mul_(1.0 / ((rms_sq ** 0.5) + 1e-12)) 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: 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: 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: else:
denom = self.update_second_moment(state, group, grad, beta2, denom_before_update=use_adopt) denom = self.update_second_moment(state, group, grad, beta2, y, denom_before_update=use_adopt)
update, num_scale = self.update_(grad, denom, group) if group['use_bias_correction'] and rho_t <= 4.0:
update = grad
else:
update = self.update_(grad, denom, group, y)
del denom del denom
if update is not None: if update is not None:
if group['use_stableadamw']: if group['use_stableadamw']:
clip_threshold = self.get_clip_threshold(group) clip_threshold = self.get_clip_threshold(group)
num_scale = max(1, self.get_rms(update, 1.0).item() / clip_threshold) rms = self.get_rms(update, 1).div(clip_threshold).clamp_min(1)
update.mul_(1 / num_scale) update.mul_(1 / rms)
z_state = state['z'] self.update_prodigy(state, group, p.grad, p)
self.update_prodigy(state, group, p.grad, z_state, 1.0)
y, z = (p.float(), z_state.float()) if stochastic else (p, z_state) weight_sum = self.update_params(y, z, update, group, dlr)
weight_sum = self.update_params(y, z, update, group)
self.smart_copy(p, y, stochastic, True) self.smart_copy(p, y, stochastic, True)
self.smart_copy(z_state, z, stochastic, True) self.smart_copy(z_state, z, stochastic, True)