pino-source-code / tests /test_pimt_invariance.py
Matthew Ford
Complete catalog recovery and enrichment
83db774
Raw
History Blame Contribute Delete
3.68 kB
from __future__ import annotations
import numpy as np
import torch
import torch.nn as nn
from pino.embeddings import OlfactoryEmbeddingEngine
from pino.heads import PIMTHeads
from pino.pimt_model import DEFAULT_EMBEDDING_DIM, PhysicsInformedMixtureTransformer
def test_ingredient_order_permutation_invariance() -> None:
"""
Shuffling the order of ingredient tokens in a formulation must leave the
transformer encoder's output unchanged, because the model uses no
positional encodings and attends over a permutation-invariant token set.
"""
engine = OlfactoryEmbeddingEngine(use_fallback=True)
smiles_list = ["CCO", "CC(=O)O", "c1ccccc1", "CC(C)Cc1ccc(C(C)C)cc1"]
embeddings = np.stack([engine.get_embedding(s) for s in smiles_list], axis=0)
embeddings = torch.from_numpy(embeddings).float()
# Synthetic physical states (T, S, 4).
timesteps = 5
states = torch.rand(timesteps, len(smiles_list), 4)
model = PhysicsInformedMixtureTransformer(
embedding_dim=DEFAULT_EMBEDDING_DIM,
state_dim=4,
hidden_dim=256,
num_heads=4,
num_layers=2,
)
model.eval()
heads = PIMTHeads(hidden_dim=256, objective_dim=138).eval()
with torch.no_grad():
latent = model(embeddings.unsqueeze(0), states.unsqueeze(0))
pred = heads(latent, states.unsqueeze(0))
perm = torch.randperm(len(smiles_list))
perm_embeddings = embeddings[perm]
perm_states = states[:, perm, :]
latent_perm = model(perm_embeddings.unsqueeze(0), perm_states.unsqueeze(0))
pred_perm = heads(latent_perm, perm_states.unsqueeze(0))
assert torch.allclose(pred["objective"], pred_perm["objective"], atol=1e-5)
for key in ("seasonality", "gender_profile", "wearability", "substantivity"):
assert torch.allclose(pred["subjective"][key], pred_perm["subjective"][key], atol=1e-5)
def test_head_output_shapes() -> None:
"""Verify both dual heads emit tensors of the expected shapes."""
model = PhysicsInformedMixtureTransformer(
embedding_dim=138, state_dim=4, hidden_dim=256, num_heads=4, num_layers=2
)
heads = PIMTHeads(hidden_dim=256, objective_dim=138)
model.eval()
heads.eval()
bsz, t_steps, n_mol = 2, 7, 6
tokens = torch.rand(bsz, n_mol, 138)
physics = torch.rand(bsz, t_steps, n_mol, 4)
# Valid tokens are False; the last token is a zero-padded placeholder marked True.
src_key_padding_mask = torch.zeros(bsz, n_mol, dtype=torch.bool)
src_key_padding_mask[:, -1] = True
with torch.no_grad():
latent = model(tokens, physics, src_key_padding_mask)
pred = heads(latent, physics, src_key_padding_mask)
assert pred["objective"].shape == (bsz, 3, 138)
assert pred["subjective"]["seasonality"].shape == (bsz, 4)
assert pred["subjective"]["gender_profile"].shape == (bsz, 1)
assert pred["subjective"]["wearability"].shape == (bsz, 2)
assert pred["subjective"]["substantivity"].shape == (bsz, 1)
def test_objective_target_curve_physics() -> None:
"""Check that the genre-specific volatility curves follow perfumery physics."""
from pino.semantics import get_objective_targets
targets = get_objective_targets("citrus_cologne", timesteps=50)
# Citrus top note should decay from start to end.
assert targets[0, 0] > targets[-1, 0]
# Wood base note should rise.
assert targets[-1, 2] > targets[0, 2]
# Lavender heart note should peak in the middle of the dry-down.
lavender_series = targets[:, 1]
peak_idx = int(np.argmax(lavender_series))
assert 5 < peak_idx < 45, "Lavender heart note did not peak in mid-dry-down"