File size: 9,053 Bytes
1289050 f7f1afd 1289050 f7f1afd 1289050 f7f1afd 1289050 | 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 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 | """IDEA 3: does the PIMT transformer itself exploit the physics channel to predict
the time-evolved (pyramid) descriptor profile?
Idea 2 showed (with GBM/linear probes on hand-built features) that the physics
trajectory DOUBLES explained variance on the drydown profile. This evaluator
tests the same claim end-to-end in the trained model: for each held-out formula,
run the objective head with the REAL physics tensor vs a ZEROED physics tensor
(physics-off ablation at inference), and score the predicted 3-tier pyramid
against the stored pyramid target -- per tier (top/mid/base) and overall.
If physics carries non-redundant signal for the time-evolved profile, physics-on
should beat physics-off, with the largest gain on the BASE tier (drydown).
Metric: per-descriptor binary cross-entropy (the training objective) and mean
cosine similarity between predicted and target tier profiles. Lower BCE / higher
cosine with physics-on than physics-off = physics helps.
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import numpy as np
import torch
from pino.heads import PIMTHeads
from pino.pimt_model import (
FragranceTrajectoryDataset,
PhysicsInformedMixtureTransformer,
objective_targets_to_pyramid,
)
from pino.upload_data import create_molecule_disjoint_split_from_records
def load_model(ckpt_path: Path, objective_dim: int):
sd = torch.load(ckpt_path, map_location="cpu", weights_only=False)
state = sd.get("model_state_dict", sd.get("model", sd))
heads_state = sd.get("heads_state_dict", sd.get("heads", {}))
state_dim = int(state["gating.physics_scaler"].shape[0])
embedding_dim = int(state["input_proj.weight"].shape[1])
hidden_dim = int(state["input_proj.weight"].shape[0])
num_layers = sum(1 for k in state if k.endswith("self_attn.in_proj_weight"))
model = PhysicsInformedMixtureTransformer(
embedding_dim=embedding_dim, state_dim=state_dim,
hidden_dim=hidden_dim, num_heads=4, num_layers=max(num_layers, 1),
)
model.load_state_dict(state, strict=False)
heads = PIMTHeads(hidden_dim=hidden_dim, objective_dim=objective_dim)
if heads_state:
heads.load_state_dict(heads_state, strict=False)
model.eval(); heads.eval()
return model, heads, {"embedding_dim": embedding_dim, "state_dim": state_dim,
"hidden_dim": hidden_dim, "num_layers": num_layers}
def tier_cosine(pred, targ):
# pred/targ: (3, 138); cosine per tier
out = []
for t in range(3):
a = pred[t]; b = targ[t]
na = np.linalg.norm(a); nb = np.linalg.norm(b)
out.append(float(np.dot(a, b) / (na * nb)) if na > 0 and nb > 0 else 0.0)
return out # [top, mid, base]
def main() -> int:
ap = argparse.ArgumentParser()
ap.add_argument("--checkpoint", required=True)
ap.add_argument("--data", required=True)
ap.add_argument("--structural-source", default="morgan")
ap.add_argument("--train-ratio", type=float, default=0.85)
ap.add_argument("--seed", type=int, default=42)
ap.add_argument("--split", choices=["molecule", "formula"], default="formula")
ap.add_argument("--arm-label", default=None)
ap.add_argument("--use-gamma", action="store_true")
ap.add_argument("--physics-mode", choices=["full", "off", "gamma_only", "headspace_only"], default="full",
help="Match the arm's training conditioning: gamma_only zeroes channels 0-1 of the 'on' tensor, headspace_only zeroes channel 2, off zeroes all. Ensures each factorial arm is evaluated with the same masked physics it was trained on.")
ap.add_argument("--max-eval", type=int, default=400)
ap.add_argument("--output", required=True)
args = ap.parse_args()
objective_dim = 575 if args.structural_source == "pom_alltags" else 138
engine_source = "openpom_256" if args.structural_source == "pom_alltags" else args.structural_source
model, heads, dims = load_model(Path(args.checkpoint), objective_dim)
state_dim = dims["state_dim"]
with open(args.data) as f:
all_records = [json.loads(line) for line in f if line.strip()]
if args.split == "molecule":
split = create_molecule_disjoint_split_from_records(all_records, train_ratio=args.train_ratio, seed=args.seed)
val_records = split["validation"]
else:
rng = np.random.default_rng(args.seed)
n = len(all_records)
n_val = max(1, int(round(n * (1.0 - args.train_ratio))))
perm = rng.permutation(n)
val_records = [all_records[i] for i in sorted(set(perm[:n_val].tolist()))]
# restrict to blends with a real trajectory + pyramid target
blends = [r for r in val_records if not r.get("is_control") and len(r.get("formula", [])) >= 3
and r.get("objective_targets") and len(r.get("trajectory", [])) > 1]
blends = blends[: args.max_eval]
ds = FragranceTrajectoryDataset(records=blends, use_embedding_fallback=True,
structural_source=engine_source, objective_dim=objective_dim,
use_gamma=args.use_gamma)
device = "cuda" if torch.cuda.is_available() else "cpu"
model.to(device); heads.to(device)
bce_on, bce_off = [], []
cos_on = np.zeros(3); cos_off = np.zeros(3)
used = 0
bce = torch.nn.functional.binary_cross_entropy
with torch.no_grad():
for i in range(len(ds)):
item = ds[i]
tokens = item["tokens"].unsqueeze(0).to(device)
physics = item["physics"].unsqueeze(0).to(device)
if physics.size(-1) != state_dim:
if physics.size(-1) < state_dim:
pad = torch.zeros(1, physics.size(1), physics.size(2), state_dim - physics.size(-1), device=device)
physics = torch.cat([physics, pad], dim=-1)
else:
physics = physics[..., :state_dim]
# mask: all real tokens present in a single (unpadded) item
S = tokens.size(1)
mask = torch.zeros(1, S, dtype=torch.bool, device=device)
# target_obj is already the (3,138) pyramid for this dataset
target = item["target_obj"].clamp(0.0, 1.0).to(device) # (3,138)
# physics ON -- with the arm's training conditioning (factorial modes
# mask channels so eval matches what the model saw during training)
if args.physics_mode == "off":
physics = torch.zeros_like(physics)
elif args.physics_mode == "gamma_only":
physics = physics.clone(); physics[..., 0] = 0.0; physics[..., 1] = 0.0
elif args.physics_mode == "headspace_only":
physics = physics.clone(); physics[..., 2] = 0.0
lat_on = model(tokens, physics)
pyr_on = heads(lat_on, physics, mask)["objective"][0] # (3,138)
# physics OFF (zeroed)
phys_zero = torch.zeros_like(physics)
lat_off = model(tokens, phys_zero)
pyr_off = heads(lat_off, phys_zero, mask)["objective"][0]
bce_on.append(float(bce(pyr_on, target).cpu()))
bce_off.append(float(bce(pyr_off, target).cpu()))
c_on = tier_cosine(pyr_on.cpu().numpy(), target.cpu().numpy())
c_off = tier_cosine(pyr_off.cpu().numpy(), target.cpu().numpy())
cos_on += np.array(c_on); cos_off += np.array(c_off)
used += 1
cos_on /= max(used, 1); cos_off /= max(used, 1)
result = {
"arm": args.arm_label or Path(args.checkpoint).stem,
"use_gamma": bool(args.use_gamma), "physics_mode": args.physics_mode, "seed": args.seed, "model_dims": dims,
"task": "3-tier descriptor pyramid (top/mid/base); physics-on vs physics-off inference ablation",
"n_eval": used,
"bce_physics_on": float(np.mean(bce_on)),
"bce_physics_off": float(np.mean(bce_off)),
"bce_delta_off_minus_on": float(np.mean(bce_off) - np.mean(bce_on)),
"cosine_top_on": float(cos_on[0]), "cosine_mid_on": float(cos_on[1]), "cosine_base_on": float(cos_on[2]),
"cosine_top_off": float(cos_off[0]), "cosine_mid_off": float(cos_off[1]), "cosine_base_off": float(cos_off[2]),
"cosine_base_delta_on_minus_off": float(cos_on[2] - cos_off[2]),
"cosine_mean_on": float(cos_on.mean()), "cosine_mean_off": float(cos_off.mean()),
"verdict": ("physics helps (on beats off)" if np.mean(bce_on) < np.mean(bce_off) else "physics does not help"),
}
out = Path(args.output); out.parent.mkdir(parents=True, exist_ok=True)
out.write_text(json.dumps(result, indent=2))
print(json.dumps({k: result[k] for k in ["arm", "n_eval", "bce_physics_on", "bce_physics_off",
"bce_delta_off_minus_on", "cosine_base_on", "cosine_base_off",
"cosine_base_delta_on_minus_off", "verdict"]}, indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())
|