From 9065054f8ddcaae524561a7110c7c0511cb658d8 Mon Sep 17 00:00:00 2001 From: sileod Date: Tue, 29 Sep 2026 13:55:16 +0200 Subject: [PATCH 1/8] Add independent Jev teacher audit stage --- src/tasksource/jev/synthetic/teacher_audit.py | 258 ++++++++++++++++++ 1 file changed, 258 insertions(+) create mode 100644 src/tasksource/jev/synthetic/teacher_audit.py diff --git a/src/tasksource/jev/synthetic/teacher_audit.py b/src/tasksource/jev/synthetic/teacher_audit.py new file mode 100644 index 0000000..c0cb8c9 --- /dev/null +++ b/src/tasksource/jev/synthetic/teacher_audit.py @@ -0,0 +1,258 @@ +"""Independent post-teacher audit for Jev-labelled synthetic bundles. + +The auditor is deliberately not shown Jev's probabilities. It answers the +state/questions independently, then this module compares that answer with the +stored teacher distribution. This catches the specific failure mode where a +teacher is highly confident on a generated example that another strong model +reads differently, without replacing Jev as the training target. + +Auditor calls are content-addressed and cached. The cache contains the raw +auditor response and its independent answers; comparison with Jev is recomputed +from the current annotations so the same audit can be reused across teacher +versions. +""" + +from __future__ import annotations + +import asyncio +import hashlib +import json +from pathlib import Path + +from . import providers +from .critic import RequestPacer +from .generate import PROMPTS_DIR, extract_json_object, prompt_hash + + +def load_teacher_audit_prompt(version: str) -> str: + return (PROMPTS_DIR / f"{version}.txt").read_text(encoding="utf-8") + + +def teacher_audit_cache_key(model: str, temperature: float, prompt: str, + bundle: dict, endpoint: str = "") -> str: + canonical = json.dumps( + {"state_id": bundle.get("state_id"), "state": bundle.get("state"), + "questions": bundle.get("questions")}, + sort_keys=True, ensure_ascii=False) + return hashlib.sha256( + f"{endpoint}\n{model}\n{temperature}\n{prompt_hash(prompt)}\n{canonical}".encode("utf-8") + ).hexdigest() + + +def _mock_independent_answers(bundle: dict) -> list[dict]: + """Deterministic offline answers for tests; never used as training targets.""" + out = [] + for question in bundle.get("questions", []): + fmt = question["format"] + if fmt == "noul": + answer = False + elif fmt == "score": + answer = 0 + else: + options = list(question.get("options", [])) + answer = options[0] if options else None + out.append({"question_id": question["question_id"], "answer": answer, + "confidence": 0.9, "issues": []}) + return out + + +def _parse_independent_answers(raw: dict, bundle: dict) -> list[dict]: + answers = raw.get("answers") + if not isinstance(answers, list): + raise ValueError("teacher-audit response needs an 'answers' list") + known = {q["question_id"] for q in bundle.get("questions", [])} + parsed = [] + for item in answers: + if not isinstance(item, dict) or item.get("question_id") not in known: + continue + try: + confidence = float(item.get("confidence", 0.0)) + except (TypeError, ValueError): + confidence = 0.0 + parsed.append({ + "question_id": item["question_id"], + "answer": item.get("answer"), + "confidence": max(0.0, min(1.0, confidence)), + "issues": list(item.get("issues", [])) if isinstance(item.get("issues", []), list) else [], + }) + return parsed + + +def _teacher_answer(question: dict, annotation: dict) -> tuple[object, float]: + probs = [float(p) for p in annotation.get("probabilities", [])] + if question["format"] == "noul": + p_yes = probs[0] + return p_yes >= 0.5, max(p_yes, 1.0 - p_yes) + index = max(range(len(probs)), key=probs.__getitem__) + if question["format"] == "score": + return index, probs[index] + return list(question.get("options", []))[index], probs[index] + + +def _normalize_auditor_answer(question: dict, answer) -> object | None: + fmt = question["format"] + if fmt == "noul": + if isinstance(answer, bool): + return answer + if isinstance(answer, str): + lowered = answer.strip().lower() + if lowered in {"true", "yes", "1"}: + return True + if lowered in {"false", "no", "0"}: + return False + return None + options = list(question.get("options", [])) + if fmt == "choice": + return answer if answer in options else None + if isinstance(answer, int) and not isinstance(answer, bool) and 0 <= answer < len(options): + return answer + if isinstance(answer, str): + stripped = answer.strip() + if stripped in options: + return options.index(stripped) + try: + index = int(stripped) + except ValueError: + return None + return index if 0 <= index < len(options) else None + return None + + +def compare_with_teacher(bundle: dict, independent_answers: list[dict], audit_cfg, + auditor_meta: dict | None = None) -> dict: + by_id = {a["question_id"]: a for a in independent_answers} + annotations = list(bundle.get("annotations", [])) + questions = list(bundle.get("questions", [])) + rows = [] + complete = len(annotations) == len(questions) + for index, question in enumerate(questions): + independent = by_id.get(question["question_id"]) + annotation = annotations[index] if index < len(annotations) else None + if independent is None or annotation is None: + complete = False + rows.append({"question_id": question["question_id"], "agrees": None, + "confident_disagreement": False, + "issues": ["missing auditor answer or teacher annotation"]}) + continue + teacher_answer, teacher_confidence = _teacher_answer(question, annotation) + auditor_answer = _normalize_auditor_answer(question, independent.get("answer")) + auditor_confidence = float(independent.get("confidence", 0.0)) + agrees = auditor_answer is not None and auditor_answer == teacher_answer + confident_disagreement = ( + auditor_answer is not None + and not agrees + and teacher_confidence >= audit_cfg.min_teacher_confidence + and auditor_confidence >= audit_cfg.min_auditor_confidence + ) + rows.append({ + "question_id": question["question_id"], + "teacher_answer": teacher_answer, + "teacher_confidence": teacher_confidence, + "auditor_answer": auditor_answer, + "auditor_confidence": auditor_confidence, + "agrees": agrees, + "confident_disagreement": confident_disagreement, + "issues": independent.get("issues", []), + }) + disagreements = sum(bool(row.get("confident_disagreement")) for row in rows) + return { + "enabled": True, + "pass": complete and disagreements == 0, + "complete": complete, + "confident_disagreements": disagreements, + "min_teacher_confidence": audit_cfg.min_teacher_confidence, + "min_auditor_confidence": audit_cfg.min_auditor_confidence, + "auditor": auditor_meta or {}, + "questions": rows, + } + + +async def _audit_one(sem, client, model: str, temperature: float, template: str, + bundle: dict, raw_dir: Path, audit_cfg, + pacer: RequestPacer | None = None, endpoint: str = "") -> dict: + key = teacher_audit_cache_key(model, temperature, template, bundle, endpoint) + cached = raw_dir / f"{key}.json" + auditor_meta = {"model": model, "endpoint": endpoint} + if cached.exists(): + record = json.loads(cached.read_text(encoding="utf-8")) + independent = record["independent_answers"] + auditor_meta.update({ + "provider": record.get("provider", ""), + "requested_model": record.get("requested_model", model), + "returned_model": record.get("returned_model", ""), + "cache_key": key, + }) + elif client is None: + independent = _mock_independent_answers(bundle) + cached.write_text(json.dumps( + {"cache_key": key, "state_id": bundle["state_id"], "provider": "mock", + "requested_model": model, "returned_model": model, + "independent_answers": independent}, + ensure_ascii=False, indent=2), encoding="utf-8") + auditor_meta.update({"provider": "mock", "requested_model": model, + "returned_model": model, "cache_key": key}) + else: + payload = {"state_id": bundle["state_id"], "state": bundle["state"], + "questions": bundle.get("questions", [])} + prompt = template.replace( + "{{BUNDLE_JSON}}", json.dumps(payload, ensure_ascii=False, indent=2)) + async with sem: + if pacer is not None: + await pacer.wait() + result = await providers.chat_complete( + client, model, [{"role": "user", "content": prompt}], + temperature=temperature, max_tokens=1600) + try: + parsed = extract_json_object(result["text"]) + independent = _parse_independent_answers(parsed, bundle) + except (ValueError, json.JSONDecodeError, TypeError): + independent = [] + cached.write_text(json.dumps( + {"cache_key": key, "state_id": bundle["state_id"], "provider": "teacher-auditor", + "requested_model": model, "returned_model": result["returned_model"], + "raw_response": result["raw"], "raw_text": result["text"], + "independent_answers": independent}, + ensure_ascii=False, indent=2), encoding="utf-8") + auditor_meta.update({"provider": "teacher-auditor", "requested_model": model, + "returned_model": result["returned_model"], "cache_key": key}) + out = dict(bundle) + out["teacher_audit"] = compare_with_teacher( + bundle, independent, audit_cfg, auditor_meta) + return out + + +async def audit_bundles_async(cfg, bundles: list[dict], raw_dir: Path) -> list[dict]: + raw_dir.mkdir(parents=True, exist_ok=True) + if not cfg.teacher_audit.enabled: + out = [] + for bundle in bundles: + copied = dict(bundle) + copied["teacher_audit"] = { + "enabled": False, "pass": True, "complete": True, + "confident_disagreements": 0, "questions": []} + out.append(copied) + return out + template = load_teacher_audit_prompt(cfg.teacher_audit.prompt_version) + provider = cfg.teacher_audit_provider() + client = None + if provider.name != "mock": + client = providers.make_client(provider, providers.require_api_key(provider)) + try: + sem = asyncio.Semaphore(max(1, cfg.generation.concurrency)) + pacer = RequestPacer(cfg.teacher_audit.requests_per_minute) if client is not None else None + return list(await asyncio.gather(*[ + _audit_one(sem, client, cfg.teacher_audit.model, cfg.teacher_audit.temperature, + template, bundle, raw_dir, cfg.teacher_audit, pacer, + f"{provider.name}@{provider.base_url.rstrip('/')}") + for bundle in bundles + ])) + finally: + if client is not None: + try: + await client.close() + except Exception: + pass + + +def audit_bundles(cfg, bundles: list[dict], raw_dir: Path) -> list[dict]: + return asyncio.run(audit_bundles_async(cfg, bundles, raw_dir)) From bbf4e546b6e8fb848d4ec40ff69e5553191e6fd9 Mon Sep 17 00:00:00 2001 From: sileod Date: Tue, 29 Sep 2026 13:55:19 +0200 Subject: [PATCH 2/8] Add independent teacher audit prompt --- .../synthetic/prompts/teacher_audit_v1.txt | 19 +++++++++++++++++++ 1 file changed, 19 insertions(+) create mode 100644 src/tasksource/jev/synthetic/prompts/teacher_audit_v1.txt diff --git a/src/tasksource/jev/synthetic/prompts/teacher_audit_v1.txt b/src/tasksource/jev/synthetic/prompts/teacher_audit_v1.txt new file mode 100644 index 0000000..740cfe4 --- /dev/null +++ b/src/tasksource/jev/synthetic/prompts/teacher_audit_v1.txt @@ -0,0 +1,19 @@ +You are an independent auditor for generated decision examples. + +You are NOT given another model's answer. Answer each question from the state +alone. Do not invent missing evidence. If the state is ambiguous or incomplete, +lower your confidence instead of forcing certainty. + +BUNDLE: +{{BUNDLE_JSON}} + +For every question: +- choice: answer with the exact option string from `options` +- score: answer with the zero-based integer index into the ordered `options` +- noul: answer with JSON true or false +- confidence: number in [0, 1] representing confidence in that answer +- issues: short list of any ambiguity, missing evidence, malformed rubric, or + other reason the answer should not be trusted + +Output STRICT JSON only: +{"answers": [{"question_id": "...", "answer": ..., "confidence": 0.0, "issues": []}]} From 00e0525fe6f9e242c6f25c0e7ea20c1728852138 Mon Sep 17 00:00:00 2001 From: sileod Date: Tue, 29 Sep 2026 13:55:21 +0200 Subject: [PATCH 3/8] Add audited synthetic Jev run config --- .../albert_deepseek_v4_flash_jev_audited.yaml | 84 +++++++++++++++++++ 1 file changed, 84 insertions(+) create mode 100644 src/tasksource/jev/synthetic/configs/albert_deepseek_v4_flash_jev_audited.yaml diff --git a/src/tasksource/jev/synthetic/configs/albert_deepseek_v4_flash_jev_audited.yaml b/src/tasksource/jev/synthetic/configs/albert_deepseek_v4_flash_jev_audited.yaml new file mode 100644 index 0000000..804b045 --- /dev/null +++ b/src/tasksource/jev/synthetic/configs/albert_deepseek_v4_flash_jev_audited.yaml @@ -0,0 +1,84 @@ +# DeepSeek generation + independent OpenAI teacher audit + pinned Jev targets. +# This is intentionally more expensive than the overnight baseline. The auditor +# never sees Jev's probabilities; it independently answers each question and only +# flags high-confidence disagreements. +run_name: deepseek_v4_flash_jev_audited +output_dir: .synthetic_runs + +provider: + name: albert + api_key_env: ALBERT_API_KEY + base_url: https://albert.api.etalab.gouv.fr/v1 + model: deepseek-v4-flash-0731 + +generation: + temperature: 0.8 + concurrency: 10 + max_output_tokens: 4000 + seed: 42 + prompt_version: generate_v2 + n_states: 4000 + +sampler: + seed: 42 + n_states: 4000 + questions_per_state: {'1': 0.6, '2': 0.25, '3': 0.1, '4': 0.05} + question_formats: {choice: 0.4, noul: 0.3, score: 0.3} + probability_mixed_formats: 0.85 + probability_all_formats_if_n_ge_3: 0.7 + +critic: + enabled: true + provider: + name: albert + api_key_env: ALBERT_API_KEY + base_url: https://albert.api.etalab.gouv.fr/v1 + model: deepseek-v4-flash-0731 + model: deepseek-v4-flash-0731 + prompt_version: critic_v1 + temperature: 0.0 + requests_per_minute: 40 + +annotator: + name: jev + version: jev-1.13 + base_url: https://openrouter.ai + api_path: /api/alpha/decisions + api_key_env: OPENROUTER_API_KEY + model: typesafe/jev-1.13 + +teacher_audit: + enabled: true + provider: + name: openai + api_key_env: OPENAI_API_KEY + base_url: https://api.openai.com/v1 + model: gpt-6-luna + model: gpt-6-luna + prompt_version: teacher_audit_v1 + temperature: 0.0 + requests_per_minute: 40 + min_teacher_confidence: 0.85 + min_auditor_confidence: 0.80 + # Keep flagged rows in the public artifact by default so filtering choices are + # auditable. Set true for a clean training-only subset. + drop_confident_disagreements: false + +selection: + buckets: + very_confident: 0.2 + confident: 0.3 + moderately_ambiguous: 0.3 + high_ambiguity: 0.15 + near_uniform: 0.05 + # When --n-target downsamples, reserve this fraction without consulting Jev + # confidence, then fill the rest with the requested ambiguity mixture. + unfiltered_fraction: 0.20 + seed_salt: jev-synthetic-selection-v2 + +split: + train: 0.85 + validation: 0.075 + test: 0.075 + ood_fraction: 0.05 + seed_salt: jev-synthetic-split-v1 From c52347272c5c9506489c8266fe51c7ba432a7b9e Mon Sep 17 00:00:00 2001 From: sileod Date: Tue, 29 Sep 2026 13:56:02 +0200 Subject: [PATCH 4/8] Configure independent teacher audits and selection reservoir --- src/tasksource/jev/synthetic/config.py | 87 ++++++++++++++++++++++---- 1 file changed, 76 insertions(+), 11 deletions(-) diff --git a/src/tasksource/jev/synthetic/config.py b/src/tasksource/jev/synthetic/config.py index 8fdb08e..3e12771 100644 --- a/src/tasksource/jev/synthetic/config.py +++ b/src/tasksource/jev/synthetic/config.py @@ -70,6 +70,33 @@ class AnnotatorConfig: model: str = "jev-mock" +@dataclass +class TeacherAuditConfig: + # Independent post-Jev sanity check. The auditor sees state/questions but + # never Jev's probabilities; comparison happens locally afterwards. + enabled: bool = False + # When absent, inherit the critic provider. For a genuinely independent + # check, explicitly choose a different provider/model from the generator. + provider: ProviderConfig | None = None + model: str = "deepseek-v4-flash-0731" + prompt_version: str = "teacher_audit_v1" + temperature: float = 0.0 + requests_per_minute: int = 40 + min_teacher_confidence: float = 0.85 + min_auditor_confidence: float = 0.80 + # False keeps flagged examples in the released artifact and records the + # disagreement. True makes the stage also emit/use audit_passed.jsonl. + drop_confident_disagreements: bool = False + + def __post_init__(self): + for name, value in ( + ("min_teacher_confidence", self.min_teacher_confidence), + ("min_auditor_confidence", self.min_auditor_confidence), + ): + if not 0 <= value <= 1: + raise ValueError(f"{name} must be in [0, 1], got {value}") + + @dataclass class SelectionConfig: # Buckets over observed Jev ambiguity (max_prob based, see select.py). @@ -81,6 +108,20 @@ class SelectionConfig: "near_uniform": 0.05, }) max_per_family: int = 0 # 0 = no cap + # When downsampling, this share is selected by a teacher-independent hash + # before confidence-bucket balancing. This leaves an unbiased reservoir for + # diagnosing selection effects. + unfiltered_fraction: float = 0.20 + seed_salt: str = "jev-synthetic-selection-v2" + + def __post_init__(self): + if not 0 <= self.unfiltered_fraction <= 1: + raise ValueError( + f"unfiltered_fraction must be in [0, 1], got {self.unfiltered_fraction}") + if not self.buckets or any(float(weight) < 0 for weight in self.buckets.values()): + raise ValueError("selection buckets must be a non-empty map of non-negative weights") + if sum(float(weight) for weight in self.buckets.values()) <= 0: + raise ValueError("selection bucket weights must have positive total mass") @dataclass @@ -105,6 +146,7 @@ class AppConfig: sampler: SamplerConfig = field(default_factory=SamplerConfig) critic: CriticConfig = field(default_factory=CriticConfig) annotator: AnnotatorConfig = field(default_factory=AnnotatorConfig) + teacher_audit: TeacherAuditConfig = field(default_factory=TeacherAuditConfig) selection: SelectionConfig = field(default_factory=SelectionConfig) split: SplitConfig = field(default_factory=SplitConfig) @@ -112,10 +154,18 @@ def critic_provider(self) -> ProviderConfig: """Resolved critic provider (explicit, else the generator provider).""" return self.critic.provider if self.critic.provider is not None else self.provider + def teacher_audit_provider(self) -> ProviderConfig: + """Resolved teacher-audit provider (explicit, else the critic provider).""" + return (self.teacher_audit.provider + if self.teacher_audit.provider is not None else self.critic_provider()) + def to_dict(self) -> dict: critic = dict(self.critic.__dict__) critic["provider"] = (self.critic.provider.__dict__ if self.critic.provider is not None else None) + teacher_audit = dict(self.teacher_audit.__dict__) + teacher_audit["provider"] = (self.teacher_audit.provider.__dict__ + if self.teacher_audit.provider is not None else None) return { "run_name": self.run_name, "output_dir": self.output_dir, @@ -124,6 +174,7 @@ def to_dict(self) -> dict: "sampler": self.sampler.__dict__, "critic": critic, "annotator": self.annotator.__dict__, + "teacher_audit": teacher_audit, "selection": self.selection.__dict__, "split": self.split.__dict__, } @@ -139,6 +190,17 @@ def _merge(base: dict, override: dict) -> dict: return out +def _nested_provider(raw, fallback: ProviderConfig, model: str) -> ProviderConfig | None: + if isinstance(raw, dict): + return ProviderConfig(**raw) + if isinstance(raw, str): + # Legacy: bare name inherits the parent connection details. + return ProviderConfig( + name=raw, api_key_env=fallback.api_key_env, + base_url=fallback.base_url, model=model) + return None + + def load_config(path: str) -> AppConfig: """Load a config.yaml into a typed AppConfig.""" with open(path, encoding="utf-8") as handle: @@ -146,25 +208,28 @@ def load_config(path: str) -> AppConfig: defaults = AppConfig().to_dict() merged = _merge(defaults, raw) provider = ProviderConfig(**merged["provider"]) - critic_raw = merged["critic"] + + critic_raw = dict(merged["critic"]) critic_provider_raw = critic_raw.pop("provider", None) - if isinstance(critic_provider_raw, dict): - critic_provider = ProviderConfig(**critic_provider_raw) - elif isinstance(critic_provider_raw, str): - # Legacy: bare name inherits the generator's connection details. - critic_provider = ProviderConfig( - name=critic_provider_raw, api_key_env=provider.api_key_env, - base_url=provider.base_url, model=critic_raw.get("model", provider.model)) - else: - critic_provider = None # inherit the generator provider at use site + critic_provider = _nested_provider( + critic_provider_raw, provider, critic_raw.get("model", provider.model)) + critic = CriticConfig(provider=critic_provider, **critic_raw) + + audit_raw = dict(merged["teacher_audit"]) + audit_provider_raw = audit_raw.pop("provider", None) + audit_provider = _nested_provider( + audit_provider_raw, critic_provider or provider, + audit_raw.get("model", (critic_provider or provider).model)) + return AppConfig( run_name=merged["run_name"], output_dir=merged["output_dir"], provider=provider, generation=GenerationConfig(**merged["generation"]), sampler=SamplerConfig(**merged["sampler"]), - critic=CriticConfig(provider=critic_provider, **critic_raw), + critic=critic, annotator=AnnotatorConfig(**merged["annotator"]), + teacher_audit=TeacherAuditConfig(provider=audit_provider, **audit_raw), selection=SelectionConfig(**merged["selection"]), split=SplitConfig(**merged["split"]), ) From 246015e952112f57f2d7adaeb29fe2bb68fe7969 Mon Sep 17 00:00:00 2001 From: sileod Date: Tue, 29 Sep 2026 13:56:33 +0200 Subject: [PATCH 5/8] Keep a teacher-independent reservoir during synthetic selection --- src/tasksource/jev/synthetic/select.py | 133 +++++++++++++++++-------- 1 file changed, 92 insertions(+), 41 deletions(-) diff --git a/src/tasksource/jev/synthetic/select.py b/src/tasksource/jev/synthetic/select.py index 6df3eb0..e95030c 100644 --- a/src/tasksource/jev/synthetic/select.py +++ b/src/tasksource/jev/synthetic/select.py @@ -1,14 +1,14 @@ """Quality / distribution selection using observed Jev ambiguity. -Buckets (by max probability, configurable in SelectionConfig): - very_confident max_prob >= 0.90 (target 20%) - confident max_prob >= 0.70 (target 30%) - moderately_ambiguous max_prob >= 0.50 (target 30%) - high_ambiguity max_prob >= 0.35 (target 15%) - near_uniform max_prob < 0.35 (target 5%) - -Selection is deterministic (stable hash ordering within buckets) and -also reports requested-vs-observed ambiguity as a quality diagnostic. +When a run is downsampled, selection deliberately has two arms: + +1. a deterministic teacher-independent reservoir, chosen only from state_id; +2. confidence-bucket balancing over the remaining examples. + +This makes it possible to measure how much Jev-confidence-based curation changes +the training distribution instead of making every retained example conditional +on the teacher's own uncertainty. Every selected bundle records its +`selection_reason` and observed `selection_bucket`. """ from __future__ import annotations @@ -39,54 +39,105 @@ def _bundle_score(bundle: dict) -> float: return sum(s["max_prob"] for s in stats) / len(stats) -def select_bundles(bundles: list[dict], selection_cfg, n_target: int | None = None) -> dict: - buckets: dict[str, list[dict]] = {} - for bundle in bundles: - buckets.setdefault(ambiguity_bucket(_bundle_score(bundle)), []).append(bundle) - for key in buckets: - buckets[key].sort(key=lambda b: hashlib.sha256(b["state_id"].encode()).hexdigest()) - total = n_target or len(bundles) - weights = selection_cfg.buckets - plan = {k: int(round(total * w)) for k, w in weights.items()} - # Fix rounding drift on the largest bucket. +def _hash_key(bundle: dict, salt: str) -> str: + return hashlib.sha256(f"{salt}:{bundle['state_id']}".encode()).hexdigest() + + +def _with_selection_metadata(bundle: dict, reason: str) -> dict: + out = dict(bundle) + out["selection_reason"] = reason + out["selection_bucket"] = ambiguity_bucket(_bundle_score(bundle)) + return out + + +def _allocation(total: int, weights: dict) -> dict[str, int]: + mass = sum(float(weight) for weight in weights.values()) + plan = {key: int(round(total * float(weight) / mass)) + for key, weight in weights.items()} drift = total - sum(plan.values()) if drift and plan: - biggest = max(plan, key=plan.get) + biggest = max(plan, key=lambda key: float(weights[key])) plan[biggest] += drift - cap = getattr(selection_cfg, "max_per_family", 0) # 0 = no cap + return plan + + +def select_bundles(bundles: list[dict], selection_cfg, + n_target: int | None = None) -> dict: + if not bundles: + return {"selected": [], "diagnostics": { + "bucket_counts": {}, "plan": {}, "unfiltered_target": 0, + "unfiltered_selected": 0, "requested_ambiguity": {}}} + + total = min(len(bundles), n_target if n_target is not None else len(bundles)) + cap = getattr(selection_cfg, "max_per_family", 0) + salt = getattr(selection_cfg, "seed_salt", "jev-synthetic-selection-v2") + unfiltered_fraction = float(getattr(selection_cfg, "unfiltered_fraction", 0.0)) per_family: Counter = Counter() + selected: list[dict] = [] + selected_ids: set[str] = set() - def admit(bundle) -> bool: + def admit(bundle: dict, reason: str) -> bool: + if bundle["state_id"] in selected_ids: + return False family = family_id(bundle) if cap and per_family[family] >= cap: return False per_family[family] += 1 + selected_ids.add(bundle["state_id"]) + selected.append(_with_selection_metadata(bundle, reason)) return True - selected: list[dict] = [] + # Arm 1: teacher-independent reservoir. This ordering never consults the + # annotation or ambiguity metadata. + unfiltered_target = int(round(total * unfiltered_fraction)) + for bundle in sorted(bundles, key=lambda b: _hash_key(b, f"{salt}:unfiltered")): + if len(selected) >= unfiltered_target: + break + admit(bundle, "unfiltered_reservoir") + unfiltered_selected = len(selected) + + # Arm 2: desired Jev ambiguity mixture over the remaining capacity. + buckets: dict[str, list[dict]] = {} + for bundle in bundles: + buckets.setdefault(ambiguity_bucket(_bundle_score(bundle)), []).append(bundle) + for key in buckets: + buckets[key].sort(key=lambda b: _hash_key(b, f"{salt}:bucket:{key}")) + + remaining_target = max(0, total - len(selected)) + plan = _allocation(remaining_target, selection_cfg.buckets) for bucket_name, count in plan.items(): taken = 0 for bundle in buckets.get(bucket_name, []): if taken >= count: break - if admit(bundle): - selected.append(bundle) + if admit(bundle, f"ambiguity_bucket:{bucket_name}"): taken += 1 - # Backfill from any bucket if a bucket is short. + + # Backfill deterministically when a requested bucket is short or a family + # cap blocks its quota. The backfill is reported distinctly. if len(selected) < total: - have = {b["state_id"] for b in selected} - for bucket_bundles in buckets.values(): - for bundle in bucket_bundles: - if len(selected) >= total: - break - if bundle["state_id"] not in have and admit(bundle): - selected.append(bundle) - have.add(bundle["state_id"]) + for bundle in sorted(bundles, key=lambda b: _hash_key(b, f"{salt}:backfill")): + if len(selected) >= total: + break + admit(bundle, "backfill") + selected.sort(key=lambda b: b["state_id"]) - by_requested = {} - for bundle in selected: - by_requested[bundle.get("ambiguity", "?")] = by_requested.get(bundle.get("ambiguity", "?"), 0) + 1 - return {"selected": selected, - "diagnostics": {"bucket_counts": {k: len(v) for k, v in buckets.items()}, - "plan": plan, - "requested_ambiguity": by_requested}} + by_requested = Counter(bundle.get("ambiguity", "?") for bundle in selected) + by_reason = Counter(bundle.get("selection_reason", "?") for bundle in selected) + final_buckets = Counter(bundle.get("selection_bucket", "?") for bundle in selected) + return { + "selected": selected, + "diagnostics": { + "input_states": len(bundles), + "target_states": total, + "bucket_counts": {key: len(value) for key, value in buckets.items()}, + "plan": plan, + "unfiltered_target": unfiltered_target, + "unfiltered_selected": unfiltered_selected, + "selection_reasons": dict(by_reason), + "selected_bucket_counts": dict(final_buckets), + "requested_ambiguity": dict(by_requested), + "max_per_family": cap, + "seed_salt": salt, + }, + } From fe5afd985b1c2531cd09af3c86a020b61b205ec7 Mon Sep 17 00:00:00 2001 From: sileod Date: Tue, 29 Sep 2026 13:58:02 +0200 Subject: [PATCH 6/8] Wire teacher audit and provenance into synthetic pipeline --- src/tasksource/jev/synthetic/run.py | 120 ++++++++++++++++++++++++---- 1 file changed, 105 insertions(+), 15 deletions(-) diff --git a/src/tasksource/jev/synthetic/run.py b/src/tasksource/jev/synthetic/run.py index 562c311..00138fe 100644 --- a/src/tasksource/jev/synthetic/run.py +++ b/src/tasksource/jev/synthetic/run.py @@ -21,6 +21,7 @@ from . import dedup as dedup_mod from . import manifest as manifest_mod from . import providers, select as select_mod +from . import teacher_audit as teacher_audit_mod from . import specs as specs_mod from . import split as split_mod from . import validate as validate_mod @@ -29,7 +30,7 @@ from .schemas import bundle_to_flat_rows, flat_to_training_row STAGES = ("specs", "generate", "validate", "critic", "dedup", - "annotate", "select", "split", "export") + "annotate", "teacher_audit", "select", "split", "export") def run_dir_for(cfg) -> Path: @@ -50,7 +51,7 @@ def _to_frame(bundles: list[dict]) -> pd.DataFrame: rows = [] for bundle in bundles: row = dict(bundle) - for key in ("questions", "annotations", "annotation_stats"): + for key in ("questions", "annotations", "annotation_stats", "teacher_audit"): if key in row and row[key] is not None: row[key] = json.dumps(row[key], ensure_ascii=False) rows.append(row) @@ -60,7 +61,7 @@ def _to_frame(bundles: list[dict]) -> pd.DataFrame: def _from_frame(frame: pd.DataFrame) -> list[dict]: bundles = [] for record in frame.to_dict(orient="records"): - for key in ("questions", "annotations", "annotation_stats"): + for key in ("questions", "annotations", "annotation_stats", "teacher_audit"): if key in record and isinstance(record[key], str): try: record[key] = json.loads(record[key]) @@ -90,7 +91,12 @@ def stage_generate(cfg, run_dir: Path) -> list[dict]: _to_frame(bundles).to_parquet(run_dir / "candidates.parquet", index=False) manifest_path = run_dir / "manifest.json" prompt_hashes = {} - for name in ("generate_v1", "critic_v1"): + prompt_versions = { + cfg.generation.prompt_version, + cfg.critic.prompt_version, + cfg.teacher_audit.prompt_version, + } + for name in sorted(prompt_versions): prompt_file = PROMPTS_DIR / f"{name}.txt" if prompt_file.exists(): prompt_hashes[name] = prompt_hash(prompt_file.read_text(encoding="utf-8")) @@ -100,7 +106,10 @@ def stage_generate(cfg, run_dir: Path) -> list[dict]: "generate_cache_key": generation_cache_key(cfg), "annotator": cfg.annotator.name, "critic_provider": cfg.critic_provider().name, - "critic_model": cfg.critic.model}, + "critic_model": cfg.critic.model, + "teacher_audit_enabled": cfg.teacher_audit.enabled, + "teacher_audit_provider": cfg.teacher_audit_provider().name, + "teacher_audit_model": cfg.teacher_audit.model}, prompt_hashes)) # Keep spec lookup for validation. (run_dir / "_spec_by_id.json").write_text(json.dumps(spec_by_id), encoding="utf-8") @@ -160,8 +169,53 @@ def stage_annotate(cfg, run_dir: Path) -> list[dict]: return annotated +def stage_teacher_audit(cfg, run_dir: Path) -> list[dict]: + """Independently answer generated questions, then compare with Jev locally.""" + bundles = _read_bundles(run_dir / "annotated.jsonl") + audited = teacher_audit_mod.audit_bundles( + cfg, bundles, run_dir / "raw" / "teacher_audit") + _write_bundles(run_dir / "audited.jsonl", audited) + _to_frame(audited).to_parquet(run_dir / "audited.parquet", index=False) + + passing = [bundle for bundle in audited + if bundle.get("teacher_audit", {}).get("pass", True)] + report = { + "enabled": cfg.teacher_audit.enabled, + "states": len(audited), + "complete": sum(bool(b.get("teacher_audit", {}).get("complete", False)) + for b in audited), + "passing": len(passing), + "confident_disagreements": sum( + int(b.get("teacher_audit", {}).get("confident_disagreements", 0)) + for b in audited), + "drop_confident_disagreements": cfg.teacher_audit.drop_confident_disagreements, + "provider": cfg.teacher_audit_provider().name, + "model": cfg.teacher_audit.model, + } + (run_dir / "teacher_audit_report.json").write_text( + json.dumps(report, indent=2), encoding="utf-8") + if cfg.teacher_audit.drop_confident_disagreements: + _write_bundles(run_dir / "audit_passed.jsonl", passing) + + manifest_path = run_dir / "manifest.json" + if manifest_path.exists(): + manifest = json.loads(manifest_path.read_text(encoding="utf-8")) + audit_config = dict(cfg.teacher_audit.__dict__) + audit_config["provider"] = cfg.teacher_audit_provider().__dict__ + manifest["teacher_audit"] = {**audit_config, "report": report} + manifest_path.write_text(json.dumps(manifest, indent=2), encoding="utf-8") + return passing if cfg.teacher_audit.drop_confident_disagreements else audited + + def stage_select(cfg, run_dir: Path, n_target: int | None = None) -> list[dict]: - annotated = _read_bundles(run_dir / "annotated.jsonl") + if (cfg.teacher_audit.enabled and cfg.teacher_audit.drop_confident_disagreements + and (run_dir / "audit_passed.jsonl").exists()): + source = run_dir / "audit_passed.jsonl" + elif (run_dir / "audited.jsonl").exists(): + source = run_dir / "audited.jsonl" + else: + source = run_dir / "annotated.jsonl" + annotated = _read_bundles(source) result = select_mod.select_bundles(annotated, cfg.selection, n_target) _write_bundles(run_dir / "selected.jsonl", result["selected"]) (run_dir / "selection_report.json").write_text( @@ -183,18 +237,50 @@ def stage_export(cfg, run_dir: Path) -> dict: bundles = _read_bundles(run_dir / "final.jsonl") flat_rows: list[dict] = [] for bundle in bundles: - annotations = {q["question_id"]: a for q, a in - zip(bundle.get("questions", []), bundle.get("annotations", []))} + annotations = { + q["question_id"]: (a, stats) + for q, a, stats in zip( + bundle.get("questions", []), + bundle.get("annotations", []), + bundle.get("annotation_stats", [])) + } + audit = bundle.get("teacher_audit", {}) + audit_by_question = { + row.get("question_id"): row for row in audit.get("questions", []) + } for flat in bundle_to_flat_rows(bundle): # a missing annotation must fail: flat_to_training_row would fill in a uniform target - target = annotations.get(flat["question_id"], {}).get("probabilities") + annotation, stats = annotations.get(flat["question_id"], (None, None)) + target = annotation.get("probabilities") if annotation else None if target is None: raise ValueError(f"no annotation for question {flat['question_id']} of state {flat['state_id']}") - flat_rows.append({**flat_to_training_row(flat, target, "synthetic/jev", - bundle.get("split", "train")), - "state_id": flat["state_id"], "question_id": flat["question_id"], - "bundle_size": flat["bundle_size"], "domain": flat["domain"], - "skill": flat.get("skill", "")}) + audit_row = audit_by_question.get(flat["question_id"], {}) + target_origin = ( + f"{annotation.get('annotator', 'unknown')}:" + f"{annotation.get('returned_model') or annotation.get('annotator_version', 'unknown')}" + ) + flat_rows.append({ + **flat_to_training_row(flat, target, "synthetic/jev", + bundle.get("split", "train")), + "state_id": flat["state_id"], + "question_id": flat["question_id"], + "bundle_size": flat["bundle_size"], + "domain": flat["domain"], + "skill": flat.get("skill", ""), + "target_origin": target_origin, + "annotator": annotation.get("annotator", ""), + "annotator_version": annotation.get("annotator_version", ""), + "annotator_model": annotation.get("returned_model", ""), + "teacher_max_prob": stats.get("max_prob") if stats else None, + "teacher_entropy": stats.get("entropy") if stats else None, + "selection_reason": bundle.get("selection_reason", ""), + "selection_bucket": bundle.get("selection_bucket", ""), + "teacher_audit_enabled": bool(audit.get("enabled", False)), + "teacher_audit_agrees": audit_row.get("agrees"), + "teacher_audit_confidence": audit_row.get("auditor_confidence"), + "teacher_audit_confident_disagreement": audit_row.get( + "confident_disagreement"), + }) flat_frame = pd.DataFrame(flat_rows) flat_frame.to_parquet(run_dir / "flat.parquet", index=False) # HF dataset dir with bundled + flat configs. @@ -208,7 +294,9 @@ def stage_export(cfg, run_dir: Path) -> dict: hf_dir / f"flat-{split_name}.parquet") (hf_dir / "README.md").write_text( "# jev-synthetic-decisions\n\nConfigs: `bundled` (one row per state, " - "`final.parquet`) and `flat` (one row per decision, `flat.parquet`).\n", + "`final.parquet`) and `flat` (one row per decision, `flat.parquet`). " + "Flat rows retain target origin, teacher uncertainty, selection arm, and " + "independent teacher-audit fields.\n", encoding="utf-8") shutil.copyfile(run_dir / "final.parquet", hf_dir / "bundled.parquet") shutil.copyfile(run_dir / "flat.parquet", hf_dir / "flat.parquet") @@ -235,6 +323,8 @@ def run_stage(cfg, run_dir: Path, stage: str, **kwargs): return stage_dedup(cfg, run_dir) if stage == "annotate": return stage_annotate(cfg, run_dir) + if stage == "teacher_audit": + return stage_teacher_audit(cfg, run_dir) if stage == "select": return stage_select(cfg, run_dir, kwargs.get("n_target")) if stage == "split": From 1aaf8ed24833a1249027649cf18a288ffad35f2f Mon Sep 17 00:00:00 2001 From: sileod Date: Tue, 29 Sep 2026 13:58:47 +0200 Subject: [PATCH 7/8] Test teacher audit and unbiased selection arm --- tests/test_jev_synthetic.py | 98 ++++++++++++++++++++++++++++++++++++- 1 file changed, 96 insertions(+), 2 deletions(-) diff --git a/tests/test_jev_synthetic.py b/tests/test_jev_synthetic.py index 61b821e..0a05c63 100644 --- a/tests/test_jev_synthetic.py +++ b/tests/test_jev_synthetic.py @@ -14,6 +14,7 @@ from tasksource.jev.synthetic import select as select_mod from tasksource.jev.synthetic import specs as specs_mod from tasksource.jev.synthetic import split as split_mod +from tasksource.jev.synthetic import teacher_audit as teacher_audit_mod from tasksource.jev.synthetic import validate as validate_mod from tasksource.jev.synthetic.config import AppConfig, load_config from tasksource.jev.synthetic.generate import ( @@ -226,6 +227,59 @@ def test_critic_results_cached_by_content_hash(self): self.assertTrue(all(not p.startswith("state_") for p in files_after_first)) +class TeacherAuditTest(unittest.TestCase): + def test_confident_teacher_disagreement_is_flagged(self): + from tasksource.jev.synthetic.config import TeacherAuditConfig + bundle = { + "state_id": "s", + "state": "A compact test state.", + "questions": [ + {"question_id": "q0", "format": "noul", "question": "Is it true?"}, + {"question_id": "q1", "format": "choice", "question": "Which?", + "options": ["a", "b"]}, + {"question_id": "q2", "format": "score", "question": "How much?", + "options": ["low", "medium", "high"]}, + ], + "annotations": [ + {"probabilities": [0.95]}, + {"probabilities": [0.9, 0.1]}, + {"probabilities": [0.05, 0.9, 0.05]}, + ], + } + independent = [ + {"question_id": "q0", "answer": False, "confidence": 0.95, "issues": []}, + {"question_id": "q1", "answer": "a", "confidence": 0.95, "issues": []}, + {"question_id": "q2", "answer": 1, "confidence": 0.95, "issues": []}, + ] + audit = teacher_audit_mod.compare_with_teacher( + bundle, independent, + TeacherAuditConfig(enabled=True, min_teacher_confidence=0.8, + min_auditor_confidence=0.8)) + self.assertFalse(audit["pass"]) + self.assertEqual(audit["confident_disagreements"], 1) + self.assertTrue(audit["questions"][0]["confident_disagreement"]) + self.assertTrue(audit["questions"][1]["agrees"]) + self.assertTrue(audit["questions"][2]["agrees"]) + + def test_mock_teacher_audit_is_cached_and_offline(self): + from tasksource.jev.synthetic.config import ProviderConfig + cfg = _cfg() + cfg.teacher_audit.enabled = True + cfg.teacher_audit.provider = ProviderConfig( + name="mock", api_key_env="UNUSED", base_url="", model="mock-auditor") + cfg.teacher_audit.model = "mock-auditor" + specs = specs_mod.sample_specs(cfg.sampler, 3) + bundles = [annot_mod.annotate_bundle(mock_realization(spec), cfg.annotator) + for spec in specs] + with tempfile.TemporaryDirectory() as tmp: + raw = Path(tmp) + first = teacher_audit_mod.audit_bundles(cfg, bundles, raw) + second = teacher_audit_mod.audit_bundles(cfg, bundles, raw) + self.assertEqual(first, second) + self.assertEqual(len(list(raw.glob("*.json"))), len(bundles)) + self.assertTrue(all(b["teacher_audit"]["enabled"] for b in first)) + + class CacheKeyTest(unittest.TestCase): def test_unrelated_settings_do_not_invalidate_generation(self): cfg = _cfg() @@ -355,7 +409,35 @@ def test_plan_sums_to_total(self): annotated = annot_mod.annotate_bundles(bundles, cfg.annotator) result = select_mod.select_bundles(annotated, cfg.selection) self.assertEqual(len(result["selected"]), len(annotated)) - self.assertEqual(sum(result["diagnostics"]["plan"].values()), len(annotated)) + diagnostics = result["diagnostics"] + self.assertEqual( + sum(diagnostics["plan"].values()) + diagnostics["unfiltered_selected"], + len(annotated)) + self.assertTrue(all("selection_reason" in b for b in result["selected"])) + self.assertTrue(all("selection_bucket" in b for b in result["selected"])) + + def test_unfiltered_reservoir_does_not_depend_on_teacher_confidence(self): + from tasksource.jev.synthetic.config import SelectionConfig + bundles = [ + {"state_id": f"s{i}", "domain": f"d{i % 3}", "scenario_type": "t", + "style": "x", "ambiguity": "clear", + "questions": [{"skill": "classification"}], + "annotation_stats": [{"max_prob": 0.95 if i % 2 else 0.4}]} + for i in range(40) + ] + cfg = SelectionConfig(unfiltered_fraction=0.5, seed_salt="test-reservoir") + first = select_mod.select_bundles(bundles, cfg, n_target=12)["selected"] + flipped = [ + {**b, "annotation_stats": [{"max_prob": 0.2 if i % 2 else 0.99}]} + for i, b in enumerate(bundles) + ] + second = select_mod.select_bundles(flipped, cfg, n_target=12)["selected"] + first_reservoir = {b["state_id"] for b in first + if b["selection_reason"] == "unfiltered_reservoir"} + second_reservoir = {b["state_id"] for b in second + if b["selection_reason"] == "unfiltered_reservoir"} + self.assertEqual(first_reservoir, second_reservoir) + self.assertEqual(len(first_reservoir), 6) class JevParsingTest(unittest.TestCase): @@ -438,7 +520,8 @@ def _config_dir(self): def test_all_configs_load(self): base = self._config_dir() - for name in ("albert_deepseek_v4_flash.yaml", "openai_luna.yaml", "mock_pilot.yaml"): + for name in ("albert_deepseek_v4_flash.yaml", "openai_luna.yaml", "mock_pilot.yaml", + "albert_deepseek_v4_flash_jev_audited.yaml"): cfg = load_config(str(base / name)) self.assertTrue(cfg.provider.model) self.assertTrue(cfg.provider.api_key_env) @@ -463,5 +546,16 @@ def test_critic_provider_independent_of_generator(self): self.assertEqual(other.critic_provider().name, "albert") + def test_audited_config_uses_independent_teacher_auditor(self): + base = self._config_dir() + cfg = load_config(str(base / "albert_deepseek_v4_flash_jev_audited.yaml")) + self.assertEqual(cfg.annotator.name, "jev") + self.assertTrue(cfg.teacher_audit.enabled) + self.assertEqual(cfg.provider.name, "albert") + self.assertEqual(cfg.teacher_audit_provider().name, "openai") + self.assertNotEqual(cfg.teacher_audit_provider().name, cfg.provider.name) + self.assertAlmostEqual(cfg.selection.unfiltered_fraction, 0.20) + + if __name__ == "__main__": unittest.main() From 2d9d8e43c476fb03fec92c1963c03d319cd71158 Mon Sep 17 00:00:00 2001 From: Damien Sileo Date: Tue, 29 Sep 2026 14:43:58 +0200 Subject: [PATCH 8/8] Harden synthetic audit and use reachable pilot auditor --- .../albert_deepseek_v4_flash_jev_audited.yaml | 12 ++++---- src/tasksource/jev/synthetic/run.py | 19 +++++++++---- src/tasksource/jev/synthetic/teacher_audit.py | 7 ++++- tests/test_jev_synthetic.py | 28 ++++++++++++++++++- 4 files changed, 53 insertions(+), 13 deletions(-) diff --git a/src/tasksource/jev/synthetic/configs/albert_deepseek_v4_flash_jev_audited.yaml b/src/tasksource/jev/synthetic/configs/albert_deepseek_v4_flash_jev_audited.yaml index 804b045..7e8be65 100644 --- a/src/tasksource/jev/synthetic/configs/albert_deepseek_v4_flash_jev_audited.yaml +++ b/src/tasksource/jev/synthetic/configs/albert_deepseek_v4_flash_jev_audited.yaml @@ -1,4 +1,4 @@ -# DeepSeek generation + independent OpenAI teacher audit + pinned Jev targets. +# DeepSeek generation + independent OpenAI-model teacher audit + pinned Jev targets. # This is intentionally more expensive than the overnight baseline. The auditor # never sees Jev's probabilities; it independently answers each question and only # flags high-confidence disagreements. @@ -50,11 +50,11 @@ annotator: teacher_audit: enabled: true provider: - name: openai - api_key_env: OPENAI_API_KEY - base_url: https://api.openai.com/v1 - model: gpt-6-luna - model: gpt-6-luna + name: openrouter + api_key_env: OPENROUTER_API_KEY + base_url: https://openrouter.ai/api/v1 + model: openai/gpt-4.1-mini + model: openai/gpt-4.1-mini prompt_version: teacher_audit_v1 temperature: 0.0 requests_per_minute: 40 diff --git a/src/tasksource/jev/synthetic/run.py b/src/tasksource/jev/synthetic/run.py index 00138fe..9bdf0cf 100644 --- a/src/tasksource/jev/synthetic/run.py +++ b/src/tasksource/jev/synthetic/run.py @@ -179,9 +179,16 @@ def stage_teacher_audit(cfg, run_dir: Path) -> list[dict]: passing = [bundle for bundle in audited if bundle.get("teacher_audit", {}).get("pass", True)] + audit_questions = [row for bundle in audited + for row in bundle.get("teacher_audit", {}).get("questions", [])] + answered = [row for row in audit_questions if row.get("agrees") is not None] report = { "enabled": cfg.teacher_audit.enabled, "states": len(audited), + "questions": len(audit_questions), + "answered_questions": len(answered), + "disagreements": sum(row.get("agrees") is False for row in answered), + "auditor_issue_questions": sum(bool(row.get("issues")) for row in audit_questions), "complete": sum(bool(b.get("teacher_audit", {}).get("complete", False)) for b in audited), "passing": len(passing), @@ -208,11 +215,13 @@ def stage_teacher_audit(cfg, run_dir: Path) -> list[dict]: def stage_select(cfg, run_dir: Path, n_target: int | None = None) -> list[dict]: - if (cfg.teacher_audit.enabled and cfg.teacher_audit.drop_confident_disagreements - and (run_dir / "audit_passed.jsonl").exists()): - source = run_dir / "audit_passed.jsonl" - elif (run_dir / "audited.jsonl").exists(): - source = run_dir / "audited.jsonl" + if cfg.teacher_audit.enabled: + source = (run_dir / "audit_passed.jsonl" if + cfg.teacher_audit.drop_confident_disagreements else + run_dir / "audited.jsonl") + if not source.exists(): + raise FileNotFoundError( + f"Teacher audit artifact {source} is missing; run the teacher_audit stage") else: source = run_dir / "annotated.jsonl" annotated = _read_bundles(source) diff --git a/src/tasksource/jev/synthetic/teacher_audit.py b/src/tasksource/jev/synthetic/teacher_audit.py index c0cb8c9..84f3679 100644 --- a/src/tasksource/jev/synthetic/teacher_audit.py +++ b/src/tasksource/jev/synthetic/teacher_audit.py @@ -17,6 +17,7 @@ import asyncio import hashlib import json +import math from pathlib import Path from . import providers @@ -69,6 +70,8 @@ def _parse_independent_answers(raw: dict, bundle: dict) -> list[dict]: confidence = float(item.get("confidence", 0.0)) except (TypeError, ValueError): confidence = 0.0 + if not math.isfinite(confidence): + confidence = 0.0 parsed.append({ "question_id": item["question_id"], "answer": item.get("answer"), @@ -137,7 +140,9 @@ def compare_with_teacher(bundle: dict, independent_answers: list[dict], audit_cf teacher_answer, teacher_confidence = _teacher_answer(question, annotation) auditor_answer = _normalize_auditor_answer(question, independent.get("answer")) auditor_confidence = float(independent.get("confidence", 0.0)) - agrees = auditor_answer is not None and auditor_answer == teacher_answer + if auditor_answer is None: + complete = False + agrees = None if auditor_answer is None else auditor_answer == teacher_answer confident_disagreement = ( auditor_answer is not None and not agrees diff --git a/tests/test_jev_synthetic.py b/tests/test_jev_synthetic.py index 0a05c63..e13c09b 100644 --- a/tests/test_jev_synthetic.py +++ b/tests/test_jev_synthetic.py @@ -11,6 +11,7 @@ from tasksource.jev.synthetic import critic as critic_mod from tasksource.jev.synthetic import dedup as dedup_mod from tasksource.jev.synthetic import providers +from tasksource.jev.synthetic import run as run_mod from tasksource.jev.synthetic import select as select_mod from tasksource.jev.synthetic import specs as specs_mod from tasksource.jev.synthetic import split as split_mod @@ -228,6 +229,31 @@ def test_critic_results_cached_by_content_hash(self): class TeacherAuditTest(unittest.TestCase): + def test_selection_requires_current_audit_artifact(self): + cfg = _cfg() + cfg.teacher_audit.enabled = True + with tempfile.TemporaryDirectory() as tmp: + run_dir = Path(tmp) + (run_dir / "annotated.jsonl").write_text("", encoding="utf-8") + with self.assertRaises(FileNotFoundError): + run_mod.stage_select(cfg, run_dir) + + def test_invalid_auditor_answer_does_not_pass(self): + from tasksource.jev.synthetic.config import TeacherAuditConfig + bundle = { + "state_id": "s", "state": "A compact test state.", + "questions": [{"question_id": "q0", "format": "choice", + "question": "Which?", "options": ["a", "b"]}], + "annotations": [{"probabilities": [0.9, 0.1]}], + } + audit = teacher_audit_mod.compare_with_teacher( + bundle, [{"question_id": "q0", "answer": "not an option", + "confidence": 0.99, "issues": []}], + TeacherAuditConfig(enabled=True)) + self.assertFalse(audit["complete"]) + self.assertFalse(audit["pass"]) + self.assertIsNone(audit["questions"][0]["agrees"]) + def test_confident_teacher_disagreement_is_flagged(self): from tasksource.jev.synthetic.config import TeacherAuditConfig bundle = { @@ -552,7 +578,7 @@ def test_audited_config_uses_independent_teacher_auditor(self): self.assertEqual(cfg.annotator.name, "jev") self.assertTrue(cfg.teacher_audit.enabled) self.assertEqual(cfg.provider.name, "albert") - self.assertEqual(cfg.teacher_audit_provider().name, "openai") + self.assertEqual(cfg.teacher_audit_provider().name, "openrouter") self.assertNotEqual(cfg.teacher_audit_provider().name, cfg.provider.name) self.assertAlmostEqual(cfg.selection.unfiltered_fraction, 0.20)