"""Training schema; labels/provenance never participate in model rendering.""" import math from typing import Any, Literal from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, model_validator from decision_schema import Question, options QUESTION = TypeAdapter(Question) KINDS = ("choice", "score", "noul") MODELS = ( "github-copilot/gpt-6-luna", "github-copilot/grok-4.7", "github-copilot/gpt-6-sol", ) def validate_question(value): q = QUESTION.validate_python(value).model_dump(exclude_none=True) if not q.get("instructions"): raise ValueError("instructions required") if q["type"] in {"choice", "score"} and len(q["criteria"]) < 2: raise ValueError("training requires at least two candidates") return q def label_key(label, kind): if kind == "noul": if isinstance(label, bool): return "true" if label else "false" if label in ("true", "false"): return label raise ValueError("noul label must be a boolean or true/false string") if isinstance(label, bool): raise ValueError("boolean is not a choice/score label") return str(label) class Target(BaseModel): model_config = ConfigDict(extra="forbid") hard_label: str | int | bool evidence: list[dict[str, Any]] = Field(min_length=1) ambiguity: Literal["none", "ambiguous", "insufficient", "conflict"] = "none" class Case(BaseModel): model_config = ConfigDict(extra="forbid") case_id: str group_id: str domain: str task_family: str language: Literal["zh", "en", "mixed"] state: str | dict[str, Any] | list[Any] questions: dict[str, dict[str, Any]] = Field(min_length=1, max_length=8) targets: dict[str, Target] @model_validator(mode="after") def check_targets(self): if set(self.questions) != set(self.targets): raise ValueError("question/target IDs differ") for qid, question in self.questions.items(): q = validate_question(question) self.questions[qid] = q key = label_key(self.targets[qid].hard_label, q["type"]) if key not in [k for k, _ in options(q)]: raise ValueError("target is not a candidate") return self def check_distribution(values): if not values or any(not math.isfinite(x) or x < 0 for x in values): raise ValueError("invalid distribution") if abs(sum(values) - 1) > 1e-5: raise ValueError("probabilities must sum to one")