File size: 5,200 Bytes
02d27c4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
"""The one contract every KODEX code obeys.

    predict(x) -> Prediction(y, uncertainty, in_domain)

A KODEX surrogate is a fast, calibrated stand-in for an expensive fusion
computation.  Three promises come with every prediction:

    y            the point estimate
    uncertainty  a calibrated 1-sigma (or class-probability spread)
    in_domain    whether the input is inside the region the surrogate is
                 trusted on; when False the caller must defer to full physics.

Provenance is carried in the type: `.as_tagged()` returns a
`kronos_toolkit.verify.tags.Tagged`, and every surrogate that stands in for a
real code is a `[T]` naming the code that retires it (`retired_by`).  Nothing
here is ever called "live"; a code is BUILT, PARTIAL, or ROADMAP.
"""
from __future__ import annotations

from dataclasses import dataclass, field
from typing import Any

try:                                    # kronos-ml depends on kronos-toolkit
    from kronos_toolkit.verify.tags import Tagged, surrogate as _surrogate_tag
except Exception:                       # pragma: no cover - toolkit always present in the env
    Tagged = None
    _surrogate_tag = None

STATUSES = ("BUILT", "PARTIAL", "ROADMAP")
PROVENANCE = ("SIM", "TWIN", "CGYRO-urep", "ANALYTIC", "REAL-ANCHOR", "n/a")


@dataclass
class Prediction:
    """A single surrogate prediction with its calibrated uncertainty and gate."""
    y: Any
    uncertainty: Any = None
    in_domain: bool = True
    note: str = ""

    def as_tuple(self):
        return (self.y, self.uncertainty, self.in_domain)

    def as_tagged(self, retired_by: str, note: str = ""):
        """Wrap as a [T] Tagged value naming the real code that retires it."""
        if Tagged is None:
            raise RuntimeError("kronos_toolkit not importable; cannot tag")
        payload = {"y": self.y, "uncertainty": self.uncertainty,
                   "in_domain": self.in_domain}
        return Tagged(payload, "T", note=note or self.note, retired_by=retired_by)


class Surrogate:
    """Base class for every KODEX code.  Subclasses set the class attributes and
    implement `_predict`; roadmap codes leave `_predict` raising and set
    status='ROADMAP'."""

    #: brand name, e.g. "KYRO"
    name: str = "SURROGATE"
    #: technical descriptor, e.g. "TRANSPORT"
    function: str = ""
    #: the real code(s) this stands in for
    real_codes: tuple = ()
    #: the real code that retires this [T] surrogate (required for BUILT/PARTIAL)
    retired_by: str = ""
    #: SIM / TWIN / CGYRO-urep / ANALYTIC / REAL-ANCHOR
    provenance: str = "SIM"
    #: BUILT / PARTIAL / ROADMAP
    status: str = "ROADMAP"
    #: rollout phase 1/2/3
    phase: int = 3
    #: de-risking register gate(s) this ties to
    gates: tuple = ()
    #: one-line description
    note: str = ""
    #: frozen input->expected-output pin for the drift regression (or None)
    ANCHOR = None

    # -- prediction -------------------------------------------------------
    def _predict(self, x) -> Prediction:
        raise NotImplementedError(
            f"{self.name} is ROADMAP — not built yet")

    def predict(self, x) -> Prediction:
        p = self._predict(x)
        if not isinstance(p, Prediction):
            # tolerate a (y, unc, in_domain) tuple
            p = Prediction(*p) if isinstance(p, tuple) else Prediction(p)
        return p

    def evaluator(self, **params) -> dict:
        """Adapter for kronos_toolkit.uq / report.Study: returns a flat dict of
        scalars {y, uncertainty, in_domain} so the existing UQ/Study machinery
        consumes a KODEX surrogate unchanged."""
        p = self.predict(params if not params else self._x_from_params(**params))
        out = {"y": _scalar(p.y), "in_domain": float(bool(p.in_domain))}
        if p.uncertainty is not None:
            out["uncertainty"] = _scalar(p.uncertainty)
        return out

    def _x_from_params(self, **params):
        """Default: pass the params dict straight through. Members override when
        they need an array/graph input."""
        return params

    # -- availability / metadata -----------------------------------------
    def available(self) -> bool:
        """True when this code can actually run here (deps + data/checkpoint on
        disk).  ROADMAP codes are never available."""
        return self.status in ("BUILT", "PARTIAL")

    def benchmark(self) -> dict:
        """Return the honest on-disk benchmark: speed / accuracy / uq calibration.
        Members override; roadmap codes return {}."""
        return {}

    def card(self) -> dict:
        return {
            "name": self.name, "function": self.function,
            "status": self.status, "phase": self.phase,
            "provenance": self.provenance, "retired_by": self.retired_by,
            "gates": list(self.gates), "note": self.note,
            "available": self.available(),
        }


def _scalar(v):
    try:
        import numpy as _np
        if isinstance(v, _np.ndarray):
            return float(v.reshape(-1)[0]) if v.size == 1 else float(_np.mean(v))
    except Exception:
        pass
    try:
        return float(v)
    except Exception:
        return v