diff --git a/src/agent/evaluation/__init__.py b/src/agent/evaluation/__init__.py index 2945533..5dd4024 100644 --- a/src/agent/evaluation/__init__.py +++ b/src/agent/evaluation/__init__.py @@ -4,12 +4,16 @@ Stage1Evaluation, Stage2Evaluation, evaluate_stage1, + evaluate_stage1_choices, evaluate_stage2, + evaluate_stage2_choices, ) __all__ = [ "Stage1Evaluation", "Stage2Evaluation", "evaluate_stage1", + "evaluate_stage1_choices", "evaluate_stage2", + "evaluate_stage2_choices", ] diff --git a/src/agent/evaluation/classification.py b/src/agent/evaluation/classification.py index 981db47..3204c9a 100644 --- a/src/agent/evaluation/classification.py +++ b/src/agent/evaluation/classification.py @@ -1,12 +1,39 @@ -"""Algorithm-independent evaluation for two-stage leaf classification.""" +"""Algorithm-independent evaluation for two-stage leaf classification. + +Canonical semantics: the final comparison object is ALWAYS the canonical +category_id (registry / corpus / ground truth). The LLM boundary speaks +choice ids instead: + + raw model output -> parse choice_id -> decode category_id + -> existing canonical correctness logic (evaluate_stage1/2) + +evaluate_stage1/2 consume the unified parser (check_stage1_output / +check_stage2_output) so the reward adapter and the evaluator share one +contract implementation. evaluate_stage1_choices / evaluate_stage2_choices +are thin adapters: they parse and decode the choice protocol BEFORE +delegating to the canonical evaluators, so choice ids never leak into +correctness logic and no canonical check logic is duplicated. +""" from __future__ import annotations from dataclasses import dataclass +import json from typing import Sequence from agent.task import LeafRegistry -from agent.task.parser import check_stage1_output, check_stage2_output +from agent.task.parser import ( + PredictionFormatError, + check_stage1_output, + check_stage2_output, + parse_stage1_output, + parse_stage2_output, +) +from agent.task.prompt_choices import ( + PromptChoiceError, + PromptChoiceRegistry, + decode_stage2_answer, +) @dataclass(frozen=True) @@ -51,6 +78,40 @@ def evaluate_stage1( ) +def evaluate_stage1_choices( + solution: str, + *, + ground_truth: str, + registry: LeafRegistry, + choices: PromptChoiceRegistry | None = None, +) -> Stage1Evaluation: + """Evaluate a choice-id Stage 1 output; decode BEFORE canonical logic. + + The model answers with global choice ids ("1".."N"); the returned + prediction is the decoded canonical category_id tuple. Invalid choice + ids / wrong counts / duplicates yield an explicit invalid result with + no name or fuzzy fallback. Delegates to evaluate_stage1 with the + decoded canonical payload, so the canonical contract logic is never + duplicated. + """ + if ground_truth not in registry.ids: + raise ValueError("ground_truth must belong to the leaf registry") + choices = choices or PromptChoiceRegistry.from_registry(registry) + try: + output = parse_stage1_output(solution) + except PredictionFormatError as exc: + return Stage1Evaluation(None, False, False, False, (str(exc),)) + try: + decoded = choices.decode_candidates(output.candidates) + except PromptChoiceError as exc: + return Stage1Evaluation(None, True, False, False, (str(exc),)) + return evaluate_stage1( + json.dumps({"candidates": list(decoded)}, ensure_ascii=False, separators=(",", ":")), + ground_truth=ground_truth, + registry=registry, + ) + + def evaluate_stage2( solution: str, *, @@ -77,3 +138,42 @@ def evaluate_stage2( contract_valid and result.output.answer == ground_truth, errors, ) + + +def evaluate_stage2_choices( + solution: str, + *, + ground_truth: str, + candidates: Sequence[str], + registry: LeafRegistry, +) -> Stage2Evaluation: + """Evaluate a local-id Stage 2 output; decode BEFORE canonical logic. + + The model answers with a LOCAL bundle id ("1".."5" in candidate order); + the returned prediction is the decoded canonical category_id. Anything + but an exact local id yields an explicit invalid result with no name or + fuzzy fallback. Delegates to evaluate_stage2 with the decoded canonical + payload, so the canonical contract logic is never duplicated. + """ + if ground_truth not in registry.ids: + raise ValueError("ground_truth must belong to the leaf registry") + if isinstance(candidates, (str, bytes)) or ( + len(candidates) != 5 + or len(set(candidates)) != 5 + or any(candidate not in registry.ids for candidate in candidates) + ): + raise ValueError("candidates must be 5 unique IDs from the leaf registry") + try: + output = parse_stage2_output(solution) + except PredictionFormatError as exc: + return Stage2Evaluation(None, False, False, False, (str(exc),)) + try: + decoded = decode_stage2_answer(output.answer, tuple(candidates)) + except PromptChoiceError as exc: + return Stage2Evaluation(output.answer, True, False, False, (str(exc),)) + return evaluate_stage2( + json.dumps({"answer": decoded}, ensure_ascii=False, separators=(",", ":")), + ground_truth=ground_truth, + candidates=candidates, + registry=registry, + ) diff --git a/src/agent/task/__init__.py b/src/agent/task/__init__.py index ce56b4c..4df4299 100644 --- a/src/agent/task/__init__.py +++ b/src/agent/task/__init__.py @@ -19,6 +19,13 @@ parse_stage1_output, parse_stage2_output, ) +from .prompt_choices import ( + PromptChoice, + PromptChoiceError, + PromptChoiceRegistry, + decode_stage2_answer, + encode_stage2_answer, +) from .prompts import ( Prompt, build_stage1_prompt, @@ -59,6 +66,11 @@ "check_stage2_output", "parse_stage1_output", "parse_stage2_output", + "PromptChoice", + "PromptChoiceError", + "PromptChoiceRegistry", + "decode_stage2_answer", + "encode_stage2_answer", "Prompt", "build_stage1_prompt", "build_stage2_prompt", diff --git a/src/agent/task/prompt_choices.py b/src/agent/task/prompt_choices.py new file mode 100644 index 0000000..7442e18 --- /dev/null +++ b/src/agent/task/prompt_choices.py @@ -0,0 +1,215 @@ +"""Prompt-facing adapter between canonical category_ids and LLM action ids. + +Boundary contract: +- canonical category_id stays the ONLY identity of the whole pipeline + (registry / corpus / ground truth / evaluation / reward). It never changes. +- choice_id is a compact, deterministic numbering ("1", "2", ... following + LeafRegistry.categories order) that exists ONLY inside prompts and model + outputs. It is decoded back to category_id immediately at the LLM boundary + and is never written back to canonical SampleTarget / CorpusCategory. +- display_name is the shortest unambiguous suffix of the category path: the + leaf name when unique in the registry, otherwise parent-qualified until + unique. No hashes, UUIDs or permanent encodings are introduced. + +Stage 2 uses LOCAL bundle ids ("1".."5") in candidate order instead of the +global choice ids; decode is positional against the candidate bundle. +""" + +from __future__ import annotations + +from collections import Counter +from dataclasses import dataclass, field +from typing import Sequence + +from .contracts import LeafCategory, LeafRegistry + + +class PromptChoiceError(ValueError): + """Raised when a choice id cannot be mapped to a canonical category_id.""" + + +@dataclass(frozen=True) +class PromptChoice: + """One prompt-facing entry: choice_id <-> category_id + display name.""" + + choice_id: str + category_id: str + display_name: str + + +_STAGE2_LOCAL_IDS = tuple(str(index) for index in range(1, 6)) + + +def _path_parts(category: LeafCategory) -> tuple[str, ...]: + """Path parts with the leaf name guaranteed as the last element.""" + parts = list(category.path) if category.path else [] + if not parts or parts[-1] != category.name: + parts = parts + [category.name] + return tuple(parts) + + +def _display_names(registry: LeafRegistry) -> tuple[str, ...]: + """Shortest unique path suffix per category; raises when impossible. + + A unique leaf name stays as-is. Duplicate leaf names are qualified with + one parent level at a time until the suffix is unique against every + other category; registries whose duplicates cannot be disambiguated + (empty paths, identical full paths) fail explicitly instead of silently + producing ambiguous display names. + """ + counts = Counter(category.name for category in registry.categories) + parts = {id(category): _path_parts(category) for category in registry.categories} + names: list[str] = [] + for category in registry.categories: + if counts[category.name] == 1: + names.append(category.name) + continue + path = parts[id(category)] + for depth in range(1, len(path) + 1): + candidate = " / ".join(path[-depth:]) + if not any( + other is not category + and " / ".join(parts[id(other)][-depth:]) == candidate + for other in registry.categories + ): + names.append(candidate) + break + else: + raise PromptChoiceError( + f"cannot build a unique display name for leaf {category.name!r} " + f"(category_id {category.category_id!r}): duplicate leaf names " + "cannot be disambiguated by path suffix" + ) + if len(set(names)) != len(names): + raise PromptChoiceError("display names must be unique across the registry") + return tuple(names) + + +@dataclass(frozen=True) +class PromptChoiceRegistry: + """Deterministic prompt-facing view over one LeafRegistry. + + choice ids are "1".."N" following LeafRegistry.categories order, so the + mapping is stable across runs and identical registries. The registry is + kept for coverage checks; canonical contracts are never modified. + """ + + registry: LeafRegistry + choices: tuple[PromptChoice, ...] + _by_choice_id: dict[str, PromptChoice] = field( + init=False, compare=False, repr=False + ) + _by_category_id: dict[str, PromptChoice] = field( + init=False, compare=False, repr=False + ) + + def __post_init__(self) -> None: + by_choice = {choice.choice_id: choice for choice in self.choices} + by_category = {choice.category_id: choice for choice in self.choices} + if len(by_choice) != len(self.choices): + raise PromptChoiceError("prompt choice ids must be unique") + if len(by_category) != len(self.choices): + raise PromptChoiceError("prompt choices must not duplicate category ids") + if set(by_category) != set(self.registry.ids): + raise PromptChoiceError("prompt choices must cover the leaf registry") + object.__setattr__(self, "_by_choice_id", by_choice) + object.__setattr__(self, "_by_category_id", by_category) + + @classmethod + def from_registry(cls, registry: LeafRegistry) -> "PromptChoiceRegistry": + display_names = _display_names(registry) + choices = tuple( + PromptChoice( + choice_id=str(index), + category_id=category.category_id, + display_name=display_name, + ) + for index, (category, display_name) in enumerate( + zip(registry.categories, display_names), start=1 + ) + ) + return cls(registry, choices) + + @property + def choice_ids(self) -> tuple[str, ...]: + return tuple(choice.choice_id for choice in self.choices) + + def contains_choice_id(self, choice_id: str) -> bool: + return choice_id in self._by_choice_id + + def contains_category_id(self, category_id: str) -> bool: + return category_id in self._by_category_id + + def choice_id_of(self, category_id: str) -> str: + choice = self._by_category_id.get(category_id) + if choice is None: + raise PromptChoiceError( + f"category_id {category_id!r} has no prompt choice" + ) + return choice.choice_id + + def category_id_of(self, choice_id: str) -> str: + choice = self._by_choice_id.get(choice_id) + if choice is None: + raise PromptChoiceError( + f"choice id {choice_id!r} is not in the prompt catalog" + ) + return choice.category_id + + def display_name_of(self, category_id: str) -> str: + choice = self._by_category_id.get(category_id) + if choice is None: + raise PromptChoiceError( + f"category_id {category_id!r} has no prompt choice" + ) + return choice.display_name + + def encode_candidates(self, category_ids: Sequence[str]) -> tuple[str, ...]: + """Canonical category_ids -> global choice ids (strict Stage 1 shape).""" + if len(category_ids) != 5 or len(set(category_ids)) != 5: + raise PromptChoiceError("stage1 requires exactly 5 unique candidates") + return tuple(self.choice_id_of(category_id) for category_id in category_ids) + + def decode_candidates(self, choice_ids: Sequence[str]) -> tuple[str, ...]: + """Global choice ids -> canonical category_ids (strict Stage 1 shape). + + Raises PromptChoiceError for anything but exactly 5 unique known + choice ids; there is deliberately no name/fuzzy fallback. + """ + if len(choice_ids) != 5: + raise PromptChoiceError("stage1 prediction must contain exactly 5 candidates") + if len(set(choice_ids)) != len(choice_ids): + raise PromptChoiceError("stage1 candidates must be unique") + return tuple(self.category_id_of(choice_id) for choice_id in choice_ids) + + +def encode_stage2_answer(category_id: str, candidates: Sequence[str]) -> str: + """Canonical answer -> local bundle id ("1".."5") in candidate order.""" + if len(candidates) != 5: + raise PromptChoiceError("stage2 requires exactly 5 candidates") + try: + return str(candidates.index(category_id) + 1) + except ValueError: + raise PromptChoiceError("stage2 answer must be one of the candidates") from None + + +def decode_stage2_answer(answer: str, candidates: Sequence[str]) -> str: + """Local bundle id ("1".."5") -> canonical category_id. + + Raises PromptChoiceError for anything but an exact local id; there is + deliberately no name/fuzzy fallback. + """ + if len(candidates) != 5: + raise PromptChoiceError("stage2 requires exactly 5 candidates") + if answer not in _STAGE2_LOCAL_IDS: + raise PromptChoiceError(f"stage2 answer {answer!r} must be one of 1..5") + return candidates[int(answer) - 1] + + +__all__ = [ + "PromptChoice", + "PromptChoiceError", + "PromptChoiceRegistry", + "encode_stage2_answer", + "decode_stage2_answer", +] diff --git a/src/agent/task/prompts.py b/src/agent/task/prompts.py index 330ca3c..220ea1c 100644 --- a/src/agent/task/prompts.py +++ b/src/agent/task/prompts.py @@ -1,10 +1,20 @@ """Deterministic two-stage HF-message prompts for leaf classification. -Stage 1 receives the FULL LeafRegistry as candidate universe, rendered as -category_id + name pairs only (descriptions/examples are deliberately kept -out of Stage 1). Stage 2 resolves candidates by category_id against the -canonical corpus (category_id/name/description/descriptions/examples); when -no corpus is provided it falls back to the registry name/description. +Prompt-facing identity is the choice protocol, never the canonical +category_id: +- Stage 1 receives the FULL LeafRegistry as candidate universe, rendered as + compact [choice_id, display_name] pairs only (descriptions/examples are + deliberately kept out of Stage 1). The model answers with global choice + ids ("1".."N" following registry order). +- Stage 2 resolves candidates by canonical category_id against the + canonical corpus (description/descriptions/examples) and renders each + candidate with a LOCAL bundle id ("1".."5"); the model answers with the + local id. When no corpus is provided it falls back to the registry + description. + +Decoding back to canonical category_id happens immediately at the LLM +boundary (evaluation adapters / SFT validator); choice ids never leak into +corpus lookup, canonical targets or reward semantics. """ from __future__ import annotations @@ -14,6 +24,7 @@ from typing import Mapping, Sequence from .contracts import CorpusCategory, LeafRegistry, TaskConfig +from .prompt_choices import PromptChoiceRegistry, encode_stage2_answer @dataclass(frozen=True) @@ -31,22 +42,23 @@ def _metadata_text(metadata: Mapping[str, object], config: TaskConfig) -> str: def build_stage1_prompt( - metadata: Mapping[str, object], registry: LeafRegistry, config: TaskConfig + metadata: Mapping[str, object], + registry: LeafRegistry, + config: TaskConfig, + choices: PromptChoiceRegistry | None = None, ) -> Prompt: + choices = choices or PromptChoiceRegistry.from_registry(registry) system = ( - "You are a leaf-category candidate retriever. Return only one JSON object, " - 'with exactly this shape: {"candidates":["category_id", "category_id", ' - '"category_id", "category_id", "category_id"]}. ' - "The candidates array must contain exactly 5 unique category_id values " - "from the registry. " - "Do not output Markdown, commentary, or any other keys." + "You are a leaf-category candidate retriever. Return exactly one JSON object " + 'with key "candidates". ' + "The value must contain exactly five unique choice ids from the catalog. " + "Do not output Markdown, commentary, canonical category ids, or any other keys." ) catalog = [ - {"category_id": category.category_id, "name": category.name} - for category in registry.categories + [choice.choice_id, choice.display_name] for choice in choices.choices ] user = ( - "Retrieve five candidate leaf categories from this registry:\n" + "Retrieve five candidate leaf categories from this catalog:\n" + json.dumps(catalog, ensure_ascii=False) + "\nField metadata:\n" + _metadata_text(metadata, config) @@ -60,13 +72,16 @@ def build_stage2_prompt( registry: LeafRegistry, config: TaskConfig, corpus: Mapping[str, CorpusCategory] | None = None, + choices: PromptChoiceRegistry | None = None, ) -> Prompt: if len(candidates) != 5 or len(set(candidates)) != 5: raise ValueError("stage2 requires exactly 5 unique candidates") if any(candidate not in registry.ids for candidate in candidates): raise ValueError("stage2 candidates must belong to the leaf registry") + choices = choices or PromptChoiceRegistry.from_registry(registry) bundle = [] - for category_id in candidates: + for index, category_id in enumerate(candidates, start=1): + display_name = choices.display_name_of(category_id) if corpus is not None: # resolve by category_id only; never by bare leaf name corpus_category = corpus.get(category_id) @@ -76,8 +91,8 @@ def build_stage2_prompt( ) bundle.append( { - "category_id": category_id, - "name": corpus_category.name, + "id": str(index), + "name": display_name, "description": corpus_category.description, "descriptions": list(corpus_category.descriptions), "examples": list(corpus_category.examples), @@ -87,14 +102,15 @@ def build_stage2_prompt( category = registry.get(category_id) bundle.append( { - "category_id": category_id, - "name": category.name, + "id": str(index), + "name": display_name, "description": category.description, } ) system = ( - "You are a leaf-category reranker. Return only one JSON object, exactly " - '{"answer":"category_id"}. The answer must be one of the five candidates. ' + "You are a leaf-category reranker. Return exactly one JSON object with key " + '"answer". ' + 'Its value must be one of the five candidate ids "1" through "5". ' "Do not output Markdown, commentary, or any other keys." ) user = ( @@ -106,13 +122,24 @@ def build_stage2_prompt( return Prompt(system, user) -def stage1_answer(candidates: Sequence[str]) -> str: - if len(candidates) != 5 or len(set(candidates)) != 5: - raise ValueError("stage1 requires exactly 5 unique candidates") - return json.dumps({"candidates": list(candidates)}, ensure_ascii=False, separators=(",", ":")) +def stage1_answer( + candidates: Sequence[str], *, choices: PromptChoiceRegistry +) -> str: + """Assistant answer for Stage 1: canonical candidates -> global choice ids.""" + choice_ids = choices.encode_candidates(candidates) + return json.dumps( + {"candidates": list(choice_ids)}, + ensure_ascii=False, + separators=(",", ":"), + ) def stage2_answer(category_id: str, candidates: Sequence[str]) -> str: + """Assistant answer for Stage 2: canonical answer -> local bundle id.""" if category_id not in candidates: raise ValueError("stage2 answer must be one of the candidates") - return json.dumps({"answer": category_id}, ensure_ascii=False, separators=(",", ":")) + return json.dumps( + {"answer": encode_stage2_answer(category_id, candidates)}, + ensure_ascii=False, + separators=(",", ":"), + ) diff --git a/src/agent/training/common.py b/src/agent/training/common.py index 5a8f39f..0708a3f 100644 --- a/src/agent/training/common.py +++ b/src/agent/training/common.py @@ -8,8 +8,9 @@ from __future__ import annotations +import hashlib from pathlib import Path -from typing import Any, Mapping +from typing import Any, Mapping, Sequence from agent.task.contracts import CorpusCategory, LeafRegistry @@ -51,13 +52,47 @@ def canonical_target( return category_id -def build_candidates(ground_truth: str, registry: LeafRegistry) -> list[str]: - """Deterministic baseline/test fixture policy: GT followed by the first - four non-GT registry IDs. This is a fixture, NOT the production Stage 1 - retrieval policy.""" +def build_candidates( + ground_truth: str, + registry: LeafRegistry, + *, + source_id: str, +) -> list[str]: + """Deterministic candidate bundle shared by SFT and RL for one sample. + + Baseline fixture policy (no hard negatives): the ground truth plus the + first four non-GT registry ids, then permuted deterministically from the + stable ``source_id``. The permutation keeps the bundle reproducible + (same source_id always yields the same ordering, across runs and across + the SFT/RL exporters) while preventing the ground truth from sitting at + a fixed position — which would otherwise leak a systematic + ``{"answer":"1"}`` bias into every Stage 2 sample. This is a fixture, + NOT the production Stage 1 retrieval policy. + """ if ground_truth not in registry.ids: - raise ValueError(f"ground-truth category_id is absent from leaf registry: {ground_truth}") - return [ground_truth] + [category_id for category_id in registry.ids if category_id != ground_truth][:4] + raise ValueError( + f"ground-truth category_id is absent from leaf registry: {ground_truth}" + ) + base = [ground_truth] + [ + category_id for category_id in registry.ids if category_id != ground_truth + ][:4] + return _permute(base, source_id) + + +def _permute(items: Sequence[str], source_id: str) -> list[str]: + """Deterministic Fisher–Yates shuffle keyed by a stable sha256 digest. + + ``hash()`` is randomized per process (PYTHONHASHSEED) and must never be + used as a reproducibility seed; sha256 over the UTF-8 source_id is + stable across runs, machines and processes. There is no runtime + randomness. + """ + digest = hashlib.sha256(source_id.encode("utf-8")).digest() + result = list(items) + for index in range(len(result) - 1, 0, -1): + swap = digest[index % len(digest)] % (index + 1) + result[index], result[swap] = result[swap], result[index] + return result def require_corpus(corpus: Mapping[str, CorpusCategory]) -> Mapping[str, CorpusCategory]: diff --git a/src/agent/training/rl/dataset.py b/src/agent/training/rl/dataset.py index 55aa06f..e6f3757 100644 --- a/src/agent/training/rl/dataset.py +++ b/src/agent/training/rl/dataset.py @@ -133,8 +133,9 @@ def export_rl_dataset( "version": "verl 0.8.0 five-field schema (data_source/prompt/ability/reward_model/extra_info)", "label_source": "canonical target.category_id (resolution_status == resolved)", "candidate_policy": ( - "baseline fixture: ground_truth_then_registry_order_first_four_" - "non_ground_truth (not the production retrieval policy)" + "baseline fixture: (ground_truth + first four registry negatives) " + "permuted deterministically by stable source_id (not position-fixed, " + "not the production retrieval policy)" ), "dataset": dataset, "metadata_fields": list(config.metadata_fields), @@ -334,8 +335,26 @@ def _validate_row( expected_prompt.user, ]: errors.append("prompt does not match registry and task contract") - if stage == "stage1" and contents and contents[1].count('"category_id"') != len(registry.categories): - errors.append("stage1 prompt must render the full leaf registry catalog") + if stage == "stage1" and contents: + user = contents[1] + if '"category_id"' in user: + errors.append("stage1 prompt must not expose canonical category ids") + try: + catalog = json.loads( + user.split("\n", 1)[1].split("\nField metadata:", 1)[0] + ) + except (json.JSONDecodeError, AttributeError): + errors.append("stage1 prompt catalog must be a JSON array") + else: + if ( + not isinstance(catalog, list) + or len(catalog) != len(registry.categories) + or any( + not (isinstance(entry, list) and len(entry) == 2) + for entry in catalog + ) + ): + errors.append("stage1 prompt must render the full leaf registry catalog") return errors diff --git a/src/agent/training/rl/sample.py b/src/agent/training/rl/sample.py index 4bf9247..0b9bdcc 100644 --- a/src/agent/training/rl/sample.py +++ b/src/agent/training/rl/sample.py @@ -124,7 +124,8 @@ def build_rl_samples( Only resolved records may be passed (the caller filters on resolution_status); the ground truth is validated as target.category_id with the registry as the final constraint. Stage 2 candidates reuse the - deterministic fixture policy (GT + first four non-GT registry IDs). + deterministic fixture policy (GT + four non-GT registry ids, permuted + deterministically by the stable source_id). """ ground_truth = canonical_target(item, index, source, registry) if ground_truth is None: @@ -133,7 +134,7 @@ def build_rl_samples( if not source_id: raise ValueError(f"item {index} in {source} has no stable id") metadata = visible_metadata(item.get("metadata", {}), config) - candidates = tuple(build_candidates(ground_truth, registry)) + candidates = tuple(build_candidates(ground_truth, registry, source_id=source_id)) stage1_prompt = build_stage1_prompt(metadata, registry, config) stage2_prompt = build_stage2_prompt( metadata, candidates, registry, config, corpus=corpus or None diff --git a/src/agent/training/sft/dataset.py b/src/agent/training/sft/dataset.py index 207f01f..1f19c6e 100644 --- a/src/agent/training/sft/dataset.py +++ b/src/agent/training/sft/dataset.py @@ -19,8 +19,9 @@ from pathlib import Path from typing import Any, Mapping -from agent.evaluation import evaluate_stage1, evaluate_stage2 +from agent.evaluation import evaluate_stage1_choices, evaluate_stage2_choices from agent.task.contracts import CorpusCategory, LeafRegistry, TaskConfig +from agent.task.prompt_choices import PromptChoiceRegistry from agent.task.prompts import ( build_stage1_prompt, build_stage2_prompt, @@ -84,13 +85,19 @@ def _row( source_id = str(item.get("id", "")).strip() if not source_id: raise ValueError(f"item {index} in {source} has no stable id") - candidates = build_candidates(ground_truth, registry) + candidates = build_candidates(ground_truth, registry, source_id=source_id) + choices = PromptChoiceRegistry.from_registry(registry) if stage == "stage1": - prompt = build_stage1_prompt(visible_metadata, registry, config) - answer = stage1_answer(candidates) + prompt = build_stage1_prompt(visible_metadata, registry, config, choices=choices) + answer = stage1_answer(candidates, choices=choices) else: prompt = build_stage2_prompt( - visible_metadata, candidates, registry, config, corpus=corpus or None + visible_metadata, + candidates, + registry, + config, + corpus=corpus or None, + choices=choices, ) answer = stage2_answer(ground_truth, candidates) return { @@ -176,9 +183,15 @@ def export_sft_dataset( report: dict[str, Any] = { "format": "verl_sft_messages_parquet", "label_source": "canonical target.category_id (resolution_status == resolved)", + "prompt_identity": ( + "choice ids in messages; ground_truth/candidates stay canonical " + "category_id (PromptChoiceRegistry, display names = shortest " + "unique path suffix)" + ), "candidate_policy": ( - "baseline fixture: ground_truth_then_registry_order_first_four_" - "non_ground_truth (not the production retrieval policy)" + "baseline fixture: (ground_truth + first four registry negatives) " + "permuted deterministically by stable source_id (not position-fixed, " + "not the production retrieval policy)" ), "metadata_fields": list(config.metadata_fields), "task_name": config.task_name, @@ -247,6 +260,7 @@ def _validate_row( corpus: Mapping[str, CorpusCategory], ) -> list[str]: errors: list[str] = [] + choices = PromptChoiceRegistry.from_registry(registry) messages = row.get("messages") if not isinstance(messages, list) or len(messages) != 3: return ["messages must contain system, user, assistant"] @@ -292,16 +306,17 @@ def _validate_row( isinstance(ground_truth, str) and ground_truth in registry.ids ) if stage == "stage1" and ground_truth_is_valid: - evaluation = evaluate_stage1( + evaluation = evaluate_stage1_choices( assistant, ground_truth=ground_truth, registry=registry, + choices=choices, ) errors.extend(f"stage1 evaluation: {error}" for error in evaluation.errors) if evaluation.prediction is None or list(evaluation.prediction) != candidates: errors.append("stage1 answer must exactly match the five candidates") elif stage == "stage2" and candidates_belong_to_registry and ground_truth_is_valid: - evaluation = evaluate_stage2( + evaluation = evaluate_stage2_choices( assistant, ground_truth=ground_truth, candidates=tuple(candidates), @@ -320,14 +335,21 @@ def _validate_row( elif all(isinstance(content, str) for content in contents): expected_prompt = None if stage == "stage1": - expected_prompt = build_stage1_prompt(visible_metadata, registry, config) + expected_prompt = build_stage1_prompt( + visible_metadata, registry, config, choices=choices + ) elif ( stage == "stage2" and candidates_are_valid and all(candidate in registry.ids for candidate in candidates) ): expected_prompt = build_stage2_prompt( - visible_metadata, candidates, registry, config, corpus=corpus or None + visible_metadata, + candidates, + registry, + config, + corpus=corpus or None, + choices=choices, ) if expected_prompt is not None and contents[:2] != [ expected_prompt.system, @@ -335,10 +357,15 @@ def _validate_row( ]: errors.append("system/user prompt does not match registry and task contract") - if isinstance(ground_truth, str) and ground_truth in registry.ids: - expected = build_candidates(ground_truth, registry) + if ( + isinstance(ground_truth, str) + and ground_truth in registry.ids + and isinstance(source_id, str) + and source_id.strip() + ): + expected = build_candidates(ground_truth, registry, source_id=source_id) if candidates != expected: - errors.append("candidates do not follow deterministic registry order") + errors.append("candidates do not follow the deterministic source-seeded bundle order") return errors diff --git a/tests/evaluation/test_classification_choices.py b/tests/evaluation/test_classification_choices.py new file mode 100644 index 0000000..92a7c83 --- /dev/null +++ b/tests/evaluation/test_classification_choices.py @@ -0,0 +1,153 @@ +"""Evaluation adapters: model outputs speak choice ids, evaluation stays canonical.""" + +from agent.evaluation import ( + evaluate_stage1_choices, + evaluate_stage2_choices, +) +from agent.task import LeafRegistry, PromptChoiceRegistry + + +def _registry() -> LeafRegistry: + return LeafRegistry.from_mapping(["A", "B", "C", "D", "E", "F"]) + + +def _choices() -> PromptChoiceRegistry: + return PromptChoiceRegistry.from_registry(_registry()) + + +# candidate canonical ids for the fixture below +CANDIDATES = ("C", "A", "B", "D", "E") + + +def test_stage1_choices_decode_and_recall_ground_truth() -> None: + # choice ids: A=1 B=2 C=3 D=4 E=5 F=6 + evaluation = evaluate_stage1_choices( + '{"candidates":["3","1","2","4","5"]}', + ground_truth="C", + registry=_registry(), + choices=_choices(), + ) + + assert evaluation.format_valid is True + assert evaluation.contract_valid is True + assert evaluation.ground_truth_recalled is True + # prediction is the DECODED canonical tuple + assert evaluation.prediction == ("C", "A", "B", "D", "E") + assert evaluation.errors == () + + +def test_stage1_choices_not_recalled_is_contract_valid() -> None: + evaluation = evaluate_stage1_choices( + '{"candidates":["1","2","4","5","6"]}', + ground_truth="C", + registry=_registry(), + choices=_choices(), + ) + + assert evaluation.contract_valid is True + assert evaluation.ground_truth_recalled is False + assert evaluation.prediction == ("A", "B", "D", "E", "F") + + +def test_stage1_choices_reject_unknown_choice_id() -> None: + evaluation = evaluate_stage1_choices( + '{"candidates":["3","1","2","4","9"]}', + ground_truth="C", + registry=_registry(), + choices=_choices(), + ) + + assert evaluation.format_valid is True + assert evaluation.contract_valid is False + assert evaluation.ground_truth_recalled is False + assert evaluation.prediction is None + assert any("not in the prompt catalog" in error for error in evaluation.errors) + + +def test_stage1_choices_reject_duplicates_and_wrong_count() -> None: + duplicate = evaluate_stage1_choices( + '{"candidates":["1","1","2","4","5"]}', + ground_truth="C", + registry=_registry(), + choices=_choices(), + ) + short = evaluate_stage1_choices( + '{"candidates":["1","2","3"]}', + ground_truth="C", + registry=_registry(), + choices=_choices(), + ) + + assert duplicate.contract_valid is False + assert any("unique" in error for error in duplicate.errors) + assert short.contract_valid is False + assert any("exactly 5" in error for error in short.errors) + + +def test_stage1_choices_reject_non_json() -> None: + evaluation = evaluate_stage1_choices( + "not json", + ground_truth="C", + registry=_registry(), + choices=_choices(), + ) + + assert evaluation.format_valid is False + assert evaluation.contract_valid is False + assert evaluation.prediction is None + + +def test_stage2_choices_decode_local_id_to_canonical() -> None: + correct = evaluate_stage2_choices( + '{"answer":"1"}', + ground_truth="C", + candidates=CANDIDATES, + registry=_registry(), + ) + wrong = evaluate_stage2_choices( + '{"answer":"2"}', + ground_truth="C", + candidates=CANDIDATES, + registry=_registry(), + ) + + assert correct.format_valid is True + assert correct.contract_valid is True + assert correct.correct is True + assert correct.prediction == "C" + assert wrong.contract_valid is True + assert wrong.correct is False + assert wrong.prediction == "A" + + +def test_stage2_choices_reject_invalid_local_ids() -> None: + outside = evaluate_stage2_choices( + '{"answer":"6"}', + ground_truth="C", + candidates=CANDIDATES, + registry=_registry(), + ) + non_numeric = evaluate_stage2_choices( + '{"answer":"A"}', + ground_truth="C", + candidates=CANDIDATES, + registry=_registry(), + ) + + assert outside.contract_valid is False + assert outside.correct is False + assert any("one of 1..5" in error for error in outside.errors) + assert non_numeric.contract_valid is False + + +def test_stage2_choices_reject_non_json() -> None: + evaluation = evaluate_stage2_choices( + "not json", + ground_truth="C", + candidates=CANDIDATES, + registry=_registry(), + ) + + assert evaluation.format_valid is False + assert evaluation.contract_valid is False + assert evaluation.prediction is None diff --git a/tests/rl/test_rl_canonical_e2e.py b/tests/rl/test_rl_canonical_e2e.py index dbd140f..4d9906c 100644 --- a/tests/rl/test_rl_canonical_e2e.py +++ b/tests/rl/test_rl_canonical_e2e.py @@ -157,12 +157,26 @@ def test_no_unresolved_samples_in_exports(exports: dict) -> None: def test_stage1_candidate_universe_is_full_registry(exports: dict) -> None: + from agent.task import PromptChoiceRegistry + for dataset, bundle in exports.items(): registry = LeafRegistry.from_path(bundle["registry_path"]) rows = pq.read_table(bundle["out"] / "train.parquet").to_pylist() stage1 = next(row for row in rows if row["extra_info"]["stage"] == "stage1") user = stage1["prompt"][1]["content"] - assert user.count('"category_id"') == len(registry.categories), dataset + # compact [choice_id, display_name] catalog; canonical ids never shown + assert '"category_id"' not in user, dataset + catalog = json.loads(user.split("\n", 1)[1].split("\nField metadata:", 1)[0]) + assert len(catalog) == len(registry.categories), dataset + assert [entry[0] for entry in catalog] == [ + str(index) for index in range(1, len(catalog) + 1) + ], dataset + display_names = [entry[1] for entry in catalog] + assert len(set(display_names)) == len(display_names), dataset + choices = PromptChoiceRegistry.from_registry(registry) + assert [choices.category_id_of(entry[0]) for entry in catalog] == list( + registry.ids + ), dataset def test_stage2_corpus_lookup_by_category_id(exports: dict) -> None: diff --git a/tests/rl/test_rl_dataset.py b/tests/rl/test_rl_dataset.py index 13abee1..e52e7be 100644 --- a/tests/rl/test_rl_dataset.py +++ b/tests/rl/test_rl_dataset.py @@ -124,11 +124,21 @@ def test_target_category_id_is_the_only_label(exported) -> None: def test_stage1_prompt_uses_full_registry(exported, registry) -> None: + from agent.task import PromptChoiceRegistry + rows = pq.read_table(exported["out"] / "train.parquet").to_pylist() stage1 = next(row for row in rows if row["extra_info"]["stage"] == "stage1") user = stage1["prompt"][1]["content"] - assert user.count('"category_id"') == len(registry.categories) - assert user.count('"name"') == len(registry.categories) + # compact [choice_id, display_name] catalog; canonical ids never shown + assert '"category_id"' not in user + catalog = json.loads(user.split("\n", 1)[1].split("\nField metadata:", 1)[0]) + assert len(catalog) == len(registry.categories) + assert [entry[0] for entry in catalog] == [ + str(index) for index in range(1, len(registry.categories) + 1) + ] + assert len({entry[1] for entry in catalog}) == len(registry.categories) + choices = PromptChoiceRegistry.from_registry(registry) + assert [choices.category_id_of(entry[0]) for entry in catalog] == list(registry.ids) def test_stage2_corpus_lookup_by_category_id(exported, corpus) -> None: @@ -408,7 +418,7 @@ def test_ground_truth_not_required_in_stage2_candidates( a reward without raising.""" import agent.training.rl.sample as rl_sample_module - def candidates_excluding_gt(ground_truth: str, reg: LeafRegistry): + def candidates_excluding_gt(ground_truth: str, reg: LeafRegistry, *, source_id: str): return [candidate for candidate in reg.ids if candidate != ground_truth][:5] monkeypatch.setattr( diff --git a/tests/sft/test_sft_canonical_e2e.py b/tests/sft/test_sft_canonical_e2e.py index 4ff8d7d..069bc75 100644 --- a/tests/sft/test_sft_canonical_e2e.py +++ b/tests/sft/test_sft_canonical_e2e.py @@ -13,7 +13,7 @@ import pyarrow.parquet as pq import pytest -from agent.task import LeafRegistry, TaskConfig +from agent.task import LeafRegistry, PromptChoiceRegistry, TaskConfig from agent.task.canonical_dataset import load_corpus_categories from agent.training.sft import export_sft_dataset, validate_sft_dataset @@ -148,15 +148,72 @@ def test_no_unresolved_samples_in_exports(exports: dict) -> None: assert details["skipped_not_resolved"] == unresolved_in_split, (dataset, split) +def _stage1_catalog(user: str) -> list: + """Extract the JSON catalog array from a Stage 1 user message.""" + body = user.split("\n", 1)[1].split("\nField metadata:", 1)[0] + return json.loads(body) + + def test_stage1_candidate_universe_is_full_registry(exports: dict) -> None: for dataset, bundle in exports.items(): registry = LeafRegistry.from_path(bundle["registry_path"]) rows = pq.read_table(bundle["out"] / "train.parquet").to_pylist() stage1 = next(row for row in rows if row["stage"] == "stage1") user = stage1["messages"][1]["content"] - # the stage-1 catalog renders every registry category as id+name - assert user.count('"category_id"') == len(registry.categories), dataset - assert f'"name": "{registry.categories[0].name}"' in user + # the stage-1 catalog renders every registry category as compact + # [choice_id, display_name] pairs; canonical ids never appear + assert '"category_id"' not in user, dataset + catalog = _stage1_catalog(user) + assert isinstance(catalog, list) and len(catalog) == len(registry.categories), dataset + assert all(isinstance(entry, list) and len(entry) == 2 for entry in catalog), dataset + assert [entry[0] for entry in catalog] == [ + str(index) for index in range(1, len(catalog) + 1) + ], dataset + display_names = [entry[1] for entry in catalog] + assert len(set(display_names)) == len(display_names), dataset + + +def test_finance_duplicate_leaf_names_are_disambiguated(exports: dict) -> None: + bundle = exports["finance"] + registry = LeafRegistry.from_path(bundle["registry_path"]) + rows = pq.read_table(bundle["out"] / "train.parquet").to_pylist() + stage1 = next(row for row in rows if row["stage"] == "stage1") + catalog = _stage1_catalog(stage1["messages"][1]["content"]) + # duplicate leaf names (e.g. 基本信息) render as parent-qualified suffixes + for category in registry.categories: + if category.name == "基本信息": + display = next( + entry[1] for entry in catalog if entry[1].endswith("基本信息") + ) + assert display != "基本信息", category.category_id + assert " / " in display + # every display name still resolves to exactly one canonical category + choices = PromptChoiceRegistry.from_registry(registry) + for entry in catalog: + assert choices.category_id_of(entry[0]) in registry.ids + + +def test_shougang_and_infra_short_codes_keep_canonical_contract(exports: dict) -> None: + """Code-strategy registries (A1-1-1 ...) must keep canonical category_ids + in ground_truth/candidates while prompts use choice ids.""" + for dataset in ("shougang", "infra"): + bundle = exports[dataset] + registry = LeafRegistry.from_path(bundle["registry_path"]) + # code strategy intact: every category carries its guanji code id + assert all(category.code for category in registry.categories), dataset + assert all(category.category_id == category.code for category in registry.categories), dataset + rows = pq.read_table(bundle["out"] / "train.parquet").to_pylist() + for row in rows: + assert row["ground_truth"] in registry.ids, (dataset, row["source_id"]) + assert all(candidate in registry.ids for candidate in row["candidates"]) + stage1 = next(row for row in rows if row["stage"] == "stage1") + assert '"category_id"' not in stage1["messages"][1]["content"] + choices = PromptChoiceRegistry.from_registry(registry) + # decoded assistant answer must round-trip to canonical candidates + decoded = choices.decode_candidates( + json.loads(stage1["messages"][-1]["content"])["candidates"] + ) + assert list(decoded) == stage1["candidates"] def test_stage2_corpus_lookup_by_category_id(exports: dict) -> None: diff --git a/tests/sft/test_sft_dataset.py b/tests/sft/test_sft_dataset.py index 2c72e8c..d518f00 100644 --- a/tests/sft/test_sft_dataset.py +++ b/tests/sft/test_sft_dataset.py @@ -124,7 +124,7 @@ def test_exporter_writes_messages_parquet_and_report(tmp_path): )["valid"] is True -def test_stage1_prompt_catalog_contains_id_and_name(tmp_path): +def test_stage1_prompt_catalog_uses_choice_ids_and_display_names(tmp_path): registry_path, config_path, _ = _write_inputs(tmp_path / "input") out = tmp_path / "out" export_sft_dataset( @@ -138,10 +138,65 @@ def test_stage1_prompt_catalog_contains_id_and_name(tmp_path): rows = pq.read_table(out / "train.parquet").to_pylist() stage1 = next(row for row in rows if row["stage"] == "stage1") user = stage1["messages"][1]["content"] - assert '"category_id": "A"' in user and '"name": "alpha data"' in user + assert '"category_id"' not in user # canonical ids never enter the prompt + assert '["1", "alpha data"]' in user + assert '"3"' in user # every choice id is rendered assert "gamma desc" not in user # descriptions stay out of Stage 1 +def test_assistant_answers_use_choice_ids_but_metadata_stays_canonical(tmp_path): + from agent.task.prompt_choices import PromptChoiceRegistry, decode_stage2_answer + + registry_path, config_path, _ = _write_inputs(tmp_path / "input") + out = tmp_path / "out" + export_sft_dataset( + tmp_path / "input" / "canonical" / "all.json", + tmp_path / "input", + out, + registry_path, + config_path, + corpus=_corpus_map(), + ) + rows = pq.read_table(out / "train.parquet").to_pylist() + stage1 = next(row for row in rows if row["stage"].startswith("stage1")) + stage2 = next(row for row in rows if row["stage"].startswith("stage2")) + choices = PromptChoiceRegistry.from_registry(LeafRegistry.from_mapping(REGISTRY)) + # messages speak choice ids; decoding restores the canonical bundle / GT + stage1_choice_ids = json.loads(stage1["messages"][-1]["content"])["candidates"] + assert choices.decode_candidates(stage1_choice_ids) == tuple(stage1["candidates"]) + stage2_local_id = json.loads(stage2["messages"][-1]["content"])["answer"] + assert decode_stage2_answer(stage2_local_id, tuple(stage2["candidates"])) == "C" + # messages speak choice ids, but metadata stays canonical category_id + assert stage1["ground_truth"] == "C" + assert stage2["ground_truth"] == "C" + assert stage1["candidates"] == stage2["candidates"] + assert all(candidate in LeafRegistry.from_mapping(REGISTRY).ids for candidate in stage2["candidates"]) + + +def test_stage2_assistant_answers_are_not_position_fixed(tmp_path): + """Phase 6: GT must not sit at local position 1 in every Stage 2 sample, + otherwise every gold is {'answer':'1'}.""" + registry_path, config_path, _ = _write_inputs(tmp_path / "input") + out = tmp_path / "out" + export_sft_dataset( + tmp_path / "input" / "canonical" / "all.json", + tmp_path / "input", + out, + registry_path, + config_path, + corpus=_corpus_map(), + ) + answers = set() + for split in ("train", "val", "test"): + rows = pq.read_table(out / f"{split}.parquet").to_pylist() + for row in rows: + if row["stage"] == "stage2": + answers.add(row["messages"][-1]["content"]) + assert answers + assert len(answers) > 1 # deterministic: row-train@2, row-val@4, row-test@4 + assert all(answer != '{"answer":"1"}' for answer in answers) + + def test_stage2_prompt_resolves_corpus_by_category_id(tmp_path): registry_path, config_path, corpus_path = _write_inputs(tmp_path / "input") out = tmp_path / "out" @@ -156,17 +211,26 @@ def test_stage2_prompt_resolves_corpus_by_category_id(tmp_path): rows = pq.read_table(out / "train.parquet").to_pylist() stage2 = next(row for row in rows if row["stage"] == "stage2") user = stage2["messages"][1]["content"] + assert '"id":"1"' in user # local bundle id, not canonical id assert '"name":"alpha data"' in user assert '"descriptions":[]' in user and '"examples":["ex1"]' in user + assert '"category_id"' not in user -def test_candidate_construction_is_deterministic_and_gt_first(): +def test_candidate_construction_is_deterministic_and_source_seeded(): from agent.training.sft import build_candidates registry = LeafRegistry.from_mapping(REGISTRY) - expected = ["D", "A", "B", "C", "E"] - assert build_candidates("D", registry) == expected - assert build_candidates("D", registry) == expected + # same source_id -> identical permuted bundle + assert build_candidates("C", registry, source_id="row-train") == build_candidates( + "C", registry, source_id="row-train" + ) + # GT present, exactly 5 unique + bundle = build_candidates("C", registry, source_id="row-train") + assert len(bundle) == 5 and len(set(bundle)) == 5 and "C" in bundle + # GT is not fixed at position 1 + assert build_candidates("C", registry, source_id="row-train").index("C") != 0 + assert build_candidates("C", registry, source_id="row-val").index("C") != 0 def test_unresolved_records_never_enter_training(tmp_path): @@ -573,8 +637,11 @@ def test_validator_returns_structured_errors_for_wrong_answers(tmp_path): corpus=_corpus_map(), ) rows = pq.read_table(out / "train.parquet").to_pylist() - rows[0]["messages"][-1]["content"] = '{"candidates":["A"]}' - rows[1]["messages"][-1]["content"] = '{"answer":"A"}' + rows[0]["messages"][-1]["content"] = '{"candidates":["1"]}' + # row-train: GT "C" sits at local position 2 (bundle E,C,D,A,B), so "3" is + # a valid-but-wrong answer (-> "D") and must trigger an "equal ground_truth" + # error, while invalid local ids stay contract errors + rows[1]["messages"][-1]["content"] = '{"answer":"3"}' pq.write_table(pa.Table.from_pylist(rows), out / "train.parquet") report = validate_sft_dataset(out, registry_path, config_path, corpus=_corpus_map()) diff --git a/tests/task/test_contracts_and_prompts.py b/tests/task/test_contracts_and_prompts.py index 4f65165..6e89f16 100644 --- a/tests/task/test_contracts_and_prompts.py +++ b/tests/task/test_contracts_and_prompts.py @@ -2,9 +2,12 @@ from agent.task import ( LeafRegistry, + PromptChoiceRegistry, TaskConfig, build_stage1_prompt, build_stage2_prompt, + stage1_answer, + stage2_answer, ) @@ -29,16 +32,21 @@ def test_prompts_have_strict_stage_contracts_and_no_chatml_tokens() -> None: registry = LeafRegistry.from_mapping(REGISTRY) config = TaskConfig.from_mapping(CONFIG) - stage1 = build_stage1_prompt(METADATA, registry, config) + choices = PromptChoiceRegistry.from_registry(registry) + stage1 = build_stage1_prompt(METADATA, registry, config, choices=choices) stage2 = build_stage2_prompt( METADATA, ["C", "A", "B", "D", "E"], registry, config, + choices=choices, ) - assert '"candidates"' in stage1.system and "exactly 5" in stage1.system - assert '"answer"' in stage2.system and "one of the five" in stage2.system + assert '"candidates"' in stage1.system and "exactly five" in stage1.system + assert '"answer"' in stage2.system and '"1" through "5"' in stage2.system + # position-bias guard: system contracts must not showcase real legal ids + assert '["1","2","3","4","5"]' not in stage1.system + assert '{"answer":"1"}' not in stage2.system prompt_text = stage1.system + stage1.user + stage2.system + stage2.user assert all( token not in prompt_text @@ -48,6 +56,79 @@ def test_prompts_have_strict_stage_contracts_and_no_chatml_tokens() -> None: assert '"field_name"' in stage1.user +def test_stage1_prompt_uses_choice_ids_and_never_canonical_ids() -> None: + registry = LeafRegistry.from_mapping(REGISTRY) + config = TaskConfig.from_mapping(CONFIG) + stage1 = build_stage1_prompt(METADATA, registry, config) + + # compact [choice_id, display_name] pairs; no category_id anywhere + assert '["1", "A"]' in stage1.user + assert '"category_id"' not in stage1.user + assert "finance:" not in stage1.user + # descriptions stay out of Stage 1 + assert "alpha data" not in stage1.user + + +def test_stage1_prompt_disambiguates_duplicate_leaf_names() -> None: + registry = LeafRegistry.from_mapping( + { + "categories": [ + { + "category_id": "finance:业务.账户信息.基本信息", + "name": "基本信息", + "path": ["业务", "账户信息", "基本信息"], + }, + { + "category_id": "finance:业务.合约协议.基本信息", + "name": "基本信息", + "path": ["业务", "合约协议", "基本信息"], + }, + {"category_id": "X", "name": "个人联系信息", "path": ["个人联系信息"]}, + {"category_id": "Y", "name": "个人财产信息", "path": ["个人财产信息"]}, + {"category_id": "Z", "name": "个人健康生理信息", "path": ["个人健康生理信息"]}, + ] + } + ) + config = TaskConfig.from_mapping(CONFIG) + stage1 = build_stage1_prompt(METADATA, registry, config) + + assert '"账户信息 / 基本信息"' in stage1.user + assert '"合约协议 / 基本信息"' in stage1.user + assert "finance:业务.账户信息.基本信息" not in stage1.user + assert '"category_id"' not in stage1.user + + +def test_stage2_prompt_uses_local_bundle_ids() -> None: + registry = LeafRegistry.from_mapping(REGISTRY) + config = TaskConfig.from_mapping(CONFIG) + stage2 = build_stage2_prompt( + METADATA, + ["C", "A", "B", "D", "E"], + registry, + config, + ) + user = stage2.user + + # local ids follow candidate order, not canonical ids + assert '"id":"1"' in user and '"name":"C"' in user + assert '"id":"5"' in user and '"name":"E"' in user + assert '"category_id"' not in user + assert '"answer"' in stage2.system + + +def test_answer_builders_use_choice_ids_and_decode_restores_canonical() -> None: + registry = LeafRegistry.from_mapping(REGISTRY) + choices = PromptChoiceRegistry.from_registry(registry) + canonical = ["C", "A", "B", "D", "E"] + + stage1 = stage1_answer(canonical, choices=choices) + assert stage1 == '{"candidates":["3","1","2","4","5"]}' + assert choices.decode_candidates(["3", "1", "2", "4", "5"]) == tuple(canonical) + + stage2 = stage2_answer("C", canonical) + assert stage2 == '{"answer":"1"}' + + def test_invalid_registry_is_rejected() -> None: with pytest.raises(ValueError, match="at least 5"): LeafRegistry.from_mapping({"categories": [{"category_id": "A"}] * 4}) diff --git a/tests/task/test_prompt_choices.py b/tests/task/test_prompt_choices.py new file mode 100644 index 0000000..31de198 --- /dev/null +++ b/tests/task/test_prompt_choices.py @@ -0,0 +1,216 @@ +"""PromptChoiceRegistry: deterministic choice ids and shortest unique suffixes.""" + +import pytest + +from agent.task import ( + LeafCategory, + LeafRegistry, + PromptChoice, + PromptChoiceError, + PromptChoiceRegistry, + decode_stage2_answer, + encode_stage2_answer, +) + + +def _registry(categories: list[dict]) -> LeafRegistry: + return LeafRegistry.from_mapping({"categories": categories}) + + +BASIC = [ + {"category_id": "A", "name": "alpha data"}, + {"category_id": "B", "name": "beta data"}, + {"category_id": "C", "name": "gamma data"}, + {"category_id": "D", "name": "delta data"}, + {"category_id": "E", "name": "epsilon data"}, + {"category_id": "F", "name": "zeta data"}, +] + + +def test_choice_ids_are_sequential_and_cover_the_registry() -> None: + registry = _registry(BASIC) + choices = PromptChoiceRegistry.from_registry(registry) + + assert choices.choice_ids == ("1", "2", "3", "4", "5", "6") + assert [choice.category_id for choice in choices.choices] == list(registry.ids) + assert len({choice.choice_id for choice in choices.choices}) == len(choices.choices) + assert all( + choices.contains_category_id(category_id) for category_id in registry.ids + ) + + +def test_bidirectional_mapping_and_roundtrip() -> None: + registry = _registry(BASIC) + choices = PromptChoiceRegistry.from_registry(registry) + + assert choices.choice_id_of("A") == "1" + assert choices.category_id_of("6") == "F" + for category in registry.categories: + choice_id = choices.choice_id_of(category.category_id) + assert choices.category_id_of(choice_id) == category.category_id + + +def test_mapping_is_deterministic() -> None: + first = PromptChoiceRegistry.from_registry(_registry(BASIC)) + second = PromptChoiceRegistry.from_registry(_registry(BASIC)) + + assert first.choices == second.choices + assert first.choice_ids == second.choice_ids + + +def test_unique_leaf_keeps_the_leaf_name() -> None: + registry = _registry(BASIC) + choices = PromptChoiceRegistry.from_registry(registry) + + assert choices.display_name_of("A") == "alpha data" + + +def test_duplicate_leaf_uses_shortest_unique_parent_suffix() -> None: + registry = _registry( + [ + { + "category_id": "finance:业务.账户信息.基本信息", + "name": "基本信息", + "path": ["业务", "账户信息", "基本信息"], + }, + { + "category_id": "finance:业务.合约协议.基本信息", + "name": "基本信息", + "path": ["业务", "合约协议", "基本信息"], + }, + {"category_id": "X", "name": "个人联系信息", "path": ["个人联系信息"]}, + {"category_id": "Y", "name": "个人财产信息", "path": ["个人财产信息"]}, + {"category_id": "Z", "name": "个人健康生理信息", "path": ["个人健康生理信息"]}, + ] + ) + choices = PromptChoiceRegistry.from_registry(registry) + + assert choices.display_name_of("finance:业务.账户信息.基本信息") == "账户信息 / 基本信息" + assert choices.display_name_of("finance:业务.合约协议.基本信息") == "合约协议 / 基本信息" + assert choices.display_name_of("X") == "个人联系信息" + + +def test_two_parent_levels_needed_until_unique() -> None: + """A/B/X and C/B/X still collide at 'B / X', so both need 3 levels.""" + registry = _registry( + [ + {"category_id": "c1", "name": "X", "path": ["A", "B", "X"]}, + {"category_id": "c2", "name": "X", "path": ["A", "C", "X"]}, + {"category_id": "c3", "name": "X", "path": ["C", "B", "X"]}, + {"category_id": "c4", "name": "Y", "path": ["Y"]}, + {"category_id": "c5", "name": "Z", "path": ["Z"]}, + ] + ) + choices = PromptChoiceRegistry.from_registry(registry) + + assert choices.display_name_of("c1") == "A / B / X" + assert choices.display_name_of("c2") == "C / X" + assert choices.display_name_of("c3") == "C / B / X" + + +def test_all_display_names_are_unique() -> None: + registry = _registry( + [ + {"category_id": "c1", "name": "基本信息", "path": ["业务", "账户信息", "基本信息"]}, + {"category_id": "c2", "name": "基本信息", "path": ["业务", "合约协议", "基本信息"]}, + {"category_id": "c3", "name": "基本信息", "path": ["经营管理", "营销服务", "基本信息"]}, + {"category_id": "c4", "name": "行为信息", "path": ["客户", "个人", "行为信息"]}, + {"category_id": "c5", "name": "行为信息", "path": ["客户", "单位", "行为信息"]}, + ] + ) + choices = PromptChoiceRegistry.from_registry(registry) + + names = [choice.display_name for choice in choices.choices] + assert len(set(names)) == len(names) + assert choices.display_name_of("c3") == "营销服务 / 基本信息" + + +def test_empty_paths_duplicate_leaves_fail_explicitly() -> None: + registry = _registry( + [ + {"category_id": "finance:业务.账户信息.基本信息", "name": "基本信息"}, + {"category_id": "finance:业务.合约协议.基本信息", "name": "基本信息"}, + {"category_id": "X", "name": "Y"}, + {"category_id": "Z", "name": "W"}, + {"category_id": "U", "name": "V"}, + ] + ) + with pytest.raises(PromptChoiceError, match="cannot build a unique display name"): + PromptChoiceRegistry.from_registry(registry) + + +def test_identical_full_paths_fail_explicitly() -> None: + registry = _registry( + [ + {"category_id": "c1", "name": "X", "path": ["A", "X"]}, + {"category_id": "c2", "name": "X", "path": ["A", "X"]}, + {"category_id": "c3", "name": "Y"}, + {"category_id": "c4", "name": "Z"}, + {"category_id": "c5", "name": "W"}, + ] + ) + with pytest.raises(PromptChoiceError, match="cannot build a unique display name"): + PromptChoiceRegistry.from_registry(registry) + + +def test_from_registry_never_mutates_leaf_categories() -> None: + registry = _registry(BASIC) + snapshot = registry.categories + PromptChoiceRegistry.from_registry(registry) + + assert registry.categories == snapshot + assert registry.categories[0].name == "alpha data" + assert registry.categories[0].category_id == "A" + + +def test_mapping_failures_raise_prompt_choice_error() -> None: + choices = PromptChoiceRegistry.from_registry(_registry(BASIC)) + + with pytest.raises(PromptChoiceError, match="no prompt choice"): + choices.choice_id_of("Z") + with pytest.raises(PromptChoiceError, match="not in the prompt catalog"): + choices.category_id_of("42") + with pytest.raises(PromptChoiceError, match="no prompt choice"): + choices.display_name_of("42") + + +def test_stage1_encode_decode_roundtrip() -> None: + choices = PromptChoiceRegistry.from_registry(_registry(BASIC)) + canonical = ("C", "A", "B", "D", "E") + + encoded = choices.encode_candidates(canonical) + assert encoded == ("3", "1", "2", "4", "5") + assert choices.decode_candidates(encoded) == canonical + + +def test_stage1_decode_rejects_wrong_shape() -> None: + choices = PromptChoiceRegistry.from_registry(_registry(BASIC)) + + with pytest.raises(PromptChoiceError, match="exactly 5"): + choices.decode_candidates(("1", "2", "3")) + with pytest.raises(PromptChoiceError, match="unique"): + choices.decode_candidates(("1", "1", "2", "4", "5")) + with pytest.raises(PromptChoiceError, match="not in the prompt catalog"): + choices.decode_candidates(("1", "2", "3", "4", "42")) + + +def test_stage2_local_ids_encode_and_decode_positionally() -> None: + candidates = ("C", "A", "B", "D", "E") + + assert encode_stage2_answer("C", candidates) == "1" + assert encode_stage2_answer("E", candidates) == "5" + assert decode_stage2_answer("1", candidates) == "C" + assert decode_stage2_answer("5", candidates) == "E" + + +def test_stage2_decode_rejects_non_local_ids() -> None: + candidates = ("C", "A", "B", "D", "E") + + with pytest.raises(PromptChoiceError, match="one of 1..5"): + decode_stage2_answer("6", candidates) + with pytest.raises(PromptChoiceError, match="one of 1..5"): + decode_stage2_answer("A", candidates) + with pytest.raises(PromptChoiceError, match="one of 1..5"): + decode_stage2_answer("01", candidates) + with pytest.raises(PromptChoiceError, match="one of the candidates"): + encode_stage2_answer("F", candidates) diff --git a/tests/training/test_common.py b/tests/training/test_common.py new file mode 100644 index 0000000..86f4c68 --- /dev/null +++ b/tests/training/test_common.py @@ -0,0 +1,77 @@ +"""Source-seeded deterministic candidate bundle (shared SFT/RL fixture policy). + +Phase 6: the baseline fixture must not fix the ground truth at position 1 +(which would leak a systematic {"answer":"1"} bias into every Stage 2 +sample). These tests pin the deterministic, source_id-seeded permutation in +agent.training.common.build_candidates. +""" + +from __future__ import annotations + +import pytest + +from agent.task import LeafRegistry +from agent.training.common import build_candidates + + +def _registry() -> LeafRegistry: + return LeafRegistry.from_mapping(["A", "B", "C", "D", "E", "F"]) + + +def test_same_source_id_yields_the_same_ordering() -> None: + registry = _registry() + assert build_candidates("C", registry, source_id="row-train") == build_candidates( + "C", registry, source_id="row-train" + ) + + +def test_different_source_ids_can_yield_different_orderings() -> None: + registry = _registry() + bundles = { + tuple(build_candidates("C", registry, source_id=sample_id)) + for sample_id in ("r1", "r2", "r3", "r4", "r5") + } + assert len(bundles) > 1 + + +def test_ground_truth_always_present_and_exactly_five_unique() -> None: + registry = _registry() + for sample_id in ("row-train", "a", "b", "c", "d"): + bundle = build_candidates("C", registry, source_id=sample_id) + assert len(bundle) == 5 + assert len(set(bundle)) == 5 + assert "C" in bundle + assert all(candidate in registry.ids for candidate in bundle) + + +def test_bundle_is_a_permutation_of_the_gt_first_slots() -> None: + registry = _registry() + base = ["C"] + [category_id for category_id in registry.ids if category_id != "C"][:4] + for sample_id in ("row-train", "row-val", "row-test", "x", "y"): + bundle = build_candidates("C", registry, source_id=sample_id) + assert sorted(bundle) == sorted(base) + + +def test_ground_truth_position_is_not_fixed_at_one() -> None: + registry = _registry() + positions = { + build_candidates("C", registry, source_id=sample_id).index("C") + for sample_id in ("row-train", "row-val", "row-test", "r1", "r2") + } + assert any(position != 0 for position in positions) + assert 0 in positions # position 1 is legal, just not universal + + +def test_all_five_positions_appear_across_large_synthetic_id_set() -> None: + registry = _registry() + positions = { + build_candidates("C", registry, source_id=f"sample-{index:04d}").index("C") + for index in range(400) + } + assert positions == {0, 1, 2, 3, 4} + + +def test_source_id_is_a_required_keyword() -> None: + registry = _registry() + with pytest.raises(TypeError): + build_candidates("C", registry) # missing required source_id