pino-source-code / scripts /diagnostic_v3_intermediate.py
mattbitzesty's picture
v6: ConcentrationAwarePyramidHead + OAV fix + curated descriptors + integration
ad424e4 unverified
Raw
History Blame Contribute Delete
6.91 kB
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()