File size: 13,651 Bytes
3383ed6
 
4f26cce
3383ed6
aa64aba
3383ed6
 
 
aa64aba
3383ed6
 
 
 
 
4f26cce
3383ed6
4f26cce
3383ed6
 
 
 
 
 
c9282b0
4f26cce
 
aa64aba
 
 
 
3383ed6
4f26cce
 
3383ed6
 
 
 
 
 
c9282b0
224d843
3383ed6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
aa64aba
 
 
 
 
 
 
 
224d843
2671b56
3383ed6
 
2671b56
3383ed6
19ac1ec
3383ed6
224d843
3383ed6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
224d843
 
aa64aba
3383ed6
 
 
224d843
 
3383ed6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
aa64aba
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3383ed6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
aa64aba
3383ed6
 
aa64aba
3383ed6
 
aa64aba
 
3383ed6
 
 
 
aa64aba
3383ed6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
aa64aba
3383ed6
 
 
 
 
 
 
 
 
 
 
 
 
 
aa64aba
 
 
 
 
 
3383ed6
 
 
 
 
 
 
 
 
 
 
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
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
#!/usr/bin/env python3
from __future__ import annotations

import argparse
import hashlib
import json
from pathlib import Path
from typing import Any
from datetime import datetime, timezone

import numpy as np
import torch
from scipy.spatial.distance import pdist
from scipy.stats import spearmanr
from torch.utils.data import DataLoader

from pino.heads import PIMTHeads
from pino.pimt_model import (
    DEFAULT_EMBEDDING_DIM,
    FragranceTrajectoryDataset,
    PhysicsInformedMixtureTransformer,
)
from pino.train import pad_trajectory_collate
from pino.upload_data import DEFAULT_TRAIN_RATIO, create_molecule_disjoint_split_from_records


PYRFUME_SOURCE = "pyrfume-cas"
SYNTHETIC_SOURCES = {"openpom-predicted", "volatility-fallback"}


DEFAULT_GENRES = ["citrus_cologne", "fougere", "floral_woody", "amber_oriental", "wildcard"]


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="Evaluate a PINO PIMT checkpoint.")
    parser.add_argument("--checkpoint", default="models/pimt_v6.pt", help="Local checkpoint path.")
    parser.add_argument("--data", default="data/empirical_dataset_v8.jsonl", help="Evaluation dataset JSONL.")
    parser.add_argument("--batch-size", type=int, default=8)
    parser.add_argument("--seed", type=int, default=42)
    parser.add_argument("--train-ratio", type=float, default=DEFAULT_TRAIN_RATIO, help="Molecule-disjoint train molecule ratio used to rebuild the 15% holdout validation set.")
    parser.add_argument("--structural-source", choices=["morgan", "openpom_256", "pom_alltags"], default="morgan", help="Structural embedding block matching the checkpoint's training arm.")
    parser.add_argument("--output", default="metrics_v6_current.json", help="Metrics JSON output path.")
    parser.add_argument(
        "--hf-repo",
        default=None,
        help="Optional HF model repo to download --hf-filename from before evaluation.",
    )
    parser.add_argument("--hf-filename", default=None, help="Checkpoint filename in --hf-repo.")
    parser.add_argument("--hf-token", default=None, help="HF token override; defaults to cached auth/env.")
    return parser.parse_args()


def resolve_checkpoint(args: argparse.Namespace) -> Path:
    checkpoint = Path(args.checkpoint)
    if checkpoint.exists():
        return checkpoint
    if not args.hf_repo:
        raise FileNotFoundError(f"Checkpoint not found: {checkpoint}")

    from huggingface_hub import hf_hub_download

    filename = args.hf_filename or checkpoint.name
    downloaded = hf_hub_download(
        repo_id=args.hf_repo,
        filename=filename,
        token=args.hf_token,
    )
    return Path(downloaded)


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


def sha256_file(path: str | Path) -> str:
    digest = hashlib.sha256()
    with Path(path).open("rb") as f:
        for chunk in iter(lambda: f.read(1024 * 1024), b""):
            digest.update(chunk)
    return digest.hexdigest()


def make_validation_dataset(records: list[dict[str, Any]], *, train_ratio: float, seed: int, structural_source: str = "morgan", objective_dim: int = 138) -> FragranceTrajectoryDataset:
    split = create_molecule_disjoint_split_from_records(records, train_ratio=train_ratio, seed=seed)
    genre_map = {genre: idx for idx, genre in enumerate(DEFAULT_GENRES)}
    return FragranceTrajectoryDataset(
        records=split["validation"],
        use_embedding_fallback=True,
        structural_source=structural_source,
        genre_map=genre_map,
        objective_dim=objective_dim,
    )


def load_vocab(path: str | Path = "data/pyrfume_vocabulary.json") -> list[str]:
    with Path(path).open("r", encoding="utf-8") as f:
        raw = json.load(f)
    if isinstance(raw, dict) and isinstance(raw.get("vocabulary"), list):
        return [str(v) for v in raw["vocabulary"]]
    if isinstance(raw, dict):
        return [str(raw[k]) for k in sorted(raw)]
    if isinstance(raw, list):
        return [str(v) for v in raw]
    return []


def infer_model_dims(state_dict: dict[str, torch.Tensor]) -> dict[str, int]:
    input_weight = state_dict["input_proj.weight"]
    hidden_dim, embedding_dim = input_weight.shape
    state_dim = state_dict["gating.physics_scaler"].shape[0]
    num_layers = max(
        int(k.split(".")[2])
        for k in state_dict
        if k.startswith("encoder.layers.") and k.endswith(".self_attn.in_proj_weight")
    ) + 1
    in_proj_rows = state_dict["encoder.layers.0.self_attn.in_proj_weight"].shape[0]
    num_heads = 8 if hidden_dim % 8 == 0 and hidden_dim >= 512 else 4
    if in_proj_rows != hidden_dim * 3:
        raise ValueError("Unexpected transformer attention weight shape")
    return {
        "hidden_dim": hidden_dim,
        "embedding_dim": embedding_dim,
        "state_dim": state_dim,
        "num_layers": num_layers,
        "num_heads": num_heads,
    }


def evaluate(args: argparse.Namespace) -> dict[str, Any]:
    torch.manual_seed(args.seed)
    np.random.seed(args.seed)

    checkpoint_path = resolve_checkpoint(args)
    ckpt = torch.load(checkpoint_path, map_location="cpu", weights_only=False)
    model_state = ckpt["model_state_dict"]
    dims = infer_model_dims(model_state)

    model = PhysicsInformedMixtureTransformer(
        embedding_dim=dims["embedding_dim"],
        state_dim=dims["state_dim"],
        hidden_dim=dims["hidden_dim"],
        num_heads=dims["num_heads"],
        num_layers=dims["num_layers"],
    )
    model.load_state_dict(model_state)
    model.eval()

    objective_dim = 575 if args.structural_source == "pom_alltags" else 138
    heads = PIMTHeads(hidden_dim=dims["hidden_dim"], objective_dim=objective_dim)
    heads_missing, heads_unexpected = heads.load_state_dict(ckpt["heads_state_dict"], strict=False)
    heads.eval()

    records = read_records(args.data)
    engine_source = "openpom_256" if args.structural_source == "pom_alltags" else args.structural_source
    val_dataset = make_validation_dataset(records, train_ratio=args.train_ratio, seed=args.seed, structural_source=engine_source, objective_dim=objective_dim)
    val_loader = DataLoader(
        val_dataset,
        batch_size=args.batch_size,
        shuffle=False,
        collate_fn=pad_trajectory_collate,
    )

    all_preds: list[np.ndarray] = []
    all_targets: list[np.ndarray] = []
    with torch.no_grad():
        for batch in val_loader:
            latent = model(
                batch["tokens"],
                batch["physics"],
                src_key_padding_mask=batch["src_key_padding_mask"],
            )
            output = heads(latent, batch["physics"], batch["src_key_padding_mask"])
            all_preds.append(output["objective"].numpy())
            all_targets.append(batch["target_obj"].numpy())

    preds = np.concatenate(all_preds, axis=0)
    targets = np.concatenate(all_targets, axis=0)

    tier_names = ["top", "middle", "base"]
    tier_metrics: dict[str, Any] = {}
    for tier_idx, tier_name in enumerate(tier_names):
        pred_tier = preds[:, tier_idx, :]
        target_tier = targets[:, tier_idx, :]
        sims = 1.0 - pdist(pred_tier, metric="cosine")
        rho_list = []
        for idx in range(len(pred_tier)):
            rho, _ = spearmanr(pred_tier[idx], target_tier[idx])
            if rho is not None and not np.isnan(rho):
                rho_list.append(float(rho))
        tier_metrics[tier_name] = {
            "mse": float(((pred_tier - target_tier) ** 2).mean()),
            "mae": float(np.abs(pred_tier - target_tier).mean()),
            "prediction_variance": float(pred_tier.var()),
            "target_variance": float(target_tier.var()),
            "mean_pairwise_cosine_similarity": float(sims.mean()) if len(sims) else None,
            "mean_spearman_rho": float(np.mean(rho_list)) if rho_list else None,
        }

    flat_rho, flat_p = spearmanr(preds.reshape(-1), targets.reshape(-1))
    profile_pred = preds.mean(axis=1)
    profile_target = targets.mean(axis=1)
    profile_rho, profile_p = spearmanr(profile_pred.reshape(-1), profile_target.reshape(-1))
    provenance_metrics: dict[str, Any] = {}
    if val_dataset.records and any("pyramid_target_provenance" in r for r in val_dataset.records):
        provenance = np.array([
            r.get("pyramid_target_provenance", [["none"] * 138 for _ in range(3)])
            for r in val_dataset.records
        ], dtype=object)
        active = targets > 0
        provenance_slices = {
            "pyrfume_cas_only": active & (provenance == PYRFUME_SOURCE),
            "openpom_or_volatility_only": active & np.isin(provenance, list(SYNTHETIC_SOURCES)),
            "full_reconstructed": active,
        }
        for name, mask in provenance_slices.items():
            n = int(mask.sum())
            if n >= 3:
                labelled_cells = mask.any(axis=0)
                sliced_preds = preds[:, labelled_cells]
                sliced_targets = np.where(mask, targets, 0.0)[:, labelled_cells]
                rho, p_value = spearmanr(sliced_preds.reshape(-1), sliced_targets.reshape(-1))
                provenance_metrics[name] = {
                    "n_positive_values": n,
                    "n_tier_descriptor_cells": int(labelled_cells.sum()),
                    "spearman_rho": float(rho),
                    "spearman_p_value": float(p_value),
                }
            else:
                provenance_metrics[name] = {
                    "n_positive_values": n,
                    "n_tier_descriptor_cells": 0,
                    "spearman_rho": None,
                    "spearman_p_value": None,
                }
        synth_n = provenance_metrics["openpom_or_volatility_only"]["n_positive_values"]
        full_n = provenance_metrics["full_reconstructed"]["n_positive_values"]
        provenance_metrics["interpretation"] = (
            "openpom-predicted + volatility-fallback are the majority of active labels; full-set rho is product/internal model-agreement, not a rigorous chemistry result."
            if full_n and synth_n / full_n > 0.5
            else "full-set rho is still split by provenance; only pyrfume-cas-only is suitable as the closest trustworthy profile subset."
        )

    vocab = load_vocab()
    top_descriptors: dict[str, Any] = {}
    for tier_idx, tier_name in enumerate(tier_names):
        pred_mean = preds[:, tier_idx, :].mean(axis=0)
        target_mean = targets[:, tier_idx, :].mean(axis=0)
        pred_top = np.argsort(pred_mean)[-5:][::-1]
        target_top = np.argsort(target_mean)[-5:][::-1]
        top_descriptors[tier_name] = {
            "predicted": [
                {"descriptor": vocab[i] if i < len(vocab) else f"d{i}", "score": float(pred_mean[i])}
                for i in pred_top
            ],
            "target": [
                {"descriptor": vocab[i] if i < len(vocab) else f"d{i}", "score": float(target_mean[i])}
                for i in target_top
            ],
        }

    return {
        "generated_at_utc": datetime.now(timezone.utc).isoformat().replace("+00:00", "Z"),
        "checkpoint": {
            "path": str(checkpoint_path),
            "sha256": sha256_file(checkpoint_path),
            "epoch": ckpt.get("epoch"),
            "val_loss": ckpt.get("val_loss"),
            "missing_head_keys_on_load": list(heads_missing),
            "unexpected_head_keys_on_load": list(heads_unexpected),
            **dims,
        },
        "data": {
            "path": args.data,
            "sha256": sha256_file(args.data),
            "n_records": len(records),
            "n_validation_samples": len(val_dataset),
        },
        "prediction": {
            "profile_mse": float(((profile_pred - profile_target) ** 2).mean()),
            "profile_mae": float(np.abs(profile_pred - profile_target).mean()),
            "profile_prediction_variance": float(profile_pred.var()),
            "tier_prediction_variance": float(preds.var()),
            "tier_metrics": tier_metrics,
        },
        "spearman": {
            "flat_spearman_rho": float(flat_rho),
            "flat_spearman_p_value": float(flat_p),
            "profile_spearman_rho": float(profile_rho),
            "profile_spearman_p_value": float(profile_p),
        },
        "provenance_spearman": provenance_metrics,
        "top_descriptors": top_descriptors,
    }


def main() -> None:
    args = parse_args()
    metrics = evaluate(args)
    output = Path(args.output)
    output.write_text(json.dumps(metrics, indent=2), encoding="utf-8")

    print(f"Checkpoint: {metrics['checkpoint']['path']}")
    print(f"Validation samples: {metrics['data']['n_validation_samples']}")
    print(f"Profile MSE: {metrics['prediction']['profile_mse']:.6f}")
    print(f"Profile Spearman rho: {metrics['spearman']['profile_spearman_rho']:.4f}")
    if metrics.get("provenance_spearman"):
        for name, vals in metrics["provenance_spearman"].items():
            if isinstance(vals, dict) and "spearman_rho" in vals:
                rho = vals["spearman_rho"]
                rho_s = f"{rho:.4f}" if rho is not None else "n/a"
                print(f"{name}: rho={rho_s} n={vals['n_positive_values']}")
    for tier, vals in metrics["prediction"]["tier_metrics"].items():
        print(
            f"{tier.title():<7} MSE={vals['mse']:.4f} "
            f"MAE={vals['mae']:.4f} Var(pred)={vals['prediction_variance']:.6f} "
            f"rho={vals['mean_spearman_rho']:.4f}"
        )
    print(f"Wrote {output}")


if __name__ == "__main__":
    main()