File size: 2,512 Bytes
f5864f8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
"""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")