import os import torch import jiwer from tqdm import tqdm import torchaudio import numpy as np from datasets import load_dataset, Audio from transformers import ( Wav2Vec2CTCTokenizer, SeamlessM4TFeatureExtractor, Wav2Vec2BertProcessor, Wav2Vec2BertForCTC, ) # ===================================================== # CONFIG # ===================================================== DEVICE = "cuda" if torch.cuda.is_available() else "cpu" TARGET_SR = 16000 BATCH_SIZE = 8 VOCAB_DIR = "/home/devbcp/Practicas/wav2vec/datasets" CHECKPOINT = "/home/devbcp/Practicas/wav2vec/w2vbert-galician/checkpoint-64000" # ===================================================== # LOAD TOKENIZER + PROCESSOR # ===================================================== tokenizer = Wav2Vec2CTCTokenizer( vocab_file=f"{VOCAB_DIR}/vocab.json", unk_token="[UNK]", pad_token="[PAD]", word_delimiter_token="|", ) feature_extractor = SeamlessM4TFeatureExtractor.from_pretrained("facebook/w2v-bert-2.0") processor = Wav2Vec2BertProcessor( feature_extractor=feature_extractor, tokenizer=tokenizer, ) # ===================================================== # LOAD MODEL # ===================================================== print(f"✔ Cargando modelo desde: {CHECKPOINT}") model = Wav2Vec2BertForCTC.from_pretrained(CHECKPOINT).to(DEVICE).eval() # ===================================================== # AUDIO LOADER # ===================================================== def load_audio(path, target_sr=16000): wav, sr = torchaudio.load(path) if wav.shape[0] > 1: wav = wav.mean(dim=0) else: wav = wav.squeeze(0) if sr != target_sr: wav = torchaudio.functional.resample(wav, sr, target_sr) return wav.numpy() import re chars_to_remove_regex = r"[\,\?\.\!\-\;\:\"\“\%\‘\”\�\'\»\«]" def normalize_text_eval(text): if text is None: return "" text = text.lower() text = re.sub(chars_to_remove_regex, "", text) text = text.replace("\u00A0", " ") # espacio raro text = text.strip() return text # ===================================================== # INFERENCE FUNCTION # ===================================================== def infer_dataset_wav2vec(ds, name, audio_mode, audio_root=None, text_key="text"): hyps, refs = [], [] for i in tqdm(range(0, len(ds), BATCH_SIZE), desc=name): batch = ds.select(range(i, min(i + BATCH_SIZE, len(ds)))) waves = [] batch_refs = [] # Load audio if audio_mode == "commonvoice": for ex in batch: path = ex.get("path") if not path: continue audio_path = os.path.join(audio_root, path) if not os.path.exists(audio_path): continue wav = load_audio(audio_path, TARGET_SR) waves.append(wav) batch_refs.append(ex[text_key]) else: for ex in batch: audio = ex.get("audio") if audio is None or audio.get("array") is None: continue waves.append(audio["array"]) batch_refs.append(ex[text_key]) if len(waves) == 0: continue # Preprocess inputs = processor( waves, sampling_rate=TARGET_SR, return_tensors="pt", padding=True, ) input_features = inputs["input_features"].to(DEVICE) # Forward with torch.no_grad(): logits = model(input_features).logits pred_ids = torch.argmax(logits, dim=-1) # Decode texts = processor.batch_decode(pred_ids, skip_special_tokens=True) # Normalizar predicciones y referencias igual que en el dataset de entrenamiento texts = [normalize_text_eval(t) for t in texts] batch_refs = [normalize_text_eval(r) for r in batch_refs] hyps.extend(texts) refs.extend(batch_refs) # Metrics wer = jiwer.wer(refs, hyps) cer = jiwer.cer(refs, hyps) print(f"{name:15} | N={len(refs):6d} | WER={wer:.4f} | CER={cer:.4f}") return refs, hyps # ===================================================== # LOAD DATASETS # ===================================================== common_voice = load_dataset( "/home/devbcp/Proyectos/00-DATASETS/ASR/CommonVoice-v23-GL", split="test", ) common_voice = common_voice.filter(lambda x: x["path"] not in [None, ""]) openslr = load_dataset( "/home/devbcp/Proyectos/00-DATASETS/ASR/OpenSLR-SpeechT-GL-EN", split="test", ).cast_column("audio", Audio(sampling_rate=TARGET_SR)) fleurs = load_dataset( "/home/devbcp/Proyectos/00-DATASETS/ASR/FLEURS-SpeechT-GL-EN", split="test", ).cast_column("audio", Audio(sampling_rate=TARGET_SR)) transcrispeech = load_dataset( "/home/devbcp/Proyectos/00-DATASETS/ASR/Transcrispeech-GL", split="test", ).cast_column("audio", Audio(sampling_rate=TARGET_SR)) falai_full = load_dataset("/home/devbcp/Proyectos/00-DATASETS/ASR/FalAI") falai_validated = falai_full["validated"].cast_column("audio", Audio(sampling_rate=TARGET_SR)) n = int(0.2 * len(falai_validated)) falai_sampled = falai_validated.shuffle(seed=42).select(range(n)) falai_split = falai_sampled.train_test_split(test_size=0.2, seed=42) val_test = falai_split["test"].train_test_split(test_size=0.5, seed=42) falai_test = val_test["test"] rg_podcast = load_dataset( "/home/devbcp/Proyectos/00-DATASETS/ASR/RG-Podcast-GL" ) rg_podcast = rg_podcast.cast_column("audio", Audio(sampling_rate=TARGET_SR)) # ===================================================== # DATASETS TO EVALUATE # ===================================================== datasets = { "FalAI": { "ds": falai_test, "audio_mode": "array", "text_key": "sentence", }, "CommonVoice": { "ds": common_voice, "audio_mode": "commonvoice", "audio_root": "/home/devbcp/Proyectos/00-DATASETS/ASR/CommonVoice-v23-GL/cv-corpus-23.0-2025-09-05/gl/clips", "text_key": "sentence", }, "OpenSLR": { "ds": openslr, "audio_mode": "array", "text_key": "text_gl", }, "FLEURS": { "ds": fleurs, "audio_mode": "array", "text_key": "text_gl", }, "Transcrispeech": { "ds": transcrispeech, "audio_mode": "array", "text_key": "text", }, "RG-Podcast": { "ds": rg_podcast["test"], "audio_mode": "array", "text_key": "text", }, } # ===================================================== # EVALUATION # ===================================================== all_refs, all_hyps = [], [] print("\n=== EVALUACIÓN POR DATASET ===\n") for name, cfg in datasets.items(): refs, hyps = infer_dataset_wav2vec( cfg["ds"], name, audio_mode=cfg["audio_mode"], audio_root=cfg.get("audio_root"), text_key=cfg["text_key"], ) all_refs.extend(refs) all_hyps.extend(hyps) wer = jiwer.wer(all_refs, all_hyps) cer = jiwer.cer(all_refs, all_hyps) print("\n=== COMBINADO ===") print(f"TOTAL | N={len(all_refs)} | WER={wer:.4f} | CER={cer:.4f}")