diff --git a/pyproject.toml b/pyproject.toml index 685ceb0..a11d076 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -11,7 +11,6 @@ dependencies = [ "matplotlib>=3.10.8", "scipy>=1.17.0", "seaborn>=0.13.2", - "sentence-transformers>=5.3.0", ] [build-system] @@ -102,8 +101,6 @@ module = [ "seaborn.*", "matplotlib", "matplotlib.*", - "sentence_transformers", - "sentence_transformers.*" ] ignore_missing_imports = true diff --git a/src/voice/__init__.py b/src/voice/__init__.py index 5551485..a3a9cc5 100644 --- a/src/voice/__init__.py +++ b/src/voice/__init__.py @@ -6,8 +6,8 @@ """ from voice.comparison import ( - make_embedding_comparison, - make_stylometric_comparison, + make_comparison, + stylometric_distribution, ) from voice.datasets import DatasetSpec, get_dataset from voice.stylometry import get_groups, get_metrics @@ -17,6 +17,6 @@ "get_metrics", "get_groups", "get_dataset", - "make_embedding_comparison", - "make_stylometric_comparison", + "stylometric_distribution", + "make_comparison", ] diff --git a/src/voice/stylometry/_defaults.py b/src/voice/_defaults.py similarity index 82% rename from src/voice/stylometry/_defaults.py rename to src/voice/_defaults.py index 3d76763..d9e5f29 100644 --- a/src/voice/stylometry/_defaults.py +++ b/src/voice/_defaults.py @@ -1,7 +1,7 @@ """ -Defaults for the stylometry module. +Defaults for the VOICE project. -This module contains defaults for stylometric analysis +This module contains defaults for the VOICE project e.g. default parameters for parameterised metrics """ @@ -14,19 +14,15 @@ class MetricGroup(str, Enum): Enumeration of stylometric metric groups. Metrics within the same group are expected to be correlated; - the group is the unit of analysis for Bonferroni correction. + groups are the unit of aggregation for the alignment score. .. attribute :: WORD_LENGTH_DISTRIBUTION Moments of the word length distribution (mean, std, skew, kurtosis). - .. attribute :: LEXICAL_RICHNESS + .. attribute :: VOCABULARY_RICHNESS - Type-token ratio and its moving-average variant. - - .. attribute :: LEGOMENA - - Hapax, dis and tri legomena ratios. + Type-token ratio, its moving-average variant, and legomena ratios. .. attribute :: FUNCTION_WORDS @@ -42,8 +38,7 @@ class MetricGroup(str, Enum): """ WORD_LENGTH_DISTRIBUTION = "word_length_distribution" - LEXICAL_RICHNESS = "lexical_richness" - LEGOMENA = "legomena" + VOCABULARY_RICHNESS = "vocabulary_richness" FUNCTION_WORDS = "function_words" CHAR_NGRAM_DIVERSITY = "char_ngram_diversity" TEXT_LENGTH = "text_length" @@ -116,22 +111,6 @@ class CalibrationDefaults: CALIBRATION_DEFAULTS: CalibrationDefaults = CalibrationDefaults() -@dataclass(frozen=True) -class ComparisonDefaults: - """ - Default parameters for comparing stylometric distributions. - - .. attribute :: alpha - - Significance level for hypothesis testing. - """ - - alpha: float = 0.05 - - -COMPARISON_DEFAULTS: ComparisonDefaults = ComparisonDefaults() - - @dataclass(frozen=True) class PlottingDefaults: """ diff --git a/src/voice/comparison/__init__.py b/src/voice/comparison/__init__.py index 2294a0a..cb46e01 100644 --- a/src/voice/comparison/__init__.py +++ b/src/voice/comparison/__init__.py @@ -5,14 +5,12 @@ - Functions for comparing distributions of stylometric metrics """ -from voice.comparison.embedding_comparison import make_embedding_comparison -from voice.comparison.stylometric_comparison import ( - make_stylometric_comparison, +from voice.comparison.comparison import ( + make_comparison, stylometric_distribution, ) __all__: list[str] = [ - "make_stylometric_comparison", - "make_embedding_comparison", "stylometric_distribution", + "make_comparison", ] diff --git a/src/voice/comparison/_utils.py b/src/voice/comparison/_utils.py index c893e47..a684a69 100644 --- a/src/voice/comparison/_utils.py +++ b/src/voice/comparison/_utils.py @@ -10,7 +10,7 @@ import numpy as np -from voice.stylometry._defaults import CALIBRATION_DEFAULTS +from voice._defaults import CALIBRATION_DEFAULTS # ----------------------------------------------------------------------------- # Calibrated percentile diff --git a/src/voice/comparison/stylometric_comparison.py b/src/voice/comparison/comparison.py similarity index 73% rename from src/voice/comparison/stylometric_comparison.py rename to src/voice/comparison/comparison.py index b5fd959..7173b68 100644 --- a/src/voice/comparison/stylometric_comparison.py +++ b/src/voice/comparison/comparison.py @@ -18,17 +18,16 @@ import numpy as np from scipy.stats import wasserstein_distance +from voice._defaults import ( + CALIBRATION_DEFAULTS, + MetricGroup, +) from voice.comparison._utils import ( bootstrap_null_distribution, calibrated_percentile, ) from voice.datasets import VoiceDataset from voice.datasets.dataset import Example -from voice.stylometry._defaults import ( - CALIBRATION_DEFAULTS, - COMPARISON_DEFAULTS, - MetricGroup, -) from voice.stylometry.metrics import get_metrics # ----------------------------------------------------------------------------- @@ -110,7 +109,7 @@ def self_wasserstein_distribution( @dataclass(frozen=True, slots=True) -class StylometricComparisonEntry: +class ComparisonEntry: """ Result for a single stylometric metric comparison. @@ -132,12 +131,7 @@ class StylometricComparisonEntry: .. attribute :: tail - Right-tail extremeness (1 - percentile). - - .. attribute :: flag - - Boolean indicating percentile <= 1 - alpha (uncorrected, - diagnostic only — decision is made at the group level). + Right tail extremeness (1 - percentile); higher means more similar. """ metric: str @@ -145,11 +139,10 @@ class StylometricComparisonEntry: wasserstein: float percentile: float tail: float - flag: bool @dataclass(frozen=True, slots=True) -class StylometricGroupEntry: +class GroupEntry: """ Aggregated result for a stylometric metric group. @@ -163,19 +156,12 @@ class StylometricGroupEntry: .. attribute :: avg_tail - Right-tail extremeness (1 - avg_percentile). - - .. attribute :: flag - - Boolean indicating avg_percentile <= 1 - alpha, i.e. the group's - average Wasserstein distance is not unusually large relative to the - self-Wasserstein null (the model is indistinguishable on this group). + Right tail extremeness (1 - avg_percentile); higher means more similar. """ group: MetricGroup avg_percentile: float avg_tail: float - flag: bool def __repr__(self) -> str: """ @@ -183,41 +169,29 @@ def __repr__(self) -> str: :return: A string representation of the group entry """ - return ( - f"StylometricGroupEntry(" - f"avg_percentile={self.avg_percentile:.4f}, " - f"flag={self.flag})" - ) + return f"GroupEntry(avg_percentile={self.avg_percentile:.4f})" @dataclass(slots=True) -class StylometricComparisonResults: +class ComparisonResults: """ Container for stylometric comparison results across multiple metrics. - The object is initialised with a subset of metric names and a - significance level `alpha`. Entries can then be added for each metric. + The object is initialised with a subset of metric names. Entries can + then be added for each metric. - The null hypothesis for each group is that the model is stylometrically - distinguishable from the reference. A group passes (flag=True) when its - average Wasserstein percentile is low enough to reject this null. - The overall result passes (all_pass=True) only when all groups pass, i.e. - an intersection-union test. + The alignment score is the mean of per-group average tail values, with + each group weighted equally regardless of how many metrics it contains. + A score of 1 indicates perfect stylistic indistinguishability; 0 indicates + maximum divergence. .. attribute :: metrics - List of stylometric metrics to run comparison for. - - .. attribute :: alpha - - Significance level for hypothesis testing. + Tuple of stylometric metric names to run comparison for. """ metrics: tuple[str, ...] - alpha: float = COMPARISON_DEFAULTS.alpha - _entries: dict[str, StylometricComparisonEntry] = field( - default_factory=dict - ) + _entries: dict[str, ComparisonEntry] = field(default_factory=dict) def __post_init__(self) -> None: """ @@ -236,18 +210,21 @@ def __post_init__(self) -> None: def __repr__(self) -> str: """ - Represent results as a group -> entry mapping. + Represent results as a group tail mapping and overall score. :return: A string representation of the results """ - return ( + group_part = ( "{" + ", ".join( - f"{group.value!r}: {entry}" - for group, entry in self.group_entries.items() + f"{group.value!r}: {tail:.4f}" + for group, tail in self.group_tails.items() ) + "}" ) + return ( + f"ComparisonResults(groups={group_part}, score={self.score:.4f})" + ) @property def num_metrics(self) -> int: @@ -255,42 +232,58 @@ def num_metrics(self) -> int: return len(self.metrics) @property - def group_entries(self) -> dict[MetricGroup, StylometricGroupEntry]: + def group_entries(self) -> dict[MetricGroup, GroupEntry]: """ Aggregated results keyed by MetricGroup. For each group, the average percentile across member metrics is - computed and used to determine group-level flags. + computed. Groups are weighted equally in the overall score. """ groups: dict[MetricGroup, list[float]] = {} for entry in self._entries.values(): groups.setdefault(entry.group, []).append(entry.percentile) - result: dict[MetricGroup, StylometricGroupEntry] = {} + result: dict[MetricGroup, GroupEntry] = {} for group, percentiles in groups.items(): avg_p = sum(percentiles) / len(percentiles) avg_tail = 1.0 - avg_p - result[group] = StylometricGroupEntry( + result[group] = GroupEntry( group=group, avg_percentile=avg_p, avg_tail=avg_tail, - flag=avg_p <= (1.0 - self.alpha), ) return result @property def score(self) -> float: - """Alignment score in [0, 1]; 1 = perfectly indistinguishable.""" + """ + Alignment score in [0, 1]; higher = more stylistically similar. + + Computed as the mean of per-group average tail values, with each + group weighted equally regardless of how many metrics it contains. + """ entries = self.group_entries if not entries: return 0.0 return sum(e.avg_tail for e in entries.values()) / len(entries) @property - def all_pass(self) -> bool: - """True if all groups pass; the model is indistinguishable overall.""" - entries = self.group_entries - return bool(entries) and all(e.flag for e in entries.values()) + def metric_tails(self) -> dict[str, float]: + """ + Per-metric tail values (1 - percentile) for diagnostic use. + + :return: Mapping of metric name to tail value + """ + return {m: e.tail for m, e in self._entries.items()} + + @property + def group_tails(self) -> dict[MetricGroup, float]: + """ + Per-group average tail values for diagnostic use. + + :return: Mapping of MetricGroup to average tail value + """ + return {g: e.avg_tail for g, e in self.group_entries.items()} def add( self, @@ -310,48 +303,45 @@ def add( """ if metric not in self.metrics: raise ValueError( - f"Metric '{metric}' not initialised in this " - f"StylometricComparisonResults." + f"Metric '{metric}' not initialised in this ComparisonResults." ) group = get_metrics()[metric].group tail = 1.0 - percentile - flag = percentile <= (1.0 - self.alpha) - entry = StylometricComparisonEntry( + entry = ComparisonEntry( metric=metric, group=group, wasserstein=float(wasserstein), percentile=float(percentile), tail=float(tail), - flag=flag, ) self._entries[metric] = entry - def get(self, metric: str) -> StylometricComparisonEntry: + def get(self, metric: str) -> ComparisonEntry: """ Retrieve the result entry for a metric. :param metric: Metric name - :return: StylometricComparisonEntry for metric + :return: ComparisonEntry for metric """ return self._entries[metric] - def get_group(self, group: MetricGroup) -> StylometricGroupEntry: + def get_group(self, group: MetricGroup) -> GroupEntry: """ Retrieve the aggregated result entry for a group. :param group: MetricGroup - :return: StylometricGroupEntry for the group + :return: GroupEntry for the group """ return self.group_entries[group] - def as_dict(self) -> dict[str, StylometricComparisonEntry]: + def as_dict(self) -> dict[str, ComparisonEntry]: """ - Return a mapping of metric name to StylometricComparisonEntry. + Return a mapping of metric name to ComparisonEntry. - :return: mapping of metric name to StylometricComparisonEntry object + :return: mapping of metric name to ComparisonEntry object """ return dict(self._entries) @@ -361,12 +351,12 @@ def as_dict(self) -> dict[str, StylometricComparisonEntry]: # ----------------------------------------------------------------------------- -def make_stylometric_comparison( +def make_comparison( completions: Sequence[Example], true_ds: VoiceDataset, *, metrics: Sequence[str] | None = None, -) -> StylometricComparisonResults: +) -> ComparisonResults: """ Compare stylometric distributions of `completions` to the true corpus. @@ -374,12 +364,12 @@ def make_stylometric_comparison( - Observed Wasserstein distance between completions and the reference completions on the same prompts. - Calibrated percentile of the observed distance under the - self-Wasserstein distribution estimated from the true training split + self-Wasserstein distribution estimated from the true training split, + using a sample size matched to the number of completions. - The returned ComparisonResults includes per-metric flags (diagnostic) and - group-level flags under H0: the model is distinguishable. A group passes - when its average Wasserstein percentile is <= 1 - alpha. The model passes - overall (all_pass=True) only when all groups pass. + The returned ComparisonResults provides per-metric and per-group tail + values, and an overall alignment score equal to the mean of group-level + average tail values (groups weighted equally). :param completions: Sequence of examples to analyse :param true_ds: VoiceDataset providing the true corpus @@ -396,25 +386,22 @@ def make_stylometric_comparison( registry = get_metrics() metric_list = list(registry.keys()) if metrics is None else list(metrics) - results = StylometricComparisonResults(metrics=tuple(metric_list)) + results = ComparisonResults(metrics=tuple(metric_list)) # Split that observed completions are generated on observed_split = completions[0].split ds_observed_split = true_ds[observed_split] for metric in metric_list: - # Observed distributions x = stylometric_distribution(completions, metric) - - # Wasserstein y = stylometric_distribution(ds_observed_split, metric) d = float(wasserstein_distance(x, y)) - # Self-W calibration distribution self_dist = self_wasserstein_distribution( true_ds, metric=metric, + sample_size=len(completions), ) p = calibrated_percentile(d, self_dist) diff --git a/src/voice/comparison/embedding_comparison.py b/src/voice/comparison/embedding_comparison.py deleted file mode 100644 index 9722a8a..0000000 --- a/src/voice/comparison/embedding_comparison.py +++ /dev/null @@ -1,231 +0,0 @@ -""" -Functions for comparing style via neural embedding similarity. - -This module contains: - - Embedding of text examples using a pretrained SentenceTransformer - - Self-similarity calibration distribution (train/train resampling) - - Comparison of generated completions against the true corpus - -All calibration is performed relative to the training split by design, -to avoid data leakage from validation/test splits. -""" - -from __future__ import annotations - -from collections.abc import Sequence -from dataclasses import dataclass - -import numpy as np -from sentence_transformers import SentenceTransformer - -from voice.comparison._utils import ( - bootstrap_null_distribution, - calibrated_percentile, -) -from voice.datasets import VoiceDataset -from voice.datasets.dataset import Example -from voice.stylometry._defaults import ( - CALIBRATION_DEFAULTS, - COMPARISON_DEFAULTS, -) - -# ----------------------------------------------------------------------------- -# Embedding extraction -# ----------------------------------------------------------------------------- - - -def embedding_matrix( - ds: Sequence[Example], - model: SentenceTransformer, -) -> np.ndarray: - """ - Embed a sequence of examples using a pretrained SentenceTransformer. - - Each example's `answer` field is embedded independently. The returned - matrix has one row per example and one column per embedding dimension, - with rows L2-normalised. - - :param ds: Sequence of examples, each expected to have an `answer` field - :param model: Pretrained SentenceTransformer model - :return: Float32 array of shape (len(ds), embedding_dim) - """ - texts = [e.answer for e in ds] - embeddings = model.encode( - texts, convert_to_numpy=True, normalize_embeddings=True - ) - return embeddings.astype(np.float32) - - -def mean_embedding(embeddings: np.ndarray) -> np.ndarray: - """ - Compute the L2-normalised mean of a set of embeddings. - - :param embeddings: Float array of shape (n, d) - :return: L2-normalised 1D array of shape (d,) - """ - mean = embeddings.mean(axis=0) - norm = np.linalg.norm(mean) - if norm == 0.0: - return mean - return mean / norm - - -def cosine_similarity(a: np.ndarray, b: np.ndarray) -> float: - """ - Compute cosine similarity between two L2-normalised vectors. - - Assumes both inputs are already L2-normalised, so this reduces - to a dot product. - - :param a: 1D array - :param b: 1D array of the same shape as `a` - :return: Cosine similarity as a float in [-1, 1] - """ - return float(np.dot(a, b)) - - -# ----------------------------------------------------------------------------- -# Calibration: self-similarity distribution -# ----------------------------------------------------------------------------- - - -def self_similarity_distribution( - ds: VoiceDataset, - model: SentenceTransformer, - *, - sample_size: int = CALIBRATION_DEFAULTS.sample_size, - num_iterations: int = CALIBRATION_DEFAULTS.num_iterations, - seed: int = CALIBRATION_DEFAULTS.seed, -) -> np.ndarray: - """ - Estimate the self-similarity distribution via resampling. - - For each iteration, we draw 2*`sample_size` unique examples from the - training split, split them into two non-overlapping groups of size - `sample_size`, and compute the cosine similarity between their - respective mean embeddings. - - :param ds: A VoiceDataset providing a `train` split - :param model: Pretrained SentenceTransformer model - :param sample_size: Size of each resampled subset - :param num_iterations: Number of Monte Carlo resampling iterations - :param seed: RNG seed for reproducibility - :return: 1D array of cosine similarity values of length `num_iterations` - """ - ds_train = ds.train - n_train = len(ds_train) - embeddings = embedding_matrix(ds_train, model) - - return bootstrap_null_distribution( - n_train, - lambda a, b: cosine_similarity( - mean_embedding(embeddings[a]), mean_embedding(embeddings[b]) - ), - sample_size=sample_size, - num_iterations=num_iterations, - seed=seed, - ) - - -# ----------------------------------------------------------------------------- -# Comparison result container -# ----------------------------------------------------------------------------- - - -@dataclass(frozen=True, slots=True) -class EmbeddingComparisonResult: - """ - Result for an embedding-based style comparison. - - .. attribute :: similarity - - Observed cosine similarity between the mean embedding of the - generated completions and the mean embedding of the true corpus. - - .. attribute :: percentile - - Calibrated percentile under the self-similarity distribution. - - .. attribute :: tail - - Right-tail extremeness (1 - percentile). - - .. attribute :: flag - - Boolean indicating percentile >= alpha, i.e. the observed similarity - is not unusually low relative to within-corpus pairs (the model is - indistinguishable from the reference under H0: distinguishable) - """ - - similarity: float - percentile: float - tail: float - flag: bool - - @property - def score(self) -> float: - """Alignment score in [0, 1]; 1 = perfectly indistinguishable.""" - return self.percentile - - -# ----------------------------------------------------------------------------- -# Comparison creation -# ----------------------------------------------------------------------------- - - -def make_embedding_comparison( - completions: Sequence[Example], - true_ds: VoiceDataset, - model: SentenceTransformer, - *, - alpha: float = COMPARISON_DEFAULTS.alpha, -) -> EmbeddingComparisonResult: - """ - Compare style of `completions` to the true corpus via embedding similarity. - - We compute the cosine similarity between the mean embedding of the - generated completions and the mean embedding of the true corpus on the - same split, then calibrate this against the self-similarity distribution - estimated from the true training split. - - H0: the model is stylistically distinguishable from the reference (its - similarity is unusually low). flag=True rejects this null; the observed - similarity is not in the low tail, concluding the model is - indistinguishable. - - :param completions: Sequence of examples to analyse - :param true_ds: VoiceDataset providing the true corpus - :param model: Pretrained SentenceTransformer model - :param alpha: Significance level for hypothesis testing - :return: EmbeddingComparisonResult - :raises ValueError: If completions is empty - """ - if not completions: - raise ValueError("completions must be non-empty.") - - # Split that observed completions are generated on - observed_split = completions[0].split - ds_observed_split = true_ds[observed_split] - - # Observed similarity - gen_embeddings = embedding_matrix(completions, model) - ref_embeddings = embedding_matrix(ds_observed_split, model) - - similarity = cosine_similarity( - mean_embedding(gen_embeddings), - mean_embedding(ref_embeddings), - ) - - # Self-similarity calibration distribution - self_dist = self_similarity_distribution(true_ds, model) - - p = calibrated_percentile(similarity, self_dist) - tail = 1.0 - p - flag = p >= alpha - - return EmbeddingComparisonResult( - similarity=similarity, - percentile=float(p), - tail=float(tail), - flag=flag, - ) diff --git a/src/voice/stylometry/metrics.py b/src/voice/stylometry/metrics.py index 71c1faa..3f4666c 100644 --- a/src/voice/stylometry/metrics.py +++ b/src/voice/stylometry/metrics.py @@ -29,7 +29,7 @@ from collections.abc import Callable from dataclasses import dataclass -from voice.stylometry._defaults import PARAMETER_DEFAULTS, MetricGroup +from voice._defaults import PARAMETER_DEFAULTS, MetricGroup from voice.stylometry._lexicons import FUNCTION_WORDS from voice.stylometry._sequence_stats import ( length_central_moment_from_units, @@ -254,7 +254,7 @@ def calculate_num_words(text: str) -> float: @metric( "hapax_legomena_ratio", - group=MetricGroup.LEGOMENA, + group=MetricGroup.VOCABULARY_RICHNESS, description="Ratio of hapax legomena", ) def calculate_hapax_legomena_ratio(text: str) -> float: @@ -282,7 +282,7 @@ def calculate_hapax_legomena_ratio(text: str) -> float: @metric( "dis_legomena_ratio", - group=MetricGroup.LEGOMENA, + group=MetricGroup.VOCABULARY_RICHNESS, description="Ratio of dis legomena", ) def calculate_dis_legomena_ratio(text: str) -> float: @@ -310,7 +310,7 @@ def calculate_dis_legomena_ratio(text: str) -> float: @metric( "tri_legomena_ratio", - group=MetricGroup.LEGOMENA, + group=MetricGroup.VOCABULARY_RICHNESS, description="Ratio of tri legomena", ) def calculate_tri_legomena_ratio(text: str) -> float: @@ -367,7 +367,7 @@ def calculate_function_word_ratio(text: str) -> float: @metric( "type_token_ratio", - group=MetricGroup.LEXICAL_RICHNESS, + group=MetricGroup.VOCABULARY_RICHNESS, description="Type-token ratio", ) def calculate_type_token_ratio(text: str) -> float: @@ -387,7 +387,7 @@ def calculate_type_token_ratio(text: str) -> float: @metric( "moving_avg_type_token_ratio", - group=MetricGroup.LEXICAL_RICHNESS, + group=MetricGroup.VOCABULARY_RICHNESS, description="Moving average type-token ratio", ) def calculate_moving_avg_type_token_ratio(text: str) -> float: diff --git a/src/voice/stylometry/plotting.py b/src/voice/stylometry/plotting.py index 1193fce..2d8e815 100644 --- a/src/voice/stylometry/plotting.py +++ b/src/voice/stylometry/plotting.py @@ -13,7 +13,7 @@ import numpy as np import seaborn as sns -from voice.stylometry._defaults import PLOTTING_DEFAULTS +from voice._defaults import PLOTTING_DEFAULTS def plot_kde( # noqa: CCR001 diff --git a/tests/voice/comparison/test_stylometric_comparison.py b/tests/voice/comparison/test_comparison.py similarity index 74% rename from tests/voice/comparison/test_stylometric_comparison.py rename to tests/voice/comparison/test_comparison.py index 4e28769..82666d3 100644 --- a/tests/voice/comparison/test_stylometric_comparison.py +++ b/tests/voice/comparison/test_comparison.py @@ -1,11 +1,11 @@ """ -Tests for voice.stylometry.comparison. +Tests for voice.comparison.comparison. Scope: - Distribution extraction - Self-Wasserstein calibration distribution - Empirical CDF calibrated percentile -- StylometricComparisonResults / StylometricComparisonEntry +- ComparisonResults / ComparisonEntry - Comparison creation """ @@ -19,19 +19,19 @@ import pytest from datasets import Dataset +from voice._defaults import MetricGroup from voice.comparison._utils import calibrated_percentile -from voice.comparison.stylometric_comparison import ( - StylometricComparisonEntry, - StylometricComparisonResults, - StylometricGroupEntry, - make_stylometric_comparison, +from voice.comparison.comparison import ( + ComparisonEntry, + ComparisonResults, + GroupEntry, + make_comparison, self_wasserstein_distribution, stylometric_distribution, ) from voice.datasets import VoiceDataset from voice.datasets._schema import Split from voice.datasets.dataset import Example, _PinnedDatasetSpec -from voice.stylometry._defaults import MetricGroup # ----------------------------------------------------------------------------- # Helpers @@ -89,14 +89,14 @@ def metric_registry(): ), "vowels": _Metric( fn=lambda s: float(sum(c in "aeiou" for c in s.lower())), - group=MetricGroup.LEXICAL_RICHNESS, + group=MetricGroup.VOCABULARY_RICHNESS, ), } @pytest.fixture() def patch_metrics(monkeypatch, metric_registry): - import voice.comparison.stylometric_comparison as comparison + import voice.comparison.comparison as comparison monkeypatch.setattr( comparison, "get_metrics", lambda: dict(metric_registry) @@ -247,7 +247,7 @@ def raising_metric(s: str) -> float: raise AssertionError("test split should not be touched") return float(len(s)) - import voice.comparison.stylometric_comparison as comparison + import voice.comparison.comparison as comparison monkeypatch.setattr( comparison, @@ -275,7 +275,7 @@ def raising_metric(s: str) -> float: def test_self_wasserstein_distribution_is_zero_for_constant_metric( patch_metrics, monkeypatch ): - import voice.comparison.stylometric_comparison as comparison + import voice.comparison.comparison as comparison monkeypatch.setattr( comparison, @@ -298,18 +298,17 @@ def test_self_wasserstein_distribution_is_zero_for_constant_metric( # ----------------------------------------------------------------------------- -# StylometricComparisonEntry / StylometricComparisonResults +# ComparisonEntry / ComparisonResults # ----------------------------------------------------------------------------- def test_comparison_entry_is_frozen_and_slots(): - e = StylometricComparisonEntry( + e = ComparisonEntry( metric="m", group=MetricGroup.TEXT_LENGTH, wasserstein=1.0, percentile=0.9, tail=0.1, - flag=True, ) assert type(e).__dataclass_params__.frozen is True with pytest.raises(dataclasses.FrozenInstanceError): @@ -317,11 +316,10 @@ def test_comparison_entry_is_frozen_and_slots(): def test_comparison_group_entry_is_frozen_and_slots(): - e = StylometricGroupEntry( + e = GroupEntry( group=MetricGroup.TEXT_LENGTH, avg_percentile=0.9, avg_tail=0.1, - flag=True, ) assert type(e).__dataclass_params__.frozen is True with pytest.raises(dataclasses.FrozenInstanceError): @@ -330,12 +328,12 @@ def test_comparison_group_entry_is_frozen_and_slots(): def test_comparison_results_rejects_unknown_metrics(patch_metrics): with pytest.raises(ValueError, match=r"Unknown metric"): - StylometricComparisonResults(metrics=("len", "nope")) + ComparisonResults(metrics=("len", "nope")) def test_comparison_results_properties_and_add_logic(patch_metrics): - # len → TEXT_LENGTH, vowels → LEXICAL_RICHNESS: 2 metrics, 2 groups - r = StylometricComparisonResults(metrics=("len", "vowels"), alpha=0.1) + # len → TEXT_LENGTH, vowels → VOCABULARY_RICHNESS: 2 metrics, 2 groups + r = ComparisonResults(metrics=("len", "vowels")) assert r.num_metrics == 2 @@ -347,24 +345,18 @@ def test_comparison_results_properties_and_add_logic(patch_metrics): assert isinstance(e.wasserstein, float) assert isinstance(e.percentile, float) assert e.tail == pytest.approx(0.1) - assert e.flag is True # 0.9 <= 1 - 0.1 r.add(metric="vowels", wasserstein=0.5, percentile=0.0) e2 = r.get("vowels") - assert e2.group == MetricGroup.LEXICAL_RICHNESS - assert e2.flag is True # 0.0 <= 0.9 + assert e2.group == MetricGroup.VOCABULARY_RICHNESS # Group-level checks ge_len = r.get_group(MetricGroup.TEXT_LENGTH) assert ge_len.avg_percentile == pytest.approx(0.9) assert ge_len.avg_tail == pytest.approx(0.1) - assert ge_len.flag is True # 0.9 <= 0.9 - ge_vowels = r.get_group(MetricGroup.LEXICAL_RICHNESS) + ge_vowels = r.get_group(MetricGroup.VOCABULARY_RICHNESS) assert ge_vowels.avg_percentile == pytest.approx(0.0) - assert ge_vowels.flag is True # 0.0 <= 0.9 - - assert r.all_pass is True d = r.as_dict() d.pop("len") @@ -372,46 +364,80 @@ def test_comparison_results_properties_and_add_logic(patch_metrics): def test_comparison_results_score_is_mean_group_tail(patch_metrics): - r = StylometricComparisonResults(metrics=("len", "vowels"), alpha=0.1) + r = ComparisonResults(metrics=("len", "vowels")) r.add(metric="len", wasserstein=1.0, percentile=0.8) r.add(metric="vowels", wasserstein=0.5, percentile=0.4) - # TEXT_LENGTH avg_tail = 0.2, LEXICAL_RICHNESS avg_tail = 0.6 → mean = 0.4 assert r.score == pytest.approx(0.4) -def test_comparison_results_all_pass_false_when_group_fails(patch_metrics): - r = StylometricComparisonResults(metrics=("len", "vowels"), alpha=0.1) - # percentile=1.0 > 1 - alpha=0.9 → flag=False → all_pass=False - r.add(metric="len", wasserstein=5.0, percentile=1.0) - r.add(metric="vowels", wasserstein=0.5, percentile=0.0) - assert r.get_group(MetricGroup.TEXT_LENGTH).flag is False - assert r.all_pass is False +def test_comparison_results_score_weights_groups_equally(patch_metrics): + # Both "len" metrics in TEXT_LENGTH, "vowels" in VOCABULARY_RICHNESS. + # Add a second TEXT_LENGTH metric to verify equal group weighting. + registry = { + "len": _Metric( + fn=lambda s: float(len(s)), group=MetricGroup.TEXT_LENGTH + ), + "len2": _Metric( + fn=lambda s: float(len(s)), group=MetricGroup.TEXT_LENGTH + ), + "vowels": _Metric( + fn=lambda s: float(sum(c in "aeiou" for c in s.lower())), + group=MetricGroup.VOCABULARY_RICHNESS, + ), + } + import voice.comparison.comparison as comparison + + comparison.get_metrics = lambda: dict(registry) # type: ignore[assignment] + + r = ComparisonResults(metrics=("len", "len2", "vowels")) + r.add(metric="len", wasserstein=1.0, percentile=0.0) # tail=1.0 + r.add(metric="len2", wasserstein=1.0, percentile=0.0) # tail=1.0 + r.add(metric="vowels", wasserstein=1.0, percentile=1.0) # tail=0.0 + assert r.score == pytest.approx(0.5) def test_comparison_results_add_rejects_metric_not_initialised(patch_metrics): - r = StylometricComparisonResults(metrics=("len",)) + r = ComparisonResults(metrics=("len",)) with pytest.raises(ValueError, match=r"not initialised"): r.add(metric="vowels", wasserstein=0.0, percentile=0.0) +def test_comparison_results_metric_tails(patch_metrics): + r = ComparisonResults(metrics=("len", "vowels")) + r.add(metric="len", wasserstein=1.0, percentile=0.8) + r.add(metric="vowels", wasserstein=0.5, percentile=0.3) + + tails = r.metric_tails + assert tails["len"] == pytest.approx(0.2) + assert tails["vowels"] == pytest.approx(0.7) + + +def test_comparison_results_group_tails(patch_metrics): + r = ComparisonResults(metrics=("len", "vowels")) + r.add(metric="len", wasserstein=1.0, percentile=0.8) + r.add(metric="vowels", wasserstein=0.5, percentile=0.3) + + tails = r.group_tails + assert tails[MetricGroup.TEXT_LENGTH] == pytest.approx(0.2) + assert tails[MetricGroup.VOCABULARY_RICHNESS] == pytest.approx(0.7) + + # ----------------------------------------------------------------------------- -# make_stylometric_comparison +# make_comparison # ----------------------------------------------------------------------------- -def test_make_stylometric_comparison_rejects_empty_completions(patch_metrics): +def test_make_comparison_rejects_empty_completions(patch_metrics): p = _pinned(splits=(Split.TRAIN,)) true_ds = VoiceDataset( datasets={Split.TRAIN: _canonical_hf_ds(["a"])}, spec=p ) with pytest.raises(ValueError, match=r"completions must be non-empty"): - make_stylometric_comparison([], true_ds) + make_comparison([], true_ds) -def test_make_stylometric_comparison_defaults_to_registry_metrics( - patch_metrics, -): +def test_make_comparison_defaults_to_registry_metrics(patch_metrics): p = _pinned(splits=(Split.TRAIN, Split.TEST)) true_ds = VoiceDataset( datasets={ @@ -425,14 +451,14 @@ def test_make_stylometric_comparison_defaults_to_registry_metrics( _ex(answer="world", split=Split.TEST), ] - import voice.comparison.stylometric_comparison as comparison + import voice.comparison.comparison as comparison def fake_self_dist(*_args, **_kwargs): return np.array([0.0, 1.0, 2.0, 3.0], dtype=float) comparison.self_wasserstein_distribution = fake_self_dist # type: ignore[assignment] - out = make_stylometric_comparison(completions, true_ds) + out = make_comparison(completions, true_ds) assert out.metrics == ("len", "vowels") for m in out.metrics: @@ -441,9 +467,7 @@ def fake_self_dist(*_args, **_kwargs): assert 0.0 <= entry.percentile <= 1.0 -def test_make_stylometric_comparison_respects_subset_metrics( - patch_metrics, monkeypatch -): +def test_make_comparison_respects_subset_metrics(patch_metrics, monkeypatch): p = _pinned(splits=(Split.TRAIN, Split.TEST)) true_ds = VoiceDataset( datasets={ @@ -454,7 +478,7 @@ def test_make_stylometric_comparison_respects_subset_metrics( ) completions = [_ex(answer="aaa", split=Split.TEST)] - import voice.comparison.stylometric_comparison as comparison + import voice.comparison.comparison as comparison monkeypatch.setattr( comparison, @@ -462,14 +486,14 @@ def test_make_stylometric_comparison_respects_subset_metrics( lambda *_args, **_kwargs: np.array([0.0, 1.0], dtype=float), ) - out = make_stylometric_comparison(completions, true_ds, metrics=["len"]) + out = make_comparison(completions, true_ds, metrics=["len"]) assert out.metrics == ("len",) assert "len" in out.as_dict() with pytest.raises(KeyError): _ = out.get("vowels") -def test_make_stylometric_comparison_uses_observed_split_from_completions( +def test_make_comparison_uses_observed_split_from_completions( patch_metrics, monkeypatch ): p = _pinned(splits=(Split.TRAIN, Split.TEST)) @@ -482,7 +506,7 @@ def test_make_stylometric_comparison_uses_observed_split_from_completions( ) completions = [_ex(answer="hello", split=Split.TEST)] - import voice.comparison.stylometric_comparison as comparison + import voice.comparison.comparison as comparison monkeypatch.setattr( comparison, @@ -499,6 +523,40 @@ def recording_getitem(self, key): monkeypatch.setattr(VoiceDataset, "__getitem__", recording_getitem) - _ = make_stylometric_comparison(completions, true_ds, metrics=["len"]) + _ = make_comparison(completions, true_ds, metrics=["len"]) assert Split.TEST in calls + + +def test_make_comparison_passes_sample_size_matching_completions( + patch_metrics, monkeypatch +): + p = _pinned(splits=(Split.TRAIN, Split.TEST)) + true_ds = VoiceDataset( + datasets={ + Split.TRAIN: _canonical_hf_ds([f"a{i}" for i in range(40)]), + Split.TEST: _canonical_hf_ds([f"b{i}" for i in range(10)]), + }, + spec=p, + ) + completions = [ + _ex(answer="hello", split=Split.TEST), + _ex(answer="world", split=Split.TEST), + _ex(answer="foo", split=Split.TEST), + ] + + import voice.comparison.comparison as comparison + + captured: list[int] = [] + + def spy_self_dist(*_args, sample_size: int = 0, **_kwargs): + captured.append(sample_size) + return np.array([0.0, 1.0], dtype=float) + + monkeypatch.setattr( + comparison, "self_wasserstein_distribution", spy_self_dist + ) + + make_comparison(completions, true_ds, metrics=["len"]) + + assert all(s == len(completions) for s in captured) diff --git a/tests/voice/comparison/test_embedding_comparison.py b/tests/voice/comparison/test_embedding_comparison.py deleted file mode 100644 index f9595cc..0000000 --- a/tests/voice/comparison/test_embedding_comparison.py +++ /dev/null @@ -1,470 +0,0 @@ -""" -Tests for voice.comparison.embedding_comparison. - -Scope: -- Embedding extraction (embedding_matrix, mean_embedding, cosine_similarity) -- Self-similarity calibration distribution -- EmbeddingComparisonResult dataclass invariants -- make_embedding_comparison end-to-end -""" - -from __future__ import annotations - -import dataclasses - -import numpy as np -import pytest -from datasets import Dataset - -from voice.comparison.embedding_comparison import ( - EmbeddingComparisonResult, - cosine_similarity, - embedding_matrix, - make_embedding_comparison, - mean_embedding, - self_similarity_distribution, -) -from voice.datasets import VoiceDataset -from voice.datasets._schema import Split -from voice.datasets.dataset import Example, _PinnedDatasetSpec - -# ----------------------------------------------------------------------------- -# Helpers / fixtures -# ----------------------------------------------------------------------------- - - -def _canonical_hf_ds(answers: list[str]) -> Dataset: - n = len(answers) - return Dataset.from_dict( - { - "system": [f"s{i}" for i in range(n)], - "question": [f"q{i}" for i in range(n)], - "answer": list(answers), - } - ) - - -def _pinned( - *, - repo_id: str = "ns/name", - revision: str = "a1b2c3d", - splits: tuple[Split, ...] = (Split.TRAIN,), -) -> _PinnedDatasetSpec: - return _PinnedDatasetSpec( - repo_id=repo_id, revision=revision, splits=splits - ) - - -def _ex( - *, - answer: str, - split: Split = Split.TRAIN, -) -> Example: - return Example( - system="s", - question="q", - answer=answer, - spec=_pinned(), - split=split, - ) - - -@pytest.fixture() -def fake_model(monkeypatch): - """ - A SentenceTransformer stand-in whose encode() returns deterministic - L2-normalised vectors derived from the input text lengths. - - Each text of length L gets a 4-dim vector proportional to - [L, 0, 0, 0], then L2-normalised → [1, 0, 0, 0] for any non-empty - string. Empty strings yield the zero vector. - """ - - class _FakeModel: - def encode( - self, - texts: list[str], - convert_to_numpy: bool = True, - normalize_embeddings: bool = True, - ) -> np.ndarray: - _ = convert_to_numpy - __ = normalize_embeddings - out = np.zeros((len(texts), 4), dtype=np.float32) - for i, t in enumerate(texts): - v = np.array([float(len(t)), 1.0, 0.0, 0.0], dtype=np.float32) - norm = np.linalg.norm(v) - if norm > 0: - v /= norm - out[i] = v - return out - - return _FakeModel() - - -# ----------------------------------------------------------------------------- -# embedding_matrix -# ----------------------------------------------------------------------------- - - -def test_embedding_matrix_shape(fake_model): - examples = [_ex(answer="hello"), _ex(answer="world"), _ex(answer="foo")] - mat = embedding_matrix(examples, fake_model) - assert mat.shape == (3, 4) - assert mat.dtype == np.float32 - - -def test_embedding_matrix_rows_are_unit_vectors(fake_model): - examples = [_ex(answer="abc"), _ex(answer="xy")] - mat = embedding_matrix(examples, fake_model) - norms = np.linalg.norm(mat, axis=1) - np.testing.assert_allclose(norms, 1.0, atol=1e-6) - - -def test_embedding_matrix_uses_answer_field(fake_model, monkeypatch): - seen: list[list[str]] = [] - orig_encode = fake_model.encode - - def recording_encode(texts, **kwargs): - seen.append(list(texts)) - return orig_encode(texts, **kwargs) - - monkeypatch.setattr(fake_model, "encode", recording_encode) - - examples = [_ex(answer="one"), _ex(answer="two")] - embedding_matrix(examples, fake_model) - - assert seen == [["one", "two"]] - - -# ----------------------------------------------------------------------------- -# mean_embedding -# ----------------------------------------------------------------------------- - - -def test_mean_embedding_returns_unit_vector(): - mat = np.array([[1.0, 0.0], [0.0, 1.0]], dtype=np.float32) - result = mean_embedding(mat) - assert result.shape == (2,) - np.testing.assert_allclose(np.linalg.norm(result), 1.0, atol=1e-6) - - -def test_mean_embedding_zero_returns_zero(): - mat = np.zeros((3, 4), dtype=np.float32) - result = mean_embedding(mat) - np.testing.assert_array_equal(result, np.zeros(4)) - - -def test_mean_embedding_single_row_returns_same_vector(): - v = np.array([[0.6, 0.8]], dtype=np.float32) - result = mean_embedding(v) - np.testing.assert_allclose(result, [0.6, 0.8], atol=1e-6) - - -# ----------------------------------------------------------------------------- -# cosine_similarity -# ----------------------------------------------------------------------------- - - -def test_cosine_similarity_identical_vectors(): - v = np.array([1.0, 0.0, 0.0]) - assert cosine_similarity(v, v) == pytest.approx(1.0) - - -def test_cosine_similarity_orthogonal_vectors(): - a = np.array([1.0, 0.0]) - b = np.array([0.0, 1.0]) - assert cosine_similarity(a, b) == pytest.approx(0.0) - - -def test_cosine_similarity_opposite_vectors(): - a = np.array([1.0, 0.0]) - b = np.array([-1.0, 0.0]) - assert cosine_similarity(a, b) == pytest.approx(-1.0) - - -def test_cosine_similarity_returns_float(): - a = np.array([1.0, 0.0]) - result = cosine_similarity(a, a) - assert isinstance(result, float) - - -# ----------------------------------------------------------------------------- -# self_similarity_distribution -# ----------------------------------------------------------------------------- - - -def test_self_similarity_distribution_shape(fake_model): - p = _pinned(splits=(Split.TRAIN,)) - vd = VoiceDataset( - datasets={ - Split.TRAIN: _canonical_hf_ds([f"text{i}" for i in range(20)]) - }, - spec=p, - ) - out = self_similarity_distribution( - vd, fake_model, sample_size=5, num_iterations=10, seed=42 - ) - assert out.shape == (10,) - assert out.dtype == float - - -def test_self_similarity_distribution_is_deterministic_given_seed(fake_model): - p = _pinned(splits=(Split.TRAIN,)) - vd = VoiceDataset( - datasets={Split.TRAIN: _canonical_hf_ds([f"t{i}" for i in range(20)])}, - spec=p, - ) - out1 = self_similarity_distribution( - vd, fake_model, sample_size=5, num_iterations=15, seed=7 - ) - out2 = self_similarity_distribution( - vd, fake_model, sample_size=5, num_iterations=15, seed=7 - ) - assert np.array_equal(out1, out2) - - out3 = self_similarity_distribution( - vd, fake_model, sample_size=5, num_iterations=15, seed=99 - ) - assert not np.array_equal(out1, out3) - - -def test_self_similarity_distribution_uses_training_split_only( - fake_model, monkeypatch -): - encode_calls: list[list[str]] = [] - orig_encode = fake_model.encode - - def recording_encode(texts, **kwargs): - encode_calls.append(list(texts)) - return orig_encode(texts, **kwargs) - - monkeypatch.setattr(fake_model, "encode", recording_encode) - - p = _pinned(splits=(Split.TRAIN, Split.TEST)) - vd = VoiceDataset( - datasets={ - Split.TRAIN: _canonical_hf_ds([f"train{i}" for i in range(20)]), - Split.TEST: _canonical_hf_ds(["SHOULD_NOT_APPEAR"]), - }, - spec=p, - ) - self_similarity_distribution( - vd, fake_model, sample_size=5, num_iterations=5 - ) - - all_texts = [t for call in encode_calls for t in call] - assert "SHOULD_NOT_APPEAR" not in all_texts - - -def test_self_similarity_distribution_values_in_range(fake_model): - p = _pinned(splits=(Split.TRAIN,)) - vd = VoiceDataset( - datasets={ - Split.TRAIN: _canonical_hf_ds([f"word{i}" for i in range(30)]) - }, - spec=p, - ) - out = self_similarity_distribution( - vd, fake_model, sample_size=5, num_iterations=20 - ) - assert np.all(out >= -1.0 - 1e-5) and np.all(out <= 1.0 + 1e-5) - - -# ----------------------------------------------------------------------------- -# EmbeddingComparisonResult -# ----------------------------------------------------------------------------- - - -def test_embedding_comparison_result_is_frozen(): - r = EmbeddingComparisonResult( - similarity=0.9, percentile=0.8, tail=0.2, flag=False - ) - assert type(r).__dataclass_params__.frozen is True - with pytest.raises(dataclasses.FrozenInstanceError): - setattr(r, "similarity", 0.5) # noqa: B010 - - -def test_embedding_comparison_result_tail_is_complement_of_percentile(): - r = EmbeddingComparisonResult( - similarity=0.7, percentile=0.6, tail=0.4, flag=False - ) - assert r.tail == pytest.approx(1.0 - r.percentile) - - -def test_embedding_comparison_result_score_equals_percentile(): - r = EmbeddingComparisonResult( - similarity=0.7, percentile=0.6, tail=0.4, flag=False - ) - assert r.score == pytest.approx(r.percentile) - - -# ----------------------------------------------------------------------------- -# make_embedding_comparison -# ----------------------------------------------------------------------------- - - -def test_make_embedding_comparison_rejects_empty_completions(fake_model): - p = _pinned(splits=(Split.TRAIN,)) - true_ds = VoiceDataset( - datasets={Split.TRAIN: _canonical_hf_ds(["a"])}, spec=p - ) - with pytest.raises(ValueError, match=r"completions must be non-empty"): - make_embedding_comparison([], true_ds, fake_model) - - -def test_make_embedding_comparison_returns_result_type( - fake_model, monkeypatch -): - import voice.comparison.embedding_comparison as ec - - monkeypatch.setattr( - ec, - "self_similarity_distribution", - lambda *_a, **_kw: np.array([0.5, 0.6, 0.7, 0.8], dtype=float), - ) - - p = _pinned(splits=(Split.TRAIN, Split.TEST)) - true_ds = VoiceDataset( - datasets={ - Split.TRAIN: _canonical_hf_ds([f"train{i}" for i in range(20)]), - Split.TEST: _canonical_hf_ds([f"ref{i}" for i in range(5)]), - }, - spec=p, - ) - completions = [_ex(answer="hello", split=Split.TEST)] - - result = make_embedding_comparison(completions, true_ds, fake_model) - - assert isinstance(result, EmbeddingComparisonResult) - assert isinstance(result.similarity, float) - assert isinstance(result.percentile, float) - assert isinstance(result.tail, float) - assert isinstance(result.flag, bool) - - -def test_make_embedding_comparison_percentile_and_tail_are_consistent( - fake_model, monkeypatch -): - import voice.comparison.embedding_comparison as ec - - monkeypatch.setattr( - ec, - "self_similarity_distribution", - lambda *_a, **_kw: np.array([0.0, 0.5, 1.0], dtype=float), - ) - - p = _pinned(splits=(Split.TRAIN, Split.TEST)) - true_ds = VoiceDataset( - datasets={ - Split.TRAIN: _canonical_hf_ds([f"t{i}" for i in range(20)]), - Split.TEST: _canonical_hf_ds(["ref"]), - }, - spec=p, - ) - completions = [_ex(answer="hello", split=Split.TEST)] - - result = make_embedding_comparison(completions, true_ds, fake_model) - - assert result.tail == pytest.approx(1.0 - result.percentile) - assert 0.0 <= result.percentile <= 1.0 - assert 0.0 <= result.tail <= 1.0 - - -def test_make_embedding_comparison_flag_false_when_low_percentile( - fake_model, monkeypatch -): - """ - Flag should be False when similarity is low; model is distinguishable. - """ - import voice.comparison.embedding_comparison as ec - - # Self-dist is always high → observed similarity will be low percentile - monkeypatch.setattr( - ec, - "self_similarity_distribution", - lambda *_a, **_kw: np.ones(100, dtype=float), - ) - - p = _pinned(splits=(Split.TRAIN, Split.TEST)) - true_ds = VoiceDataset( - datasets={ - Split.TRAIN: _canonical_hf_ds([f"t{i}" for i in range(20)]), - Split.TEST: _canonical_hf_ds(["ref"]), - }, - spec=p, - ) - completions = [_ex(answer="x", split=Split.TEST)] - - result = make_embedding_comparison( - completions, true_ds, fake_model, alpha=0.01 - ) - assert result.flag is False - - -def test_make_embedding_comparison_flag_true_when_high_percentile( - fake_model, monkeypatch -): - """ - Flag should be True when similarity is high; model is indistinguishable. - """ - import voice.comparison.embedding_comparison as ec - - # Self-dist is always low → observed similarity will be high percentile - monkeypatch.setattr( - ec, - "self_similarity_distribution", - lambda *_a, **_kw: np.zeros(100, dtype=float), - ) - - p = _pinned(splits=(Split.TRAIN, Split.TEST)) - true_ds = VoiceDataset( - datasets={ - Split.TRAIN: _canonical_hf_ds([f"t{i}" for i in range(20)]), - Split.TEST: _canonical_hf_ds(["ref"]), - }, - spec=p, - ) - completions = [_ex(answer="hello world", split=Split.TEST)] - - result = make_embedding_comparison( - completions, true_ds, fake_model, alpha=0.05 - ) - assert result.flag is True - - -def test_make_embedding_comparison_uses_observed_split( - fake_model, monkeypatch -): - """ - The reference embeddings should come from the same split as completions. - """ - import voice.comparison.embedding_comparison as ec - - monkeypatch.setattr( - ec, - "self_similarity_distribution", - lambda *_a, **_kw: np.array([0.5], dtype=float), - ) - - calls: list[object] = [] - orig_getitem = VoiceDataset.__getitem__ - - def recording_getitem(self, key): - calls.append(key) - return orig_getitem(self, key) - - monkeypatch.setattr(VoiceDataset, "__getitem__", recording_getitem) - - p = _pinned(splits=(Split.TRAIN, Split.TEST)) - true_ds = VoiceDataset( - datasets={ - Split.TRAIN: _canonical_hf_ds([f"t{i}" for i in range(20)]), - Split.TEST: _canonical_hf_ds([f"r{i}" for i in range(5)]), - }, - spec=p, - ) - completions = [_ex(answer="hello", split=Split.TEST)] - - make_embedding_comparison(completions, true_ds, fake_model) - - assert Split.TEST in calls diff --git a/tests/voice/stylometry/test_metrics.py b/tests/voice/stylometry/test_metrics.py index c4f49ac..69fc947 100644 --- a/tests/voice/stylometry/test_metrics.py +++ b/tests/voice/stylometry/test_metrics.py @@ -17,7 +17,7 @@ import pytest import voice.stylometry.metrics as metrics -from voice.stylometry._defaults import MetricGroup +from voice._defaults import MetricGroup # ----------------------------------------------------------------------------- # Helpers @@ -153,12 +153,12 @@ def test_metric_groups_cover_all_enum_values(): "skew_word_length": MetricGroup.WORD_LENGTH_DISTRIBUTION, "kurtosis_word_length": MetricGroup.WORD_LENGTH_DISTRIBUTION, "num_words": MetricGroup.TEXT_LENGTH, - "hapax_legomena_ratio": MetricGroup.LEGOMENA, - "dis_legomena_ratio": MetricGroup.LEGOMENA, - "tri_legomena_ratio": MetricGroup.LEGOMENA, + "hapax_legomena_ratio": MetricGroup.VOCABULARY_RICHNESS, + "dis_legomena_ratio": MetricGroup.VOCABULARY_RICHNESS, + "tri_legomena_ratio": MetricGroup.VOCABULARY_RICHNESS, "function_word_ratio": MetricGroup.FUNCTION_WORDS, - "type_token_ratio": MetricGroup.LEXICAL_RICHNESS, - "moving_avg_type_token_ratio": MetricGroup.LEXICAL_RICHNESS, + "type_token_ratio": MetricGroup.VOCABULARY_RICHNESS, + "moving_avg_type_token_ratio": MetricGroup.VOCABULARY_RICHNESS, "char_3gram_type_token_ratio": MetricGroup.CHAR_NGRAM_DIVERSITY, "char_4gram_type_token_ratio": MetricGroup.CHAR_NGRAM_DIVERSITY, "char_5gram_type_token_ratio": MetricGroup.CHAR_NGRAM_DIVERSITY,