CoolFace
Apppublic

dskill/DiffRhythm

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
2likes
trainer.py351 linesDownload Raw Back to model
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