added disc loss balancer
This commit is contained in:
@@ -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))
|
||||
)
|
||||
|
||||
+33
-27
@@ -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
|
||||
|
||||
|
||||
|
||||
+43
-18
@@ -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<least_loss:
|
||||
least_loss = total_loss
|
||||
@@ -674,7 +699,7 @@ def train_and_evaluate(
|
||||
|
||||
with open(loss_file,"w") as f:
|
||||
json.dump(dict(least_loss=least_loss.item(),best_model_name=best_model_name,epoch=epoch,steps=global_step,
|
||||
loss_weights = dict(balancer.ema_weights),
|
||||
loss_weights = dict(**balancer_g.ema_weights,**balancer_d.ema_weights),
|
||||
scalar_dict={
|
||||
"total/loss/all": commons.serialize_tensor(loss_gen_all+loss_disc_all),
|
||||
"total/loss/gen_all": commons.serialize_tensor(loss_gen_all),
|
||||
|
||||
Reference in New Issue
Block a user