| from core.tasks.base_task import BaseTask |
| from core.tasks.task_registry import register_task |
| from transformers import Trainer, TrainingArguments |
| from transformers.trainer import find_batch_size, EvalLoopOutput |
| from torch.utils.data import DataLoader |
| from core.datasets.cosyvoice_dataset import CosyVoiceDataset, CosyVoiceCollator |
| from core.datasets.samplers import DistributedDynamicBatchSampler |
| from models.cosyvoice.cosyvoice2 import CosyVoice2Model |
| from omegaconf import OmegaConf |
| from transformers import EvalPrediction |
| from models.cosyvoice.utils.common import np_accuracy, IGNORE_ID, th_accuracy |
| import torch.distributed as dist |
| import torch |
| import numpy as np |
| import os |
|
|
|
|
| def compute_metrics(eval_pred): |
| predictions, labels = eval_pred |
| acc = th_accuracy( |
| predictions, |
| labels, |
| ignore_label=IGNORE_ID, |
| ) |
| return {"accuracy": acc} |
|
|
|
|
| class CustomTrainer(Trainer): |
| def __init__(self, sampler_cfg=None, *args, **kwargs): |
| super().__init__(*args, **kwargs) |
| self.sampler_cfg = sampler_cfg |
|
|
| def evaluation_loop( |
| self, |
| dataloader, |
| description, |
| prediction_loss_only=None, |
| ignore_keys=None, |
| metric_key_prefix="eval", |
| ): |
| args = self.args |
| model = self._wrap_model(self.model, training=False, dataloader=dataloader) |
|
|
| if not self.is_in_train: |
| if args.fp16_full_eval: |
| model = model.to(dtype=torch.float16, device=args.device) |
| elif args.bf16_full_eval: |
| model = model.to(dtype=torch.bfloat16, device=args.device) |
|
|
| model.eval() |
|
|
| total_acc = 0.0 |
| total_count = 0 |
| total_loss = 0.0 |
|
|
| observed_num_examples = 0 |
|
|
| for step, inputs in enumerate(dataloader): |
| |
| with torch.no_grad(): |
| outputs = model(**inputs) |
| loss = outputs.get("loss", None) |
| acc = outputs.get("acc", None) |
| batch_size = find_batch_size(inputs) or 1 |
| observed_num_examples += batch_size |
|
|
| if loss is not None: |
| total_loss += loss.detach().float().item() * batch_size |
|
|
| if acc is not None: |
| total_acc += float(acc.item()) * batch_size |
| total_count += batch_size |
|
|
| |
| del outputs |
| torch.cuda.empty_cache() |
|
|
| |
| if self.accelerator.num_processes > 1: |
| total_loss = self.accelerator.reduce( |
| torch.tensor(total_loss, device=args.device), reduction="sum" |
| ).item() |
| total_acc = self.accelerator.reduce( |
| torch.tensor(total_acc, device=args.device), reduction="sum" |
| ).item() |
| observed_num_examples = self.accelerator.reduce( |
| torch.tensor(observed_num_examples, device=args.device), reduction="sum" |
| ).item() |
|
|
| metrics = {} |
|
|
| if observed_num_examples > 0: |
| metrics[f"{metric_key_prefix}_loss"] = total_loss / observed_num_examples |
| metrics[f"{metric_key_prefix}_accuracy"] = total_acc / total_count |
| metrics[f"{metric_key_prefix}_step"] = self.state.global_step |
|
|
| return EvalLoopOutput( |
| predictions=None, |
| label_ids=None, |
| metrics=metrics, |
| num_samples=observed_num_examples, |
| ) |
|
|
| def get_train_dataloader(self): |
| if self.train_dataset is None: |
| raise ValueError("Trainer: training requires a train_dataset.") |
|
|
| |
| sampler = DistributedDynamicBatchSampler( |
| self.train_dataset.get_lengths(), **self.sampler_cfg |
| ) |
|
|
| return DataLoader( |
| self.train_dataset, |
| batch_sampler=sampler, |
| collate_fn=self.data_collator, |
| num_workers=self.args.dataloader_num_workers, |
| pin_memory=self.args.dataloader_pin_memory, |
| ) |
|
|
| def compute_loss( |
| self, |
| model, |
| inputs, |
| return_outputs=False, |
| num_items_in_batch=None, |
| ): |
| outputs = model(**inputs) |
| loss = outputs["loss"] |
|
|
| acc = outputs.get("acc", None) |
|
|
| |
| if ( |
| self.state.global_step > 0 |
| and self.args.logging_steps > 0 |
| and self.state.global_step % self.args.logging_steps == 0 |
| ): |
| |
| self.log({"train_accuracy": acc.item(), "step": self.state.global_step}) |
|
|
| if return_outputs: |
| return loss, outputs |
| return loss |
|
|
|
|
| @register_task("cosyvoice2") |
| class CosyVoice2Task(BaseTask): |
| def build_dataset(self): |
| dataset_cfg = self.config["datasets"] |
| dataset_cfg = OmegaConf.to_container(dataset_cfg, resolve=True) |
| train_split = dataset_cfg.pop("train_file") |
| valid_split = dataset_cfg.pop("valid_file") |
| train_dataset = CosyVoiceDataset(**dataset_cfg, split=train_split) |
| eval_dataset = CosyVoiceDataset(**dataset_cfg, split=valid_split) |
| print(f"train dataset size: {len(train_dataset)}") |
| print(f"eval dataset size: {len(eval_dataset)}") |
| return train_dataset, eval_dataset |
|
|
| def build_collator(self): |
| collator_cfg = self.config.get("collator", {}) |
| return CosyVoiceCollator(**collator_cfg) |
|
|
| def build_model(self): |
| model_cfg = OmegaConf.to_container(self.config.get("model", {}), resolve=True) |
| pre_ckpt = model_cfg.pop("pretrained_path", "") |
| model = CosyVoice2Model(**model_cfg) |
| if os.path.exists(pre_ckpt): |
| state = torch.load( |
| pre_ckpt, |
| map_location="cpu", |
| ) |
| model.model.load_state_dict(state) |
| print(f"load pretrained ckpt: {pre_ckpt}") |
| return model |
|
|
| def build_training_args(self): |
| args_cfg = self.config.get("trainer", {}) |
| sampler_cfg = self.config.get("sampler", {}) |
| return TrainingArguments(**args_cfg), sampler_cfg |
|
|
| def build_trainer(self): |
| trainer = CustomTrainer( |
| model=self.model, |
| args=self.training_args, |
| sampler_cfg=self.sampler_cfg, |
| train_dataset=self.train_dataset, |
| eval_dataset=self.eval_dataset, |
| data_collator=self.collator, |
| compute_metrics=compute_metrics, |
| ) |
| return trainer |
|
|
| def run(self): |
| |
| self.train_dataset, self.eval_dataset = self.build_dataset() |
| self.collator = self.build_collator() |
| self.model = self.build_model() |
| self.training_args, self.sampler_cfg = self.build_training_args() |
| self.trainer = self.build_trainer() |
|
|
| |
| print(f"train args: {self.trainer.args}", flush=True) |
| self.trainer.train() |
|
|