Files
Jyotisha/scripts/research/probe_supply_after_six.py
T
Jesse_ChenandCursor b063c66835
Independent Staging Quality Gate / validate (push) Successful in 11m55s
Independent Staging Quality Gate / publish (push) Successful in 8m53s
docs(research): measure dated probe supply after six answers
Offline holdout replay shows refresh-only R3+R4 add discriminative dated probes; R1/R2 do not meet the gate.

Co-authored-by: Cursor <cursoragent@cursor.com>
2026-09-13 12:59:46 +08:00

1225 lines
48 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
"""Offline measurement: can refresh emit more dated probes after six answers?
Does not change production `event_probes.py` defaults. Relaxations are applied
only inside this script via temporary module patches.
"""
from __future__ import annotations
import argparse
import json
import statistics
import sys
import traceback
from calendar import monthrange
from contextlib import contextmanager, nullcontext
from dataclasses import dataclass
from datetime import date
from pathlib import Path
from typing import Any, Iterator, Sequence
from uuid import NAMESPACE_URL, uuid5
ROOT = Path(__file__).resolve().parents[2]
if str(ROOT) not in sys.path:
sys.path.insert(0, str(ROOT))
from scripts.rectification.candidate_contrast import ( # noqa: E402
cluster_contexts_by_signature,
distinguish_contract_errors,
)
from scripts.rectification.contracts import normalize_rectification_request # noqa: E402
from scripts.rectification.event_probes import ( # noqa: E402
DOMAIN_CATALOG,
EXISTENCE_NEARBY_YEARS,
MAX_BOUNDARY_CANDIDATES_PER_DOMAIN,
MIN_BOUNDARY_DAYS,
PROBE_PHASE_CANDIDATE_DISCRIMINATOR,
_annotate_nearby_ledger,
_answer_priors_for,
_apply_prior_ranking,
_best_probe_per_year,
_boundary_windows,
_differing_layers,
_dominant_existence_prior,
_evaluation_order,
_event_years,
_existence_blocked_years,
_layer_value,
_narayana_cache_key,
_narayana_start_dates,
_partition_ranked_probes,
_probe_caps,
_probe_domains,
_probe_sort_key,
_quality_distinguish_probes,
_remaining_contexts,
_representative_pairs,
_scoreable,
_static_contexts,
_try_activation_probe,
_vim_cache_key,
_vim_start_dates,
asked_years_for_domain,
discriminating_event_probes,
)
from scripts.rectification.case_holdout import holdout_domain_years, holdout_event_ids # noqa: E402
from scripts.rectification.refinement_packet import window_scan # noqa: E402
from scripts.rectification.scoring_service import ( # noqa: E402
build_event_contribution_matrix,
score_from_matrix,
scoreable_request,
)
HOLDOUT_MANIFEST = ROOT / "references" / "real_case_calibration" / "minute_rectification_holdout_v3.json"
TODAY = date(2026, 9, 13)
ASK_COUNT = 6
REMAINING_CAP = 5
STRONG_CONFLICT_ELIMINATION_COUNT = 3
SCORE_DELTA = {
"support": 2.0,
"weak_support": 1.0,
"neutral": 0.0,
"weak_conflict": -1.0,
"conflict": -2.0,
}
KIND_BY_DOMAIN = {
"education": "education_milestone",
"career": "career_change",
"relationship": "relationship_commitment",
"relocation": "relocation",
"health_pressure": "self_health_event",
"health": "self_health_event",
"finance": "finance_change",
"family": "family_event",
}
EXTRA_FAMILIES: dict[str, tuple[tuple[str, str], ...]] = {
"career": (
("career_entry", "入职"),
("promotion", "升职"),
("career_exit", "离职"),
),
"relationship": (
("relationship_commitment", "结婚"),
("relationship_separation", "分手"),
),
}
@dataclass(frozen=True)
class Relaxation:
r1: bool = False
r2: bool = False
r3: bool = False
r4: bool = False
@property
def name(self) -> str:
labels = [
label
for label, on in (("R1", self.r1), ("R2", self.r2), ("R3", self.r3), ("R4", self.r4))
if on
]
return "+".join(labels) if labels else "baseline"
def all_relaxations() -> list[Relaxation]:
rows: list[Relaxation] = []
for mask in range(16):
rows.append(Relaxation(
r1=bool(mask & 1),
r2=bool(mask & 2),
r3=bool(mask & 4),
r4=bool(mask & 8),
))
return rows
def _clock(value: str) -> int:
return int(value[:2]) * 60 + int(value[3:5])
def _hhmm_from_minutes(value: int) -> str:
wrapped = value % 1440
return f"{wrapped // 60:02d}:{wrapped % 60:02d}"
def _shift_clock(value: str, delta: int) -> str:
return _hhmm_from_minutes(_clock(value) + delta)
def _expand_date(raw: str, precision: str) -> tuple[str, str]:
text = str(raw or "").strip()
if precision == "day":
day = date.fromisoformat(text)
return day.isoformat(), day.isoformat()
if precision == "month":
month = date.fromisoformat(f"{text}-01") if len(text) == 7 else date.fromisoformat(text[:10]).replace(day=1)
last = monthrange(month.year, month.month)[1]
return month.isoformat(), month.replace(day=last).isoformat()
year = int(text[:4])
return f"{year}-01-01", f"{year}-12-31"
def _event_kind(event: dict[str, Any]) -> str:
domain = str(event.get("domain") or "")
description = str(event.get("description") or "").lower()
if domain == "relationship" and not any(
token in description for token in ("married", "wedding", "wife", "husband")
):
return "relationship_start"
return KIND_BY_DOMAIN[domain]
def request_from_case(case: dict[str, Any]) -> dict[str, Any]:
birth = case["birth"]
true_time = str(birth["time"])[:5]
radius = int(case.get("candidate_radius_minutes") or 10)
events = []
for event in case.get("events") or []:
precision = str(event.get("precision") or "year")
start, end = _expand_date(str(event.get("date") or ""), precision)
domain = str(event.get("domain") or "")
events.append({
"id": str(uuid5(NAMESPACE_URL, str(event.get("id") or ""))),
"domain": domain,
"event_kind": _event_kind(event),
"date_start": start,
"date_end": end,
"precision": precision,
"summary": str(event.get("description") or event.get("id") or domain)[:200],
})
return normalize_rectification_request({
"birth_date": str(birth["date"]),
"start_time": _shift_clock(true_time, -radius),
"end_time": _shift_clock(true_time, radius),
"lat": float(birth["latitude"]),
"lon": float(birth["longitude"]),
"tz": float(birth["timezone_offset"]),
"events": events,
}, today=TODAY)
def _scan_for(built: dict[str, Any]) -> dict[str, Any]:
return window_scan(built)
def family_slug(kind: str) -> str:
return kind.replace("_", "-")
@contextmanager
def patched_boundary_days(days: int) -> Iterator[None]:
import scripts.rectification.event_probes as ep
previous = ep.MIN_BOUNDARY_DAYS
ep.MIN_BOUNDARY_DAYS = days
try:
yield
finally:
ep.MIN_BOUNDARY_DAYS = previous
@contextmanager
def catalog_family(domain: str, kind: str, family: str) -> Iterator[None]:
original = DOMAIN_CATALOG[domain]
DOMAIN_CATALOG[domain] = {**original, "kind": kind, "event_family": family}
try:
yield
finally:
DOMAIN_CATALOG[domain] = original
def union_boundary_dates(
reps: Sequence[dict[str, Any]],
*,
birth_date: str,
lo: int,
hi: int,
include_pratyantar: bool,
varga_narayana: bool,
) -> list[date]:
vim_cache: dict[tuple[Any, ...], list[date]] = {}
narayana_cache: dict[tuple[Any, ...], list[date] | None] = {}
dates_by_key: dict[tuple[int, int], date] = {}
for left, right in _representative_pairs(reps):
left_moon = float(left["planet_longitudes"]["Moon"])
right_moon = float(right["planet_longitudes"]["Moon"])
left_vim_key = (*_vim_cache_key(birth_date, left_moon, lo, hi), include_pratyantar)
right_vim_key = (*_vim_cache_key(birth_date, right_moon, lo, hi), include_pratyantar)
if left_vim_key not in vim_cache:
vim_cache[left_vim_key] = _vim_start_dates(
birth_date, left_moon, lo, hi, include_pratyantar=include_pratyantar,
)
if right_vim_key not in vim_cache:
vim_cache[right_vim_key] = _vim_start_dates(
birth_date, right_moon, lo, hi, include_pratyantar=include_pratyantar,
)
windows = list(_boundary_windows(vim_cache[left_vim_key], vim_cache[right_vim_key]))
layers: list[str | None] = [None]
if varga_narayana:
layers.extend(["d9", "d10"])
for layer in layers:
if layer is None:
left_asc = int(left["ascendant_index"])
right_asc = int(right["ascendant_index"])
else:
left_raw = _layer_value(left, layer)
right_raw = _layer_value(right, layer)
if not isinstance(left_raw, int) or not isinstance(right_raw, int):
continue
left_asc = left_raw
right_asc = right_raw
left_nara_key = (*_narayana_cache_key(
left_asc, left["planet_longitudes"], birth_date, lo, hi,
), layer)
right_nara_key = (*_narayana_cache_key(
right_asc, right["planet_longitudes"], birth_date, lo, hi,
), layer)
if left_nara_key not in narayana_cache:
narayana_cache[left_nara_key] = _narayana_start_dates(
left_asc, left["planet_longitudes"], birth_date, lo, hi,
)
if right_nara_key not in narayana_cache:
narayana_cache[right_nara_key] = _narayana_start_dates(
right_asc, right["planet_longitudes"], birth_date, lo, hi,
)
left_narayana = narayana_cache[left_nara_key]
right_narayana = narayana_cache[right_nara_key]
if left_narayana is not None and right_narayana is not None:
windows.extend(_boundary_windows(left_narayana, right_narayana))
for item in windows:
dates_by_key.setdefault((item.year, item.month), item)
return sorted(dates_by_key.values(), key=lambda item: (item.year, item.month))
def best_probe_per_year_family(rows: Sequence[dict[str, Any]]) -> list[dict[str, Any]]:
by_key: dict[tuple[int, str], dict[str, Any]] = {}
for row in sorted(rows, key=_probe_sort_key):
key = (int(row["year"]), str(row.get("event_family") or ""))
if key not in by_key:
by_key[key] = row
return list(by_key.values())
def partition_probes(
probes: Sequence[dict[str, Any]],
*,
max_probes: int,
include_family: bool,
) -> list[dict[str, Any]]:
if not include_family:
public, _dropped = _partition_ranked_probes(probes, max_probes=max_probes)
return public
public: list[dict[str, Any]] = []
seen: set[tuple[str, int, int, str, str]] = set()
for row in probes:
if row.get("source") == "known_event_quality" and not (
row.get("role") == "distinguish" and row.get("target_evidence_id")
):
continue
if row.get("phase") != PROBE_PHASE_CANDIDATE_DISCRIMINATOR:
continue
if distinguish_contract_errors(row):
continue
ranked = row if "raw_split_gain" in row else _apply_prior_ranking(dict(row))
priors = ranked.get("answer_priors") or _answer_priors_for(ranked)
if _dominant_existence_prior(ranked, priors):
continue
if not isinstance(ranked.get("year"), int) or int(ranked["year"]) <= 0:
continue
key = (
str(ranked["domain"]),
int(ranked["year"]),
int(ranked.get("month") or 0),
str(ranked["source"]),
str(ranked.get("event_family") or ""),
)
if key in seen or "points" in str(ranked):
continue
seen.add(key)
public.append(ranked)
if len(public) >= max_probes:
break
public.sort(key=_probe_sort_key)
return public
def annotate_family(row: dict[str, Any], *, kind: str, family: str) -> dict[str, Any]:
payload = dict(row)
payload["event_family"] = family
payload["event_kind"] = kind
month = payload.get("month")
source = str(payload.get("source") or "dasha_boundary")
domain = str(payload.get("domain") or "")
year = int(payload["year"])
slug = family_slug(kind)
if isinstance(month, int) and 1 <= month <= 12:
payload["semantic_key"] = f"{domain}.{year}.{month:02d}.{source}.{slug}"
else:
payload["semantic_key"] = f"{domain}.{year}.{source}.{slug}"
return payload
def generate_refresh_probes(
request: dict[str, Any],
built: dict[str, Any],
*,
remaining_times: Sequence[str],
relax: Relaxation,
today: date,
) -> list[dict[str, Any]]:
birth_date = str(request.get("birth_date") or "").strip()
birth_year = int(birth_date[:4])
events = [item for item in (request.get("events") or []) if isinstance(item, dict)]
asked_probe_keys = [
str(item).strip()
for item in (request.get("asked_probe_keys") or [])
if isinstance(item, str) and str(item).strip()
]
holdout_keys = holdout_domain_years(events)
remaining = _remaining_contexts(built, remaining_times)
if len(remaining) < 2:
return []
clusters = cluster_contexts_by_signature(remaining)
if len(clusters) < 2:
return []
reps = [cluster["representative"] for cluster in clusters if _scoreable(cluster["representative"])]
if len(reps) < 2:
reps = [item for item in remaining if _scoreable(item)]
if len(reps) < 2:
return []
max_probes, max_per_domain = _probe_caps(refresh=True, remaining_count=len(remaining))
from scripts.rectification.candidate_contrast import candidate_set_version
set_version = candidate_set_version([cluster["times"] for cluster in clusters])
remaining_layers = _differing_layers(remaining)
scan = _scan_for(built)
if not remaining_layers:
remaining_layers = {
layer for layer in (
"d1", "d9", "d10", "d4", "d5", "d24", "d7", "d12", "d2", "d11", "d30",
)
if scan.get(f"{layer}_candidates_differ")
}
domains = _probe_domains(
remaining_layers,
events,
d1_differs="d1" in remaining_layers or bool(scan.get("d1_candidates_differ")),
)
lo, hi = birth_year + 5, min(today.year, birth_year + 80)
with patched_boundary_days(30 if relax.r3 else MIN_BOUNDARY_DAYS):
boundary_dates = union_boundary_dates(
reps,
birth_date=birth_date,
lo=lo,
hi=hi,
include_pratyantar=relax.r4,
varga_narayana=relax.r4,
) if domains else []
probes: list[dict[str, Any]] = []
from scripts.rectification.event_probes import _evaluate_contexts, _domain_year_floor
for domain in domains:
if domain not in DOMAIN_CATALOG:
continue
known_years = _event_years(events, domain)
asked_years = asked_years_for_domain(asked_probe_keys, domain)
event_blocked = set(known_years) if relax.r1 else _existence_blocked_years(domain, known_years)
if relax.r2 and domain in EXTRA_FAMILIES:
blocked_years = event_blocked
elif relax.r1:
blocked_years = event_blocked | asked_years
else:
blocked_years = _existence_blocked_years(domain, known_years | asked_years)
domain_lo = _domain_year_floor(birth_year, domain, lo)
eligible = [
item
for item in boundary_dates
if domain_lo <= item.year <= hi
and item.year not in blocked_years
and f"{domain}:{item.year}" not in holdout_keys
]
families: tuple[tuple[str, str] | None, ...]
if relax.r2 and domain in EXTRA_FAMILIES:
families = EXTRA_FAMILIES[domain]
else:
families = (None,)
found: list[dict[str, Any]] = []
for family in families:
if relax.r2 and family is not None:
asked_for_family = {
year
for year in asked_years
if any(
key.startswith(f"{domain}.{year}.") and family_slug(family[0]) in key
for key in asked_probe_keys
)
}
family_blocked = set(blocked_years) | asked_for_family
family_eligible = [
item for item in boundary_dates
if domain_lo <= item.year <= hi
and item.year not in family_blocked
and f"{domain}:{item.year}" not in holdout_keys
]
else:
family_eligible = eligible
evaluated = 0
sample_size = min(MAX_BOUNDARY_CANDIDATES_PER_DOMAIN, len(family_eligible))
ctx = catalog_family(domain, family[0], family[1]) if family else nullcontext()
with ctx:
for at in _evaluation_order(family_eligible, MAX_BOUNDARY_CANDIDATES_PER_DOMAIN):
if evaluated >= sample_size and found:
break
row = _evaluate_contexts(
reps,
birth_date=birth_date,
domain=domain,
year=at.year,
month=at.month,
source="dasha_boundary",
clusters=clusters,
set_version=set_version,
)
evaluated += 1
if row is None or distinguish_contract_errors(row):
continue
if not isinstance(row.get("year"), int) or int(row["year"]) <= 0:
continue
if family is not None:
row = annotate_family(row, kind=family[0], family=family[1])
found.append(row)
if evaluated > sample_size:
break
kept = (
best_probe_per_year_family(found)[:max_per_domain]
if relax.r2
else _best_probe_per_year(found)[:max_per_domain]
)
if len(kept) < max_per_domain:
activation = _try_activation_probe(
reps=reps,
birth_date=birth_date,
birth_year=birth_year,
domain=domain,
now=today,
blocked_years=blocked_years | asked_years,
holdout_keys=set(holdout_keys),
clusters=clusters,
set_version=set_version,
)
if activation is not None:
activation_key = (
str(activation["domain"]),
int(activation["year"]),
int(activation.get("month") or 0),
str(activation["source"]),
)
existing = {
(str(row["domain"]), int(row["year"]), int(row.get("month") or 0), str(row["source"]))
for row in kept
}
if activation_key not in existing:
kept.append(activation)
probes.extend(kept)
probes.extend(_quality_distinguish_probes(
events,
clusters,
set_version=set_version,
holdout_ids=set(holdout_event_ids(events)),
holdout_keys=set(holdout_keys),
))
_annotate_nearby_ledger(probes, events)
probes.sort(key=_probe_sort_key)
return partition_probes(probes, max_probes=max_probes, include_family=relax.r2)
def outcome_groups(probe: dict[str, Any]) -> tuple[set[str], set[str]]:
yes: set[str] = set()
no: set[str] = set()
for row in probe.get("expected_outcomes") or []:
if not isinstance(row, dict):
continue
answer = str(row.get("answer_class") or "")
supports = {str(item)[:5] for item in (row.get("supports") or []) if str(item)}
if answer in {"yes", "weak_yes"}:
yes |= supports
elif answer == "no":
no |= supports
return yes, no
def inherit_direction(minute: str, yes: set[str], no: set[str]) -> str:
point = _clock(minute)
known = sorted(yes | no, key=_clock)
left = next((item for item in reversed(known) if _clock(item) < point), None)
right = next((item for item in known if _clock(item) > point), None)
if not left or not right:
return "neutral"
yes_times = [_clock(item) for item in yes]
no_times = [_clock(item) for item in no]
in_yes = left in yes and right in yes and yes_times and min(yes_times) <= point <= max(yes_times)
in_no = left in no and right in no and no_times and min(no_times) <= point <= max(no_times)
if in_yes == in_no:
return "neutral"
return "support" if in_yes else "conflict"
def true_side(probe: dict[str, Any], true_time: str) -> str:
yes, no = outcome_groups(probe)
if true_time in yes:
return "yes"
if true_time in no:
return "no"
direction = inherit_direction(true_time, yes, no)
if direction == "support":
return "yes"
if direction == "conflict":
return "no"
return "unknown"
def separates_true(
probe: dict[str, Any],
true_time: str,
remaining: Sequence[str],
clusters: Sequence[dict[str, Any]],
) -> bool:
target = displayed_true(true_time, remaining, clusters) or true_time
side = true_side(probe, target)
if side == "unknown":
side = true_side(probe, true_time)
if side == "unknown":
return False
yes, no = outcome_groups(probe)
remaining_set = {str(item)[:5] for item in remaining}
yes &= remaining_set
no &= remaining_set
if not yes or not no:
return False
if side == "yes":
return target in yes or inherit_direction(target, yes, no) == "support"
return target in no or inherit_direction(target, yes, no) == "conflict"
def split_tuple(probe: dict[str, Any], remaining: Sequence[str]) -> tuple[str, ...]:
yes, no = outcome_groups(probe)
remaining_set = {str(item)[:5] for item in remaining}
yes_kept = tuple(sorted(yes & remaining_set))
no_kept = tuple(sorted(no & remaining_set))
return yes_kept + ("|",) + no_kept
def apply_answer(
scores: dict[str, float],
conflicts: dict[str, int],
eliminated: set[str],
probe: dict[str, Any],
answer: str,
candidate_times: Sequence[str],
) -> tuple[dict[str, float], dict[str, int], set[str]]:
yes, no = outcome_groups(probe)
next_scores = dict(scores)
next_conflicts = dict(conflicts)
next_eliminated = set(eliminated)
if answer not in {"yes", "weak_yes", "no"}:
return next_scores, next_conflicts, next_eliminated
for time in candidate_times:
if time in next_eliminated:
continue
if time in yes:
raw = "support"
elif time in no:
raw = "conflict"
else:
raw = inherit_direction(time, yes, no)
if answer == "no":
raw = {"support": "conflict", "conflict": "support", "neutral": "neutral"}[raw]
if answer == "weak_yes":
direction = {"support": "weak_support", "conflict": "weak_conflict", "neutral": "neutral"}[raw]
else:
direction = raw
if direction == "conflict":
next_conflicts[time] = next_conflicts.get(time, 0) + 1
next_scores[time] = next_scores.get(time, 0.0) + SCORE_DELTA[direction]
newly = [
time for time in candidate_times
if time not in next_eliminated and next_conflicts.get(time, 0) >= STRONG_CONFLICT_ELIMINATION_COUNT
]
still_active = [
time for time in candidate_times
if time not in next_eliminated and time not in newly
]
survivor = None
if not still_active and newly:
survivor = sorted(newly, key=lambda item: (-next_scores.get(item, 0.0), item))[0]
for time in newly:
if time != survivor:
next_eliminated.add(time)
return next_scores, next_conflicts, next_eliminated
def optimal_answer(probe: dict[str, Any], true_time: str) -> str | None:
side = true_side(probe, true_time)
if side in {"yes", "no"}:
return side
return None
def range_width(times: Sequence[str]) -> int | None:
clocks = sorted(_clock(str(item)[:5]) for item in times if str(item)[:5])
if not clocks:
return None
return clocks[-1] - clocks[0] + 1
def displayed_true(true_time: str, remaining: Sequence[str], clusters: Sequence[dict[str, Any]]) -> str | None:
remaining_set = {str(item)[:5] for item in remaining}
if true_time in remaining_set:
return true_time
cluster = cluster_of(true_time, clusters)
if cluster is None:
return None
cluster_times = {str(item)[:5] for item in cluster.get("times") or []}
for time in remaining:
if str(time)[:5] in cluster_times:
return str(time)[:5]
return None
def top1_hit(
scores: dict[str, float],
active: Sequence[str],
true_time: str,
clusters: Sequence[dict[str, Any]],
) -> bool:
target = displayed_true(true_time, active, clusters)
if target is None:
return False
ranked = sorted(active, key=lambda item: (-scores.get(item, 0.0), item))
if not ranked:
return False
best = scores.get(ranked[0], 0.0)
leaders = [item for item in active if scores.get(item, 0.0) == best]
return leaders == [target]
def cluster_of(true_time: str, clusters: Sequence[dict[str, Any]]) -> dict[str, Any] | None:
for cluster in clusters:
times = [str(item)[:5] for item in cluster.get("times") or []]
if true_time in times:
return cluster
return None
def remaining_after_six(
*,
all_times: Sequence[str],
scores: dict[str, float],
eliminated: set[str],
clusters: Sequence[dict[str, Any]],
true_time: str,
) -> tuple[list[str], str, bool]:
active = [time for time in all_times if time not in eliminated]
true_alive = true_time in active
reps: list[str] = []
for cluster in clusters:
members = [str(item)[:5] for item in cluster.get("times") or [] if str(item)[:5] in active]
if not members:
continue
rep = str(cluster.get("representative_time") or "")[:5]
if not rep or rep not in members:
rep = max(members, key=lambda item: (scores.get(item, 0.0), item))
reps.append(rep)
if len(reps) <= REMAINING_CAP:
return reps, "natural", true_alive
picked: list[str] = []
true_cluster = cluster_of(true_time, clusters)
if true_alive and true_cluster is not None:
true_rep = str(true_cluster.get("representative_time") or "")[:5]
members = [str(item)[:5] for item in true_cluster.get("times") or [] if str(item)[:5] in active]
if true_rep not in members and members:
true_rep = max(members, key=lambda item: (scores.get(item, 0.0), item))
if true_rep in reps:
picked.append(true_rep)
ranked = sorted(reps, key=lambda item: (-scores.get(item, 0.0), item))
for time in ranked:
if time not in picked:
picked.append(time)
if len(picked) >= REMAINING_CAP:
break
return picked, "capped_to_5", true_alive
def asked_key(probe: dict[str, Any]) -> str:
return str(probe.get("semantic_key") or "")
def compact_probe(
probe: dict[str, Any],
true_time: str,
remaining: Sequence[str],
clusters: Sequence[dict[str, Any]],
) -> dict[str, Any]:
yes, no = outcome_groups(probe)
side = true_side(probe, true_time)
return {
"semantic_key": asked_key(probe),
"domain": probe.get("domain"),
"year": probe.get("year"),
"month": probe.get("month"),
"source": probe.get("source"),
"event_family": probe.get("event_family"),
"information_gain": probe.get("information_gain"),
"true_side": side,
"separates": separates_true(probe, true_time, remaining, clusters),
"yes_count": len(yes),
"no_count": len(no),
}
def measure_variant(
*,
probes: Sequence[dict[str, Any]],
asked_keys: set[str],
remaining: Sequence[str],
true_time: str,
base_scores: dict[str, float],
base_conflicts: dict[str, int],
base_eliminated: set[str],
all_times: Sequence[str],
clusters: Sequence[dict[str, Any]],
) -> dict[str, Any]:
new_probes = [
probe for probe in probes
if asked_key(probe) and asked_key(probe) not in asked_keys
]
discriminative = [
probe for probe in new_probes
if separates_true(probe, true_time, remaining, clusters)
]
unique_splits = {
split_tuple(probe, remaining)
for probe in discriminative
}
scores = dict(base_scores)
conflicts = dict(base_conflicts)
eliminated = set(base_eliminated)
applied = 0
skipped = 0
for probe in new_probes:
answer = optimal_answer(probe, true_time)
if answer is None:
skipped += 1
continue
scores, conflicts, eliminated = apply_answer(
scores, conflicts, eliminated, probe, answer, all_times,
)
applied += 1
active = [time for time in remaining if time not in eliminated]
if not active:
active = [time for time in all_times if time not in eliminated]
return {
"new_probe_count": len(new_probes),
"discriminative_count": len(discriminative),
"unique_split_count": len(unique_splits),
"applied_optimal": applied,
"skipped_unknown_side": skipped,
"top1_hit": top1_hit(scores, active, true_time, clusters),
"range_width": range_width(active),
"active_count": len(active),
"true_alive": true_time not in eliminated,
"probes": [compact_probe(probe, true_time, remaining, clusters) for probe in new_probes],
}
def score_case(case: dict[str, Any]) -> dict[str, Any]:
true_time = str(case["birth"]["time"])[:5]
request = request_from_case(case)
scoring = scoreable_request(request)
built = build_event_contribution_matrix(scoring)
rows = score_from_matrix(scoring, built)
all_times = [str(row["time"])[:5] for row in rows]
prior = {str(row["time"])[:5]: float(row["score"] or 0) for row in rows}
contexts = _static_contexts(built)
clusters = cluster_contexts_by_signature(contexts)
scan = _scan_for(built)
initial = discriminating_event_probes(
{**request, "refresh_probes": False},
built,
scan=scan,
candidate_times=all_times,
representative_time=true_time,
today=TODAY,
)
asked = initial[:ASK_COUNT]
scores = dict(prior)
conflicts = {time: 0 for time in all_times}
eliminated: set[str] = set()
for probe in asked:
answer = optimal_answer(probe, true_time)
if answer is None:
continue
scores, conflicts, eliminated = apply_answer(
scores, conflicts, eliminated, probe, answer, all_times,
)
remaining, remaining_mode, true_alive = remaining_after_six(
all_times=all_times,
scores=scores,
eliminated=eliminated,
clusters=clusters,
true_time=true_time,
)
asked_keys = {asked_key(probe) for probe in asked if asked_key(probe)}
refresh_request = {
**request,
"refresh_probes": True,
"asked_probe_keys": sorted(asked_keys),
"column_times": remaining,
}
baseline_after_six = {
"asked_count": len(asked),
"initial_probe_count": len(initial),
"remaining_count": len(remaining),
"remaining_mode": remaining_mode,
"true_alive": true_alive,
"top1_hit": top1_hit(scores, remaining, true_time, clusters),
"range_width": range_width(remaining),
"asked_keys": sorted(asked_keys),
"remaining_times": remaining,
}
variants: dict[str, Any] = {}
for relax in all_relaxations():
if relax.name == "baseline":
probes = discriminating_event_probes(
refresh_request,
built,
scan=scan,
candidate_times=remaining,
representative_time=remaining[0] if remaining else true_time,
today=TODAY,
)
else:
probes = generate_refresh_probes(
refresh_request,
built,
remaining_times=remaining,
relax=relax,
today=TODAY,
)
variants[relax.name] = measure_variant(
probes=probes,
asked_keys=asked_keys,
remaining=remaining,
true_time=true_time,
base_scores=scores,
base_conflicts=conflicts,
base_eliminated=eliminated,
all_times=all_times,
clusters=clusters,
)
variants[relax.name]["refresh_probe_count"] = len(probes)
variants[relax.name]["refresh_keys"] = [asked_key(item) for item in probes]
return {
"case_id": case["case_id"],
"true_time": true_time,
"radius": int(case.get("candidate_radius_minutes") or 10),
"cluster_count": len(clusters),
"after_six": baseline_after_six,
"variants": variants,
}
def summarize(cases: list[dict[str, Any]]) -> dict[str, Any]:
eligible = [
row for row in cases
if not row.get("error")
and int((row.get("after_six") or {}).get("remaining_count") or 0) >= 2
and int((row.get("after_six") or {}).get("asked_count") or 0) >= ASK_COUNT
]
table: dict[str, Any] = {}
names = [relax.name for relax in all_relaxations()]
after_six_hit = [
bool(row["after_six"]["top1_hit"]) for row in eligible
]
after_six_width = [
int(row["after_six"]["range_width"])
for row in eligible
if row["after_six"].get("range_width") is not None
]
for name in names:
new_counts = [int(row["variants"][name]["new_probe_count"]) for row in eligible]
disc = [int(row["variants"][name]["discriminative_count"]) for row in eligible]
unique = [int(row["variants"][name]["unique_split_count"]) for row in eligible]
hits = [bool(row["variants"][name]["top1_hit"]) for row in eligible]
widths = [
int(row["variants"][name]["range_width"])
for row in eligible
if row["variants"][name].get("range_width") is not None
]
hit_rate = round(sum(hits) / len(hits), 4) if hits else None
base_hit = round(sum(after_six_hit) / len(after_six_hit), 4) if after_six_hit else None
table[name] = {
"n": len(eligible),
"mean_new_probes": round(statistics.mean(new_counts), 3) if new_counts else 0.0,
"mean_discriminative": round(statistics.mean(disc), 3) if disc else 0.0,
"mean_unique_splits": round(statistics.mean(unique), 3) if unique else 0.0,
"cases_with_any_discriminative": sum(item > 0 for item in disc),
"top1_after_six": base_hit,
"top1_after_replay": hit_rate,
"top1_delta": None if hit_rate is None or base_hit is None else round(hit_rate - base_hit, 4),
"mean_width_after_six": round(statistics.mean(after_six_width), 2) if after_six_width else None,
"mean_width_after_replay": round(statistics.mean(widths), 2) if widths else None,
}
baseline = table.get("baseline") or {}
for name, row in table.items():
if name == "baseline":
row["extra_new_vs_baseline"] = 0.0
row["extra_unique_vs_baseline"] = 0.0
continue
row["extra_new_vs_baseline"] = round(
float(row["mean_new_probes"]) - float(baseline.get("mean_new_probes") or 0), 3,
)
row["extra_unique_vs_baseline"] = round(
float(row["mean_unique_splits"]) - float(baseline.get("mean_unique_splits") or 0), 3,
)
return {
"case_count": len(cases),
"eligible_count": len(eligible),
"errors": [row["case_id"] for row in cases if row.get("error")],
"variants": table,
}
def decide(summary: dict[str, Any]) -> dict[str, Any]:
eligible = int(summary.get("eligible_count") or 0)
if eligible < 10:
return {
"verdict": "uncertain",
"reason": f"only {eligible} cases asked 6 dated probes and still had a remaining set",
"implement": False,
}
winners: list[str] = []
for name, row in summary["variants"].items():
if name == "baseline":
continue
extra = float(row.get("extra_unique_vs_baseline") or 0)
delta = row.get("top1_delta")
if extra >= 1.0 and delta is not None and delta >= 0:
winners.append(name)
if winners:
return {
"verdict": "benefit",
"reason": "at least one relaxation added >=1 unique true-cluster split vs production refresh without dropping top-1",
"implement": True,
"variants": winners,
}
return {
"verdict": "no_benefit",
"reason": "no relaxation added >=1 unique true-cluster split vs production refresh with non-falling top-1",
"implement": False,
}
def render_markdown(report: dict[str, Any]) -> str:
summary = report["summary"]
decision = report["decision"]
lines = [
"# 六题之后刷新出题供给测量(2026-09-13",
"",
"- 任务:`docs/tasks/TASK-rectification-probe-supply-research-20260913.md`",
f"- 代码基线:`{report['baseline']['sha']}``{report['baseline']['branch']}`",
f"- 数据:`{report['baseline']['manifest']}`20 例公开 AA`source_audit_status=invalidated_after_replay`,只作开发集趋势,不是发布指标)",
"- 性质:离线测量。生产 `event_probes.py` 默认行为未改。",
"",
"## 方法",
"",
"1. 每例用 holdout 公开事件对声明分钟 ±`candidate_radius_minutes` 打分,生成生产路径带年月题。",
"2. 按信息增益取前 6 道,用对真实分钟最优的 yes/no 回放(真实分钟落在 yes 组答 yes,落在 no 组答 no)。",
"3. 剩余活动簇代表若多于 5 个,保留真实簇并按后验截到 5 个,模拟「六题后剩余 ≤5 候选」。",
"4. 对该剩余集按 `refresh_probes=true` 再生成题:baseline 走生产函数;R1–R4 只在本脚本里临时改封锁年、事件家族、`MIN_BOUNDARY_DAYS`、Vimshottari 第三级与 D9/D10 上升 Narayana。",
"5. 新题按最优答案继续回放,比较头名命中与剩余范围宽度。",
"",
"有分辨力:真实分钟(或其簇代表)能落到 yes 或 no,且另一组里还有剩余候选。相同 yes/no 划分只计一次 unique split。收益门槛是相对生产刷新 **多出** ≥1 道独立划分,且头名不降。",
"",
"## 放宽项",
"",
"| 代号 | 改法 |",
"| --- | --- |",
"| R1 | 刷新时只封已问/已知的确切年,去掉 ±1 年封锁 |",
"| R2 | 事业按入职/升职/离职、感情按结婚/分手,同域同年可再出一道 |",
"| R3 | 刷新阶段 `MIN_BOUNDARY_DAYS` 45→30 |",
"| R4 | 边界加 Vimshottari pratyantar,以及 D9/D10 上升的 Narayana |",
"",
f"问满 6 道且剩余 ≥2 因而进入刷新测量的例子:{summary['eligible_count']}/{summary['case_count']}。",
"",
"## 总表",
"",
"| 方案 | 平均新增题 | 相对生产多出题 | 平均有分辨力 | 平均独立划分 | 相对生产多出划分 | 有分辨力的例子 | 六题后头名 | 回放后头名 | 头名差 | 六题后宽度 | 回放后宽度 |",
"| --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: |",
]
for name, row in summary["variants"].items():
lines.append(
"| {name} | {mean_new_probes} | {extra_new_vs_baseline} | {mean_discriminative} | {mean_unique_splits} | {extra_unique_vs_baseline} | {cases_with_any_discriminative}/{n} | {top1_after_six} | {top1_after_replay} | {top1_delta} | {mean_width_after_six} | {mean_width_after_replay} |".format(
name=name,
**{key: row[key] for key in (
"mean_new_probes", "extra_new_vs_baseline", "mean_discriminative",
"mean_unique_splits", "extra_unique_vs_baseline",
"cases_with_any_discriminative", "n", "top1_after_six",
"top1_after_replay", "top1_delta", "mean_width_after_six",
"mean_width_after_replay",
)},
)
)
verdict_text = {
"benefit": "有收益,可另立实现单。",
"no_benefit": "无收益,关闭方案 B。",
"uncertain": "不确定,还缺足够进入「六题后剩余 ≤5」的例子。",
}[decision["verdict"]]
lines.extend([
"",
"## 结论",
"",
f"**{verdict_text}**",
"",
f"- 判定:`{decision['verdict']}`",
f"- 原因:{decision['reason']}",
f"- 立实现单:{'是' if decision['implement'] else '否'}",
])
if decision.get("variants"):
lines.append(f"- 达到门槛的组合:{', '.join(decision['variants'])}")
if any(name in {"R3", "R4"} or "R3" in name.split("+") or "R4" in name.split("+") for name in decision["variants"]):
r3 = summary["variants"].get("R3") or {}
r4 = summary["variants"].get("R4") or {}
r34 = summary["variants"].get("R3+R4") or {}
r1 = summary["variants"].get("R1") or {}
r2 = summary["variants"].get("R2") or {}
lines.extend([
"",
"## 建议实现范围",
"",
"立实现单只做 **刷新阶段 R3 + R4**",
"",
f"- R3:刷新时 `MIN_BOUNDARY_DAYS` 45→30。单独达到门槛(相对生产 +{r3.get('extra_unique_vs_baseline')} 独立划分、+{r3.get('extra_new_vs_baseline')} 题),改动最小。",
f"- R4:刷新边界加 Vimshottari pratyantar 与 D9/D10 上升 Narayana。单独 +{r4.get('extra_unique_vs_baseline')} 独立划分、+{r4.get('extra_new_vs_baseline')} 题。`_vim_start_dates` 已有 `include_pratyantar`,生产默认仍为 false。",
f"- 合做 R3+R4+{r34.get('extra_unique_vs_baseline')} 独立划分、+{r34.get('extra_new_vs_baseline')} 题;头名 {r34.get('top1_after_replay')} 不降,宽度 {r34.get('mean_width_after_six')}{r34.get('mean_width_after_replay')}。",
f"- **不做 R1**:只多 {r1.get('extra_unique_vs_baseline')} 划分。",
f"- **不做 R2**:多出的题大多是同一 yes/no 划分的家族复题;单独只多 {r2.get('extra_unique_vs_baseline')} 划分。",
"- 即使 R3+R4,仍有例子问满 6 道后刷新还是 0 题。定向补事(方案 A)仍要保留。",
])
lines.extend([
"",
"## 分例",
"",
"| 例子 | 问了几道 | 剩余 | 截断 | 真实还在 | baseline 新题 | 最好组合 | 该组合新题 | 该组合分辨力 |",
"| --- | ---: | ---: | --- | --- | ---: | --- | ---: | ---: |",
])
for row in report["cases"]:
if row.get("error"):
lines.append(f"| {row['case_id']} | error | | | | | | | {row.get('error')} |")
continue
after = row["after_six"]
best_name = "baseline"
best = row["variants"]["baseline"]
for name, variant in row["variants"].items():
if int(variant["discriminative_count"]) > int(best["discriminative_count"]):
best_name = name
best = variant
lines.append(
f"| {row['case_id']} | {after['asked_count']} | {after['remaining_count']} | {after['remaining_mode']} | {after['true_alive']} | {row['variants']['baseline']['new_probe_count']} | {best_name} | {best['new_probe_count']} | {best['discriminative_count']} |"
)
lines.extend([
"",
"## 边界",
"",
"- 最优答案是相对已公布 Rodden AA 分钟的上界,不是真实用户会怎么答。",
"- holdout 已被 `invalidated_after_replay`,不得写成发布准确率。",
"- R2 若只换事件家族文案、分盘激活规则不变,yes/no 划分可能与默认家族重复。",
"- 未把研究脚本接到生产路径。",
"",
])
return "\n".join(lines) + "\n"
def git_sha() -> str:
head = ROOT / ".git"
try:
import subprocess
return subprocess.check_output(
["git", "rev-parse", "HEAD"],
cwd=ROOT,
text=True,
).strip()
except Exception:
return str(head)
def run(limit: int | None = None, case_id: str | None = None) -> dict[str, Any]:
payload = json.loads(HOLDOUT_MANIFEST.read_text(encoding="utf-8"))
cases = [item for item in payload.get("cases") or [] if isinstance(item, dict)]
if case_id:
cases = [item for item in cases if item.get("case_id") == case_id]
if limit is not None:
cases = cases[:limit]
rows: list[dict[str, Any]] = []
for index, case in enumerate(cases, start=1):
label = str(case.get("case_id") or index)
print(f"[{index}/{len(cases)}] {label}", flush=True)
try:
rows.append(score_case(case))
except Exception as exc:
rows.append({
"case_id": str(case.get("case_id") or ""),
"error": f"{type(exc).__name__}: {exc}",
"traceback": traceback.format_exc(),
})
summary = summarize(rows)
decision = decide(summary)
return {
"scope": "probe_supply_after_six",
"today": TODAY.isoformat(),
"baseline": {
"sha": git_sha(),
"branch": "codex/rectification-probe-supply-research-20260913",
"manifest": str(HOLDOUT_MANIFEST.relative_to(ROOT)),
"benchmark_id": payload.get("benchmark_id"),
"min_boundary_days_default": MIN_BOUNDARY_DAYS,
"existence_nearby_years": dict(EXISTENCE_NEARBY_YEARS),
},
"summary": summary,
"decision": decision,
"cases": rows,
}
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--from-json", type=Path, default=None)
parser.add_argument("--limit", type=int, default=None)
parser.add_argument("--case-id", type=str, default=None)
parser.add_argument("--json-out", type=Path, default=ROOT / "docs/research/probe_supply_after_six_2026_09_13.json")
parser.add_argument("--md-out", type=Path, default=ROOT / "docs/research/probe_supply_after_six_2026_09_13.md")
args = parser.parse_args()
if args.from_json:
previous = json.loads(args.from_json.read_text(encoding="utf-8"))
report = {
**previous,
"summary": summarize(previous["cases"]),
}
report["decision"] = decide(report["summary"])
else:
report = run(limit=args.limit, case_id=args.case_id)
args.json_out.parent.mkdir(parents=True, exist_ok=True)
args.json_out.write_text(json.dumps(report, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
args.md_out.write_text(render_markdown(report), encoding="utf-8")
print(json.dumps({
"decision": report["decision"],
"eligible_count": report["summary"]["eligible_count"],
"errors": report["summary"]["errors"],
"json_out": str(args.json_out),
"md_out": str(args.md_out),
}, ensure_ascii=False, indent=2))
return 0 if not report["summary"]["errors"] else 1
if __name__ == "__main__":
raise SystemExit(main())