diff --git a/flux_train_comfy.py b/flux_train_comfy.py index 60502a2..6251ccc 100644 --- a/flux_train_comfy.py +++ b/flux_train_comfy.py @@ -669,6 +669,9 @@ class FluxTrainer: if not args.apply_t5_attn_mask: t5_attn_mask = None + if args.bypass_flux_guidance: + flux_utils.bypass_flux_guidance(flux) + with accelerator.autocast(): # YiYi notes: divide it by 1000 for now because we scale it by 1000 in the transformer model (we should not keep it but I want to keep the inputs same for the model for testing) model_pred = flux( @@ -685,6 +688,9 @@ class FluxTrainer: # unpack latents model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width) + if args.bypass_flux_guidance: + flux_utils.restore_flux_guidance(flux) + # apply model prediction type model_pred, weighting = flux_train_utils.apply_model_prediction_type(args, model_pred, noisy_model_input, sigmas) diff --git a/flux_train_network_comfy.py b/flux_train_network_comfy.py index 7494acd..8cbbd9e 100644 --- a/flux_train_network_comfy.py +++ b/flux_train_network_comfy.py @@ -356,6 +356,9 @@ class FluxNetworkTrainer(NetworkTrainer): """ return model_pred + + if args.bypass_flux_guidance: + flux_utils.bypass_flux_guidance(unet) model_pred = call_dit( img=packed_noisy_model_input, @@ -371,6 +374,9 @@ class FluxNetworkTrainer(NetworkTrainer): # unpack latents model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width) + if args.bypass_flux_guidance: #for flex + flux_utils.restore_flux_guidance(unet) + # apply model prediction type model_pred, weighting = flux_train_utils.apply_model_prediction_type(args, model_pred, noisy_model_input, sigmas) diff --git a/library/flux_train_utils.py b/library/flux_train_utils.py index 6a2506a..c518221 100644 --- a/library/flux_train_utils.py +++ b/library/flux_train_utils.py @@ -578,3 +578,8 @@ def add_flux_train_arguments(parser: argparse.ArgumentParser): default=3.0, help="Discrete flow shift for the Euler Discrete Scheduler, default is 3.0. / Euler Discrete Schedulerの離散フローシフト、デフォルトは3.0。", ) + parser.add_argument( + "--bypass_flux_guidance" + , action="store_true" + , help="bypass flux guidance module for Flex.1-Alpha Training" + ) diff --git a/library/flux_utils.py b/library/flux_utils.py index 3034b07..62c58bd 100644 --- a/library/flux_utils.py +++ b/library/flux_utils.py @@ -21,7 +21,14 @@ MODEL_VERSION_FLUX_V1 = "flux1" MODEL_NAME_DEV = "dev" MODEL_NAME_SCHNELL = "schnell" +# bypass guidance +def bypass_flux_guidance(transformer): + transformer.params.guidance_embed = False +# restore the forward function +def restore_flux_guidance(transformer): + transformer.params.guidance_embed = True + def analyze_checkpoint_state(ckpt_path: str) -> Tuple[bool, bool, Tuple[int, int], List[str]]: """ チェックポイントの状態を分析し、DiffusersかBFLか、devかschnellか、ブロック数を計算して返す。 diff --git a/nodes.py b/nodes.py index c43b8cb..c53e04a 100644 --- a/nodes.py +++ b/nodes.py @@ -396,15 +396,16 @@ class OptimizerConfigProdigyPlusScheduleFree: "split_groups": ("BOOLEAN",{"default": True, "tooltip": "Track individual adaptation values for each parameter group."}), #"beta3": ("FLOAT",{"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.0001, "tooltip": " Coefficient for computing the Prodigy stepsize using running averages. If set to None, uses the value of square root of beta2 (default: None)."}), #"beta4": ("FLOAT",{"default": 0, "min": 0.0, "max": 1.0, "step": 0.0001, "tooltip": "Coefficient for updating the learning rate from Prodigy's adaptive stepsize. Smooths out spikes in learning rate adjustments. If set to None, beta1 is used instead. (default 0, which disables smoothing and uses original Prodigy behaviour)."}), - "use_bias_correction": ("BOOLEAN",{"default": False, "tooltip": "Turn on Adafactor-style bias correction, which scales beta2 directly."}), + "use_bias_correction": ("BOOLEAN",{"default": False, "tooltip": "Use the RAdam variant of schedule-free"}), "min_snr_gamma": ("FLOAT",{"default": 5.0, "min": 0.0, "step": 0.01, "tooltip": "gamma for reducing the weight of high loss timesteps. Lower numbers have stronger effect. 5 is recommended by the paper"}), "use_stableadamw": ("BOOLEAN",{"default": True, "tooltip": "Scales parameter updates by the root-mean-square of the normalised gradient, in essence identical to Adafactor's gradient scaling. Set to False if the adaptive learning rate never improves."}), "use_cautious" : ("BOOLEAN",{"default": False, "tooltip": "Experimental. Perform 'cautious' updates, as proposed in https://arxiv.org/pdf/2411.16085. Modifies the update to isolate and boost values that align with the current gradient."}), "use_adopt": ("BOOLEAN",{"default": False, "tooltip": "Experimental. Performs a modified step where the second moment is updated after the parameter update, so as not to include the current gradient in the denominator. This is a partial implementation of ADOPT (https://arxiv.org/abs/2411.02853), as we don't have a first moment to use for the update."}), "use_grams": ("BOOLEAN",{"default": False, "tooltip": "Perform 'grams' updates, as proposed in https://arxiv.org/abs/2412.17107. Modifies the update using sign operations that align with the current gradient. Note that we do not have access to a first moment, so this deviates from the paper (we apply the sign directly to the update). May have a limited effect."}), "stochastic_rounding": ("BOOLEAN",{"default": True, "tooltip": "Use stochastic rounding for bfloat16 weights"}), + "use_orthograd": ("BOOLEAN",{"default": False, "tooltip": "Experimental. Updates weights using the component of the gradient that is orthogonal to the current weight direction, as described in (https://arxiv.org/pdf/2501.04697). Can help prevent overfitting and improve generalisation."}), + "use_focus ": ("BOOLEAN",{"default": False, "tooltip": "Experimental. Modifies the update step to better handle noise at large step sizes. (https://arxiv.org/abs/2501.12243). This method is incompatible with factorisation, Muon and Adam-atan2."}), "extra_optimizer_args": ("STRING",{"multiline": True, "default": "", "tooltip": "additional optimizer args"}), - }, } @@ -497,6 +498,7 @@ class InitFluxLoRATraining: dataset_toml = toml.dumps(json.loads(dataset_config)) parser = train_network_setup_parser() + flux_train_utils.add_flux_train_arguments(parser) if additional_args is not None: print(f"additional_args: {additional_args}") args, _ = parser.parse_known_args(args=shlex.split(additional_args)) diff --git a/prodigyplusschedulefree/core_optimiser.py b/prodigyplusschedulefree/core_optimiser.py index e0aeb95..a50664a 100644 --- a/prodigyplusschedulefree/core_optimiser.py +++ b/prodigyplusschedulefree/core_optimiser.py @@ -10,16 +10,20 @@ class CoreOptimiser(torch.optim.Optimizer): use_bias_correction=False, d0=1e-6, d_coef=1.0, prodigy_steps=0, + use_speed=False, eps=1e-8, split_groups=True, split_groups_mean=True, factored=True, + factored_fp32=True, fused_back_pass=False, use_stableadamw=True, use_muon_pp=False, use_cautious=False, use_grams=False, use_adopt=False, + use_orthograd=False, + use_focus=False, stochastic_rounding=True): if not 0.0 < d0: @@ -50,6 +54,17 @@ class CoreOptimiser(torch.optim.Optimizer): print(f"[{self.__class__.__name__}] 'use_grams' has been disabled (mutually exclusive with 'use_cautious').") use_grams = False + if use_focus: + if factored: + print(f"[{self.__class__.__name__}] 'factored' has been disabled (incompatible with 'use_focus').") + factored = False + if use_muon_pp: + print(f"[{self.__class__.__name__}] 'use_muon_pp' has been disabled (incompatible with 'use_focus').") + use_muon_pp = False + if eps is None: + print(f"[{self.__class__.__name__}] Adam-atan2 ('eps=None') has been disabled (incompatible with 'use_focus').") + # We skip the Adam-atan2 branch entirely when FOCUS is enabled. + defaults = dict(lr=lr, betas=betas, beta3=beta3, eps=eps, weight_decay=weight_decay, @@ -58,15 +73,19 @@ class CoreOptimiser(torch.optim.Optimizer): k=1, train_mode=True, weight_sum=0, prodigy_steps=prodigy_steps, + use_speed=use_speed, use_bias_correction=use_bias_correction, d_numerator=0.0, d_denom=0, factored=factored, + factored_fp32=factored_fp32, use_stableadamw=use_stableadamw, use_muon_pp=use_muon_pp, use_cautious=use_cautious, use_grams=use_grams, use_adopt=use_adopt, + use_orthograd=use_orthograd, + use_focus=use_focus, stochastic_rounding=stochastic_rounding) super().__init__(params, defaults) @@ -112,11 +131,22 @@ class CoreOptimiser(torch.optim.Optimizer): @torch.no_grad() def get_sliced_tensor(self, tensor, slice_p=11): return tensor.ravel()[::slice_p] - + + @torch.no_grad() + def check_running_values_for_group(self, p, group): + if not self.split_groups: + group = self.param_groups[0] + + if group['running_d_numerator'].device != p.device: + group['running_d_numerator'] = group['running_d_numerator'].to(p.device) + if group['running_d_denom'].device != p.device: + group['running_d_denom'] = group['running_d_denom'].to(p.device) + @torch.no_grad() def get_running_values_for_group(self, group): if not self.split_groups: group = self.param_groups[0] + return group['running_d_numerator'], group['running_d_denom'] @torch.no_grad() @@ -135,10 +165,11 @@ class CoreOptimiser(torch.optim.Optimizer): @torch.no_grad() def newton_schulz_(self, G, steps=6, eps=1e-7): # Inline reshaping step within the method itself. - X = G.view(G.size(0), -1) + G_shape = G.shape + G = G.view(G.size(0), -1) a, b, c = (3.4445, -4.7750, 2.0315) - X = X.to(dtype=torch.bfloat16, copy=True) + X = G.to(dtype=torch.bfloat16, copy=True) if G.size(0) > G.size(1): X = X.T @@ -153,11 +184,34 @@ class CoreOptimiser(torch.optim.Optimizer): # Gradient scaling adaptation from: https://github.com/leloykun/adaptive-muon X = torch.einsum('ij,ij->', G.type_as(X), X).clamp(-1.0, 1.0) * X - G.copy_(X.view_as(G)) + G.copy_(X) del X - return G + return G.view(G_shape) + # Implementation from: https://github.com/LucasPrietoAl/grokking-at-the-edge-of-numerical-stability/blob/main/orthograd.py + def orthograd(self, p, grad): + G_shape = grad.shape + w = p.view(-1) + g = grad.view(-1) + + proj = torch.dot(w, g) / (torch.dot(w, w) + 1e-30) + g_orth = g.to(dtype=torch.float32, copy=True).sub_(w, alpha=proj) + g_orth_scaled = g_orth.mul_(g.norm(2) / (g_orth.norm(2) + 1e-30)) + + return g_orth_scaled.view(G_shape) + + def orthograd_(self, p, grad): + G_shape = grad.shape + w = p.view(-1) + g = grad.view(-1) + + proj = torch.dot(w, g) / (torch.dot(w, w) + 1e-30) + g_orth = g.sub_(w, alpha=proj) + g_orth_scaled = g_orth.mul_(g.norm(2) / (g_orth.norm(2) + 1e-30)) + + return g_orth_scaled.view(G_shape) + # Implementation by Nerogar. From: https://github.com/pytorch/pytorch/issues/120376#issuecomment-1974828905 def copy_stochastic_(self, target, source): # create a random 16 bit integer @@ -224,36 +278,40 @@ class CoreOptimiser(torch.optim.Optimizer): if needs_init: grad = p.grad - dtype = torch.bfloat16 if p.dtype == torch.float32 else p.dtype + dtype = torch.bfloat16 if grad.dtype == torch.float32 else grad.dtype sliced_data = self.get_sliced_tensor(p) - # NOTE: We don't initialise z/exp_avg here -- subclass needs to do that. - state['muon'] = group['use_muon_pp'] and len(grad.shape) >= 2 - - if state['muon']: - state["rms_sq"] = 0 + if group['use_focus']: + state['exp_avg_sq'] = torch.zeros_like(grad, memory_format=torch.preserve_format).detach() + state['muon'] = False else: - factored_dims = self.factored_dims( - grad.shape, - factored=group['factored'], - min_dim_size_to_factor=32 - ) + # NOTE: We don't initialise z/exp_avg here -- subclass needs to do that. + state['muon'] = group['use_muon_pp'] and len(grad.shape) >= 2 - if factored_dims is not None: - dc, dr = factored_dims - row_shape = list(p.grad.shape) - row_shape[dr] = 1 - col_shape = list(p.grad.shape) - col_shape[dc] = 1 - reduce_dc = dc - 1 if dc > dr else dc - # Store reduction variables so we don't have to recalculate each step. - # Always store second moment low ranks in fp32 to avoid precision issues. Memory difference - # between bf16/fp16 and fp32 is negligible here. - state["exp_avg_sq"] = [torch.zeros(row_shape, dtype=torch.float32, device=p.device).detach(), - torch.zeros(col_shape, dtype=torch.float32, device=p.device).detach(), - dr, dc, reduce_dc] + if state['muon']: + state["rms_sq"] = None if group['use_speed'] else torch.tensor(0.0, dtype=dtype, device=p.device) else: - state['exp_avg_sq'] = torch.zeros_like(p, memory_format=torch.preserve_format).detach() + factored_dims = self.factored_dims( + grad.shape, + factored=group['factored'], + min_dim_size_to_factor=32 + ) + + if factored_dims is not None: + # Store reduction variables so we don't have to recalculate each step. + dc, dr = factored_dims + row_shape = list(grad.shape) + row_shape[dr] = 1 + col_shape = list(grad.shape) + col_shape[dc] = 1 + reduce_dc = dc - 1 if dc > dr else dc + + factored_dtype = torch.float32 if group['factored_fp32'] else grad.dtype + state["exp_avg_sq"] = [torch.zeros(row_shape, dtype=factored_dtype, device=p.device).detach(), + torch.zeros(col_shape, dtype=factored_dtype, device=p.device).detach(), + dr, dc, reduce_dc] + else: + state['exp_avg_sq'] = torch.zeros_like(grad, memory_format=torch.preserve_format).detach() # If the initial weights are zero, don't bother storing them. if p.any() > 0: @@ -266,7 +324,7 @@ class CoreOptimiser(torch.optim.Optimizer): return state, needs_init @torch.no_grad() - def update_d_and_reset(self, group): + def update_d_stats_and_reset(self, group): k = group['k'] prodigy_steps = group['prodigy_steps'] @@ -274,8 +332,6 @@ class CoreOptimiser(torch.optim.Optimizer): return d, d0 = group['d'], group['d0'] - d_prev = group['d_prev'] - d_coef = group['d_coef'] beta3 = group['beta3'] running_d_numerator, running_d_denom = self.get_running_values_for_group(group) @@ -283,8 +339,6 @@ class CoreOptimiser(torch.optim.Optimizer): d_numerator = group['d_numerator'] d_numerator *= beta3 - d_prev = d - d_numerator_item = running_d_numerator.item() d_denom_item = running_d_denom.item() @@ -299,47 +353,65 @@ class CoreOptimiser(torch.optim.Optimizer): else: d_numerator += d_numerator_item - d_hat = math.atan2(d_coef * d_numerator, d_denom_item) - d = max(d, d_hat) - - group['d'] = d - group['d_prev'] = d_prev group['d_numerator'] = d_numerator group['d_denom'] = d_denom_item running_d_numerator.zero_() running_d_denom.zero_() - def on_start_step(self): + @torch.no_grad() + def calculate_d(self, group): + k = group['k'] + prodigy_steps = group['prodigy_steps'] + + if prodigy_steps > 0 and k >= prodigy_steps: + return + + d = group['d'] + d_hat = math.atan2(group['d_coef'] * group['d_numerator'], group['d_denom']) + + if group['use_speed']: + d_hat = max(d, d_hat) + d = min(d_hat, d ** 0.975) + else: + d = max(d, d_hat) + + group['d_prev'] = group['d'] + group['d'] = d + + def on_start_step(self, p, group): if self.parameters_to_process is None or self.parameters_to_process == 0: # Optimiser hasn't run yet (or is starting a new step), so initialise. self.parameters_to_process = sum(len(group['params']) for group in self.param_groups) - + # Check running values are on-device. + self.check_running_values_for_group(p, group) + def on_end_step(self): self.parameters_to_process -= 1 if self.parameters_to_process == 0: # Update d for next optimiser step. if self.split_groups: - i = 0 - for group in self.param_groups: + for i, group in enumerate(self.param_groups): if group['prodigy_steps'] > 0 and group['k'] == group['prodigy_steps']: print(f"[{self.__class__.__name__}] Prodigy stepsize adaptation disabled after {group['k']} steps for param_group {i}.") - self.update_d_and_reset(group) + self.update_d_stats_and_reset(group) + + for group in self.param_groups: + self.calculate_d(group) group['weight_sum'] = group.get('running_weight_sum', 0) group['k'] += 1 - i += 1 self.shared_d = self.get_d_mean() else: # When groups aren't split, calculate d for the first group (which collects stats for all groups in non-split mode), # then copy to all other groups. first_group = self.param_groups[0] - self.update_d_and_reset(first_group) + self.update_d_stats_and_reset(first_group) + self.calculate_d(first_group) - i = 0 - for group in self.param_groups: + for i, group in enumerate(self.param_groups): if group['prodigy_steps'] > 0 and group['k'] == group['prodigy_steps']: print(f"[{self.__class__.__name__}] Prodigy stepsize adaptation disabled after {group['k']} steps for param_group {i}.") @@ -348,63 +420,69 @@ class CoreOptimiser(torch.optim.Optimizer): group['d_denom'] = first_group['d_denom'] group['weight_sum'] = group.get('running_weight_sum', 0) group['k'] += 1 - i += 1 def get_dlr(self, group): - return (self.shared_d if self.split_groups and self.shared_d else group['d']) * group['lr'] - - def update_prodigy(self, state, group, grad, data, num_scale): - # num_scale is used to compensate the numerator calculations when - # clipping/scaling is applied to the incoming update. If we don't - # do this, it will dampen Prodigy's 'd' predictions. + dlr = (self.shared_d if self.split_groups and self.shared_d else group['d']) * group['lr'] + return dlr * group.get('rect', 1.0) + def update_prodigy(self, state, group, grad, data): k = group['k'] prodigy_steps = group['prodigy_steps'] if prodigy_steps <= 0 or k < prodigy_steps: beta3 = group['beta3'] - d, d0 = group['d'], group['d0'] + 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_data = self.get_sliced_tensor(data) running_d_numerator, running_d_denom = self.get_running_values_for_group(group) - s = state['s'] - x0_minus = state['p0'] - sliced_data - running_d_numerator.add_(torch.dot(sliced_grad, x0_minus), alpha=d_k * num_scale) + x0_dot = torch.dot(sliced_grad, x0_minus) - s.mul_(beta3).add_(sliced_grad, alpha=d_k) + if group['use_speed']: + d_update *= group['d0'] + x0_dot, sliced_grad = x0_dot.sign(), sliced_grad.sign() + + s = state['s'] + s.mul_(beta3).add_(sliced_grad, alpha=d_update) + + running_d_numerator.add_(x0_dot, alpha=d_update) running_d_denom.add_(s.abs().sum()) del x0_minus elif 's' in state: # Free the memory used by Prodigy, as we no longer need it. del state['s'] del state['p0'] - def update_(self, num, denom, group): - eps = group['eps'] + 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 - if eps is None: - # Approximate scaling for a regular Adam-style update. - b = self.get_clip_threshold(group) - a = 1 / math.atan(1 / b) + # Original form. + # update = torch.sign(num) + gamma * torch.sign(w - denom) - # Adam-atan2. Use atan2 rather than epsilon and division - # for parameter updates (https://arxiv.org/abs/2407.05872). - # Has the nice property of "clipping" the gradient as well. - update = num.atan2_(denom.mul_(b)).mul_(a) + denom = denom.sub_(w).sign_().mul_(-gamma) + update = num.sign_().add_(denom) else: - update = num.div_(denom.add_(eps)) + eps = group['eps'] - return update, 1.0 + if eps is None: + # Approximate scaling for a regular Adam-style update. + b = self.get_clip_threshold(group) - def get_denom(self, state): + # Adam-atan2. Use atan2 rather than epsilon and division + # for parameter updates (https://arxiv.org/abs/2407.05872). + # Has the nice property of "clipping" the gradient as well. + update = num.atan2_(denom.mul_(b)).mul_(b) + else: + update = num.div_(denom.add_(eps)) + + return update + + def get_denom(self, state, group): exp_avg_sq = state['exp_avg_sq'] # Adam EMA updates @@ -415,53 +493,63 @@ class CoreOptimiser(torch.optim.Optimizer): row_factor = row_var.div(row_col_mean).sqrt_() col_factor = col_var.sqrt() denom = row_factor * col_factor + elif group['use_focus']: + denom = exp_avg_sq.clone() else: denom = exp_avg_sq.sqrt() return denom - def update_first_moment(self, state, group, grad): + def update_first_moment(self, state, group, grad, beta1): exp_avg = state['exp_avg'] - beta1, _ = group['betas'] + d_k = group['d_prev'] / group['d'] - return exp_avg.mul_(beta1).add_(grad, alpha=1 - beta1) + return exp_avg.mul_(beta1 * d_k).add_(grad, weight=1 - beta1) - def update_second_moment(self, state, group, grad, beta2, return_denom=True, denom_before_update=False): + def update_second_moment(self, state, group, grad, beta2, w, return_denom=True, denom_before_update=False): exp_avg_sq = state['exp_avg_sq'] + d_k = group['d_prev'] / group['d'] denom = None if return_denom and denom_before_update: - denom = self.get_denom(state) - - # Adam EMA updates - if isinstance(exp_avg_sq, list): - row_var, col_var, dr, dc, _ = exp_avg_sq + denom = self.get_denom(state, group) - 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 - ) + # Adam EMA updates + if group['use_focus']: + exp_avg_sq.mul_(beta2 * d_k * d_k).add_(w, alpha=1 - beta2) else: - exp_avg_sq.mul_(beta2).addcmul_(grad, grad, value=1 - beta2) + if isinstance(exp_avg_sq, list): + row_var, col_var, dr, dc, _ = exp_avg_sq + + row_var.mul_(beta2 * d_k * d_k).add_( + grad.norm(dim=dr, keepdim=True).square_().mul_(1 / grad.shape[dr]), + alpha=1 - beta2 + ) + col_var.mul_(beta2 * d_k * d_k).add_( + grad.norm(dim=dc, keepdim=True).square_().mul_(1 / grad.shape[dc]), + alpha=1 - beta2 + ) + else: + exp_avg_sq.mul_(beta2 * d_k * d_k).addcmul_(grad, grad, value=1 - beta2) if return_denom and denom is None: - denom = self.get_denom(state) + denom = self.get_denom(state, group) return denom def get_rms(self, tensor, eps=1e-8): - return tensor.norm().div(tensor.numel() ** 0.5).clamp_min(eps) + return tensor.norm(2).div(tensor.numel() ** 0.5).clamp_min(eps) def rms_(self, tensor, eps): return tensor.div_(self.get_rms(tensor, eps)) def get_clip_threshold(self, group): - return max(1, 8 * (0.99 ** (group['k'] - 1))) + # Prodigy works best with unscaled gradients during early steps. + if not group['use_speed'] and group['d'] <= group['d0']: + return 50 + + return 1 def try_hook_kohya_fbp(self): self.kohya_original_patch_adafactor_fused = None diff --git a/prodigyplusschedulefree/prodigy_plus_schedulefree.py b/prodigyplusschedulefree/prodigy_plus_schedulefree.py index 7b384b5..c286bc1 100644 --- a/prodigyplusschedulefree/prodigy_plus_schedulefree.py +++ b/prodigyplusschedulefree/prodigy_plus_schedulefree.py @@ -62,7 +62,11 @@ class ProdigyPlusScheduleFree(CoreOptimiser): If False, weight_decay will have a much stronger effect. (default: True). use_bias_correction (boolean): - Turn on Adafactor-style bias correction, which scales beta2 directly. (default: False). + Use the RAdam variant of schedule-free (https://github.com/facebookresearch/schedule_free/blob/main/schedulefree/radam_schedulefree.py). + This combines bias correction with automatic warmup. Please note this will significantly dampen Prodigy's adaptive stepsize + calculations -- it can take up to 10 times longer to start adjusting the learning rate. This can be mitigated somewhat by enabling + SPEED (use_speed=True). + (default: False). d0 (float): Initial estimate for Prodigy. Also serves as the minimum learning rate. (default: 1e-6). @@ -74,6 +78,11 @@ class ProdigyPlusScheduleFree(CoreOptimiser): Freeze Prodigy stepsize adjustments after a certain optimiser step and releases all state memory required by Prodigy. (default: 0) + use_speed (boolean): + Highly experimental. Signed Prodigy with ExponEntial D. This decouples the adaptive stepsize calculations from + the magnitude of the weights and gradient. This can provide faster, more accurate LRs in some scenarios, + but may fail in situations where the optimal LR is very close to (or less than) d0. + (default: False): split_groups (boolean): Track individual adaptation values for each parameter group. For example, if training a text encoder beside a Unet. Note this can have a significant impact on training dynamics. @@ -88,6 +97,12 @@ class ProdigyPlusScheduleFree(CoreOptimiser): Use factored approximation of the second moment, similar to Adafactor. Reduces memory usage. Disable if training results in NaNs or the learning rate fails to grow. (default: True) + factored_fp32 (boolean): + Force the use of float32 for the factored second moment. Because the factorisation is an approximation, it can + be beneficial to use high precision to avoid stability issues. However, if you're training in lower precision + for short durations, setting this to False will slightly reduce memory usage. + Ignored if factored is False. + (default: True) fused_back_pass (boolean): Stops the optimiser from running the normal step method. Set to True if using fused backward pass. Really only needed for scripts and UIs that call the regular step method even when using fused backward pass (OneTrainer). @@ -118,6 +133,16 @@ class ProdigyPlusScheduleFree(CoreOptimiser): so as not to include the current gradient in the denominator. This is a partial implementation of ADOPT (https://arxiv.org/abs/2411.02853), as we don't have a first moment to use for the update. (default: False) + use_orthograd (boolean): + Experimental. Updates weights using the component of the gradient that is orthogonal to the current + weight direction, as described in "Grokking at the Edge of Numerical Stability" (https://arxiv.org/pdf/2501.04697). + Can help prevent overfitting and improve generalisation. + (default: False) + use_focus (boolean): + Experimental. Modifies the update step to better handle noise at large step sizes. From + "FOCUS: First-Order Concentrated Update Scheme" (https://arxiv.org/abs/2501.12243). This method is + incompatible with factorisation, Muon and Adam-atan2. + (default: False) stochastic_rounding (boolean): Use stochastic rounding for bfloat16 weights (https://github.com/pytorch/pytorch/issues/120376). Brings bfloat16 training performance close to that of float32. @@ -130,27 +155,32 @@ class ProdigyPlusScheduleFree(CoreOptimiser): use_bias_correction=False, d0=1e-6, d_coef=1.0, prodigy_steps=0, + use_speed=False, eps=1e-8, split_groups=True, split_groups_mean=True, factored=True, + factored_fp32=True, fused_back_pass=False, use_stableadamw=True, use_muon_pp=False, use_cautious=False, use_grams=False, use_adopt=False, + use_orthograd=False, + use_focus=False, stochastic_rounding=True): super().__init__(params=params, lr=lr, betas=betas, beta3=beta3, weight_decay=weight_decay, weight_decay_by_lr=weight_decay_by_lr, use_bias_correction=use_bias_correction, - d0=d0, d_coef=d_coef, prodigy_steps=prodigy_steps, + d0=d0, d_coef=d_coef, prodigy_steps=prodigy_steps, use_speed=use_speed, eps=eps, split_groups=split_groups, - split_groups_mean=split_groups_mean, factored=factored, + split_groups_mean=split_groups_mean, factored=factored, factored_fp32=factored_fp32, fused_back_pass=fused_back_pass, use_stableadamw=use_stableadamw, use_muon_pp=use_muon_pp, use_cautious=use_cautious, use_grams=use_grams, - use_adopt=use_adopt, stochastic_rounding=stochastic_rounding) + use_adopt=use_adopt, use_orthograd=use_orthograd, use_focus=use_focus, + stochastic_rounding=stochastic_rounding) @torch.no_grad() def eval(self): @@ -188,9 +218,7 @@ class ProdigyPlusScheduleFree(CoreOptimiser): return state @torch.no_grad() - def update_params(self, y, z, update, group): - dlr = self.get_dlr(group) - + def update_params(self, y, z, update, group, dlr): beta1, _ = group['betas'] decay = group['weight_decay'] @@ -208,21 +236,23 @@ class ProdigyPlusScheduleFree(CoreOptimiser): y.sub_(y, alpha=decay * xy_step) z.sub_(y, alpha=decay) - if group['use_cautious']: - # "Cautious Optimizer (C-Optim): Improving Training with One Line of Code": https://github.com/kyleliang919/c-optim - # ScheduleFree implementation by nhamanasu: https://github.com/facebookresearch/schedule_free/pull/54 + cautious, grams = group['use_cautious'], group['use_grams'] + + if cautious or grams: u = (y - z).mul_(ckp1).add_(update, alpha=dlr * xy_step) z.sub_(update, alpha=dlr) - mask = (u * update > 0).to(update.dtype) - mask.mul_(mask.numel() / (mask.sum() + 1)) - u.mul_(mask) + + 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) + elif grams: + # "Grams: Gradient Descent with Adaptive Momentum Scaling": https://arxiv.org/abs/2412.17107 + u.abs_().mul_(update.sign_()) + 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) @@ -233,7 +263,7 @@ class ProdigyPlusScheduleFree(CoreOptimiser): @torch.no_grad() def step_param(self, p, group): - self.on_start_step() + self.on_start_step(p, group) if not group['train_mode']: raise Exception("Not in train mode!") @@ -241,49 +271,68 @@ class ProdigyPlusScheduleFree(CoreOptimiser): weight_sum = group['weight_sum'] if p.grad is not None: - grad = p.grad.to(dtype=torch.float32, copy=True) - - use_adopt = group['use_adopt'] + use_adopt = group['use_adopt'] stochastic = group['stochastic_rounding'] _, beta2 = group['betas'] k = group['k'] state = self.initialise_state(p, group) + z_state = state['z'] + y, z = (p.float(), z_state.float()) if stochastic else (p, z_state) + + grad = self.orthograd(z_state, p.grad) if group['use_orthograd'] else p.grad.to(dtype=torch.float32, copy=True) + dlr = self.get_dlr(group) + + if group['use_bias_correction']: + beta2_t = beta2 ** k + bias_correction2 = 1 - beta2_t + + # maximum length of the approximated SMA + rho_inf = 2 / (1 - beta2) - 1 + # compute the length of the approximated SMA + rho_t = rho_inf - 2 * k * beta2_t / bias_correction2 + rect = ( + ((rho_t - 4) * (rho_t - 2) * rho_inf / ((rho_inf - 4) * (rho_inf - 2) * rho_t)) ** 0.5 + if rho_t > 4.0 + else 0.0 + ) + dlr *= rect + beta2 = 1 - (1 - beta2) / (1 - beta2_t) + update = None if state['muon']: grad = self.newton_schulz_(grad) - grad_rms = self.get_rms(grad).item() ** 2 - - rms_sq = (state["rms_sq"] * beta2) + (grad_rms * (1 - beta2)) - state["rms_sq"] = rms_sq - - update = grad.mul_(1.0 / ((rms_sq ** 0.5) + 1e-12)) - else: - if group['use_bias_correction']: - # Adafactor / PaLM beta2 decay. Clip beta2 as per Scaling ViT paper. - beta2 = min(beta2, 1 - k ** -0.8) - beta2 = (1 - beta2) / (1 - beta2 ** k) - - if use_adopt and group['k'] == 1: - self.update_second_moment(state, group, grad, 0, return_denom=False) + if group['use_speed']: + grad_rms = state['rms_sq'] + if grad_rms is None: + grad_rms = state['rms_sq'] = 1 / self.get_rms(grad) + update = grad.mul_(grad_rms) else: - denom = self.update_second_moment(state, group, grad, beta2, denom_before_update=use_adopt) - update, num_scale = self.update_(grad, denom, group) + d_k = group['d_prev'] / group['d'] + rms_sq = state["rms_sq"].mul_(beta2 * d_k * d_k).add_(self.get_rms(grad).square(), alpha=1 - beta2) + update = grad.mul_(1 / rms_sq.sqrt().add(1e-8)) + else: + if use_adopt and group['k'] == 1: + self.update_second_moment(state, group, grad, 0, y, return_denom=False) + else: + denom = self.update_second_moment(state, group, grad, beta2, y, denom_before_update=use_adopt) + if group['use_bias_correction'] and rho_t <= 4.0: + update = grad + else: + update = self.update_(grad, denom, group, y) del denom if update is not None: if group['use_stableadamw']: clip_threshold = self.get_clip_threshold(group) - num_scale = max(1, self.get_rms(update, 1.0).item() / clip_threshold) - update.mul_(1 / num_scale) + rms = self.get_rms(update, 1).div(clip_threshold).clamp_min(1) + update.mul_(1 / rms) - z_state = state['z'] - self.update_prodigy(state, group, p.grad, z_state, 1.0) + self.update_prodigy(state, group, p.grad, p) - y, z = (p.float(), z_state.float()) if stochastic else (p, z_state) - weight_sum = self.update_params(y, z, update, group) + weight_sum = self.update_params(y, z, update, group, dlr) self.smart_copy(p, y, stochastic, True) self.smart_copy(z_state, z, stochastic, True)