File size: 7,132 Bytes
456264c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
"""Leakage-resistant genre benchmark with an elementary composition baseline."""

from __future__ import annotations

import hashlib
import json
from collections.abc import Sequence
from pathlib import Path
from typing import Any

import numpy as np

SOLVENTS = {"64-17-5", "ethanol", "water", "7732-18-5"}


def _identity(component: dict[str, Any]) -> str:
    return str(component.get("cas") or component.get("smiles") or component.get("name") or "").strip().lower()


def active_compounds(record: dict[str, Any]) -> set[str]:
    """Return non-solvent component identities used for leakage accounting."""
    return {identity for component in record.get("formula", []) if (identity := _identity(component)) and identity not in SOLVENTS}


def formula_fingerprint(record: dict[str, Any], precision: int = 6) -> str:
    """Stable, order-invariant identity for exact/near-exact formula duplicates."""
    parts = sorted(
        (_identity(component), round(float(component.get("weight_fraction", 0.0)), precision))
        for component in record.get("formula", [])
    )
    return hashlib.sha256(json.dumps(parts, separators=(",", ":")).encode()).hexdigest()


def audit_split(train: Sequence[dict[str, Any]], test: Sequence[dict[str, Any]]) -> dict[str, Any]:
    """Fail closed on molecule or formula leakage between train and test."""
    train_compounds = set().union(*(active_compounds(row) for row in train)) if train else set()
    test_compounds = set().union(*(active_compounds(row) for row in test)) if test else set()
    train_formulas = {formula_fingerprint(row) for row in train}
    test_formulas = {formula_fingerprint(row) for row in test}
    compound_overlap = sorted(train_compounds & test_compounds)
    formula_overlap = sorted(train_formulas & test_formulas)
    result = {
        "passed": not compound_overlap and not formula_overlap,
        "train_records": len(train),
        "test_records": len(test),
        "compound_overlap": compound_overlap,
        "formula_overlap": formula_overlap,
    }
    if not result["passed"]:
        raise ValueError(
            f"leaky split: {len(compound_overlap)} shared compounds, "
            f"{len(formula_overlap)} shared formulas"
        )
    return result


def composition_features(records: Sequence[dict[str, Any]]) -> np.ndarray:
    """Compute label-free mass-distribution statistics; no identities are encoded."""
    features: list[list[float]] = []
    for row in records:
        formula = row.get("formula", [])
        solvent_mass = sum(float(c.get("weight_fraction", 0.0)) for c in formula if _identity(c) in SOLVENTS)
        weights = np.asarray(
            [float(c.get("weight_fraction", 0.0)) for c in formula if _identity(c) not in SOLVENTS], dtype=float
        )
        weights = weights[weights > 0]
        if not len(weights):
            weights = np.zeros(1)
        total = float(weights.sum())
        proportions = weights / total if total else weights
        sorted_weights = np.sort(weights)[::-1]
        top = np.pad(sorted_weights, (0, max(0, 3 - len(sorted_weights))))[:3]
        entropy = float(-(proportions[proportions > 0] * np.log(proportions[proportions > 0])).sum())
        hhi = float(np.square(proportions).sum())
        q25, median, q75 = np.quantile(weights, [0.25, 0.5, 0.75])
        features.append([
            len(weights), solvent_mass, total, float(weights.mean()), float(weights.std()),
            float(q25), float(median), float(q75), float(weights.max()), entropy, hhi,
            1.0 / hhi if hhi else 0.0, *top.tolist(),
        ])
    return np.asarray(features, dtype=float)


def nearest_centroid_predict(train_x: np.ndarray, train_y: np.ndarray, test_x: np.ndarray) -> np.ndarray:
    """Standardized nearest-centroid classifier with training-only statistics."""
    mean = train_x.mean(axis=0)
    scale = train_x.std(axis=0)
    scale[scale == 0] = 1.0
    train_z, test_z = (train_x - mean) / scale, (test_x - mean) / scale
    labels = np.unique(train_y)
    centroids = np.vstack([train_z[train_y == label].mean(axis=0) for label in labels])
    distances = np.square(test_z[:, None, :] - centroids[None, :, :]).sum(axis=2)
    return labels[np.argmin(distances, axis=1)]


def paired_bootstrap_delta(
    learned_correct: np.ndarray,
    baseline_correct: np.ndarray,
    *,
    seed: int = 0,
    samples: int = 10_000,
) -> dict[str, float]:
    """Bootstrap the paired accuracy advantage on the identical held-out rows."""
    if learned_correct.shape != baseline_correct.shape or not learned_correct.size:
        raise ValueError("paired non-empty correctness arrays are required")
    differences = learned_correct.astype(float) - baseline_correct.astype(float)
    rng = np.random.default_rng(seed)
    draws = rng.choice(differences, size=(samples, len(differences)), replace=True).mean(axis=1)
    low, high = np.quantile(draws, [0.025, 0.975])
    return {"delta": float(differences.mean()), "ci95_low": float(low), "ci95_high": float(high)}


def benchmark_split(
    train: Sequence[dict[str, Any]],
    test: Sequence[dict[str, Any]],
    learned_train: np.ndarray,
    learned_test: np.ndarray,
    *,
    seed: int = 0,
    bootstrap_samples: int = 10_000,
) -> dict[str, Any]:
    """Compare learned representations with the composition baseline on one split."""
    leakage = audit_split(train, test)
    train_y = np.asarray([row["genre"] for row in train])
    test_y = np.asarray([row["genre"] for row in test])
    baseline_pred = nearest_centroid_predict(composition_features(train), train_y, composition_features(test))
    learned_pred = nearest_centroid_predict(np.asarray(learned_train), train_y, np.asarray(learned_test))
    baseline_correct = baseline_pred == test_y
    learned_correct = learned_pred == test_y
    return {
        "leakage_audit": leakage,
        "n_test": len(test),
        "baseline_accuracy": float(baseline_correct.mean()),
        "learned_accuracy": float(learned_correct.mean()),
        "paired_advantage": paired_bootstrap_delta(
            learned_correct, baseline_correct, seed=seed, samples=bootstrap_samples
        ),
    }


def decide(results: Sequence[dict[str, Any]], *, minimum_splits: int = 3, margin: float = 0.0) -> dict[str, Any]:
    """Apply the preregistered rule: every split CI must clear the margin."""
    valid = [r for r in results if r.get("leakage_audit", {}).get("passed")]
    beats = len(valid) >= minimum_splits and all(r["paired_advantage"]["ci95_low"] > margin for r in valid)
    return {
        "decision": "learned_representation_supported" if beats else "negative_result",
        "reason": (
            "learned representation reliably beats composition baseline"
            if beats else "learned representation did not reliably beat composition baseline"
        ),
        "valid_splits": len(valid),
        "minimum_splits": minimum_splits,
        "required_margin": margin,
    }


def load_jsonl(path: str | Path) -> list[dict[str, Any]]:
    with Path(path).open(encoding="utf-8") as handle:
        return [json.loads(line) for line in handle if line.strip()]