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()
|