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) # Build token list: find nearest registry token for each Shalimar ingredient 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()