diff --git a/custom_nodes/rvc_nodes.py b/custom_nodes/rvc_nodes.py index dfb467a..5773d3c 100644 --- a/custom_nodes/rvc_nodes.py +++ b/custom_nodes/rvc_nodes.py @@ -359,7 +359,7 @@ class RVCTrainParamsNode: return { "optional": dict( batch_size=("INT",dict(default=4,min=1,max=64,step=1)), - c_gen=("FLOAT",dict(default=1.,min=0.,max=100.,step=.1)), + c_adv=("FLOAT",dict(default=1.,min=0.,max=100.,step=.1)), c_mel=("FLOAT",dict(default=45.,min=0.,max=100.,step=.1)), c_kl=("FLOAT",dict(default=1.,min=0.,max=100.,step=.1)), c_fm=("FLOAT",dict(default=2.,min=0.,max=100.,step=.1)), @@ -372,7 +372,7 @@ class RVCTrainParamsNode: use_pareto=("BOOLEAN",dict(default=False)), fast_mode=("BOOLEAN",dict(default=False)), fp16_run=("BOOLEAN",dict(default=True)), - learning_rate=("FLOAT",dict(default=2e-4,min=1e-8,max=1,step=1e-8)) + learning_rate=("FLOAT",dict(default=1e-4,min=1e-8,max=1,step=1e-8)) ) } @@ -421,7 +421,7 @@ class RVCTrainModelNode: train_index=("BOOLEAN",{"default": True}), retrain=("BOOLEAN",{"default": False}), save_best_model=("BOOLEAN",{"default": True}), - best_model_threshold=("INT",dict(default=10,min=1,max=100)), + best_model_threshold=("INT",dict(default=30,min=1,max=100)), log_every_epoch=("FLOAT",dict(default=1.,min=0.,max=2.,step=.1)), num_workers=("INT",dict(default=1,min=1,max=16)) ) diff --git a/lib/train/losses.py b/lib/train/losses.py index 9392650..45b0741 100644 --- a/lib/train/losses.py +++ b/lib/train/losses.py @@ -53,7 +53,8 @@ class LossBalancer: """Updates the EMA of the weights.""" if not self.ema_weights: self.ema_weights = dict(**new_weights) else: self.ema_weights = { - k: self.weights_decay * self.ema_weights.get(k, 1.) + (1 - self.weights_decay) * new_weights[k] + k: np.nan_to_num( + self.weights_decay * self.ema_weights.get(k, 1.) + (1 - self.weights_decay) * new_weights[k],nan=self.epsilon) for k in new_weights } return dict(**self.ema_weights) @@ -63,26 +64,26 @@ class LossBalancer: if not self.historical_losses: self.historical_losses = dict(**new_losses) else: for k,v in new_losses.items(): - if np.nan_to_num(v,nan=-1)==-1: print(f"{k=} {v=}") - self.historical_losses[k]=self.loss_decay*self.historical_losses.get(k, v) + (1 - self.loss_decay)*v + self.historical_losses[k]=np.nan_to_num( + self.loss_decay*self.historical_losses.get(k, v) + (1 - self.loss_decay)*v,nan=self.epsilon) return dict(**self.historical_losses) def calculate_loss_slope(self, key: str, current_loss: torch.Tensor): """Calculates the slope of the loss using the current loss and its EMA.""" ema_loss = self.historical_losses.get(key, current_loss)+self.epsilon slope = (current_loss-ema_loss)/ema_loss # relative loss change - return slope + return slope.abs() def calculate_gradients(self, key: str, current_loss: torch.Tensor, input: torch.Tensor): """Calculates the gradient norm of the current loss wrt the model params.""" self.model.zero_grad() # Clear previous gradients input.requires_grad_(True) + current_loss.requires_grad_(True) # Backward pass to output layer output_params = torch.autograd.grad(current_loss, [input], retain_graph=True, allow_unused=True, materialize_grads=True)[0] # Compute L2 gradient norm - # grad_norm = output_params.view(output_params.size(0),-1).norm(2, dim=-1).mean() if output_params.ndim>1: grad_norm = output_params.view(output_params.size(0), -1).norm(2, dim=-1).mean() else: grad_norm = output_params.norm(2) @@ -111,12 +112,14 @@ class LossBalancer: weights[i] = remaining_weight / (len(sorted_data) - top_20_count) # Normalize the weights - total_weight = sum(weights) - normalized_weights = [w / total_weight for w in weights] + normalized_weights = {k: weights[i]*v for i,(k,v) in enumerate(sorted_data)} + total_weight = sum(normalized_weights.values()) + if total_weight>self.epsilon: scale = sum(data.values())/total_weight + else: scale=self.epsilon # Combine the keys with their normalized weights - normalized_data = {sorted_data[i][0]: normalized_weights[i]*len(data) for i in range(len(sorted_data))} - + normalized_data = {k: v*scale for k,v in normalized_weights.items()} + return normalized_data def on_train_batch_start(self, losses: dict, input: Optional[torch.Tensor]=None): @@ -160,17 +163,17 @@ class LossBalancer: if not valid_losses: return torch.tensor(0.0) # If all losses are skipped # Calculate loss weights based on gradient magnitudes - if self.use_pareto: #20% hardest tasks gets 80% weight - normalized_weights = self.pareto_normalizer(gradients) - else: - total_gradient = sum(gradients.values()) + self.epsilon - normalized_weights = {k: w/total_gradient*len(gradients) for k, w in gradients.items()} + total_gradient = sum(gradients.values()) + self.epsilon + normalized_weights = {k: w/total_gradient*len(gradients) for k, w in gradients.items()} # Update EMA weights - self.update_ema_weights(normalized_weights) + if self.use_pareto and len(normalized_weights)>1: #20% hardest tasks gets 80% weight + normalized_weights = self.pareto_normalizer(normalized_weights) + normalized_weights = self.update_ema_weights(normalized_weights) + balanced_loss = 0 for k, loss in valid_losses.items(): - balanced_loss += self.ema_weights.get(k, 1.) * loss + balanced_loss += normalized_weights.get(k, 1.) * loss # update historical losses self.update_historical_losses({k: loss.item() for k, loss in valid_losses.items()}) @@ -576,8 +579,8 @@ class MultiScaleMelSpectrogramLoss(torch.nn.Module): "window_type": s.window_type, } - x_mels = self.mel_spectrogram(x, **kwargs) - y_mels = self.mel_spectrogram(y, **kwargs) + x_mels = self.mel_spectrogram(x, **kwargs).nan_to_num(self.clamp_eps) + y_mels = self.mel_spectrogram(y, **kwargs).nan_to_num(self.clamp_eps) x_logmels = torch.log10(x_mels.pow(self.pow)+self.clamp_eps).nan_to_num(self.clamp_eps) y_logmels = torch.log10(y_mels.pow(self.pow)+self.clamp_eps).nan_to_num(self.clamp_eps) @@ -598,24 +601,27 @@ def discriminator_loss( disc_generated_outputs: List[torch.Tensor] ): loss = 0 - r_losses = [] - g_losses = [] + # r_losses = [] + # g_losses = [] + disc_losses = [] for dr, dg in zip(disc_real_outputs, disc_generated_outputs): r_loss = torch.mean((1 - dr) ** 2) g_loss = torch.mean(dg**2) - loss += r_loss + g_loss - r_losses.append(r_loss.item()) - g_losses.append(g_loss.item()) - return loss, r_losses, g_losses + L = r_loss + g_loss + loss += L + # r_losses.append(r_loss.item()) + # g_losses.append(g_loss.item()) + disc_losses.append(L) + return loss, disc_losses def generator_loss(disc_outputs: List[torch.Tensor]): loss = 0 gen_losses = [] for dg in disc_outputs: - l = torch.mean((1 - dg) ** 2) - gen_losses.append(l) - loss += l + L = torch.mean((1 - dg) ** 2) + gen_losses.append(L) + loss += L return loss, gen_losses diff --git a/training_cli.py b/training_cli.py index 77b7157..e087897 100644 --- a/training_cli.py +++ b/training_cli.py @@ -3,6 +3,7 @@ import json import os import shutil import traceback +from typing import Tuple import numpy as np from tqdm import tqdm @@ -285,10 +286,10 @@ def run(rank, n_gpus, hps, device): try: balancer_state = g_kwargs["balancer"] logger.info(f"Using existing balancer: {balancer_state}") - balancer = LossBalancer(net_g,**balancer_state) + balancer_g = LossBalancer(net_g,**balancer_state) except Exception as e: logger.error(f"Failed to load balancer state: {e}") - balancer = LossBalancer( + balancer_g = LossBalancer( net_g, weights_decay=.5 / (1 + np.exp(-10 * (epoch_str / hps.total_epoch - 0.16)))+.5, #sigmoid scaled ema .8 at 20% epoch loss_decay=.8, @@ -297,7 +298,7 @@ def run(rank, n_gpus, hps, device): use_pareto=hps.train.get("use_pareto",False), use_norm=not hps.train.get("fast_mode",False), initial_weights=dict( - loss_gen=hps.train.get("c_gen",1.), + loss_gen=hps.train.get("c_adv",1.), loss_fm=hps.train.get("c_fm",2.), loss_mel=hps.train.get("c_mel",45.), loss_kl=hps.train.get("c_kl",1.), @@ -305,6 +306,25 @@ def run(rank, n_gpus, hps, device): tsi_loss=hps.train.get("c_tsi",0.), tefs_loss=hps.train.get("c_tefs",0.), )) + + try: + balancer_state = d_kwargs["balancer"] + logger.info(f"Using existing balancer: {balancer_state}") + balancer_d = LossBalancer(net_d,**balancer_state) + except Exception as e: + logger.error(f"Failed to load balancer state: {e}") + balancer_d = LossBalancer( + net_d, + weights_decay=.5 / (1 + np.exp(-10 * (epoch_str / hps.total_epoch - 0.16)))+.5, #sigmoid scaled ema .8 at 20% epoch + loss_decay=.8, + epsilon=hps.train.eps, + active=hps.train.get("use_balancer",False), + use_pareto=hps.train.get("use_pareto",False), + use_norm=not hps.train.get("fast_mode",False), + initial_weights=dict( + loss_disc=hps.train.get("c_adv",1.), + gradient_penalty=hps.train.get("c_gp",0.), + )) cache = [] for epoch in range(epoch_str, hps.train.epochs + 1): train_loader.batch_sampler.set_epoch(epoch) @@ -321,7 +341,7 @@ def run(rank, n_gpus, hps, device): logger, [writer, writer_eval], cache, - balancer + [balancer_g, balancer_d] ) else: train_and_evaluate( @@ -336,14 +356,14 @@ def run(rank, n_gpus, hps, device): None, None, cache, - balancer + [balancer_g, balancer_d] ) scheduler_g.step() scheduler_d.step() def train_and_evaluate( - rank, epoch, hps, nets, optims, _, scaler, loaders, logger, writers, cache, balancer: "LossBalancer" + rank, epoch, hps, nets, optims, _, scaler, loaders, logger, writers, cache, balancer: Tuple["LossBalancer","LossBalancer"] ): net_g, net_d = nets optim_g, optim_d = optims @@ -355,6 +375,7 @@ def train_and_evaluate( net_g.train() net_d.train() + balancer_g, balancer_d = balancer # Prepare data iterator if hps.if_cache_data_in_gpu: @@ -519,8 +540,11 @@ def train_and_evaluate( with autocast(enabled=False): gradient_penalty = gradient_norm_loss(wave,gen_wave, net_d, eps=hps.train.eps)*hps.train.c_gp if hps.train.get("c_gp",0.)>0 else 0 - loss_disc, losses_disc_r, losses_disc_g = discriminator_loss(y_d_hat_r, y_d_hat_g) - loss_disc_all = loss_disc+gradient_penalty + loss_disc, losses_disc = discriminator_loss(y_d_hat_r, y_d_hat_g) + loss_disc_all = balancer_d.on_train_batch_start(dict( + loss_disc=loss_disc, + gradient_penalty=gradient_penalty + ),input=y_hat) optim_d.zero_grad() scaler.scale(loss_disc_all).backward() @@ -548,7 +572,7 @@ def train_and_evaluate( ) loss_gen, losses_gen = generator_loss(y_d_hat_g) aux_loss = harmonic_loss + tefs_loss + tsi_loss - loss_gen_all = balancer.on_train_batch_start(dict( + loss_gen_all = balancer_g.on_train_batch_start(dict( loss_gen=loss_gen, loss_fm=loss_fm, loss_mel=loss_mel, @@ -594,10 +618,9 @@ def train_and_evaluate( "gradient/grad_norm_gen": grad_norm_g, "gradient/gradient_penalty": gradient_penalty, **{f"loss/g/{i}": v for i, v in enumerate(losses_gen)}, - **{f"loss/d_r/{i}": v for i, v in enumerate(losses_disc_r)}, - **{f"loss/d_g/{i}": v for i, v in enumerate(losses_disc_g)}, - **{f"balancer/weights/{k}": v for k, v in balancer.ema_weights.items()}, - **{f"balancer/losses/{k}": v for k, v in balancer.historical_losses.items()} + **{f"loss/d/{i}": v for i, v in enumerate(losses_disc)}, + **{f"balancer_g/weights/{k}": v for k, v in balancer_g.ema_weights.items()}, + **{f"balancer_d/weights/{k}": v for k, v in balancer_d.ema_weights.items()}, } image_dict = { @@ -637,14 +660,15 @@ def train_and_evaluate( hps.train.learning_rate, epoch, os.path.join(hps.model_dir, f"G_{saved_epoch}.pth"), - balancer=balancer.to_dict() + balancer=balancer_g.to_dict() ) utils.save_checkpoint( net_d, optim_d, hps.train.learning_rate, epoch, - os.path.join(hps.model_dir, f"D_{saved_epoch}.pth") + os.path.join(hps.model_dir, f"D_{saved_epoch}.pth"), + balancer=balancer_d.to_dict() ) if hps.save_every_weights: ckpt = net_g.module.state_dict() if hasattr(net_g, "module") else net_g.state_dict() @@ -653,12 +677,13 @@ def train_and_evaluate( logger.info(f"saving ckpt {save_name}: {status}") if rank == 0: - total_loss = balancer.weighted_ema_loss + loss_disc_all + total_loss = balancer_g.weighted_ema_loss + balancer_d.weighted_ema_loss logger.info(f"====> Epoch {epoch} ({total_loss=:.3f}): {global_step=} {lr=:.2E} {epoch_recorder.record()}") logger.info(f"|| {loss_disc_all=:.3f}: {loss_disc=:.3f}, {gradient_penalty=:.3f}") logger.info(f"|| {loss_gen_all=:.3f}: {loss_gen=:.3f}, {loss_fm=:.3f}, {loss_mel=:.3f}, {loss_kl=:.3f}") logger.info(f"|| {aux_loss=:.3f}: {harmonic_loss=:.3f}, {tefs_loss=:.3f}, {tsi_loss=:.3f}") - balancer.on_epoch_end(.5 / (1 + np.exp(-10 * (epoch / hps.total_epoch - 0.16)))+.5) #sigmoid scaling of ema + balancer_g.on_epoch_end(.5 / (1 + np.exp(-10 * (epoch / hps.total_epoch - 0.16)))+.5) #sigmoid scaling of ema + balancer_d.on_epoch_end(.5 / (1 + np.exp(-10 * (epoch / hps.total_epoch - 0.16)))+.5) #sigmoid scaling of ema if total_loss