| import sys |
| from pathlib import Path |
| import json |
| import numpy as np |
| import torch |
|
|
| sys.path.insert(0, "src") |
|
|
| from pino.pimt_model import PhysicsInformedMixtureTransformer, DEFAULT_EMBEDDING_DIM |
| from pino.heads import PIMTHeads |
| from pino.embeddings import get_openpom_embedding |
| from pino.registry import AromaRegistry |
| from huggingface_hub import hf_hub_download |
|
|
| SHALIMAR_1990 = [ |
| ("AMBREINE PURE", 6.00), ("BENZOAT BENZYL", 23.00), ("BENZOIN SIAM RESINOID", 21.50), |
| ("BENZYL ACETATE", 10.00), ("BERGAMOTE ITALIAN OIL", 255.00), ("CARDAMOM GUATEMALA OIL", 5.00), |
| ("CASTOREUM GIVCO", 116.00), ("CITRONELLOL", 0.50), ("CITRONELLYL ACETATE", 0.10), |
| ("CIVETTE SYNTHETIC", 0.50), ("CORIANDER SEED OIL", 20.00), ("COUMARIN", 72.00), |
| ("DIMETHYL BENZYL CARBINYL ACETATE", 0.10), ("DIPROPYLENE GLYCOL", 320.00), ("ESTRAGOLE", 5.00), |
| ("GERANIUM OIL CHINE", 5.00), ("GERANIUM OIL EGYPT", 0.10), ("GUAIACWOOD OIL", 0.05), |
| ("HONEY BEE ABSOLUTE (Cire D'abeille)", 0.10), ("IRONE ALPHA", 5.00), ("LAVANDER OIL 40/42", 5.20), |
| ("LEMON SICILIAN OIL", 1.30), ("LINALYL ACETATE", 25.40), ("MUSK KETONE (MUSK XYLOL <0,1%)", 4.00), |
| ("MYRRH RESINOID", 0.20), ("OPOPONAX RESINOID", 1.00), ("ORANGE SEET BRASIL OIL", 26.20), |
| ("PETITGRAIN BIGARADE OIL", 5.00), ("PHENYL ETHYL ACETATE", 0.10), ("PHENYLETHYL ALCOHOL", 1.50), |
| ("POLYSANTOL", 25.00), ("RHODINOL 70", 0.50), ("ROSALVA", 0.05), ("ROSE ACETATE", 0.20), |
| ("ROSE OIL TURKISH", 4.30), ("SANDALOR", 30.00), ("VANILLIN", 95.00), ("VETIVER OIL JAVA", 8.00), |
| ] |
|
|
| NAME_TO_CAS = { |
| "AMBREINE PURE": "NATURAL:68916-26-7", "BENZOAT BENZYL": "120-51-4", "BENZOIN SIAM RESINOID": "NATURAL:9000-05-9", |
| "BENZYL ACETATE": "140-11-4", "BERGAMOTE ITALIAN OIL": "NATURAL:8007-35-0", "CARDAMOM GUATEMALA OIL": "NATURAL:8000-66-6", |
| "CASTOREUM GIVCO": "NATURAL:8023-83-4", "CITRONELLOL": "106-22-9", "CITRONELLYL ACETATE": "150-84-5", |
| "CIVETTE SYNTHETIC": "NATURAL:68916-26-7", "CORIANDER SEED OIL": "NATURAL:8008-52-4", "COUMARIN": "91-64-5", |
| "DIMETHYL BENZYL CARBINYL ACETATE": "151-05-3", "DIPROPYLENE GLYCOL": "25265-71-8", "ESTRAGOLE": "140-67-0", |
| "GERANIUM OIL CHINE": "NATURAL:8000-46-2", "GERANIUM OIL EGYPT": "NATURAL:8000-46-2", "GUAIACWOOD OIL": "NATURAL:8016-23-7", |
| "HONEY BEE ABSOLUTE (Cire D'abeille)": "NATURAL:8029-66-5", "IRONE ALPHA": "79-69-6", "LAVANDER OIL 40/42": "NATURAL:8000-28-0", |
| "LEMON SICILIAN OIL": "NATURAL:8008-56-8", "LINALYL ACETATE": "115-95-7", "MUSK KETONE (MUSK XYLOL <0,1%)": "81-14-1", |
| "MYRRH RESINOID": "NATURAL:8016-37-3", "OPOPONAX RESINOID": "NATURAL:9000-78-6", "ORANGE SEET BRASIL OIL": "NATURAL:8008-57-9", |
| "PETITGRAIN BIGARADE OIL": "NATURAL:8014-17-3", "PHENYL ETHYL ACETATE": "103-45-7", "PHENYLETHYL ALCOHOL": "60-12-8", |
| "POLYSANTOL": "NATURAL:66068-84-6", "RHODINOL 70": "141-25-3", "ROSALVA": "128-50-7", "ROSE ACETATE": "141-12-8", |
| "ROSE OIL TURKISH": "NATURAL:8007-01-0", "SANDALOR": "NATURAL:66068-84-6", "VANILLIN": "121-33-5", "VETIVER OIL JAVA": "NATURAL:8016-96-4", |
| } |
|
|
|
|
| def build_inputs(): |
| registry = AromaRegistry() |
| records = registry.all_records() |
| vocab_emb = [] |
| for cas, rec in records.items(): |
| try: |
| emb = get_openpom_embedding(rec.get("smiles", "") or "", dim=DEFAULT_EMBEDDING_DIM, cas=cas) |
| vocab_emb.append(emb) |
| except Exception: |
| vocab_emb.append(np.zeros(DEFAULT_EMBEDDING_DIM, dtype=np.float32)) |
| vocab_emb = np.array(vocab_emb, dtype=np.float32) |
|
|
| |
| token_to_cas = {i: cas for i, cas in enumerate(records.keys())} |
|
|
| tokens, weights = [], [] |
| for name, w in SHALIMAR_1990: |
| cas = NAME_TO_CAS[name] |
| emb = get_openpom_embedding("", dim=DEFAULT_EMBEDDING_DIM, cas=cas) |
| norm = np.linalg.norm(emb) |
| emb = emb / (norm + 1e-8) |
| sims = vocab_emb @ emb |
| best_idx = int(np.argmax(sims)) |
| tokens.append(best_idx) |
| weights.append(w) |
| total = sum(weights) |
| weights = [w / total for w in weights] |
| return tokens, weights |
|
|
|
|
| def load_model(device, untrained=False): |
| model = PhysicsInformedMixtureTransformer( |
| embedding_dim=DEFAULT_EMBEDDING_DIM, state_dim=2, hidden_dim=256, num_heads=4, num_layers=4 |
| ).to(device) |
| heads = PIMTHeads(hidden_dim=256, objective_dim=138).to(device) |
| if not untrained: |
| checkpoint_path = Path("models/pimt_v3.pt") |
| if not checkpoint_path.exists(): |
| hf_hub_download(repo_id="mattbitzesty/pino-pimt-checkpoint", filename="pimt_v3.pt", repo_type="model", local_dir="models") |
| checkpoint = torch.load(checkpoint_path, map_location=device) |
| model.load_state_dict(checkpoint["model_state_dict"]) |
| heads.load_state_dict(checkpoint["heads_state_dict"]) |
| model.eval() |
| heads.eval() |
| return model, heads |
|
|
|
|
| def run_forward(model, heads, tokens, weights, device): |
| activations = {} |
|
|
| def make_hook(name): |
| def fn(m, i, o): |
| activations[name] = o.detach() if not isinstance(o, tuple) else o[0].detach() |
| return fn |
|
|
| model.gating.register_forward_hook(make_hook("film")) |
| model.encoder.register_forward_hook(make_hook("encoder")) |
|
|
| tokens = torch.tensor([tokens], dtype=torch.long, device=device) |
| weights = torch.tensor([weights], dtype=torch.float32, device=device) |
| with torch.no_grad(): |
| out = model(tokens, weights) |
| h = heads(out) |
| return out, h, activations |
|
|
|
|
| def main(): |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| tokens, weights = build_inputs() |
|
|
| print("=== TRAINED v3 ===") |
| model, heads = load_model(device, untrained=False) |
| out, h, acts = run_forward(model, heads, tokens, weights, device) |
| print(f"physics state: shape {out['physics_state'].shape}") |
| print(" per-channel mean time-variance:", out["physics_state"][0].var(dim=0).mean(dim=0).cpu().numpy()) |
| for k in ["film", "encoder"]: |
| v = acts[k] |
| print(f"{k}: shape {v.shape}") |
| print(f" mean per-token time-variance: {v[0].var(dim=0).mean().item():.6e}") |
| print(f" max per-token time-variance: {v[0].var(dim=0).max().item():.6e}") |
| print(f"objective: shape {h['objective'].shape}") |
| print(f" mean per-dim time-variance: {h['objective'][0].var(dim=0).mean().item():.6e}") |
| print(f" max per-dim time-variance: {h['objective'][0].var(dim=0).max().item():.6e}") |
|
|
| print("\n=== UNTRAINED ===") |
| model2, heads2 = load_model(device, untrained=True) |
| out2, h2, acts2 = run_forward(model2, heads2, tokens, weights, device) |
| for k in ["film", "encoder"]: |
| v = acts2[k] |
| print(f"{k}: mean per-token time-variance: {v[0].var(dim=0).mean().item():.6e}") |
| print(f"objective mean per-dim time-variance: {h2['objective'][0].var(dim=0).mean().item():.6e}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|