Team Ai
Apppublic

pratik-250620/MultiModal-Coherence-AI

sourceHugging Facemitupdated 8mo agoView on Hugging Face
2likes
human_correlation.py399 linesDownload Raw Back to validation
1"""2Human Correlation Analysis3 4Analyzes correlation between MSCI scores and human judgments.5This addresses RQ3: "Does MSCI correlate with human judgments of multimodal coherence?"6 7Key analyses:8- Spearman rank correlation (for ordinal human ratings)9- Pearson correlation (for continuous relationship)10- Per-dimension correlations (text-image, text-audio, image-audio)11- Agreement analysis12"""13 14from __future__ import annotations15 16import json17from dataclasses import dataclass18from pathlib import Path19from typing import Any, Dict, List, Optional, Tuple20import numpy as np21from scipy import stats22 23 24@dataclass25class CorrelationResult:26    """Result of a correlation analysis."""27    variable1: str28    variable2: str29    spearman_rho: float30    spearman_p: float31    pearson_r: float32    pearson_p: float33    n: int34    ci_lower: float35    ci_upper: float36    significant: bool37    interpretation: str38 39    def to_dict(self) -> Dict[str, Any]:40        """Convert to dictionary."""41        return {42            "variable1": self.variable1,43            "variable2": self.variable2,44            "spearman_rho": self.spearman_rho,45            "spearman_p": self.spearman_p,46            "pearson_r": self.pearson_r,47            "pearson_p": self.pearson_p,48            "n": self.n,49            "ci_95": [self.ci_lower, self.ci_upper],50            "significant": self.significant,51            "interpretation": self.interpretation,52        }53 54 55class HumanCorrelationAnalyzer:56    """57    Analyzes correlation between MSCI and human judgments.58 59    RQ3: "Does MSCI correlate with human judgments of multimodal coherence?"60    H0: ρ(MSCI, human) ≤ 061    H1: ρ(MSCI, human) > 062    """63 64    def __init__(self, alpha: float = 0.05):65        self.alpha = alpha66 67    def compute_correlation(68        self,69        msci_scores: List[float],70        human_scores: List[float],71        var1_name: str = "MSCI",72        var2_name: str = "Human",73    ) -> CorrelationResult:74        """75        Compute correlation with confidence interval.76 77        Args:78            msci_scores: MSCI scores79            human_scores: Human coherence scores (normalized to 0-1)80            var1_name: Name for first variable81            var2_name: Name for second variable82 83        Returns:84            CorrelationResult with all statistics85        """86        if len(msci_scores) != len(human_scores):87            raise ValueError("Score lists must have same length")88 89        n = len(msci_scores)90        if n < 3:91            return CorrelationResult(92                variable1=var1_name,93                variable2=var2_name,94                spearman_rho=0.0,95                spearman_p=1.0,96                pearson_r=0.0,97                pearson_p=1.0,98                n=n,99                ci_lower=-1.0,100                ci_upper=1.0,101                significant=False,102                interpretation="Insufficient data (N < 3)",103            )104 105        # Spearman correlation (better for ordinal human ratings)106        spearman = stats.spearmanr(msci_scores, human_scores)107 108        # Pearson correlation109        pearson = stats.pearsonr(msci_scores, human_scores)110 111        # Confidence interval for Spearman (using Fisher z-transformation)112        z = np.arctanh(spearman.correlation)113        se_z = 1 / np.sqrt(n - 3)114        z_crit = stats.norm.ppf(1 - self.alpha / 2)115        ci_lower = np.tanh(z - z_crit * se_z)116        ci_upper = np.tanh(z + z_crit * se_z)117 118        # Significance (one-tailed test: ρ > 0)119        significant = spearman.pvalue / 2 < self.alpha and spearman.correlation > 0120 121        # Interpretation122        interpretation = self._interpret_correlation(123            spearman.correlation, spearman.pvalue / 2, significant124        )125 126        return CorrelationResult(127            variable1=var1_name,128            variable2=var2_name,129            spearman_rho=float(spearman.correlation),130            spearman_p=float(spearman.pvalue),131            pearson_r=float(pearson.statistic),132            pearson_p=float(pearson.pvalue),133            n=n,134            ci_lower=float(ci_lower),135            ci_upper=float(ci_upper),136            significant=significant,137            interpretation=interpretation,138        )139 140    def _interpret_correlation(141        self,142        rho: float,143        p_one_tailed: float,144        significant: bool,145    ) -> str:146        """Generate interpretation of correlation."""147        if not significant:148            if p_one_tailed >= self.alpha:149                return f"No significant positive correlation (ρ={rho:.3f}, p={p_one_tailed:.4f})"150            else:151                return f"Significant negative correlation (unexpected; ρ={rho:.3f})"152 153        abs_rho = abs(rho)154        if abs_rho >= 0.7:155            strength = "strong"156        elif abs_rho >= 0.5:157            strength = "moderate-strong"158        elif abs_rho >= 0.3:159            strength = "moderate"160        else:161            strength = "weak"162 163        return f"Significant {strength} positive correlation (ρ={rho:.3f}, p={p_one_tailed:.4f})"164 165    def analyze_from_human_eval(166        self,167        human_eval_path: Path,168        msci_scores: Optional[Dict[str, float]] = None,169    ) -> Dict[str, Any]:170        """171        Analyze correlation from human evaluation session.172 173        Args:174            human_eval_path: Path to human evaluation session JSON175            msci_scores: Optional dict of sample_id -> MSCI score176 177        Returns:178            Comprehensive correlation analysis179        """180        from src.evaluation.human_eval_schema import EvaluationSession181 182        session = EvaluationSession.load(Path(human_eval_path))183 184        # Build sample ID -> MSCI mapping from session if not provided185        if msci_scores is None:186            msci_scores = {}187            for sample in session.samples:188                if sample.msci_score is not None:189                    msci_scores[sample.sample_id] = sample.msci_score190 191        # Collect paired data192        pairs: List[Dict[str, Any]] = []193 194        for eval in session.evaluations:195            if eval.is_rerating:196                continue197            if eval.sample_id not in msci_scores:198                continue199 200            pairs.append({201                "sample_id": eval.sample_id,202                "msci": msci_scores[eval.sample_id],203                "human_weighted": eval.weighted_score(),204                "human_overall": eval.overall_coherence / 5.0,  # Normalize205                "human_ti": eval.text_image_coherence / 5.0,206                "human_ta": eval.text_audio_coherence / 5.0,207                "human_ia": eval.image_audio_coherence / 5.0,208            })209 210        if len(pairs) < 3:211            return {212                "error": "Insufficient paired data",213                "n_pairs": len(pairs),214            }215 216        # Extract arrays217        msci = [p["msci"] for p in pairs]218        human_weighted = [p["human_weighted"] for p in pairs]219        human_overall = [p["human_overall"] for p in pairs]220        human_ti = [p["human_ti"] for p in pairs]221        human_ta = [p["human_ta"] for p in pairs]222        human_ia = [p["human_ia"] for p in pairs]223 224        # Compute correlations225        results = {226            "n_pairs": len(pairs),227            "overall_correlation": self.compute_correlation(228                msci, human_weighted, "MSCI", "Human Weighted Score"229            ).to_dict(),230            "overall_rating_correlation": self.compute_correlation(231                msci, human_overall, "MSCI", "Human Overall Rating"232            ).to_dict(),233            "per_dimension": {234                "text_image": self.compute_correlation(235                    msci, human_ti, "MSCI", "Human Text-Image"236                ).to_dict(),237                "text_audio": self.compute_correlation(238                    msci, human_ta, "MSCI", "Human Text-Audio"239                ).to_dict(),240                "image_audio": self.compute_correlation(241                    msci, human_ia, "MSCI", "Human Image-Audio"242                ).to_dict(),243            },244        }245 246        # RQ3 verdict247        main_corr = results["overall_correlation"]248        results["rq3_verdict"] = self._rq3_verdict(main_corr)249 250        return results251 252    def _rq3_verdict(self, correlation: Dict[str, Any]) -> Dict[str, Any]:253        """Generate RQ3 verdict from correlation result."""254        rho = correlation["spearman_rho"]255        p = correlation["spearman_p"]256        significant = correlation["significant"]257 258        if significant and rho > 0.3:259            verdict = "SUPPORTED"260            explanation = (261                f"MSCI shows significant positive correlation with human judgments "262                f"(ρ={rho:.3f}, p={p/2:.4f}). MSCI is a valid proxy for human-perceived coherence."263            )264        elif significant and rho > 0:265            verdict = "WEAKLY SUPPORTED"266            explanation = (267                f"MSCI shows significant but weak correlation with human judgments "268                f"(ρ={rho:.3f}). MSCI captures some aspects of human-perceived coherence."269            )270        elif not significant and rho > 0:271            verdict = "NOT SUPPORTED"272            explanation = (273                f"No significant correlation between MSCI and human judgments "274                f"(ρ={rho:.3f}, p={p/2:.4f}). MSCI may not reliably reflect human perception."275            )276        else:277            verdict = "CONTRADICTED"278            explanation = (279                f"Unexpected negative correlation (ρ={rho:.3f}). "280                f"MSCI may be inversely related to human perception."281            )282 283        return {284            "verdict": verdict,285            "explanation": explanation,286            "threshold_met": significant and rho > 0.3,287            "rho": rho,288            "p_value": p / 2,  # One-tailed289        }290 291    def analyze_disagreements(292        self,293        pairs: List[Dict[str, Any]],294        threshold: float = 0.2,295    ) -> Dict[str, Any]:296        """297        Analyze cases where MSCI and human judgments disagree.298 299        Args:300            pairs: List of dicts with 'msci' and 'human_weighted' keys301            threshold: Disagreement threshold (normalized)302 303        Returns:304            Analysis of disagreement patterns305        """306        disagreements = []307 308        for pair in pairs:309            msci = pair.get("msci", 0)310            human = pair.get("human_weighted", 0)311            diff = msci - human312 313            if abs(diff) > threshold:314                disagreements.append({315                    "sample_id": pair.get("sample_id"),316                    "msci": msci,317                    "human": human,318                    "difference": diff,319                    "type": "MSCI_overestimates" if diff > 0 else "MSCI_underestimates",320                })321 322        n_total = len(pairs)323        n_disagree = len(disagreements)324 325        overestimates = [d for d in disagreements if d["type"] == "MSCI_overestimates"]326        underestimates = [d for d in disagreements if d["type"] == "MSCI_underestimates"]327 328        return {329            "n_total": n_total,330            "n_disagreements": n_disagree,331            "disagreement_rate": n_disagree / n_total if n_total > 0 else 0,332            "n_overestimates": len(overestimates),333            "n_underestimates": len(underestimates),334            "mean_overestimate": (335                np.mean([d["difference"] for d in overestimates])336                if overestimates else 0337            ),338            "mean_underestimate": (339                np.mean([abs(d["difference"]) for d in underestimates])340                if underestimates else 0341            ),342            "samples": disagreements,343        }344 345    def generate_report(346        self,347        analysis_results: Dict[str, Any],348        output_path: Optional[Path] = None,349    ) -> Dict[str, Any]:350        """351        Generate comprehensive human correlation report.352 353        Args:354            analysis_results: Results from analyze_from_human_eval355            output_path: Optional path to save report356 357        Returns:358            Complete correlation report359        """360        report = {361            "analysis_type": "MSCI-Human Correlation Analysis",362            "research_question": "RQ3: Does MSCI correlate with human judgments?",363            "hypothesis": {364                "H0": "ρ(MSCI, human) ≤ 0",365                "H1": "ρ(MSCI, human) > 0",366                "threshold": "ρ > 0.3 for meaningful validity",367            },368            "results": analysis_results,369        }370 371        # Add recommendations based on results372        verdict = analysis_results.get("rq3_verdict", {})373        if verdict.get("verdict") == "SUPPORTED":374            report["recommendations"] = [375                "MSCI can be used as a proxy for human coherence judgments",376                "Consider using MSCI for automated evaluation at scale",377            ]378        elif verdict.get("verdict") == "WEAKLY SUPPORTED":379            report["recommendations"] = [380                "MSCI provides some signal but should not be sole metric",381                "Consider combining MSCI with other metrics or human spot-checks",382                "Investigate which dimensions MSCI captures well vs poorly",383            ]384        else:385            report["recommendations"] = [386                "MSCI may not reliably reflect human perception",387                "Consider revising MSCI weights or embedding approach",388                "Human evaluation remains necessary for validation",389                "Investigate failure modes to improve MSCI",390            ]391 392        if output_path:393            output_path = Path(output_path)394            output_path.parent.mkdir(parents=True, exist_ok=True)395            with output_path.open("w", encoding="utf-8") as f:396                json.dump(report, f, indent=2, ensure_ascii=False)397 398        return report399