File size: 6,913 Bytes
ad424e4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
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()