Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_017eEAG8HD3mm8gsKXgk8uU8
821 lines
33 KiB
Python
821 lines
33 KiB
Python
"""Scaffold (R-0) for TASK-rectification-scoring-research-20260929.
|
||
|
||
Offline only. Nothing here changes a production default or a frozen scoring
|
||
file. The module records, for every (case, candidate minute, event, sample
|
||
date), which scoring rules fire — the eleven production rules from
|
||
`active_rectification_event_engine._score_event` plus the auxiliary
|
||
transit / Ashtakavarga / Shadbala rules — and a set of *extra* observations
|
||
that production computes but does not score (KP cusp sub-lords, Pranapada /
|
||
Hora / Ghati / Bhava lagna, D60). Likelihood-ratio weights are estimated per
|
||
feature with leave-one-**case**-out folds, and a weighted row provider feeds
|
||
the same contribution-matrix / probe / six-question replay machinery that the
|
||
09-14 / 09-26 / 09-29 studies used.
|
||
|
||
Recording provider vs production: the provider returns exactly the production
|
||
evidence (same rule ids, same points), so `build_event_contribution_matrix`
|
||
yields the same matrix as the production path. `reconcile_rows` asserts that
|
||
per case before any experiment (task hard line 4).
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
import math
|
||
import statistics
|
||
import sys
|
||
from collections import defaultdict
|
||
from dataclasses import dataclass, field
|
||
from datetime import date, datetime
|
||
from pathlib import Path
|
||
from random import Random
|
||
from typing import Any, Callable, Sequence
|
||
|
||
ROOT = Path(__file__).resolve().parents[2]
|
||
if str(ROOT) not in sys.path:
|
||
sys.path.insert(0, str(ROOT))
|
||
|
||
from scripts.active_rectification_event_engine import ( # noqa: E402
|
||
DOMAIN_CONFIG,
|
||
OBSERVATION_ONLY_LAYERS,
|
||
_active_narayana,
|
||
_active_vimshottari,
|
||
_ashtakavarga_auxiliary,
|
||
_controlled_transit_rules,
|
||
_event_datetime,
|
||
_relative_house,
|
||
_shadbala_verified_components_auxiliary,
|
||
_varga_chart,
|
||
_varga_house,
|
||
)
|
||
import scripts.rectification.event_probes as event_probes # noqa: E402
|
||
from scripts.rectification.candidate_contrast import cluster_contexts_by_signature # noqa: E402
|
||
from scripts.rectification.case_holdout import holdout_event_ids # noqa: E402
|
||
from scripts.rectification.contracts import is_scoreable_event # 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,
|
||
)
|
||
from scripts.research.cluster_width_lib import ( # noqa: E402
|
||
SEPARATION_LEAD,
|
||
delivery_from_public,
|
||
merge_adjacent_traced,
|
||
metrics_bundle,
|
||
public_from_clusters,
|
||
raw_signature_clusters,
|
||
shannon_entropy,
|
||
still_valid_public,
|
||
)
|
||
from scripts.research.precision_gate_lib import ( # noqa: E402
|
||
VargaPolicy,
|
||
score_event_with_policy,
|
||
)
|
||
from scripts.research.probe_supply_after_six import ( # noqa: E402
|
||
ASK_COUNT,
|
||
SCORE_DELTA,
|
||
apply_answer,
|
||
inherit_direction,
|
||
optimal_answer,
|
||
outcome_groups,
|
||
)
|
||
import varga # noqa: E402
|
||
|
||
TODAY = date(2026, 9, 16) # same probe "today" as the 09-16 / 09-29 replays
|
||
LAPLACE_ALPHA = 1.0
|
||
EPS = 1e-9
|
||
EXTRA_PREFIXES = ("kp_", "pranapada_", "hora_lagna_", "ghati_lagna_", "bhava_lagna_")
|
||
D60_SUFFIX = "_domain_varga_d60"
|
||
IGNORED_RULE_PREFIXES = ("event_kind:", "event_kind_profile:", "no_domain_activation")
|
||
CONSTANT_RULES = frozenset({"occupation_auxiliary_not_primary", "appearance_auxiliary_not_primary"})
|
||
VARIANTS = ("A1", "A2", "A3")
|
||
VARIANT_SPEC = {
|
||
# (uses extra features, allows negative / absence weights)
|
||
"A1": (False, False),
|
||
"A2": (True, False),
|
||
"A3": (True, True),
|
||
}
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# small helpers
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def hhmm(value: object) -> str | None:
|
||
text = str(value or "")[:5]
|
||
return text if len(text) == 5 and text[2] == ":" else None
|
||
|
||
|
||
def clock(value: str) -> int:
|
||
return int(value[:2]) * 60 + int(value[3:5])
|
||
|
||
|
||
def is_extra_feature(name: str) -> bool:
|
||
return name.startswith(EXTRA_PREFIXES) or name.endswith(D60_SUFFIX)
|
||
|
||
|
||
def is_scoring_rule(name: str) -> bool:
|
||
if name.startswith(IGNORED_RULE_PREFIXES) or name in CONSTANT_RULES:
|
||
return False
|
||
return not is_extra_feature(name)
|
||
|
||
|
||
def median(values: Sequence[float]) -> float | None:
|
||
items = [float(item) for item in values if item is not None]
|
||
return statistics.median(items) if items else None
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# R-0 · feature recording provider (production-identical evidence + extras)
|
||
# ---------------------------------------------------------------------------
|
||
|
||
@dataclass
|
||
class FeatureStore:
|
||
"""(event_id, sample_date) -> HH:MM -> {rules, extras}."""
|
||
|
||
rows: dict[tuple[str, str], dict[str, dict[str, Any]]] = field(default_factory=dict)
|
||
|
||
def put(self, event_id: str, sample: str, time: str, rules: Sequence[str], extras: Sequence[str]) -> None:
|
||
self.rows.setdefault((event_id, sample), {})[time] = {
|
||
"rules": sorted(set(rules)),
|
||
"extras": sorted(set(extras)),
|
||
}
|
||
|
||
def samples_for(self, event_id: str) -> list[str]:
|
||
return sorted(sample for (item, sample) in self.rows if item == event_id)
|
||
|
||
|
||
def d60_charts(contexts: Sequence[dict[str, Any]]) -> dict[str, dict[str, Any] | None]:
|
||
out: dict[str, dict[str, Any] | None] = {}
|
||
for context in contexts:
|
||
stamp = context["candidate_at"].strftime("%H:%M")
|
||
try:
|
||
computed = varga.calc_all_vargas(
|
||
context["planet_longitudes"], float(context["ascendant_longitude"]), divisions=[60],
|
||
)
|
||
out[stamp] = _varga_chart(computed, "D60")
|
||
except (KeyError, TypeError, ValueError):
|
||
out[stamp] = None
|
||
return out
|
||
|
||
|
||
def extra_features(
|
||
context: dict[str, Any],
|
||
target_houses: tuple[int, ...],
|
||
vimshottari: tuple[str, str, str],
|
||
d60: dict[str, Any] | None,
|
||
) -> list[str]:
|
||
feature = context.get("feature") if isinstance(context.get("feature"), dict) else {}
|
||
ascendant_index = int(context["ascendant_index"])
|
||
kp_houses = ((feature.get("kp_cusps") or {}).get("houses") or {})
|
||
target_sub = {str((kp_houses.get(str(h)) or {}).get("sub_lord") or "") for h in target_houses}
|
||
target_star = {str((kp_houses.get(str(h)) or {}).get("nakshatra_lord") or "") for h in target_houses}
|
||
asc_sub = str((kp_houses.get("1") or {}).get("sub_lord") or "")
|
||
out: list[str] = []
|
||
for lord, label in zip(vimshottari, ("md", "ad", "pd")):
|
||
if lord in target_sub:
|
||
out.append(f"kp_target_cusp_sublord_is_vim_{label}")
|
||
if lord in target_star:
|
||
out.append(f"kp_target_cusp_starlord_is_vim_{label}")
|
||
if asc_sub and lord == asc_sub:
|
||
out.append(f"kp_asc_sublord_is_vim_{label}")
|
||
if isinstance(d60, dict) and _varga_house(d60, lord) in target_houses:
|
||
out.append(f"vim_{label}{D60_SUFFIX}")
|
||
for key, name in (
|
||
("pranapada_sign_index", "pranapada"),
|
||
("hora_sign_index", "hora_lagna"),
|
||
("ghati_sign_index", "ghati_lagna"),
|
||
("bhava_sign_index", "bhava_lagna"),
|
||
):
|
||
index = feature.get(key)
|
||
if isinstance(index, int) and _relative_house(index, ascendant_index) in target_houses:
|
||
out.append(f"{name}_in_target_house")
|
||
return out
|
||
|
||
|
||
def recorded_row(
|
||
request: dict[str, Any],
|
||
context: dict[str, Any],
|
||
*,
|
||
store: FeatureStore,
|
||
transit_cache: dict[tuple[Any, ...], dict[str, Any]],
|
||
d60_by_time: dict[str, dict[str, Any] | None],
|
||
) -> dict[str, Any]:
|
||
"""Production `_candidate_row` for one legacy (single event, one date) request, plus extras."""
|
||
candidate_at = context["candidate_at"]
|
||
chart = context["chart"]
|
||
planet_longitudes = context["planet_longitudes"]
|
||
ascendant_index = context["ascendant_index"]
|
||
arudha_padas = context["arudha_padas"]
|
||
varga_charts = context["varga_charts"]
|
||
moon_longitude = planet_longitudes["Moon"]
|
||
stamp = candidate_at.strftime("%H:%M")
|
||
policy = VargaPolicy(name="V0", window_minutes=0.0)
|
||
evidence: list[dict[str, Any]] = []
|
||
missing: list[str] = []
|
||
for event in request["events"]:
|
||
event_at = _event_datetime(event)
|
||
prefixes, target_houses = DOMAIN_CONFIG[event["domain"]]
|
||
selected = {prefix: varga_charts.get(prefix) for prefix in prefixes}
|
||
if any(item is None for item in selected.values()):
|
||
missing.extend(prefixes)
|
||
continue
|
||
try:
|
||
vimshottari = _active_vimshottari(
|
||
candidate_at.date().isoformat(), moon_longitude, event_at, context.get("vimshottari_timeline"),
|
||
)
|
||
except (KeyError, TypeError, ValueError):
|
||
missing.append("Vimshottari_MD_AD_PD")
|
||
continue
|
||
try:
|
||
narayana = _active_narayana(
|
||
ascendant_index, planet_longitudes, candidate_at, event_at, context.get("narayana_periods"),
|
||
)
|
||
except (KeyError, TypeError, ValueError):
|
||
missing.append("Narayana_MD_AD")
|
||
continue
|
||
if narayana[0] is None or narayana[1] is None:
|
||
missing.append("Narayana_MD_AD")
|
||
continue
|
||
row = score_event_with_policy(
|
||
candidate_time=stamp, event=event, natal_chart=chart,
|
||
varga_by_prefix={k: v for k, v in selected.items() if isinstance(v, dict)},
|
||
vimshottari=vimshottari, narayana=narayana, arudha_padas=arudha_padas, policy=policy,
|
||
)
|
||
weight = 1.0 # legacy requests carry precision "day"; the matrix applies the real weight
|
||
transit_rules = _controlled_transit_rules(
|
||
request, event, ascendant_index, target_houses, transit_cache,
|
||
)
|
||
if transit_rules:
|
||
row["rule_ids"].extend(transit_rules)
|
||
row["points"] = round(row["points"] + 0.25 * len(transit_rules) * weight, 4)
|
||
av_rules, av_points = _ashtakavarga_auxiliary(
|
||
chart, ascendant_index, target_houses, context.get("ashtakavarga_result"),
|
||
)
|
||
if av_rules:
|
||
row["rule_ids"].extend(av_rules)
|
||
row["points"] = round(row["points"] + av_points * weight, 4)
|
||
sb_rules, sb_points = _shadbala_verified_components_auxiliary(
|
||
chart, candidate_at.hour + candidate_at.minute / 60, vimshottari, context.get("shadbala_result"),
|
||
)
|
||
if sb_rules:
|
||
row["rule_ids"].extend(sb_rules)
|
||
row["points"] = round(row["points"] + sb_points * weight, 4)
|
||
evidence.append(row)
|
||
store.put(
|
||
str(event["id"]), str(event["date"]), stamp, row["rule_ids"],
|
||
extra_features(context, target_houses, vimshottari, d60_by_time.get(stamp)),
|
||
)
|
||
return {
|
||
"time": stamp,
|
||
"score": round(sum(item["points"] for item in evidence), 4),
|
||
"evidence": evidence,
|
||
"missing_layers": sorted(set(
|
||
missing + [layer for layer in context["feature"]["blocked_layers"] if layer not in OBSERVATION_ONLY_LAYERS]
|
||
)),
|
||
}
|
||
|
||
|
||
def make_recording_provider(
|
||
contexts: Sequence[dict[str, Any]],
|
||
store: FeatureStore,
|
||
d60_by_time: dict[str, dict[str, Any] | None],
|
||
) -> Callable[[dict[str, Any]], list[dict[str, Any]]]:
|
||
transit_cache: dict[tuple[Any, ...], dict[str, Any]] = {}
|
||
|
||
def provider(request: dict[str, Any]) -> list[dict[str, Any]]:
|
||
return [
|
||
recorded_row(request, context, store=store, transit_cache=transit_cache, d60_by_time=d60_by_time)
|
||
for context in contexts
|
||
]
|
||
|
||
return provider
|
||
|
||
|
||
def score_map(rows: Sequence[dict[str, Any]]) -> dict[str, float]:
|
||
return {str(row["time"])[:5]: float(row.get("score") or 0) for row in rows}
|
||
|
||
|
||
def reconcile_rows(production: Sequence[dict[str, Any]], research: Sequence[dict[str, Any]]) -> dict[str, Any]:
|
||
"""Per-minute score and rule-id comparison. `changed == 0` is the R-0 gate."""
|
||
prod = {str(r["time"])[:5]: r for r in production}
|
||
rese = {str(r["time"])[:5]: r for r in research}
|
||
changed: list[dict[str, Any]] = []
|
||
for time, row in prod.items():
|
||
other = rese.get(time)
|
||
if other is None:
|
||
changed.append({"time": time, "reason": "missing"})
|
||
continue
|
||
if abs(float(row.get("score") or 0) - float(other.get("score") or 0)) > EPS:
|
||
changed.append({"time": time, "reason": "score", "production": row.get("score"), "research": other.get("score")})
|
||
continue
|
||
prod_rules = {
|
||
(item["event_id"], tuple(sorted(item["rule_ids"]))) for item in row.get("evidence") or []
|
||
}
|
||
rese_rules = {
|
||
(item["event_id"], tuple(sorted(item["rule_ids"]))) for item in other.get("evidence") or []
|
||
}
|
||
if prod_rules != rese_rules:
|
||
changed.append({"time": time, "reason": "rule_ids"})
|
||
return {"candidates": len(prod), "changed": len(changed), "details": changed[:5]}
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# feature matrices and likelihood-ratio weights
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def event_feature_matrix(
|
||
store: FeatureStore,
|
||
event_id: str,
|
||
times: Sequence[str],
|
||
*,
|
||
include_extras: bool,
|
||
) -> dict[str, dict[str, float]]:
|
||
"""HH:MM -> feature -> fraction of sample dates on which it fired."""
|
||
samples = store.samples_for(event_id)
|
||
out: dict[str, dict[str, float]] = {}
|
||
if not samples:
|
||
return {time: {} for time in times}
|
||
for time in times:
|
||
counts: dict[str, int] = defaultdict(int)
|
||
for sample in samples:
|
||
cell = store.rows.get((event_id, sample), {}).get(time)
|
||
if not cell:
|
||
continue
|
||
names = [name for name in cell["rules"] if is_scoring_rule(name)]
|
||
if include_extras:
|
||
names += list(cell["extras"])
|
||
for name in names:
|
||
counts[name] += 1
|
||
out[time] = {name: round(count / len(samples), 6) for name, count in counts.items()}
|
||
return out
|
||
|
||
|
||
def training_event_ids(request: dict[str, Any]) -> list[str]:
|
||
holdout = holdout_event_ids(request["events"])
|
||
return [str(e["id"]) for e in request["events"] if is_scoreable_event(e) and str(e["id"]) not in holdout]
|
||
|
||
|
||
@dataclass
|
||
class LRTable:
|
||
alpha: float
|
||
positives: float
|
||
negatives: float
|
||
present: dict[str, float] # log LR of the feature firing
|
||
absent: dict[str, float] # log LR of the feature not firing (naive Bayes complement)
|
||
support: dict[str, dict[str, float]]
|
||
|
||
def weight(self, name: str, *, allow_negative: bool) -> tuple[float, float]:
|
||
"""(present weight, absent weight) under the variant policy."""
|
||
present = self.present.get(name, 0.0)
|
||
absent = self.absent.get(name, 0.0)
|
||
if allow_negative:
|
||
return present, absent
|
||
return max(present, 0.0), 0.0
|
||
|
||
|
||
def estimate_lr(examples: Sequence[tuple[dict[str, float], bool]], *, alpha: float = LAPLACE_ALPHA) -> LRTable:
|
||
"""Per-feature log likelihood ratios from (feature fractions, is-truth-cluster) examples."""
|
||
positives = float(sum(1 for _values, truth in examples if truth))
|
||
negatives = float(sum(1 for _values, truth in examples if not truth))
|
||
names: set[str] = set()
|
||
pos_sum: dict[str, float] = defaultdict(float)
|
||
neg_sum: dict[str, float] = defaultdict(float)
|
||
for values, truth in examples:
|
||
for name, value in values.items():
|
||
names.add(name)
|
||
if truth:
|
||
pos_sum[name] += float(value)
|
||
else:
|
||
neg_sum[name] += float(value)
|
||
present: dict[str, float] = {}
|
||
absent: dict[str, float] = {}
|
||
support: dict[str, dict[str, float]] = {}
|
||
for name in sorted(names):
|
||
p_pos = (pos_sum[name] + alpha) / (positives + 2 * alpha)
|
||
p_neg = (neg_sum[name] + alpha) / (negatives + 2 * alpha)
|
||
present[name] = round(math.log(p_pos / p_neg), 6)
|
||
absent[name] = round(math.log((1 - p_pos) / (1 - p_neg)), 6)
|
||
support[name] = {
|
||
"fires_truth": round(pos_sum[name], 4),
|
||
"fires_other": round(neg_sum[name], 4),
|
||
"p_truth": round(p_pos, 6),
|
||
"p_other": round(p_neg, 6),
|
||
}
|
||
return LRTable(alpha=alpha, positives=positives, negatives=negatives, present=present, absent=absent, support=support)
|
||
|
||
|
||
def bootstrap_lr(
|
||
per_case_examples: dict[str, list[tuple[dict[str, float], bool]]],
|
||
*,
|
||
resamples: int = 200,
|
||
seed: int = 20260929,
|
||
alpha: float = LAPLACE_ALPHA,
|
||
) -> dict[str, dict[str, float]]:
|
||
"""Case-level bootstrap 5–95 % interval of the present-log-LR per feature."""
|
||
ids = sorted(per_case_examples)
|
||
rng = Random(f"{seed}:lr-bootstrap")
|
||
draws: dict[str, list[float]] = defaultdict(list)
|
||
for _ in range(resamples):
|
||
picked = [rng.choice(ids) for _ in ids]
|
||
table = estimate_lr([ex for cid in picked for ex in per_case_examples[cid]], alpha=alpha)
|
||
for name, value in table.present.items():
|
||
draws[name].append(value)
|
||
out: dict[str, dict[str, float]] = {}
|
||
for name, values in draws.items():
|
||
ordered = sorted(values)
|
||
lo = ordered[int(0.05 * (len(ordered) - 1))]
|
||
hi = ordered[int(0.95 * (len(ordered) - 1))]
|
||
out[name] = {"p05": round(lo, 4), "p95": round(hi, 4), "draws": len(ordered)}
|
||
return out
|
||
|
||
|
||
def weighted_points(values: dict[str, float], table: LRTable, *, allow_negative: bool, universe: Sequence[str]) -> float:
|
||
total = 0.0
|
||
for name in universe:
|
||
present, absent = table.weight(name, allow_negative=allow_negative)
|
||
x = float(values.get(name, 0.0))
|
||
total += x * present + (1.0 - x) * absent
|
||
return total
|
||
|
||
|
||
def make_weighted_provider(
|
||
store: FeatureStore,
|
||
table: LRTable,
|
||
*,
|
||
variant: str,
|
||
scale: float = 1.0,
|
||
baseline_rows_by_sample: dict[tuple[str, str], dict[str, dict[str, Any]]] | None = None,
|
||
) -> Callable[[dict[str, Any]], list[dict[str, Any]]]:
|
||
"""Row provider that scores each (event, sample date) as Σ feature weights.
|
||
|
||
Production rule ids are kept on the evidence so the matrix's kind factor
|
||
and technique layers behave as in production; only the points change.
|
||
"""
|
||
include_extras, allow_negative = VARIANT_SPEC[variant]
|
||
universe = [
|
||
name for name in sorted(table.present)
|
||
if (include_extras or not is_extra_feature(name))
|
||
]
|
||
|
||
def provider(request: dict[str, Any]) -> list[dict[str, Any]]:
|
||
rows: list[dict[str, Any]] = []
|
||
for event in request["events"]:
|
||
key = (str(event["id"]), str(event["date"]))
|
||
cells = store.rows.get(key, {})
|
||
for time in sorted(cells, key=clock):
|
||
cell = cells[time]
|
||
names = [name for name in cell["rules"] if is_scoring_rule(name)]
|
||
if include_extras:
|
||
names += list(cell["extras"])
|
||
values = {name: 1.0 for name in names}
|
||
points = round(scale * weighted_points(values, table, allow_negative=allow_negative, universe=universe), 4)
|
||
rows.append({
|
||
"time": time,
|
||
"score": points,
|
||
"evidence": [{
|
||
"event_id": event["id"], "domain": event["domain"], "candidate_time": time,
|
||
"rule_ids": list(cell["rules"]), "points": points,
|
||
}],
|
||
"missing_layers": [],
|
||
})
|
||
return rows
|
||
|
||
return provider
|
||
|
||
|
||
def loo_folds(case_ids: Sequence[str]) -> list[tuple[str, list[str]]]:
|
||
ids = list(case_ids)
|
||
return [(held, [other for other in ids if other != held]) for held in ids]
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# truth labels, priors, replay
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def truth_cluster_times(contexts: Sequence[dict[str, Any]], true_time: str) -> list[str]:
|
||
for cluster in raw_signature_clusters(contexts):
|
||
times = [str(item)[:5] for item in cluster.get("times") or []]
|
||
if true_time[:5] in times:
|
||
return sorted(times, key=clock)
|
||
return [true_time[:5]]
|
||
|
||
|
||
def public_for(rows: Sequence[dict[str, Any]], contexts: Sequence[dict[str, Any]]) -> list[dict[str, Any]]:
|
||
raw = raw_signature_clusters(contexts)
|
||
by_time = {stamp: row for row in rows if (stamp := hhmm(row.get("time")))}
|
||
merged, _trace = merge_adjacent_traced(raw, by_time)
|
||
public = public_from_clusters(merged, rows)
|
||
for row in public:
|
||
row["score"] = float(row.get("score") or 0)
|
||
return public
|
||
|
||
|
||
def percent_proportional(scores: dict[str, float]) -> dict[str, float]:
|
||
total = sum(max(value, 0.0) for value in scores.values())
|
||
if total <= 0:
|
||
return {key: 0.0 for key in scores}
|
||
return {key: float(round(max(value, 0.0) / total * 100)) for key, value in scores.items()}
|
||
|
||
|
||
def softmax_probabilities(scores: dict[str, float], temperature: float = 1.0) -> dict[str, float]:
|
||
if not scores:
|
||
return {}
|
||
temp = float(temperature) if temperature and temperature > 0 else 1.0
|
||
peak = max(scores.values())
|
||
weights = {key: math.exp((value - peak) / temp) for key, value in scores.items()}
|
||
total = sum(weights.values())
|
||
return {key: weight / total for key, weight in weights.items()}
|
||
|
||
|
||
def percent_softmax(scores: dict[str, float], temperature: float = 1.0) -> dict[str, float]:
|
||
return {key: float(round(p * 100)) for key, p in softmax_probabilities(scores, temperature).items()}
|
||
|
||
|
||
TEMPERATURE_GRID = (0.25, 0.5, 1.0, 2.0, 4.0, 8.0, 16.0, 32.0, 64.0, 128.0, 256.0)
|
||
|
||
|
||
def truth_mass(public: Sequence[dict[str, Any]], probabilities: dict[str, float], truth_times: Sequence[str]) -> float:
|
||
truth = {str(t)[:5] for t in truth_times}
|
||
mass = 0.0
|
||
for row in public:
|
||
stamp = str(row["time"])[:5]
|
||
members = {str(t)[:5] for t in (row.get("cluster_times") or [stamp])}
|
||
if members & truth:
|
||
mass += probabilities.get(stamp, 0.0)
|
||
return mass
|
||
|
||
|
||
def fit_temperature(
|
||
training: Sequence[tuple[Sequence[dict[str, Any]], Sequence[str]]],
|
||
grid: Sequence[float] = TEMPERATURE_GRID,
|
||
) -> float:
|
||
"""Platt-style single temperature: maximise Σ log P(truth cluster) on training cases only."""
|
||
best_t, best_ll = 1.0, None
|
||
for temp in grid:
|
||
ll = 0.0
|
||
for public, truth_times in training:
|
||
base = {str(r["time"])[:5]: float(r.get("score") or 0) for r in public}
|
||
mass = truth_mass(public, softmax_probabilities(base, temp), truth_times)
|
||
ll += math.log(max(mass, 1e-9))
|
||
if best_ll is None or ll > best_ll + 1e-12:
|
||
best_t, best_ll = float(temp), ll
|
||
return best_t
|
||
|
||
|
||
def prior_for(public: Sequence[dict[str, Any]], *, mode: str, softmax: bool, temperature: float = 1.0) -> dict[str, float]:
|
||
base = {str(row["time"])[:5]: float(row.get("score") or 0) for row in public}
|
||
if mode == "raw":
|
||
return base
|
||
return percent_softmax(base, temperature) if softmax else percent_proportional(base)
|
||
|
||
|
||
def apply_scaled_answer(
|
||
scores: dict[str, float],
|
||
conflicts: dict[str, int],
|
||
eliminated: set[str],
|
||
probe: dict[str, Any],
|
||
answer: str,
|
||
candidate_times: Sequence[str],
|
||
factor: float,
|
||
) -> tuple[dict[str, float], dict[str, int], set[str]]:
|
||
"""`apply_answer` at card strength (factor == 1); weaker evidence scales the
|
||
delta and never counts a strong conflict (so it can never eliminate)."""
|
||
if factor >= 1.0 - EPS:
|
||
return apply_answer(scores, conflicts, eliminated, probe, answer, candidate_times)
|
||
yes, no = outcome_groups(probe)
|
||
next_scores = dict(scores)
|
||
for time in candidate_times:
|
||
if time in 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]
|
||
next_scores[time] = next_scores.get(time, 0.0) + SCORE_DELTA[raw] * factor
|
||
return next_scores, dict(conflicts), set(eliminated)
|
||
|
||
|
||
def replay(
|
||
*,
|
||
public: Sequence[dict[str, Any]],
|
||
prior: dict[str, float],
|
||
probes: Sequence[dict[str, Any]],
|
||
true_time: str,
|
||
window_times: Sequence[str],
|
||
pre_answers: Sequence[tuple[dict[str, Any], str, float]] = (),
|
||
answers: Sequence[str | None] | None = None,
|
||
) -> dict[str, Any]:
|
||
"""Prior → optional pre-answers (probe, answer, factor) → six probes answered from truth."""
|
||
reps = [str(row["time"])[:5] for row in public]
|
||
scores = {time: float(prior.get(time, 0.0)) for time in reps}
|
||
conflicts = {time: 0 for time in reps}
|
||
eliminated: set[str] = set()
|
||
for probe, answer, factor in pre_answers:
|
||
scores, conflicts, eliminated = apply_scaled_answer(scores, conflicts, eliminated, probe, answer, reps, factor)
|
||
asked = list(probes)[:ASK_COUNT]
|
||
given = list(answers) if answers is not None else [optimal_answer(p, true_time) for p in asked]
|
||
answered = 0
|
||
for probe, answer in zip(asked, given):
|
||
if answer is None:
|
||
continue
|
||
answered += 1
|
||
scores, conflicts, eliminated = apply_answer(scores, conflicts, eliminated, probe, answer, reps)
|
||
posterior = [{**row, "score": scores.get(str(row["time"])[:5], row.get("score") or 0)} for row in public]
|
||
valid = still_valid_public(posterior, scores, eliminated, lead=SEPARATION_LEAD)
|
||
delivery = delivery_from_public(valid)
|
||
alive = [row for row in posterior if str(row["time"])[:5] not in eliminated]
|
||
metrics = metrics_bundle(
|
||
public=alive, true_time=true_time, window_times=window_times,
|
||
delivery_times=delivery["times"], delivery_width=delivery["width"], independent=True,
|
||
entropy_scores=[scores.get(str(row["time"])[:5], 0.0) for row in alive],
|
||
)
|
||
truth_eliminated = any(
|
||
true_time in (row.get("cluster_times") or [str(row.get("time"))[:5]]) and str(row["time"])[:5] in eliminated
|
||
for row in public
|
||
)
|
||
return {
|
||
"top1": bool(metrics["top1_hit"]),
|
||
"coverage": bool(metrics["coverage"]),
|
||
"width": delivery["width"],
|
||
"tie": bool(metrics["tie"]),
|
||
"squeezed": bool(metrics["truth_squeezed"]),
|
||
"eliminated": len(eliminated),
|
||
"truth_eliminated": bool(truth_eliminated),
|
||
"questions": answered,
|
||
"entropy": round(shannon_entropy(max(s, 0.0) for t, s in scores.items() if t not in eliminated), 4),
|
||
}
|
||
|
||
|
||
def engine_top1(public: Sequence[dict[str, Any]], true_time: str) -> bool:
|
||
from scripts.research.cluster_width_lib import top1_from_public
|
||
return bool(top1_from_public(public, true_time))
|
||
|
||
|
||
def within_case_range(rows: Sequence[dict[str, Any]]) -> float:
|
||
scores = [float(row.get("score") or 0) for row in rows]
|
||
return (max(scores) - min(scores)) if scores else 0.0
|
||
|
||
|
||
def summarize(rows: Sequence[dict[str, Any]]) -> dict[str, Any]:
|
||
"""Same keys as precision_gate_sweep.summarize so `gate_verdict` can read it."""
|
||
if not rows:
|
||
return {"n": 0, "top1": None, "coverage": None, "width_median": None, "tie": None, "squeezed": 0,
|
||
"engine_top1": None, "refresh_mean": 0.0, "entropy": None, "truth_eliminated": 0}
|
||
n = len(rows)
|
||
return {
|
||
"n": n,
|
||
"top1": round(sum(1 for r in rows if r.get("top1")) / n, 4),
|
||
"coverage": round(sum(1 for r in rows if r.get("coverage")) / n, 4),
|
||
"width_median": median([r.get("width") for r in rows]),
|
||
"tie": round(sum(1 for r in rows if r.get("tie")) / n, 4),
|
||
"squeezed": sum(1 for r in rows if r.get("squeezed")),
|
||
"engine_top1": round(sum(1 for r in rows if r.get("engine_top1")) / n, 4),
|
||
"refresh_mean": 0.0,
|
||
"entropy": round(sum(float(r.get("entropy") or 0) for r in rows) / n, 4),
|
||
"truth_eliminated": sum(1 for r in rows if r.get("truth_eliminated")),
|
||
}
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# probes on representatives (typed-event convention), absence, precision follow-up
|
||
# ---------------------------------------------------------------------------
|
||
|
||
def production_probes(request: dict[str, Any], built: dict[str, Any], times: Sequence[str], true_time: str) -> list[dict[str, Any]]:
|
||
payload = {**request, "refresh_probes": False, "asked_probe_keys": []}
|
||
return event_probes.discriminating_event_probes(
|
||
payload, built, scan=window_scan(built), candidate_times=list(times),
|
||
representative_time=true_time, today=TODAY,
|
||
)
|
||
|
||
|
||
@dataclass
|
||
class Representatives:
|
||
reps: list[dict[str, Any]]
|
||
clusters: list[dict[str, Any]]
|
||
set_version: str
|
||
|
||
|
||
def representatives_for(contexts: Sequence[dict[str, Any]]) -> Representatives:
|
||
full = [item for item in contexts if event_probes._context_time(item)]
|
||
full.sort(key=lambda item: event_probes._clock(str(event_probes._context_time(item))))
|
||
clusters = cluster_contexts_by_signature(full)
|
||
reps = [c["representative"] for c in clusters if event_probes._scoreable(c["representative"])]
|
||
if len(reps) < 2:
|
||
reps = [item for item in full if event_probes._scoreable(item)]
|
||
version = event_probes.candidate_set_version([c["times"] for c in clusters])
|
||
return Representatives(reps=reps, clusters=clusters, set_version=version)
|
||
|
||
|
||
def split_probe(
|
||
representatives: Representatives,
|
||
*,
|
||
birth_date: str,
|
||
domain: str,
|
||
year: int,
|
||
month: int | None,
|
||
source: str,
|
||
) -> dict[str, Any] | None:
|
||
canonical = event_probes.canonical_domain(domain)
|
||
if canonical not in event_probes.DOMAIN_CATALOG:
|
||
return None
|
||
return event_probes._evaluate_contexts(
|
||
representatives.reps, birth_date=birth_date, domain=canonical, year=year, month=month,
|
||
source=source, clusters=representatives.clusters, set_version=representatives.set_version,
|
||
)
|
||
|
||
|
||
def event_year(event: dict[str, Any]) -> int | None:
|
||
raw = str(event.get("date_start") or event.get("date") or "")
|
||
return int(raw[:4]) if raw[:4].isdigit() else None
|
||
|
||
|
||
def event_month(event: dict[str, Any]) -> int | None:
|
||
raw = str(event.get("date_start") or event.get("date") or "")
|
||
return int(raw[5:7]) if len(raw) >= 7 and raw[4] == "-" else None
|
||
|
||
|
||
def absent_years(events: Sequence[dict[str, Any]]) -> dict[str, list[int]]:
|
||
"""Per domain: years strictly inside the span of that domain's dated events with no event."""
|
||
years: dict[str, set[int]] = defaultdict(set)
|
||
for event in events:
|
||
year = event_year(event)
|
||
if year is None:
|
||
continue
|
||
years[str(event.get("domain") or "")].add(year)
|
||
out: dict[str, list[int]] = {}
|
||
for domain, known in years.items():
|
||
if len(known) < 2:
|
||
continue
|
||
lo, hi = min(known), max(known)
|
||
gaps = [year for year in range(lo + 1, hi) if year not in known]
|
||
if gaps:
|
||
out[domain] = gaps
|
||
return dict(sorted(out.items()))
|
||
|
||
|
||
def drop_events(events: Sequence[dict[str, Any]], count: int, *, seed: str) -> list[dict[str, Any]]:
|
||
rng = Random(seed)
|
||
if count <= 0 or count >= len(events):
|
||
return list(events)
|
||
dropped = set(rng.sample(range(len(events)), count))
|
||
return [event for index, event in enumerate(events) if index not in dropped]
|
||
|
||
|
||
def degrade_to_year(events: Sequence[dict[str, Any]]) -> list[dict[str, Any]]:
|
||
"""The 'spoken' form of a case: every day/month event becomes year precision."""
|
||
out = []
|
||
for event in events:
|
||
item = dict(event)
|
||
if str(item.get("precision")) in {"day", "month"}:
|
||
item["precision"] = "year"
|
||
item["date"] = str(item.get("date") or "")[:4]
|
||
out.append(item)
|
||
return out
|
||
|
||
|
||
def upgrade_to_month(events: Sequence[dict[str, Any]], event_ids: set[str], originals: dict[str, dict[str, Any]]) -> list[dict[str, Any]]:
|
||
out = []
|
||
for event in events:
|
||
item = dict(event)
|
||
original = originals.get(str(item.get("id")))
|
||
if str(item.get("id")) in event_ids and original is not None and len(str(original.get("date") or "")) >= 7:
|
||
item["precision"] = "month"
|
||
item["date"] = str(original["date"])[:7]
|
||
out.append(item)
|
||
return out
|
||
|
||
|
||
def shift_day_events_by_months(events: Sequence[dict[str, Any]], *, seed: str, choices: Sequence[int] = (-3, -2, -1, 1, 2, 3)) -> list[dict[str, Any]]:
|
||
rng = Random(seed)
|
||
out = []
|
||
for event in events:
|
||
item = dict(event)
|
||
raw = str(item.get("date") or "")
|
||
if str(item.get("precision")) == "day" and len(raw) >= 10:
|
||
year, month, day = int(raw[:4]), int(raw[5:7]), int(raw[8:10])
|
||
index = year * 12 + (month - 1) + rng.choice(list(choices))
|
||
new_year, new_month = divmod(index, 12)
|
||
item["date"] = f"{new_year:04d}-{new_month + 1:02d}-{min(day, 28):02d}"
|
||
item["shift_months"] = index - (year * 12 + month - 1)
|
||
out.append(item)
|
||
return out
|
||
|
||
|
||
def calibration_bins(points: Sequence[tuple[float, bool]], edges: Sequence[float] = (0.0, 0.05, 0.1, 0.2, 0.3, 0.5, 0.7, 1.0001)) -> list[dict[str, Any]]:
|
||
out = []
|
||
for lo, hi in zip(edges, edges[1:]):
|
||
inside = [(p, t) for p, t in points if lo <= p < hi]
|
||
out.append({
|
||
"bin": f"[{lo:.2f},{min(hi, 1.0):.2f})",
|
||
"n": len(inside),
|
||
"predicted_mean": round(sum(p for p, _ in inside) / len(inside), 4) if inside else None,
|
||
"observed_rate": round(sum(1 for _, t in inside if t) / len(inside), 4) if inside else None,
|
||
})
|
||
return out
|
||
|
||
|
||
__all__ = [name for name in dir() if not name.startswith("__")]
|