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): # forward 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() # ===== DDP:只在 rank 0 汇总 ===== 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.") # 使用自定义 DynamicBatchSampler 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) # ⭐ 关键:只在 logging_steps 那一步算 + log if ( self.state.global_step > 0 and self.args.logging_steps > 0 and self.state.global_step % self.args.logging_steps == 0 ): # ⚠️ 用 self.log,而不是 print 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", {}) # 透传到 Trainer 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()