EchoLoc / rendering /universal_tts /core /tasks /cosyvoice2_task.py
zsy814's picture
Initial EchoLoc code release
d8bfe4a verified
Raw
History Blame Contribute Delete
7 kB
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()