update prodigyplusschedulefree, support Flex with arg bypass_flux_guidance
https://github.com/kohya-ss/sd-scripts/pull/1893
This commit is contained in:
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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"
|
||||||
|
)
|
||||||
|
|||||||
@@ -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]]:
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -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))
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user