dskill/DiffRhythm
2
1from __future__ import annotations2 3import os4import gc5from tqdm import tqdm6import wandb7 8import torch9from torch.optim import AdamW10from torch.optim.lr_scheduler import LinearLR, SequentialLR, ConstantLR11 12from accelerate import Accelerator13from accelerate.utils import DistributedDataParallelKwargs14from diffrhythm.dataset.custom_dataset_align2f5 import LanceDiffusionDataset15 16from torch.utils.data import DataLoader, DistributedSampler17 18from ema_pytorch import EMA19 20from diffrhythm.model import CFM21from diffrhythm.model.utils import exists, default22 23import time24 25# from apex.optimizers.fused_adam import FusedAdam26 27# trainer28 29 30class Trainer:31 def __init__(32 self,33 model: CFM,34 args,35 epochs,36 learning_rate,37 #dataloader,38 num_warmup_updates=20000,39 save_per_updates=1000,40 checkpoint_path=None,41 batch_size=32,42 batch_size_type: str = "sample",43 max_samples=32,44 grad_accumulation_steps=1,45 max_grad_norm=1.0,46 noise_scheduler: str | None = None,47 duration_predictor: torch.nn.Module | None = None,48 wandb_project="test_e2-tts",49 wandb_run_name="test_run",50 wandb_resume_id: str = None,51 last_per_steps=None,52 accelerate_kwargs: dict = dict(),53 ema_kwargs: dict = dict(),54 bnb_optimizer: bool = False,55 reset_lr: bool = False,56 use_style_prompt: bool = False,57 grad_ckpt: bool = False58 ):59 self.args = args60 61 ddp_kwargs = DistributedDataParallelKwargs(find_unused_parameters=False, )62 63 logger = "wandb" if wandb.api.api_key else None64 #logger = None65 print(f"Using logger: {logger}")66 # print("-----------1-------------")67 import tbe.common68 # print("-----------2-------------")69 self.accelerator = Accelerator(70 log_with=logger,71 kwargs_handlers=[ddp_kwargs],72 gradient_accumulation_steps=grad_accumulation_steps,73 **accelerate_kwargs,74 )75 # print("-----------3-------------")76 77 if logger == "wandb":78 if exists(wandb_resume_id):79 init_kwargs = {"wandb": {"resume": "allow", "name": wandb_run_name, "id": wandb_resume_id}}80 else:81 init_kwargs = {"wandb": {"resume": "allow", "name": wandb_run_name}}82 self.accelerator.init_trackers(83 project_name=wandb_project,84 init_kwargs=init_kwargs,85 config={86 "epochs": epochs,87 "learning_rate": learning_rate,88 "num_warmup_updates": num_warmup_updates,89 "batch_size": batch_size,90 "batch_size_type": batch_size_type,91 "max_samples": max_samples,92 "grad_accumulation_steps": grad_accumulation_steps,93 "max_grad_norm": max_grad_norm,94 "gpus": self.accelerator.num_processes,95 "noise_scheduler": noise_scheduler,96 },97 )98 99 self.precision = self.accelerator.state.mixed_precision100 self.precision = self.precision.replace("no", "fp32")101 print("!!!!!!!!!!!!!!!!!", self.precision)102 103 self.model = model104 #self.model = torch.compile(model)105 106 #self.dataloader = dataloader107 108 if self.is_main:109 self.ema_model = EMA(model, include_online_model=False, **ema_kwargs)110 111 self.ema_model.to(self.accelerator.device)112 if self.accelerator.state.distributed_type in ["DEEPSPEED", "FSDP"]:113 self.ema_model.half()114 115 self.epochs = epochs116 self.num_warmup_updates = num_warmup_updates117 self.save_per_updates = save_per_updates118 self.last_per_steps = default(last_per_steps, save_per_updates * grad_accumulation_steps)119 self.checkpoint_path = default(checkpoint_path, "ckpts/test_e2-tts")120 121 self.max_samples = max_samples122 self.grad_accumulation_steps = grad_accumulation_steps123 self.max_grad_norm = max_grad_norm124 125 self.noise_scheduler = noise_scheduler126 127 self.duration_predictor = duration_predictor128 129 self.reset_lr = reset_lr130 131 self.use_style_prompt = use_style_prompt132 133 self.grad_ckpt = grad_ckpt134 135 if bnb_optimizer:136 import bitsandbytes as bnb137 138 self.optimizer = bnb.optim.AdamW8bit(model.parameters(), lr=learning_rate)139 else:140 self.optimizer = AdamW(model.parameters(), lr=learning_rate)141 #self.optimizer = FusedAdam(model.parameters(), lr=learning_rate)142 143 #self.model = torch.compile(self.model)144 if self.accelerator.state.distributed_type == "DEEPSPEED":145 self.accelerator.state.deepspeed_plugin.deepspeed_config['train_micro_batch_size_per_gpu'] = batch_size146 147 self.get_dataloader()148 self.get_scheduler()149 # self.get_constant_scheduler()150 151 self.model, self.optimizer, self.scheduler, self.train_dataloader = self.accelerator.prepare(self.model, self.optimizer, self.scheduler, self.train_dataloader)152 153 def get_scheduler(self):154 warmup_steps = (155 self.num_warmup_updates * self.accelerator.num_processes156 ) # consider a fixed warmup steps while using accelerate multi-gpu ddp157 total_steps = len(self.train_dataloader) * self.epochs / self.grad_accumulation_steps158 decay_steps = total_steps - warmup_steps159 warmup_scheduler = LinearLR(self.optimizer, start_factor=1e-8, end_factor=1.0, total_iters=warmup_steps)160 decay_scheduler = LinearLR(self.optimizer, start_factor=1.0, end_factor=1e-8, total_iters=decay_steps)161 # constant_scheduler = ConstantLR(self.optimizer, factor=1, total_iters=decay_steps)162 self.scheduler = SequentialLR(163 self.optimizer, schedulers=[warmup_scheduler, decay_scheduler], milestones=[warmup_steps]164 )165 166 def get_constant_scheduler(self):167 total_steps = len(self.train_dataloader) * self.epochs / self.grad_accumulation_steps168 self.scheduler = ConstantLR(self.optimizer, factor=1, total_iters=total_steps)169 170 def get_dataloader(self):171 prompt_path = self.args.prompt_path.split('|')172 lrc_path = self.args.lrc_path.split('|')173 latent_path = self.args.latent_path.split('|')174 ldd = LanceDiffusionDataset(*LanceDiffusionDataset.init_data(self.args.dataset_path), \175 max_frames=self.args.max_frames, min_frames=self.args.min_frames, \176 align_lyrics=self.args.align_lyrics, lyrics_slice=self.args.lyrics_slice, \177 use_style_prompt=self.args.use_style_prompt, parse_lyrics=self.args.parse_lyrics,178 lyrics_shift=self.args.lyrics_shift, downsample_rate=self.args.downsample_rate, \179 skip_empty_lyrics=self.args.skip_empty_lyrics, tokenizer_type=self.args.tokenizer_type, precision=self.precision, \180 start_time=time.time(), pure_prob=self.args.pure_prob)181 182 # start_time = time.time()183 self.train_dataloader = DataLoader(184 dataset=ldd,185 batch_size=self.args.batch_size, # 每个批次的样本数186 shuffle=True, # 是否随机打乱数据187 num_workers=4, # 用于加载数据的子进程数188 pin_memory=True, # 加速GPU训练189 collate_fn=ldd.custom_collate_fn,190 persistent_workers=True191 )192 193 194 @property195 def is_main(self):196 return self.accelerator.is_main_process197 198 def save_checkpoint(self, step, last=False):199 self.accelerator.wait_for_everyone()200 if self.is_main:201 checkpoint = dict(202 model_state_dict=self.accelerator.unwrap_model(self.model).state_dict(),203 optimizer_state_dict=self.accelerator.unwrap_model(self.optimizer).state_dict(),204 ema_model_state_dict=self.ema_model.state_dict(),205 scheduler_state_dict=self.scheduler.state_dict(),206 step=step,207 )208 if not os.path.exists(self.checkpoint_path):209 os.makedirs(self.checkpoint_path)210 if last:211 self.accelerator.save(checkpoint, f"{self.checkpoint_path}/model_last.pt")212 print(f"Saved last checkpoint at step {step}")213 else:214 self.accelerator.save(checkpoint, f"{self.checkpoint_path}/model_{step}.pt")215 216 def load_checkpoint(self):217 if (218 not exists(self.checkpoint_path)219 or not os.path.exists(self.checkpoint_path)220 or not os.listdir(self.checkpoint_path)221 ):222 return 0223 224 self.accelerator.wait_for_everyone()225 if "model_last.pt" in os.listdir(self.checkpoint_path):226 latest_checkpoint = "model_last.pt"227 else:228 latest_checkpoint = sorted(229 [f for f in os.listdir(self.checkpoint_path) if f.endswith(".pt")],230 key=lambda x: int("".join(filter(str.isdigit, x))),231 )[-1]232 233 checkpoint = torch.load(f"{self.checkpoint_path}/{latest_checkpoint}", map_location="cpu")234 235 ### **1. 过滤 `ema_model` 的不匹配参数**236 if self.is_main:237 ema_dict = self.ema_model.state_dict()238 ema_checkpoint_dict = checkpoint["ema_model_state_dict"]239 240 filtered_ema_dict = {241 k: v for k, v in ema_checkpoint_dict.items()242 if k in ema_dict and ema_dict[k].shape == v.shape # 仅加载 shape 匹配的参数243 }244 245 print(f"Loading {len(filtered_ema_dict)} / {len(ema_checkpoint_dict)} ema_model params")246 self.ema_model.load_state_dict(filtered_ema_dict, strict=False)247 248 ### **2. 过滤 `model` 的不匹配参数**249 model_dict = self.accelerator.unwrap_model(self.model).state_dict()250 checkpoint_model_dict = checkpoint["model_state_dict"]251 252 filtered_model_dict = {253 k: v for k, v in checkpoint_model_dict.items()254 if k in model_dict and model_dict[k].shape == v.shape # 仅加载 shape 匹配的参数255 }256 257 print(f"Loading {len(filtered_model_dict)} / {len(checkpoint_model_dict)} model params")258 self.accelerator.unwrap_model(self.model).load_state_dict(filtered_model_dict, strict=False)259 260 ### **3. 加载优化器、调度器和步数**261 if "step" in checkpoint:262 if self.scheduler and not self.reset_lr:263 self.scheduler.load_state_dict(checkpoint["scheduler_state_dict"])264 step = checkpoint["step"]265 else:266 step = 0267 268 del checkpoint269 gc.collect()270 print("Checkpoint loaded at step", step)271 return step272 273 def train(self, resumable_with_seed: int = None):274 train_dataloader = self.train_dataloader275 276 start_step = self.load_checkpoint()277 global_step = start_step278 279 if resumable_with_seed > 0:280 orig_epoch_step = len(train_dataloader)281 skipped_epoch = int(start_step // orig_epoch_step)282 skipped_batch = start_step % orig_epoch_step283 skipped_dataloader = self.accelerator.skip_first_batches(train_dataloader, num_batches=skipped_batch)284 else:285 skipped_epoch = 0286 287 for epoch in range(skipped_epoch, self.epochs):288 self.model.train()289 if resumable_with_seed > 0 and epoch == skipped_epoch:290 progress_bar = tqdm(291 skipped_dataloader,292 desc=f"Epoch {epoch+1}/{self.epochs}",293 unit="step",294 disable=not self.accelerator.is_local_main_process,295 initial=skipped_batch,296 total=orig_epoch_step,297 smoothing=0.15298 )299 else:300 progress_bar = tqdm(301 train_dataloader,302 desc=f"Epoch {epoch+1}/{self.epochs}",303 unit="step",304 disable=not self.accelerator.is_local_main_process,305 smoothing=0.15306 )307 308 for batch in progress_bar:309 with self.accelerator.accumulate(self.model):310 text_inputs = batch["lrc"]311 mel_spec = batch["latent"].permute(0, 2, 1)312 mel_lengths = batch["latent_lengths"]313 style_prompt = batch["prompt"]314 style_prompt_lens = batch["prompt_lengths"]315 start_time = batch["start_time"]316 317 loss, cond, pred = self.model(318 mel_spec, text=text_inputs, lens=mel_lengths, noise_scheduler=self.noise_scheduler,319 style_prompt=style_prompt if self.use_style_prompt else None,320 style_prompt_lens=style_prompt_lens if self.use_style_prompt else None,321 grad_ckpt=self.grad_ckpt, start_time=start_time322 )323 self.accelerator.backward(loss)324 325 if self.max_grad_norm > 0 and self.accelerator.sync_gradients:326 self.accelerator.clip_grad_norm_(self.model.parameters(), self.max_grad_norm)327 328 self.optimizer.step()329 self.scheduler.step()330 self.optimizer.zero_grad()331 332 if self.is_main:333 self.ema_model.update()334 335 global_step += 1336 337 if self.accelerator.is_local_main_process:338 self.accelerator.log({"loss": loss.item(), "lr": self.scheduler.get_last_lr()[0]}, step=global_step)339 340 progress_bar.set_postfix(step=str(global_step), loss=loss.item())341 342 if global_step % (self.save_per_updates * self.grad_accumulation_steps) == 0:343 self.save_checkpoint(global_step)344 345 if global_step % self.last_per_steps == 0:346 self.save_checkpoint(global_step, last=True)347 348 self.save_checkpoint(global_step, last=True)349 350 self.accelerator.end_training()351 