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