File size: 3,949 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
import torch, json, numpy as np
from pathlib import Path
import sys
sys.path.insert(0, 'src')

from pino.embeddings import OlfactoryEmbeddingEngine
from pino.heads import PIMTHeads
from pino.pimt_model import DEFAULT_EMBEDDING_DIM, PhysicsInformedMixtureTransformer
from pino.registry import AromaRegistry
from pino.thermo.naturals import NATURAL_PROFILES

device = torch.device('cpu')
checkpoint = torch.load('models/pimt_v1.pt', map_location=device)
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)
model.load_state_dict(checkpoint['model_state_dict'])
heads.load_state_dict(checkpoint['heads_state_dict'])
model.eval(); heads.eval()

with open('data/empirical_dataset_v1.jsonl') as f:
    for line in f:
        rec = json.loads(line)
        if any(str(i.get('cas','')).startswith('NATURAL:') for i in rec.get('formula',[])):
            sample = rec
            break

engine = OlfactoryEmbeddingEngine()
registry = AromaRegistry()

tokens = []
for item in sample['formula']:
    cas = str(item['cas']).strip()
    bare = cas.replace('NATURAL:', '')
    rec = registry.get(bare)
    smiles = rec.get('smiles', '') if rec else ''
    z = engine.get_embedding(smiles, cas=cas)
    tokens.append(torch.from_numpy(z))
tokens = torch.stack(tokens, dim=0).unsqueeze(0).float().to(device)

physics = np.zeros((len(sample['trajectory']), len(sample['formula']), 2), dtype=np.float32)
const_to_token = {}
for idx, item in enumerate(sample['formula']):
    raw_cas = str(item['cas']).strip()
    bare = raw_cas.replace('NATURAL:', '')
    profile = NATURAL_PROFILES.get(bare) or NATURAL_PROFILES.get(f'NATURAL:{bare}')
    if profile:
        for const_cas in profile['constituents']:
            const_to_token[const_cas] = idx
    else:
        const_to_token[bare] = idx
for t, step in enumerate(sample['trajectory']):
    for const_cas, token_idx in const_to_token.items():
        physics[t, token_idx, 0] += step['x_liquid'].get(const_cas, 0.0)
        physics[t, token_idx, 1] += step['OAV'].get(const_cas, 0.0)
physics[:,:,1] = np.log10(np.maximum(physics[:,:,1], 1e-10))
physics_t = torch.from_numpy(physics).unsqueeze(0).float().to(device)
mask = torch.zeros(1, tokens.size(1), dtype=torch.bool, device=device)

# Probe gating
gated = model.gating(tokens, physics_t)
print('Gated output shape:', gated.shape)
print('Gated variance across time:', gated.var(dim=1).mean().item())
print('Gated first timestep mean/std:', gated[0,0].mean().item(), gated[0,0].std().item())
print('Gated last timestep mean/std:', gated[0,-1].mean().item(), gated[0,-1].std().item())

# Probe input projection
proj = model.input_proj(gated)
print('Projected variance across time:', proj.var(dim=1).mean().item())

# Probe encoder output
b, t, s, _ = proj.shape
x = proj.reshape(b * t, s, 256)
enc_out = model.encoder(x, src_key_padding_mask=mask.unsqueeze(1).expand(-1, t, -1).reshape(b * t, s))
enc_out = enc_out.reshape(b, t, s, 256)
print('Encoder output variance across time:', enc_out.var(dim=1).mean().item())

# Compare with zero physics
physics0 = torch.zeros_like(physics_t)
gated0 = model.gating(tokens, physics0)
proj0 = model.input_proj(gated0)
x0 = proj0.reshape(b * t, s, 256)
enc_out0 = model.encoder(x0, src_key_padding_mask=mask.unsqueeze(1).expand(-1, t, -1).reshape(b * t, s))
enc_out0 = enc_out0.reshape(b, t, s, 256)
print('Diff encoder output vs zero physics (mean):', (enc_out - enc_out0).abs().mean().item())
print('Diff encoder output vs zero physics (max):', (enc_out - enc_out0).abs().max().item())

# Distribution of physics values
print('Physics tensor shape:', physics_t.shape)
print('Physics x_liquid range:', physics_t[0,:,:,0].min().item(), physics_t[0,:,:,0].max().item())
print('Physics log10OAV range:', physics_t[0,:,:,1].min().item(), physics_t[0,:,:,1].max().item())