import pandas as pd from jiwer import wer, cer from inference import batch_correct_texts, get_model_name import logging from typing import Dict, List, Tuple from datasets import load_dataset as hf_load_dataset logger = logging.getLogger(__name__) def load_dataset(csv_path: str = None, hf_dataset_id: str = None, split: str = "test") -> Tuple[List[str], List[str]]: """ Load dataset from CSV or HuggingFace Args: csv_path: Path to CSV file with 'prediction' and 'reference' columns (optional) hf_dataset_id: HuggingFace dataset ID (optional) split: Dataset split to use if loading from HuggingFace (default: "test") Returns: Tuple of (dyslexic_sentences, clean_sentences) """ # Load from HuggingFace dataset if specified if hf_dataset_id: logger.info(f"Loading HuggingFace dataset: {hf_dataset_id} (split: {split})") try: dataset = hf_load_dataset(hf_dataset_id, split=split) logger.info(f"Loaded {len(dataset)} samples from HuggingFace") # Extract dyslexic and clean sentences dyslexic = dataset['dyslexic_sentence'] clean = dataset['clean_sentence'] return dyslexic, clean except Exception as e: logger.error(f"Error loading HuggingFace dataset: {e}") raise # Load from CSV if specified elif csv_path: logger.info(f"Loading CSV dataset: {csv_path}") df = pd.read_csv(csv_path) # Try to find the right column names if 'dyslexic_sentence' in df.columns and 'clean_sentence' in df.columns: dyslexic = df['dyslexic_sentence'].tolist() clean = df['clean_sentence'].tolist() elif 'prediction' in df.columns and 'reference' in df.columns: dyslexic = df['prediction'].tolist() clean = df['reference'].tolist() else: raise ValueError(f"CSV must have columns: 'dyslexic_sentence'/'clean_sentence' or 'prediction'/'reference'. Found: {df.columns.tolist()}") logger.info(f"Loaded {len(dyslexic)} samples from CSV") return dyslexic, clean else: raise ValueError("Either csv_path or hf_dataset_id must be provided") def calculate_accuracy(predicted: List[str], reference: List[str]) -> float: """Calculate exact match accuracy""" matches = sum(1 for p, r in zip(predicted, reference) if p.strip() == r.strip()) return (matches / len(predicted)) * 100 if predicted else 0 def calculate_wer(predicted: List[str], reference: List[str]) -> float: """Calculate Word Error Rate""" total_wer = 0 valid_count = 0 for p, r in zip(predicted, reference): try: w = wer(r.split(), p.split()) total_wer += w valid_count += 1 except Exception as e: logger.warning(f"Error calculating WER: {e}") continue return (total_wer / valid_count) if valid_count > 0 else 0 def calculate_cer(predicted: List[str], reference: List[str]) -> float: """Calculate Character Error Rate""" total_cer = 0 valid_count = 0 for p, r in zip(predicted, reference): try: c = cer(r, p) total_cer += c valid_count += 1 except Exception as e: logger.warning(f"Error calculating CER: {e}") continue return (total_cer / valid_count) if valid_count > 0 else 0 def evaluate_model( model_id: str, predictions: List[str], references: List[str] ) -> Dict[str, float]: """ Evaluate a single model on the dataset Args: model_id: HuggingFace model ID predictions: List of input texts to correct references: List of reference (correct) texts Returns: Dictionary with metrics """ logger.info(f"Evaluating model: {model_id}") try: corrected_texts = batch_correct_texts(predictions, model_id) accuracy = calculate_accuracy(corrected_texts, references) wer_score = calculate_wer(corrected_texts, references) cer_score = calculate_cer(corrected_texts, references) return { 'model_name': get_model_name(model_id), 'model_id': model_id, 'accuracy': accuracy, 'wer': wer_score, 'cer': cer_score } except Exception as e: logger.error(f"Error evaluating {model_id}: {e}") return { 'model_name': get_model_name(model_id), 'model_id': model_id, 'accuracy': 0, 'wer': float('inf'), 'cer': float('inf'), 'error': str(e) } def evaluate_all_models( model_ids: List[str], csv_path: str = None, hf_dataset_id: str = None, split: str = "test", sample_size: int = None ) -> pd.DataFrame: """ Evaluate all models Args: model_ids: List of model IDs csv_path: Path to CSV dataset (optional) hf_dataset_id: HuggingFace dataset ID (optional) split: Dataset split to use if loading from HuggingFace (default: "test") sample_size: If set, evaluate only on a sample Returns: DataFrame with evaluation results """ # Determine which dataset source to use if hf_dataset_id: logger.info(f"Using HuggingFace dataset: {hf_dataset_id}") predictions, references = load_dataset(hf_dataset_id=hf_dataset_id, split=split) elif csv_path: logger.info(f"Using CSV dataset: {csv_path}") predictions, references = load_dataset(csv_path=csv_path) else: raise ValueError("Either csv_path or hf_dataset_id must be provided") if sample_size and sample_size < len(predictions): logger.info(f"Using sample of {sample_size} from {len(predictions)} total samples") predictions = predictions[:sample_size] references = references[:sample_size] results = [] for i, model_id in enumerate(model_ids, 1): logger.info(f"[{i}/{len(model_ids)}] Evaluating {model_id}") result = evaluate_model(model_id, predictions, references) results.append(result) return pd.DataFrame(results) def get_best_models(results_df: pd.DataFrame, metric: str = 'accuracy', top_k: int = 5): """Get top-k models by metric""" if metric == 'accuracy': sorted_df = results_df.sort_values('accuracy', ascending=False) elif metric == 'wer': sorted_df = results_df.sort_values('wer', ascending=True) elif metric == 'cer': sorted_df = results_df.sort_values('cer', ascending=True) else: sorted_df = results_df.sort_values('accuracy', ascending=False) return sorted_df.head(top_k)