proiect-pvmd / src /train_parakeet.py
alexandrubent's picture
Initial project upload
a039392 verified
Raw
History Blame Contribute Delete
11.1 kB
"""
Fine-tuning script pentru gabrielpirlo/Sped_ParakeetRomanian_110M_TDT-CTC
pe dataset-ul datadriven-company/TTS-Romanian folosind HuggingFace Transformers.
"""
import os
import sys
import warnings
from dataclasses import dataclass, field
from typing import Any, Dict, List, Optional, Union
import torch
import numpy as np
from datasets import load_from_disk, load_dataset, Audio
from transformers import (
AutoModelForCTC,
AutoProcessor,
AutoFeatureExtractor,
AutoTokenizer,
TrainingArguments,
Trainer,
EarlyStoppingCallback,
)
from transformers.trainer_utils import get_last_checkpoint
import evaluate
@dataclass
class DataConfig:
"""Configurare pentru prelucrarea datelor."""
dataset_name: str = "datadriven-company/TTS-Romanian"
dataset_path: Optional[str] = None # Cale locală dacă e deja descărcat
train_split: str = "train"
eval_split: str = "validation"
audio_column: str = "audio"
text_column: str = "text"
preprocessing_num_workers: int = 4
max_duration_seconds: float = 30.0
# Filtre opționale (vor fi setate de utilizator)
min_text_length: int = 5
max_text_length: int = 500
@dataclass
class ModelConfig:
"""Configurare pentru model."""
model_name_or_path: str = "gabrielpirlo/Sped_ParakeetRomanian_110M_TDT-CTC"
processor_name_or_path: Optional[str] = None
@dataclass
class TrainingConfig:
"""Configurare pentru antrenament."""
output_dir: str = "/mnt/parakeet-training/outputs/parakeet-romanian-tts"
num_train_epochs: int = 10
per_device_train_batch_size: int = 8
per_device_eval_batch_size: int = 8
gradient_accumulation_steps: int = 2
learning_rate: float = 5e-5
warmup_steps: int = 500
eval_strategy: str = "steps"
eval_steps: int = 500
save_strategy: str = "steps"
save_steps: int = 500
logging_steps: int = 100
save_total_limit: int = 3
load_best_model_at_end: bool = True
metric_for_best_model: str = "wer"
greater_is_better: bool = False
fp16: bool = True
dataloader_num_workers: int = 4
remove_unused_columns: bool = False
seed: int = 42
report_to: List[str] = field(default_factory=lambda: ["tensorboard"])
def load_data_and_model(data_cfg: DataConfig, model_cfg: ModelConfig):
"""Încarcă dataset-ul și modelul."""
print(f"[INFO] Încărcare model: {model_cfg.model_name_or_path}")
processor = AutoProcessor.from_pretrained(
model_cfg.processor_name_or_path or model_cfg.model_name_or_path
)
model = AutoModelForCTC.from_pretrained(model_cfg.model_name_or_path)
print(f"[INFO] Încărcare dataset: {data_cfg.dataset_name}")
if data_cfg.dataset_path and os.path.exists(data_cfg.dataset_path):
dataset = load_from_disk(data_cfg.dataset_path)
else:
dataset = load_dataset(data_cfg.dataset_name)
# Asigurăm sampling rate-ul corect
if data_cfg.audio_column in dataset[data_cfg.train_split].column_names:
dataset = dataset.cast_column(data_cfg.audio_column, Audio(sampling_rate=16000))
return dataset, model, processor
def prepare_dataset(batch, processor, audio_column: str = "audio", text_column: str = "text"):
"""Preprocesează un batch de date."""
# Extragem features audio
audio = batch[audio_column]
# Procesăm audio
inputs = processor(
audio["array"],
sampling_rate=audio["sampling_rate"],
return_tensors="pt"
)
batch["input_values"] = inputs.input_values[0]
# Procesăm text (labels)
with processor.as_target_processor():
batch["labels"] = processor(batch[text_column]).input_ids
return batch
def compute_metrics(pred, processor):
"""Calculează WER și CER."""
pred_logits = pred.predictions
pred_ids = np.argmax(pred_logits, axis=-1)
pred.label_ids[pred.label_ids == -100] = processor.tokenizer.pad_token_id
pred_str = processor.batch_decode(pred_ids)
label_str = processor.batch_decode(pred.label_ids, group_tokens=False)
wer_metric = evaluate.load("wer")
cer_metric = evaluate.load("cer")
wer = wer_metric.compute(predictions=pred_str, references=label_str)
cer = cer_metric.compute(predictions=pred_str, references=label_str)
return {"wer": wer, "cer": cer}
def main():
import argparse
parser = argparse.ArgumentParser()
parser.add_argument("--dataset_name", default="datadriven-company/TTS-Romanian")
parser.add_argument("--model_name", default="gabrielpirlo/Sped_ParakeetRomanian_110M_TDT-CTC")
parser.add_argument("--output_dir", default="/data/outputs/parakeet-romanian-tts")
parser.add_argument("--num_epochs", type=int, default=10)
parser.add_argument("--batch_size", type=int, default=8)
parser.add_argument("--learning_rate", type=float, default=5e-5)
parser.add_argument("--local_dataset_path", default=None)
args = parser.parse_args()
data_cfg = DataConfig(
dataset_name=args.dataset_name,
dataset_path=args.local_dataset_path,
)
model_cfg = ModelConfig(model_name_or_path=args.model_name)
train_cfg = TrainingConfig(
output_dir=args.output_dir,
num_train_epochs=args.num_epochs,
per_device_train_batch_size=args.batch_size,
learning_rate=args.learning_rate,
)
# Verificare GPU
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"[INFO] Device: {device}")
if device == "cuda":
print(f"[INFO] GPU: {torch.cuda.get_device_name(0)}")
print(f"[INFO] GPU Memory: {torch.cuda.get_device_properties(0).total_memory / 1024**3:.1f} GB")
# Încărcare date și model
dataset, model, processor = load_data_and_model(data_cfg, model_cfg)
print(f"[INFO] Dataset splits: {list(dataset.keys())}")
# Aplicăm preprocesarea
print("[INFO] Preprocesare dataset...")
dataset = dataset.map(
lambda x: prepare_dataset(x, processor, data_cfg.audio_column, data_cfg.text_column),
remove_columns=dataset[data_cfg.train_split].column_names,
num_proc=data_cfg.preprocessing_num_workers,
batched=False,
)
# Verificăm existența split-urilor
train_dataset = dataset.get(data_cfg.train_split)
eval_dataset = dataset.get(data_cfg.eval_split)
if eval_dataset is None:
print("[WARNING] Nu există split de validare. Se folosește un procent din train.")
split = train_dataset.train_test_split(test_size=0.1, seed=42)
train_dataset = split["train"]
eval_dataset = split["test"]
print(f"[INFO] Train samples: {len(train_dataset)}")
print(f"[INFO] Eval samples: {len(eval_dataset)}")
# Configurare training
training_args = TrainingArguments(
output_dir=train_cfg.output_dir,
num_train_epochs=train_cfg.num_train_epochs,
per_device_train_batch_size=train_cfg.per_device_train_batch_size,
per_device_eval_batch_size=train_cfg.per_device_eval_batch_size,
gradient_accumulation_steps=train_cfg.gradient_accumulation_steps,
learning_rate=train_cfg.learning_rate,
warmup_steps=train_cfg.warmup_steps,
evaluation_strategy=train_cfg.eval_strategy,
eval_steps=train_cfg.eval_steps,
save_strategy=train_cfg.save_strategy,
save_steps=train_cfg.save_steps,
logging_steps=train_cfg.logging_steps,
save_total_limit=train_cfg.save_total_limit,
load_best_model_at_end=train_cfg.load_best_model_at_end,
metric_for_best_model=train_cfg.metric_for_best_model,
greater_is_better=train_cfg.greater_is_better,
fp16=train_cfg.fp16,
dataloader_num_workers=train_cfg.dataloader_num_workers,
remove_unused_columns=train_cfg.remove_unused_columns,
seed=train_cfg.seed,
report_to=train_cfg.report_to,
)
# Data collator pentru CTC
from dataclasses import dataclass
from typing import Any, Dict, List, Union
@dataclass
class DataCollatorCTCWithPadding:
processor: Any
padding: Union[bool, str] = True
max_length: Optional[int] = None
max_length_labels: Optional[int] = None
pad_to_multiple_of: Optional[int] = None
pad_to_multiple_of_labels: Optional[int] = None
def __call__(self, features: List[Dict[str, Union[List[int], torch.Tensor]]]) -> Dict[str, torch.Tensor]:
input_features = [{"input_values": feature["input_values"]} for feature in features]
label_features = [{"input_ids": feature["labels"]} for feature in features]
batch = self.processor.pad(
input_features,
padding=self.padding,
max_length=self.max_length,
pad_to_multiple_of=self.pad_to_multiple_of,
return_tensors="pt",
)
with self.processor.as_target_processor():
labels_batch = self.processor.pad(
label_features,
padding=self.padding,
max_length=self.max_length_labels,
pad_to_multiple_of=self.pad_to_multiple_of_labels,
return_tensors="pt",
)
labels = labels_batch["input_ids"].masked_fill(labels_batch.attention_mask.ne(1), -100)
batch["labels"] = labels
return batch
data_collator = DataCollatorCTCWithPadding(processor=processor, padding=True)
# Inițializare Trainer
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
tokenizer=processor.feature_extractor,
data_collator=data_collator,
compute_metrics=lambda p: compute_metrics(p, processor),
callbacks=[EarlyStoppingCallback(early_stopping_patience=3)],
)
# Verificare checkpoint existent
last_checkpoint = None
if os.path.isdir(train_cfg.output_dir) and len(os.listdir(train_cfg.output_dir)) > 0:
last_checkpoint = get_last_checkpoint(train_cfg.output_dir)
if last_checkpoint:
print(f"[INFO] Continuare de la checkpoint: {last_checkpoint}")
# Antrenament
print("[INFO] Începere antrenament...")
train_result = trainer.train(resume_from_checkpoint=last_checkpoint)
# Salvare finală
trainer.save_model()
processor.save_pretrained(train_cfg.output_dir)
# Metrici finale
metrics = train_result.metrics
trainer.save_metrics("train", metrics)
# Evaluare finală
print("[INFO] Evaluare finală...")
eval_metrics = trainer.evaluate()
trainer.save_metrics("eval", eval_metrics)
print(f"\n[REZULTATE]")
print(f" WER final: {eval_metrics.get('eval_wer', 'N/A')}")
print(f" CER final: {eval_metrics.get('eval_cer', 'N/A')}")
print(f" Model salvat în: {train_cfg.output_dir}")
if __name__ == "__main__":
main()