diff --git a/lib/infer_pack/commons.py b/lib/infer_pack/commons.py index bf4f867..bb2c61f 100644 --- a/lib/infer_pack/commons.py +++ b/lib/infer_pack/commons.py @@ -1,4 +1,5 @@ import math +import numpy as np import torch from torch.nn import functional as F @@ -272,4 +273,29 @@ def clip_grad_value_(parameters, clip_value, norm_type=2, batch_size=1): def serialize_tensor(tensor): if hasattr(tensor,"item"): return tensor.item() - else: return float(f"{float(tensor):.3f}") \ No newline at end of file + else: return float(f"{float(tensor):.3f}") + +def sigmoid_value(step, total_steps=100, start_value=0, end_value=1, midpoint=.5): + """ + Calculate the scaled sigmoid function. + + Parameters: + step (int): The current training epoch. + total_steps (int): The total number of training epochs. + start_value (float): The starting value of the clipping range. + end_value (float): The ending value of the clipping range. + midpoint (float): Percentage where sigmoid function is 50% between start and end value + steepness (int): How fast the values change per step + + Returns: + float: The calculated scaled sigmoid value. + """ + # Sigmoid function parameters + + + # Calculate the sigmoid value + sigmoid_value = 1 / (1 + np.power(10,-np.log10(total_steps) * (step / total_steps - midpoint))) + + # Scale the sigmoid output to the desired range + value = start_value + (end_value - start_value) * sigmoid_value + return value \ No newline at end of file diff --git a/training_cli.py b/training_cli.py index 1403717..54fc342 100644 --- a/training_cli.py +++ b/training_cli.py @@ -122,10 +122,9 @@ def train_model(hps: "utils.HParams"): def run(rank, n_gpus, hps, device): print(f"{__name__=}") - global global_step, least_loss, loss_file, best_model_name, gradient_clip_value, MultiscaleMelLoss + global global_step, least_loss, loss_file, best_model_name, MultiscaleMelLoss global_step = 0 loss_file = os.path.join(hps.model_dir,"losses.json") - gradient_clip_value = np.log(hps.train.c_gp/hps.train.learning_rate) if hps.train.get("c_gp",0.)>0 else None if os.path.isfile(loss_file): with open(loss_file,"r") as f: @@ -325,7 +324,7 @@ def run(rank, n_gpus, hps, device): 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 + weights_decay=commons.sigmoid_value(global_step,total_steps=10000,start_value=.5, end_value=.999, midpoint=.2), loss_decay=.8, epsilon=hps.train.eps, active=hps.train.get("use_balancer",False), @@ -386,6 +385,7 @@ def train_and_evaluate( net_g.train() net_d.train() balancer_g, balancer_d = balancer + gradient_clip_value = commons.sigmoid_value(global_step, total_steps=10000, start_value=1, end_value=500, midpoint=.2) # Prepare data iterator if hps.if_cache_data_in_gpu: @@ -696,8 +696,11 @@ def train_and_evaluate( 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_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 + + #sigmoid scaling of ema + weights_decay = commons.sigmoid_value(global_step,total_steps=10000,start_value=.5, end_value=.999, midpoint=.2) + balancer_g.on_epoch_end(weights_decay) + balancer_d.on_epoch_end(weights_decay) if loss_gen_all