#!/usr/bin/env python3 """Deep audit of PINO v6 dataset trajectory physics quality.""" import json import math import statistics from collections import defaultdict, Counter DATASET = "/home/hermes/pino/data/empirical_dataset_v6.jsonl" def load(): records = [] with open(DATASET) as f: for line in f: records.append(json.loads(line)) return records def analyze(): records = load() print(f"Total records: {len(records)}") # Collect all CAS keys and their trajectory appearances all_cas = set() cas_traj = defaultdict(list) # cas -> list of (formula_id, [OAV values across 49 steps]) for r in records: fid = r['formula_id'] n_steps = len(r['trajectory']) # Build per-CAS OAV series cas_series = defaultdict(lambda: [None]*n_steps) for i, step in enumerate(r['trajectory']): for cas, oav in step.get('OAV', {}).items(): cas_series[cas][i] = oav for cas, series in cas_series.items(): all_cas.add(cas) cas_traj[cas].append((fid, series)) print(f"Distinct CAS keys in trajectories: {len(all_cas)}") # Categorize CAS keys natural_cas = [c for c in all_cas if c.startswith('NATURAL:')] smiles_cas = [c for c in all_cas if c.startswith('SMILES:')] plain_cas = [c for c in all_cas if not c.startswith('NATURAL:') and not c.startswith('SMILES:')] print(f" Plain CAS: {len(plain_cas)}") print(f" SMILES: prefixed: {len(smiles_cas)}") print(f" NATURAL: prefixed: {len(natural_cas)}") # ========================================================= # (1) FLAT TRAJECTORIES: peak at step 0, no decay # ========================================================= # For each (cas, formula) pair, determine if trajectory is "flat" # Flat = peak OAV at step 0 (or near step 0), and end/peak ratio ~ 1 or barely decays flat_count = 0 flat_examples = [] non_flat_count = 0 total_pairs = 0 zero_traj = 0 # all zeros for cas, occurrences in cas_traj.items(): for fid, series in occurrences: total_pairs += 1 valid = [v for v in series if v is not None] if not valid or all(v == 0 for v in valid): zero_traj += 1 continue peak_val = max(series) peak_idx = series.index(peak_val) end_val = series[-1] start_val = series[0] # Flat: peak at step 0 AND end/peak > 0.95 (minimal decay) # Also catch: start == peak == end (truly flat) if peak_idx == 0 and end_val > 0: ratio = end_val / peak_val if ratio > 0.95: flat_count += 1 if len(flat_examples) < 10: flat_examples.append({ 'cas': cas, 'formula': fid, 'start': start_val, 'peak': peak_val, 'end': end_val, 'ratio': ratio }) else: non_flat_count += 1 elif peak_idx <= 2 and start_val > 0: ratio = end_val / peak_val if peak_val > 0 else 0 if ratio > 0.95: flat_count += 1 if len(flat_examples) < 10: flat_examples.append({ 'cas': cas, 'formula': fid, 'start': start_val, 'peak': peak_val, 'end': end_val, 'ratio': ratio, 'peak_idx': peak_idx }) else: non_flat_count += 1 else: non_flat_count += 1 print(f"\n{'='*70}") print(f"(1) FLAT OAV TRAJECTORIES") print(f"{'='*70}") print(f"Total (cas, formula) pairs analyzed: {total_pairs}") print(f" Zero trajectories (all OAV=0): {zero_traj}") print(f" Flat trajectories (peak@t0, end/peak>0.95): {flat_count}") print(f" Non-flat (decaying): {non_flat_count}") print(f" Flat % of non-zero: {flat_count/(flat_count+non_flat_count)*100:.1f}%" if (flat_count+non_flat_count) else "N/A") print(f"\n Examples of flat trajectories:") for ex in flat_examples[:5]: print(f" {ex}") # ========================================================= # (2) OAV INCREASES OVER TIME (peak not at step 0) # ========================================================= increasing_count = 0 increasing_examples = [] for cas, occurrences in cas_traj.items(): for fid, series in occurrences: valid = [v for v in series if v is not None] if not valid or all(v == 0 for v in valid): continue peak_val = max(series) peak_idx = series.index(peak_val) start_val = series[0] # Increasing = peak well after step 0 AND peak > 1.1 * start if peak_idx > 5 and start_val > 0 and peak_val > 1.1 * start_val: increasing_count += 1 if len(increasing_examples) < 15: # Find peak and characterize end_val = series[-1] increasing_examples.append({ 'cas': cas, 'formula': fid, 'start': start_val, 'peak': peak_val, 'peak_idx': peak_idx, 'end': end_val, 'start_to_peak_ratio': peak_val/start_val }) elif peak_idx > 5 and start_val == 0 and peak_val > 0: # Starts at 0, increases later increasing_count += 1 if len(increasing_examples) < 15: increasing_examples.append({ 'cas': cas, 'formula': fid, 'start': start_val, 'peak': peak_val, 'peak_idx': peak_idx, 'end': series[-1], 'note': 'starts at 0' }) print(f"\n{'='*70}") print(f"(2) OAV INCREASING OVER TIME (physically suspicious)") print(f"{'='*70}") print(f"Count: {increasing_count}") print(f"Examples:") for ex in increasing_examples[:10]: print(f" {ex}") # ========================================================= # (3) OAV DECAY RATIO (end/peak) DISTRIBUTION PER TIER # ========================================================= # Tier = ingredient count in formula decay_by_tier = defaultdict(list) for r in records: n_ing = len(r['formula']) fid = r['formula_id'] n_steps = len(r['trajectory']) cas_series = defaultdict(lambda: [None]*n_steps) for i, step in enumerate(r['trajectory']): for cas, oav in step.get('OAV', {}).items(): cas_series[cas][i] = oav for cas, series in cas_series.items(): valid = [v for v in series if v is not None] if not valid or all(v == 0 for v in valid): continue peak_val = max(series) end_val = series[-1] if peak_val > 0: ratio = end_val / peak_val decay_by_tier[n_ing].append(ratio) print(f"\n{'='*70}") print(f"(3) OAV DECAY RATIO (end/peak) DISTRIBUTION PER TIER (n_ingredients)") print(f"{'='*70}") for tier in sorted(decay_by_tier.keys()): ratios = decay_by_tier[tier] if not ratios: continue pcts = {p: sorted(ratios)[int(len(ratios)*p)] for p in [0.1,0.25,0.5,0.75,0.9]} mean_r = statistics.mean(ratios) print(f" Tier n={tier:3d} (n={len(ratios):5d}): " f"p10={pcts[0.1]:.4f} p25={pcts[0.25]:.4f} median={pcts[0.5]:.4f} " f"p75={pcts[0.75]:.4f} p90={pcts[0.9]:.4f} mean={mean_r:.4f}") # ========================================================= # (4) VP/OT REGRESSION PHYSICAL REASONABLENESS # ========================================================= # Check if trajectories are monotonically decreasing after peak # Also check for NaN/Inf values nan_inf_count = 0 non_monotonic_after_peak = 0 monotonic_after_peak = 0 sudden_jumps = 0 for cas, occurrences in cas_traj.items(): for fid, series in occurrences: valid = [v for v in series if v is not None] if not valid or all(v == 0 for v in valid): continue # Check NaN/Inf for v in series: if v is not None and (isinstance(v, float) and (math.isnan(v) or math.isinf(v))): nan_inf_count += 1 break # Check monotonicity after peak peak_val = max(series) peak_idx = series.index(peak_val) after_peak = series[peak_idx:] non_none_after = [v for v in after_peak if v is not None] if len(non_none_after) > 2: is_mono_dec = all(non_none_after[i] >= non_none_after[i+1] for i in range(len(non_none_after)-1)) if is_mono_dec: monotonic_after_peak += 1 else: non_monotonic_after_peak += 1 # Check for sudden jumps (>2x in adjacent steps) for i in range(len(non_none_after)-1): if non_none_after[i] > 0 and non_none_after[i+1] > non_none_after[i]*2: sudden_jumps += 1 break print(f"\n{'='*70}") print(f"(4) VP/OT REGRESSION PHYSICAL REASONABLENESS") print(f"{'='*70}") print(f"Records with NaN/Inf in OAV: {nan_inf_count}") print(f"Monotonic decay after peak: {monotonic_after_peak}") print(f"Non-monotonic after peak: {non_monotonic_after_peak}") print(f"With sudden jumps (>2x adjacent): {sudden_jumps}") print(f"Monotonicity rate: {monotonic_after_peak/(monotonic_after_peak+non_monotonic_after_peak)*100:.1f}%") # ========================================================= # (5) NATURAL: vs PLAIN CAS COMPARISON # ========================================================= def summarize_group(cas_list, label): ratios = [] flat = 0 increasing = 0 peak_at_0 = 0 peak_after_5 = 0 total = 0 zero_all = 0 peak_vals = [] for cas in cas_list: for fid, series in cas_traj[cas]: total += 1 valid = [v for v in series if v is not None] if not valid or all(v == 0 for v in valid): zero_all += 1 continue peak_val = max(series) peak_idx = series.index(peak_val) end_val = series[-1] peak_vals.append(peak_val) if peak_idx == 0: peak_at_0 += 1 if end_val / peak_val > 0.95 if peak_val > 0 else False: flat += 1 if peak_idx > 5 and peak_val > 1.1 * series[0]: increasing += 1 if peak_idx > 5: peak_after_5 += 1 if peak_val > 0: ratios.append(end_val/peak_val) print(f"\n {label}: {len(cas_list)} CAS keys, {total} trajectory pairs") print(f" Zero trajectories: {zero_all} ({zero_all/total*100:.1f}%)") print(f" Peak at step 0: {peak_at_0} ({peak_at_0/total*100:.1f}%)") print(f" Peak after step 5: {peak_after_5} ({peak_after_5/total*100:.1f}%)") print(f" Flat (peak@0 + no decay): {flat}") print(f" Increasing (peak>5 + rise): {increasing}") if ratios: med = statistics.median(ratios) mean = statistics.mean(ratios) print(f" End/peak ratio: median={med:.4f}, mean={mean:.4f}") if peak_vals: med_peak = statistics.median(peak_vals) mean_peak = statistics.mean(peak_vals) print(f" Peak OAV: median={med_peak:.2f}, mean={mean_peak:.2f}") print(f"\n{'='*70}") print(f"(5) NATURAL: vs PLAIN CAS COMPARISON") print(f"{'='*70}") summarize_group(plain_cas, "PLAIN CAS") summarize_group(smiles_cas, "SMILES: PREFIXED") summarize_group(natural_cas, "NATURAL: PREFIXED") # Specific NATURAL examples print(f"\n Sample NATURAL: CAS keys:") for c in sorted(natural_cas)[:20]: print(f" {c}") print(f"\n NATURAL: trajectory examples (first 5):") nat_count = 0 for cas in sorted(natural_cas): for fid, series in cas_traj[cas]: valid = [v for v in series if v is not None] if valid and not all(v == 0 for v in valid): peak_val = max(series) peak_idx = series.index(peak_val) end_val = series[-1] print(f" {cas} in {fid}: start={series[0]:.2f} peak={peak_val:.2f}@{peak_idx} end={end_val:.2f}") nat_count += 1 if nat_count >= 10: break if nat_count >= 10: break # ========================================================= # ADDITIONAL: Overall trajectory shape distribution # ========================================================= shapes = Counter() for cas, occurrences in cas_traj.items(): for fid, series in occurrences: valid = [v for v in series if v is not None] if not valid or all(v == 0 for v in valid): shapes['zero'] += 1 continue peak_val = max(series) peak_idx = series.index(peak_val) end_val = series[-1] start_val = series[0] if peak_idx == 0: if peak_val > 0 and end_val/peak_val > 0.95: shapes['flat_peak0'] += 1 else: shapes['decay_from_0'] += 1 elif peak_idx > 5: if start_val > 0 and peak_val > 1.1 * start_val: shapes['increasing'] += 1 else: shapes['peak_late_flat'] += 1 else: shapes['peak_early'] += 1 print(f"\n{'='*70}") print(f"OVERALL TRAJECTORY SHAPE DISTRIBUTION") print(f"{'='*70}") for shape, cnt in shapes.most_common(): print(f" {shape}: {cnt}") # ========================================================= # ADDITIONAL: Tier distribution of records # ========================================================= tier_dist = Counter() for r in records: tier_dist[len(r['formula'])] += 1 print(f"\nFormula size distribution:") for sz in sorted(tier_dist.keys()): print(f" n={sz:3d}: {tier_dist[sz]:4d} formulas") if __name__ == '__main__': analyze()