File size: 6,901 Bytes
82ddb20 | 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 | #!/usr/bin/env python3
"""Enrich empirical_dataset_v1.jsonl with physics-derived time-varying targets.
For each record we compute how the OAV budget is distributed across top / heart /
base notes at every timestep. The objective target is then a time-varying blend
of the corresponding descriptor profiles (citrus/fresh/green for top notes,
floral/rose/jasmin for heart notes, woody/musk/amber for base notes). This
produces a target that is physically grounded in the diffusion trajectory and
varies over time because top notes fade and base notes emerge.
If Pyrfume annotations are available for an ingredient, the note class is
reconciled with the Pyrfume descriptor profile by using the ingredient's OAV
peak time as the primary signal.
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from typing import Any
import numpy as np
DESCRIPTOR_DIM = 138
# Descriptor profiles for the three fragrance notes in the Pyrfume vocabulary.
TOP_PROFILE = ["citrus", "fresh", "fruity", "green", "aldehydic", "lemon", "grapefruit", "orange"]
HEART_PROFILE = ["floral", "rose", "jasmin", "jasmine", "muguet", "lily", "neroli", "ylang", "hyacinth", "lavender"]
BASE_PROFILE = ["woody", "musk", "amber", "vanilla", "balsamic", "sweet", "cedar", "leathery", "tobacco"]
def build_profile_vectors(vocab: list[str]) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
def _vec(words: list[str]) -> np.ndarray:
arr = np.zeros(DESCRIPTOR_DIM, dtype=np.float32)
for w in words:
if w in vocab:
arr[vocab.index(w)] = 1.0
return arr / (arr.sum() + 1e-8)
return _vec(TOP_PROFILE), _vec(HEART_PROFILE), _vec(BASE_PROFILE)
def classify_ingredients(oav: np.ndarray) -> list[str]:
"""Classify each ingredient as top/heart/base by its OAV peak location."""
t_steps = oav.shape[0]
notes: list[str] = []
for s in range(oav.shape[1]):
traj = oav[:, s]
if traj.max() < 1e-12:
notes.append("heart")
continue
peak_t = int(np.argmax(traj))
frac = peak_t / (t_steps - 1) if t_steps > 1 else 0.5
if frac < 0.33:
notes.append("top")
elif frac > 0.66:
notes.append("base")
else:
notes.append("heart")
return notes
def enrich_record(
record: dict[str, Any],
top_vec: np.ndarray,
heart_vec: np.ndarray,
base_vec: np.ndarray,
) -> dict[str, Any]:
record = dict(record)
formula = record.get("formula", [])
trajectory = record.get("trajectory", [])
if not formula or not trajectory:
return record
t_steps = len(trajectory)
s_ing = len(formula)
oav = np.zeros((t_steps, s_ing), dtype=np.float32)
for t, step in enumerate(trajectory):
oav_step = step.get("OAV", {})
for s, comp in enumerate(formula):
cas = comp.get("cas", "")
lookup_cas = cas[len("NATURAL:"):] if cas.startswith("NATURAL:") else cas
oav[t, s] = oav_step.get(lookup_cas, 0.0)
notes = classify_ingredients(oav)
# Build a time-varying descriptor matrix per ingredient.
ingredient_matrix = np.zeros((t_steps, s_ing, DESCRIPTOR_DIM), dtype=np.float32)
for s, note in enumerate(notes):
if note == "top":
profile = top_vec
elif note == "base":
profile = base_vec
else:
profile = heart_vec
# The ingredient's descriptor contribution is modulated by its OAV at
# each timestep so that it fades in/out of the mixture profile.
traj = oav[:, s]
max_oav = traj.max()
if max_oav < 1e-12:
weights = np.zeros(t_steps, dtype=np.float32)
else:
weights = traj / max_oav
ingredient_matrix[:, s, :] = weights[:, None] * profile[None, :]
# Mixture target per timestep is the sum over ingredients.
targets = ingredient_matrix.sum(axis=1) # (T, 138)
targets = np.clip(targets, 0.0, 1.0)
# Normalize each timestep to a unit-length descriptor direction while
# preserving the relative descriptor weights within the timestep.
norms = np.linalg.norm(targets, axis=1, keepdims=True)
targets = np.divide(targets, norms, out=np.zeros_like(targets), where=norms > 1e-8)
record["objective_targets"] = targets.astype(np.float32).tolist()
return record
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--input", default="data/empirical_dataset_v1.jsonl")
parser.add_argument("--output", default="data/empirical_dataset_v2.jsonl")
parser.add_argument("--vocab", default="data/pyrfume_vocabulary.json")
args = parser.parse_args()
input_path = Path(args.input)
output_path = Path(args.output)
vocab_data = json.loads(Path(args.vocab).read_text())
vocab = vocab_data["vocabulary"] if isinstance(vocab_data, dict) else vocab_data
assert len(vocab) == DESCRIPTOR_DIM, f"Unexpected vocab size {len(vocab)}"
print(f"Loaded vocab: {len(vocab)} descriptors")
top_vec, heart_vec, base_vec = build_profile_vectors(vocab)
records = [json.loads(line) for line in input_path.open()]
print(f"Enriching {len(records)} records...")
out_file = output_path.open("w")
variance_list: list[float] = []
note_counts = {"top": 0, "heart": 0, "base": 0}
for rec in records:
enriched = enrich_record(rec, top_vec, heart_vec, base_vec)
obj = np.array(enriched["objective_targets"])
variance_list.append(float(obj.var(axis=0).mean()))
out_file.write(json.dumps(enriched, default=float) + "\n")
# Count notes for diagnostics.
t_steps = len(rec.get("trajectory", []))
if t_steps:
s_ing = len(rec.get("formula", []))
oav = np.zeros((t_steps, s_ing), dtype=np.float32)
for t, step in enumerate(rec["trajectory"]):
oav_step = step.get("OAV", {})
for s, comp in enumerate(rec["formula"]):
cas = comp.get("cas", "")
lookup_cas = cas[len("NATURAL:"):] if cas.startswith("NATURAL:") else cas
oav[t, s] = oav_step.get(lookup_cas, 0.0)
for note in classify_ingredients(oav):
note_counts[note] += 1
out_file.close()
variances = np.array(variance_list)
print(f"Wrote {len(records)} records to {output_path}")
print(f"Note distribution: {note_counts}")
print(f"Mean objective variance: {variances.mean():.6f}")
print(f"Median objective variance: {np.median(variances):.6f}")
print(f"Max objective variance: {variances.max():.6f}")
print(f"Fraction > 1e-3: {(variances > 1e-3).mean():.3f}")
print(f"Fraction > 1e-2: {(variances > 1e-2).mean():.3f}")
print(f"Fraction > 5e-2: {(variances > 5e-2).mean():.3f}")
if __name__ == "__main__":
main()
|