""" 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()