@pi-in-go/pigpen-pi-typesafe-api 0.1.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/CREDITS.md +14 -0
- package/LICENSE +22 -0
- package/README.md +30 -0
- package/ask.go +62 -0
- package/ask_test.go +76 -0
- package/auth.go +249 -0
- package/auth_test.go +131 -0
- package/backends.go +336 -0
- package/backends_test.go +404 -0
- package/batch.go +202 -0
- package/batch_test.go +202 -0
- package/battery_test.go +41 -0
- package/calibrate.go +354 -0
- package/calibrate_test.go +186 -0
- package/client.go +615 -0
- package/client_test.go +490 -0
- package/credentials.go +252 -0
- package/credentials_test.go +216 -0
- package/doc.go +14 -0
- package/errors.go +143 -0
- package/evaluation.go +86 -0
- package/evaluation_schema.json +264 -0
- package/gaps_test.go +77 -0
- package/go.mod +9 -0
- package/go.sum +2 -0
- package/helpers_test.go +169 -0
- package/hostmodel/hostmodel.go +87 -0
- package/json.go +299 -0
- package/json_test.go +92 -0
- package/ownmodel_test.go +79 -0
- package/package.json +40 -0
- package/port/PORT.md +6 -0
- package/provenance.json +18 -0
- package/review_test.go +23 -0
- package/schema.go +473 -0
- package/schema_test.go +262 -0
- package/testdata/tools/typebox-messages.mts +5 -0
- package/testdata/typebox-messages.json +285 -0
- package/twin_test.go +28 -0
- package/ui/fakehost_test.go +548 -0
- package/ui/keyprompt.go +115 -0
- package/ui/login.go +106 -0
- package/ui/twin_test.go +28 -0
- package/ui/ui_test.go +285 -0
- package/usage.go +366 -0
- package/usage_test.go +139 -0
package/batch_test.go
ADDED
|
@@ -0,0 +1,202 @@
|
|
|
1
|
+
package pitypesafe
|
|
2
|
+
|
|
3
|
+
import (
|
|
4
|
+
"context"
|
|
5
|
+
"encoding/json"
|
|
6
|
+
"errors"
|
|
7
|
+
"io"
|
|
8
|
+
"net/http"
|
|
9
|
+
"path/filepath"
|
|
10
|
+
"reflect"
|
|
11
|
+
"sync"
|
|
12
|
+
"sync/atomic"
|
|
13
|
+
"testing"
|
|
14
|
+
"time"
|
|
15
|
+
|
|
16
|
+
"github.com/MichaelKinsy/pigpen/components/typesafe/libraries/typesafe"
|
|
17
|
+
)
|
|
18
|
+
|
|
19
|
+
// noulAnswering answers every question with a fixed Noul probability, so merging is observable per question id.
|
|
20
|
+
func noulAnswering(counter *atomic.Int32) doerFunc {
|
|
21
|
+
return func(r *http.Request) (*http.Response, error) {
|
|
22
|
+
counter.Add(1)
|
|
23
|
+
body, _ := io.ReadAll(r.Body)
|
|
24
|
+
var req struct {
|
|
25
|
+
Questions json.RawMessage `json:"questions"`
|
|
26
|
+
}
|
|
27
|
+
_ = json.Unmarshal(body, &req)
|
|
28
|
+
questions, _ := typesafe.ParseQuestions(req.Questions)
|
|
29
|
+
answers := map[string]any{}
|
|
30
|
+
for _, q := range questions {
|
|
31
|
+
answers[q.Name] = map[string]any{"type": "noul", "noul": 0.75}
|
|
32
|
+
}
|
|
33
|
+
return jsonResponse(200, map[string]any{"model": "jev-test", "answers": answers, "usage": map[string]any{"input_tokens": 10, "output_tokens": 0}}, nil), nil
|
|
34
|
+
}
|
|
35
|
+
}
|
|
36
|
+
|
|
37
|
+
func batchClient(t *testing.T, name string, opts Options, counter *atomic.Int32) *TypeSafe {
|
|
38
|
+
t.Helper()
|
|
39
|
+
isolate(t)
|
|
40
|
+
t.Setenv("TYPESAFE_API_KEY", "batch-test-key")
|
|
41
|
+
opts.Ledger = OpenUsageLedger(LedgerOptions{Path: filepath.Join(t.TempDir(), name+".json")})
|
|
42
|
+
opts.HTTPClient = noulAnswering(counter)
|
|
43
|
+
c, err := New(opts)
|
|
44
|
+
if err != nil {
|
|
45
|
+
t.Fatal(err)
|
|
46
|
+
}
|
|
47
|
+
return c
|
|
48
|
+
}
|
|
49
|
+
|
|
50
|
+
func TestBatch(t *testing.T) {
|
|
51
|
+
tw(t, "batch", "fan-out bounds concurrency, preserves order, and never throws", func(t *testing.T) {
|
|
52
|
+
var inFlight, peak atomic.Int32
|
|
53
|
+
results := FanOut(context.Background(), []int{1, 2, 3, 4, 5, 6, 7}, func(_ context.Context, v int, _ int) (int, error) {
|
|
54
|
+
n := inFlight.Add(1)
|
|
55
|
+
for {
|
|
56
|
+
p := peak.Load()
|
|
57
|
+
if n <= p || peak.CompareAndSwap(p, n) {
|
|
58
|
+
break
|
|
59
|
+
}
|
|
60
|
+
}
|
|
61
|
+
time.Sleep(5 * time.Millisecond)
|
|
62
|
+
inFlight.Add(-1)
|
|
63
|
+
if v == 4 {
|
|
64
|
+
return 0, errors.New("bad 4")
|
|
65
|
+
}
|
|
66
|
+
return v * 2, nil
|
|
67
|
+
}, FanOutOptions{Concurrency: 3})
|
|
68
|
+
if peak.Load() != 3 {
|
|
69
|
+
t.Errorf("peak = %d", peak.Load())
|
|
70
|
+
}
|
|
71
|
+
var got []any
|
|
72
|
+
for i, r := range results {
|
|
73
|
+
if r.Index != i {
|
|
74
|
+
t.Errorf("index %d = %d", i, r.Index)
|
|
75
|
+
}
|
|
76
|
+
if r.OK {
|
|
77
|
+
got = append(got, r.Value)
|
|
78
|
+
} else {
|
|
79
|
+
got = append(got, "failed")
|
|
80
|
+
}
|
|
81
|
+
}
|
|
82
|
+
if !reflect.DeepEqual(got, []any{2, 4, 6, "failed", 10, 12, 14}) || results[3].Skipped {
|
|
83
|
+
t.Fatalf("results = %v", got)
|
|
84
|
+
}
|
|
85
|
+
})
|
|
86
|
+
tw(t, "batch", "fan-out stops launching after a stop rule and marks the rest skipped", func(t *testing.T) {
|
|
87
|
+
var started []int
|
|
88
|
+
results := FanOut(context.Background(), []int{0, 1, 2, 3, 4, 5}, func(_ context.Context, v int, _ int) (int, error) {
|
|
89
|
+
started = append(started, v)
|
|
90
|
+
if v == 2 {
|
|
91
|
+
return 0, newError(CodeBudget, "cap reached")
|
|
92
|
+
}
|
|
93
|
+
return v, nil
|
|
94
|
+
}, FanOutOptions{Concurrency: 1, StopOn: func(err error) bool { return hasCode(err, CodeBudget) }})
|
|
95
|
+
skipped := 0
|
|
96
|
+
for _, r := range results {
|
|
97
|
+
if !r.OK && r.Skipped {
|
|
98
|
+
skipped++
|
|
99
|
+
}
|
|
100
|
+
}
|
|
101
|
+
if !reflect.DeepEqual(started, []int{0, 1, 2}) || skipped != 3 || results[2].Skipped || results[2].OK {
|
|
102
|
+
t.Fatalf("started=%v skipped=%d", started, skipped)
|
|
103
|
+
}
|
|
104
|
+
})
|
|
105
|
+
tw(t, "batch", "fan-out stops launching once the caller's signal aborts", func(t *testing.T) {
|
|
106
|
+
ctx, cancel := context.WithCancel(context.Background())
|
|
107
|
+
results := FanOut(ctx, []int{0, 1, 2, 3}, func(_ context.Context, v int, _ int) (int, error) {
|
|
108
|
+
if v == 1 {
|
|
109
|
+
cancel()
|
|
110
|
+
}
|
|
111
|
+
return v, nil
|
|
112
|
+
}, FanOutOptions{Concurrency: 1})
|
|
113
|
+
if !results[0].OK || !results[1].OK || results[2].OK || !results[2].Skipped || !hasCode(results[2].Err, CodeAborted) {
|
|
114
|
+
t.Fatalf("results = %+v", results)
|
|
115
|
+
}
|
|
116
|
+
})
|
|
117
|
+
tw(t, "batch", "the splitter is pure: a request that already fits is one chunk, unchanged", func(t *testing.T) {
|
|
118
|
+
req := sampleRequest()
|
|
119
|
+
if got := ChunkRequest(req, 0); len(got) != 1 || !reflect.DeepEqual(got[0], req) {
|
|
120
|
+
t.Fatalf("chunks = %+v", got)
|
|
121
|
+
}
|
|
122
|
+
// No validation here: an invalid request is passed through and fails once, per chunk, at admission.
|
|
123
|
+
invalid := typesafe.SystemOneRequest{State: typesafe.Text("synthetic")}
|
|
124
|
+
if got := ChunkRequest(invalid, 0); len(got) != 1 || !reflect.DeepEqual(got[0], invalid) {
|
|
125
|
+
t.Fatalf("invalid = %+v", got)
|
|
126
|
+
}
|
|
127
|
+
})
|
|
128
|
+
tw(t, "batch", "an invalid request is one settled failure, not a thrown error", func(t *testing.T) {
|
|
129
|
+
var calls atomic.Int32
|
|
130
|
+
c := batchClient(t, "invalid", Options{}, &calls)
|
|
131
|
+
invalid := typesafe.SystemOneRequest{State: typesafe.Text("synthetic")}
|
|
132
|
+
batch := c.EvaluateMany(context.Background(), []typesafe.SystemOneRequest{invalid}, BatchOptions{})
|
|
133
|
+
if calls.Load() != 0 || batch.OK || len(batch.Results) != 1 || batch.Results[0].OK || !hasCode(batch.Results[0].Err, CodeValidation) || batch.Results[0].Skipped {
|
|
134
|
+
t.Fatalf("batch = %+v", batch)
|
|
135
|
+
}
|
|
136
|
+
if _, err := c.Evaluate(context.Background(), invalid); !hasCode(err, CodeValidation) {
|
|
137
|
+
t.Fatalf("err = %v", err)
|
|
138
|
+
}
|
|
139
|
+
})
|
|
140
|
+
tw(t, "batch", "more questions than one request may carry are split in order", func(t *testing.T) {
|
|
141
|
+
questions := manyQuestions(DefaultMaxQuestions + 3)
|
|
142
|
+
chunks := ChunkRequest(typesafe.SystemOneRequest{State: typesafe.Text("many"), Questions: questions}, 0)
|
|
143
|
+
if len(chunks) != 2 || len(chunks[0].Questions) != DefaultMaxQuestions || len(chunks[1].Questions) != 3 {
|
|
144
|
+
t.Fatalf("chunks = %d", len(chunks))
|
|
145
|
+
}
|
|
146
|
+
var ids, want []string
|
|
147
|
+
for _, c := range chunks {
|
|
148
|
+
for _, q := range c.Questions {
|
|
149
|
+
ids = append(ids, q.Name)
|
|
150
|
+
}
|
|
151
|
+
if c.State.Data() != "many" {
|
|
152
|
+
t.Error("the state is repeated in each chunk")
|
|
153
|
+
}
|
|
154
|
+
}
|
|
155
|
+
for _, q := range questions {
|
|
156
|
+
want = append(want, q.Name)
|
|
157
|
+
}
|
|
158
|
+
if !reflect.DeepEqual(ids, want) {
|
|
159
|
+
t.Fatalf("order = %v", ids)
|
|
160
|
+
}
|
|
161
|
+
})
|
|
162
|
+
tw(t, "batch", "evaluateAll spends one request per chunk and merges answers, usage, and model", func(t *testing.T) {
|
|
163
|
+
var calls atomic.Int32
|
|
164
|
+
c := batchClient(t, "many", Options{}, &calls)
|
|
165
|
+
batch := c.EvaluateAll(context.Background(), typesafe.SystemOneRequest{State: typesafe.Text("many"), Questions: manyQuestions(DefaultMaxQuestions + 2)}, BatchOptions{Concurrency: 2})
|
|
166
|
+
if calls.Load() != 2 || !batch.OK || batch.Failures != 0 || batch.Skipped != 0 || len(batch.Answers) != DefaultMaxQuestions+2 || batch.Model != "jev-test" || batch.Usage.InputTokens != 20 || c.GetUsage().RequestsSucceeded != 2 {
|
|
167
|
+
t.Fatalf("batch = %+v calls=%d", batch, calls.Load())
|
|
168
|
+
}
|
|
169
|
+
if len(batch.AnswerOrder) != DefaultMaxQuestions+2 || batch.AnswerOrder[0] != "q0" || batch.AnswerOrder[DefaultMaxQuestions+1] != "q"+itoa(DefaultMaxQuestions+1) {
|
|
170
|
+
t.Errorf("order = %v", batch.AnswerOrder)
|
|
171
|
+
}
|
|
172
|
+
})
|
|
173
|
+
tw(t, "batch", "evaluateMany reports per-request failures and stops submitting after a budget error", func(t *testing.T) {
|
|
174
|
+
var calls atomic.Int32
|
|
175
|
+
c := batchClient(t, "budget", Options{MaxRequests: 1}, &calls)
|
|
176
|
+
req := sampleRequest()
|
|
177
|
+
batch := c.EvaluateMany(context.Background(), []typesafe.SystemOneRequest{req, req, req}, BatchOptions{Concurrency: 1})
|
|
178
|
+
// One request succeeded, one failed on the session cap, and the third was never submitted.
|
|
179
|
+
if calls.Load() != 1 || batch.OK || batch.Failures != 2 || batch.Skipped != 1 || !batch.Results[0].OK || len(batch.Answers) != 1 || batch.ElapsedMs < 0 {
|
|
180
|
+
t.Fatalf("batch = %+v calls=%d", batch, calls.Load())
|
|
181
|
+
}
|
|
182
|
+
})
|
|
183
|
+
tw(t, "batch", "a daily cap stops the batch with the cap named", func(t *testing.T) {
|
|
184
|
+
var calls atomic.Int32
|
|
185
|
+
c := batchClient(t, "day-cap", Options{MaxRequestsPerDay: 1}, &calls)
|
|
186
|
+
req := sampleRequest()
|
|
187
|
+
batch := c.EvaluateMany(context.Background(), []typesafe.SystemOneRequest{req, req}, BatchOptions{})
|
|
188
|
+
// JavaScript runs the first request to its cap check before the second starts; two goroutines race for the
|
|
189
|
+
// one allowed request, so which of the two is refused is not fixed. Exactly one is, and it names the cap.
|
|
190
|
+
var refused *Settled[*Evaluation]
|
|
191
|
+
for i := range batch.Results {
|
|
192
|
+
if !batch.Results[i].OK {
|
|
193
|
+
refused = &batch.Results[i]
|
|
194
|
+
}
|
|
195
|
+
}
|
|
196
|
+
if calls.Load() != 1 || refused == nil || batch.Failures != 1 || !hasCode(refused.Err, CodeBudget) || !contains(refused.Err.Error(), "daily request cap") {
|
|
197
|
+
t.Fatalf("batch = %+v calls=%d", batch, calls.Load())
|
|
198
|
+
}
|
|
199
|
+
})
|
|
200
|
+
}
|
|
201
|
+
|
|
202
|
+
var _ sync.Mutex
|
package/battery_test.go
ADDED
|
@@ -0,0 +1,41 @@
|
|
|
1
|
+
package pitypesafe
|
|
2
|
+
|
|
3
|
+
import (
|
|
4
|
+
"encoding/json"
|
|
5
|
+
"os"
|
|
6
|
+
"regexp"
|
|
7
|
+
"testing"
|
|
8
|
+
)
|
|
9
|
+
|
|
10
|
+
// TestTypeBoxMessagesForCommonFailures compares admission errors with the messages the original's TypeBox
|
|
11
|
+
// validator produced for the same requests (port/oracle/battery.json through parseEvaluationRequest, see
|
|
12
|
+
// testdata/typebox-messages.json). Union-branch wording for deep question failures is not reproduced; those rows
|
|
13
|
+
// are marked "-" in the file and skipped.
|
|
14
|
+
func TestTypeBoxMessagesForCommonFailures(t *testing.T) {
|
|
15
|
+
raw, err := os.ReadFile("testdata/typebox-messages.json")
|
|
16
|
+
if err != nil {
|
|
17
|
+
t.Fatal(err)
|
|
18
|
+
}
|
|
19
|
+
var rows []struct {
|
|
20
|
+
Request json.RawMessage `json:"request"`
|
|
21
|
+
Message string `json:"message"`
|
|
22
|
+
Compared bool `json:"compared"`
|
|
23
|
+
}
|
|
24
|
+
if err := json.Unmarshal(raw, &rows); err != nil {
|
|
25
|
+
t.Fatal(err)
|
|
26
|
+
}
|
|
27
|
+
strip := regexp.MustCompile(` Expected \{.*$`)
|
|
28
|
+
for _, r := range rows {
|
|
29
|
+
if !r.Compared {
|
|
30
|
+
continue
|
|
31
|
+
}
|
|
32
|
+
_, err := ParseEvaluationRequest(mustTree(t, string(r.Request)))
|
|
33
|
+
got := "OK"
|
|
34
|
+
if err != nil {
|
|
35
|
+
got = strip.ReplaceAllString(err.Error(), "")
|
|
36
|
+
}
|
|
37
|
+
if got != r.Message {
|
|
38
|
+
t.Errorf("%s:\n got %s\n want %s", r.Request, got, r.Message)
|
|
39
|
+
}
|
|
40
|
+
}
|
|
41
|
+
}
|
package/calibrate.go
ADDED
|
@@ -0,0 +1,354 @@
|
|
|
1
|
+
package pitypesafe
|
|
2
|
+
|
|
3
|
+
import (
|
|
4
|
+
"context"
|
|
5
|
+
"fmt"
|
|
6
|
+
"math"
|
|
7
|
+
"sort"
|
|
8
|
+
"strings"
|
|
9
|
+
)
|
|
10
|
+
|
|
11
|
+
// A judge-tuning kit: label a set of cases, score them with Jev, and read off AUC and threshold behaviour. It
|
|
12
|
+
// carries no domain knowledge: a case is anything a scorer can turn into a number, so the same toolkit fits an
|
|
13
|
+
// action guard, a triage rule, or a prose check.
|
|
14
|
+
|
|
15
|
+
// ScoredSample is one labelled observation: the truth about the case, and the number the judge assigned to it.
|
|
16
|
+
type ScoredSample struct {
|
|
17
|
+
Label bool
|
|
18
|
+
Score float64
|
|
19
|
+
// ID is an optional name; it is used in the missed and flagged listings.
|
|
20
|
+
ID string
|
|
21
|
+
}
|
|
22
|
+
|
|
23
|
+
// ThresholdRow holds outcome counts for one threshold. Precision and Recall are nil when their denominator is empty.
|
|
24
|
+
type ThresholdRow struct {
|
|
25
|
+
Threshold float64
|
|
26
|
+
Flagged int
|
|
27
|
+
TP, FP int
|
|
28
|
+
FN, TN int
|
|
29
|
+
Precision *float64
|
|
30
|
+
Recall *float64
|
|
31
|
+
// FlagRate is the share of all cases the threshold selects.
|
|
32
|
+
FlagRate float64
|
|
33
|
+
}
|
|
34
|
+
|
|
35
|
+
// Calibration is the result of Calibrate.
|
|
36
|
+
type Calibration struct {
|
|
37
|
+
Name string
|
|
38
|
+
Scored int
|
|
39
|
+
Positives int
|
|
40
|
+
Negatives int
|
|
41
|
+
Errors int
|
|
42
|
+
// AUC is the rank-based AUC (Mann–Whitney, ties count half); nil when one class is empty.
|
|
43
|
+
AUC *float64
|
|
44
|
+
Rows []ThresholdRow
|
|
45
|
+
// Recommended is the lowest threshold meeting the requested floors, when one exists.
|
|
46
|
+
Recommended *ThresholdRow
|
|
47
|
+
// Missed are positive cases the recommended threshold misses.
|
|
48
|
+
Missed []ScoredSample
|
|
49
|
+
// Flagged are negative cases the recommended threshold flags.
|
|
50
|
+
Flagged []ScoredSample
|
|
51
|
+
}
|
|
52
|
+
|
|
53
|
+
// CalibrateOptions configure Calibrate.
|
|
54
|
+
type CalibrateOptions struct {
|
|
55
|
+
// Thresholds to evaluate. Default: every distinct score, ascending (at most 64 rows).
|
|
56
|
+
Thresholds []float64
|
|
57
|
+
// MinPrecision and MinRecall are floors for the recommendation; nil means no floor.
|
|
58
|
+
MinPrecision *float64
|
|
59
|
+
MinRecall *float64
|
|
60
|
+
// Errors counts cases that could not be scored; reported and excluded from the metrics.
|
|
61
|
+
Errors int
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
// AUC is the rank-based probability that a positive outranks a negative; nil when a class is empty.
|
|
65
|
+
func AUC(samples []ScoredSample) *float64 {
|
|
66
|
+
var pos, neg []float64
|
|
67
|
+
for _, s := range samples {
|
|
68
|
+
if s.Label {
|
|
69
|
+
pos = append(pos, s.Score)
|
|
70
|
+
} else {
|
|
71
|
+
neg = append(neg, s.Score)
|
|
72
|
+
}
|
|
73
|
+
}
|
|
74
|
+
if len(pos) == 0 || len(neg) == 0 {
|
|
75
|
+
return nil
|
|
76
|
+
}
|
|
77
|
+
wins := 0.0
|
|
78
|
+
for _, p := range pos {
|
|
79
|
+
for _, n := range neg {
|
|
80
|
+
switch {
|
|
81
|
+
case p > n:
|
|
82
|
+
wins++
|
|
83
|
+
case p == n:
|
|
84
|
+
wins += 0.5
|
|
85
|
+
}
|
|
86
|
+
}
|
|
87
|
+
}
|
|
88
|
+
v := wins / float64(len(pos)*len(neg))
|
|
89
|
+
return &v
|
|
90
|
+
}
|
|
91
|
+
|
|
92
|
+
func ratio(num, den int) *float64 {
|
|
93
|
+
if den == 0 {
|
|
94
|
+
return nil
|
|
95
|
+
}
|
|
96
|
+
v := float64(num) / float64(den)
|
|
97
|
+
return &v
|
|
98
|
+
}
|
|
99
|
+
|
|
100
|
+
// MetricsAt counts the outcomes at one threshold: a case is flagged when its score is at least the threshold.
|
|
101
|
+
func MetricsAt(samples []ScoredSample, threshold float64) ThresholdRow {
|
|
102
|
+
row := ThresholdRow{Threshold: threshold}
|
|
103
|
+
for _, s := range samples {
|
|
104
|
+
flagged := s.Score >= threshold
|
|
105
|
+
switch {
|
|
106
|
+
case flagged && s.Label:
|
|
107
|
+
row.TP++
|
|
108
|
+
case flagged:
|
|
109
|
+
row.FP++
|
|
110
|
+
case s.Label:
|
|
111
|
+
row.FN++
|
|
112
|
+
default:
|
|
113
|
+
row.TN++
|
|
114
|
+
}
|
|
115
|
+
}
|
|
116
|
+
row.Flagged = row.TP + row.FP
|
|
117
|
+
row.Precision = ratio(row.TP, row.TP+row.FP)
|
|
118
|
+
row.Recall = ratio(row.TP, row.TP+row.FN)
|
|
119
|
+
if len(samples) > 0 {
|
|
120
|
+
row.FlagRate = float64(row.Flagged) / float64(len(samples))
|
|
121
|
+
}
|
|
122
|
+
return row
|
|
123
|
+
}
|
|
124
|
+
|
|
125
|
+
// Sweep evaluates every threshold.
|
|
126
|
+
func Sweep(samples []ScoredSample, thresholds []float64) []ThresholdRow {
|
|
127
|
+
rows := make([]ThresholdRow, len(thresholds))
|
|
128
|
+
for i, t := range thresholds {
|
|
129
|
+
rows[i] = MetricsAt(samples, t)
|
|
130
|
+
}
|
|
131
|
+
return rows
|
|
132
|
+
}
|
|
133
|
+
|
|
134
|
+
// DefaultThresholds are the distinct scores, ascending, as a threshold grid: every point where the counts can
|
|
135
|
+
// change. limit (64 when not positive) caps the rows by sampling evenly.
|
|
136
|
+
func DefaultThresholds(samples []ScoredSample, limit int) []float64 {
|
|
137
|
+
if limit <= 0 {
|
|
138
|
+
limit = 64
|
|
139
|
+
}
|
|
140
|
+
seen := map[float64]bool{}
|
|
141
|
+
var distinct []float64
|
|
142
|
+
for _, s := range samples {
|
|
143
|
+
if !seen[s.Score] {
|
|
144
|
+
seen[s.Score] = true
|
|
145
|
+
distinct = append(distinct, s.Score)
|
|
146
|
+
}
|
|
147
|
+
}
|
|
148
|
+
sort.Float64s(distinct)
|
|
149
|
+
if len(distinct) <= limit {
|
|
150
|
+
return distinct
|
|
151
|
+
}
|
|
152
|
+
step := float64(len(distinct)-1) / float64(limit-1)
|
|
153
|
+
out := make([]float64, limit)
|
|
154
|
+
for i := range out {
|
|
155
|
+
out[i] = distinct[int(math.Floor(float64(i)*step+0.5))]
|
|
156
|
+
}
|
|
157
|
+
return out
|
|
158
|
+
}
|
|
159
|
+
|
|
160
|
+
// PickThreshold returns the lowest threshold that clears the precision and recall floors. With no floors, it
|
|
161
|
+
// returns the best F1 among the rows; nil when nothing clears them.
|
|
162
|
+
func PickThreshold(rows []ThresholdRow, minPrecision, minRecall *float64) *ThresholdRow {
|
|
163
|
+
if minPrecision == nil && minRecall == nil {
|
|
164
|
+
var best *ThresholdRow
|
|
165
|
+
bestF1 := -1.0
|
|
166
|
+
for i := range rows {
|
|
167
|
+
row := rows[i]
|
|
168
|
+
if row.Precision == nil || row.Recall == nil {
|
|
169
|
+
continue
|
|
170
|
+
}
|
|
171
|
+
f1 := 0.0
|
|
172
|
+
if *row.Precision+*row.Recall != 0 {
|
|
173
|
+
f1 = 2 * *row.Precision * *row.Recall / (*row.Precision + *row.Recall)
|
|
174
|
+
}
|
|
175
|
+
if f1 > bestF1 {
|
|
176
|
+
bestF1 = f1
|
|
177
|
+
best = &rows[i]
|
|
178
|
+
}
|
|
179
|
+
}
|
|
180
|
+
return best
|
|
181
|
+
}
|
|
182
|
+
var candidates []ThresholdRow
|
|
183
|
+
for _, row := range rows {
|
|
184
|
+
p, r := 0.0, 0.0
|
|
185
|
+
if row.Precision != nil {
|
|
186
|
+
p = *row.Precision
|
|
187
|
+
}
|
|
188
|
+
if row.Recall != nil {
|
|
189
|
+
r = *row.Recall
|
|
190
|
+
}
|
|
191
|
+
if (minPrecision == nil || p >= *minPrecision) && (minRecall == nil || r >= *minRecall) {
|
|
192
|
+
candidates = append(candidates, row)
|
|
193
|
+
}
|
|
194
|
+
}
|
|
195
|
+
if len(candidates) == 0 {
|
|
196
|
+
return nil
|
|
197
|
+
}
|
|
198
|
+
sort.SliceStable(candidates, func(i, j int) bool { return candidates[i].Threshold < candidates[j].Threshold })
|
|
199
|
+
return &candidates[0]
|
|
200
|
+
}
|
|
201
|
+
|
|
202
|
+
// Calibrate labels, scores, and reads the numbers: AUC, the threshold sweep, and one recommendation.
|
|
203
|
+
func Calibrate(name string, samples []ScoredSample, opts CalibrateOptions) Calibration {
|
|
204
|
+
thresholds := opts.Thresholds
|
|
205
|
+
if thresholds == nil {
|
|
206
|
+
thresholds = DefaultThresholds(samples, 0)
|
|
207
|
+
}
|
|
208
|
+
rows := Sweep(samples, thresholds)
|
|
209
|
+
recommendation := PickThreshold(rows, opts.MinPrecision, opts.MinRecall)
|
|
210
|
+
c := Calibration{Name: name, Scored: len(samples), Errors: opts.Errors, AUC: AUC(samples), Rows: rows, Recommended: recommendation}
|
|
211
|
+
for _, s := range samples {
|
|
212
|
+
if s.Label {
|
|
213
|
+
c.Positives++
|
|
214
|
+
} else {
|
|
215
|
+
c.Negatives++
|
|
216
|
+
}
|
|
217
|
+
}
|
|
218
|
+
if recommendation != nil {
|
|
219
|
+
t := recommendation.Threshold
|
|
220
|
+
for _, s := range samples {
|
|
221
|
+
if s.Label && s.Score < t {
|
|
222
|
+
c.Missed = append(c.Missed, s)
|
|
223
|
+
}
|
|
224
|
+
if !s.Label && s.Score >= t {
|
|
225
|
+
c.Flagged = append(c.Flagged, s)
|
|
226
|
+
}
|
|
227
|
+
}
|
|
228
|
+
}
|
|
229
|
+
return c
|
|
230
|
+
}
|
|
231
|
+
|
|
232
|
+
func percent(v *float64) string {
|
|
233
|
+
if v == nil {
|
|
234
|
+
return "-"
|
|
235
|
+
}
|
|
236
|
+
return fmt.Sprintf("%.0f%%", math.Floor(*v*100+0.5))
|
|
237
|
+
}
|
|
238
|
+
|
|
239
|
+
func listing(samples []ScoredSample) string {
|
|
240
|
+
parts := make([]string, len(samples))
|
|
241
|
+
for i, s := range samples {
|
|
242
|
+
if s.ID != "" {
|
|
243
|
+
parts[i] = s.ID
|
|
244
|
+
} else {
|
|
245
|
+
parts[i] = fmt.Sprintf("%.2f", s.Score)
|
|
246
|
+
}
|
|
247
|
+
}
|
|
248
|
+
text := strings.Join(parts, ", ")
|
|
249
|
+
return truncateUTF16(text, 300)
|
|
250
|
+
}
|
|
251
|
+
|
|
252
|
+
// FormatCalibration renders plain text, no colour: safe to write to a report file or a log.
|
|
253
|
+
func FormatCalibration(c Calibration) string {
|
|
254
|
+
errs := ""
|
|
255
|
+
if c.Errors > 0 {
|
|
256
|
+
errs = fmt.Sprintf(", %d errors", c.Errors)
|
|
257
|
+
}
|
|
258
|
+
auc := "-"
|
|
259
|
+
if c.AUC != nil {
|
|
260
|
+
auc = fmt.Sprintf("%.3f", *c.AUC)
|
|
261
|
+
}
|
|
262
|
+
lines := []string{
|
|
263
|
+
fmt.Sprintf("%s: %d scored, %d positives, %d negatives%s", c.Name, c.Scored, c.Positives, c.Negatives, errs),
|
|
264
|
+
"AUC " + auc,
|
|
265
|
+
"threshold flagged TP FP FN TN precision recall",
|
|
266
|
+
}
|
|
267
|
+
for _, r := range c.Rows {
|
|
268
|
+
lines = append(lines, fmt.Sprintf("%9.2f %7d %2d %2d %2d %2d %9s %6s", r.Threshold, r.Flagged, r.TP, r.FP, r.FN, r.TN, percent(r.Precision), percent(r.Recall)))
|
|
269
|
+
}
|
|
270
|
+
if c.Recommended == nil {
|
|
271
|
+
lines = append(lines, "recommended: none (no threshold clears the floors)")
|
|
272
|
+
} else {
|
|
273
|
+
r := c.Recommended
|
|
274
|
+
lines = append(lines, fmt.Sprintf("recommended %.2f: precision %s, recall %s, flags %s", r.Threshold, percent(r.Precision), percent(r.Recall), percent(&r.FlagRate)))
|
|
275
|
+
}
|
|
276
|
+
if len(c.Missed) > 0 {
|
|
277
|
+
lines = append(lines, fmt.Sprintf("missed positives (%d): %s", len(c.Missed), listing(c.Missed)))
|
|
278
|
+
}
|
|
279
|
+
if len(c.Flagged) > 0 {
|
|
280
|
+
lines = append(lines, fmt.Sprintf("flagged negatives (%d): %s", len(c.Flagged), listing(c.Flagged)))
|
|
281
|
+
}
|
|
282
|
+
return strings.Join(lines, "\n")
|
|
283
|
+
}
|
|
284
|
+
|
|
285
|
+
// ReplayCase is one labelled replay case: the truth, plus whatever the scorer needs to judge it.
|
|
286
|
+
type ReplayCase[T any] struct {
|
|
287
|
+
ID string
|
|
288
|
+
Label bool
|
|
289
|
+
Data T
|
|
290
|
+
}
|
|
291
|
+
|
|
292
|
+
// ReplayResult is the outcome of one replayed case.
|
|
293
|
+
type ReplayResult[T any] struct {
|
|
294
|
+
ID string
|
|
295
|
+
Label bool
|
|
296
|
+
Data T
|
|
297
|
+
// Score is the judge's number; valid only when Scored.
|
|
298
|
+
Score float64
|
|
299
|
+
Scored bool
|
|
300
|
+
// Error is the failure message, empty on success. It carries no upstream body.
|
|
301
|
+
Error string
|
|
302
|
+
// Skipped is true when the case was never submitted (cancellation or a stopped batch).
|
|
303
|
+
Skipped bool
|
|
304
|
+
}
|
|
305
|
+
|
|
306
|
+
// ReplayOptions configure Replay.
|
|
307
|
+
type ReplayOptions struct {
|
|
308
|
+
// Concurrency is the number of cases in flight at once. Default: DefaultConcurrency.
|
|
309
|
+
Concurrency int
|
|
310
|
+
// StopOn stops launching new cases once it returns true for a failure, for example a budget error.
|
|
311
|
+
StopOn func(error) bool
|
|
312
|
+
// DescribeError turns a scorer error into the reported message. Default: the error's own message.
|
|
313
|
+
DescribeError func(error) string
|
|
314
|
+
}
|
|
315
|
+
|
|
316
|
+
// Replay runs labelled cases through a scorer with bounded concurrency, keeping order and capturing per-case
|
|
317
|
+
// failures. The scorer is usually one Jev question; an error is recorded rather than aborting the run, so one
|
|
318
|
+
// bad case cannot destroy a long calibration. Results feed straight into SamplesOf and Calibrate.
|
|
319
|
+
func Replay[T any](ctx context.Context, cases []ReplayCase[T], score func(ctx context.Context, data T, index int) (float64, error), opts ReplayOptions) []ReplayResult[T] {
|
|
320
|
+
settled := FanOut(ctx, cases, func(ctx context.Context, c ReplayCase[T], index int) (float64, error) {
|
|
321
|
+
return score(ctx, c.Data, index)
|
|
322
|
+
}, FanOutOptions{Concurrency: opts.Concurrency, StopOn: opts.StopOn})
|
|
323
|
+
out := make([]ReplayResult[T], len(settled))
|
|
324
|
+
for i, r := range settled {
|
|
325
|
+
c := cases[i]
|
|
326
|
+
res := ReplayResult[T]{ID: c.ID, Label: c.Label, Data: c.Data}
|
|
327
|
+
if r.OK {
|
|
328
|
+
res.Score, res.Scored = r.Value, true
|
|
329
|
+
} else {
|
|
330
|
+
res.Skipped = r.Skipped
|
|
331
|
+
if opts.DescribeError != nil {
|
|
332
|
+
res.Error = opts.DescribeError(r.Err)
|
|
333
|
+
} else if r.Err != nil && r.Err.Error() != "" {
|
|
334
|
+
res.Error = r.Err.Error()
|
|
335
|
+
} else {
|
|
336
|
+
res.Error = "The scorer failed."
|
|
337
|
+
}
|
|
338
|
+
}
|
|
339
|
+
out[i] = res
|
|
340
|
+
}
|
|
341
|
+
return out
|
|
342
|
+
}
|
|
343
|
+
|
|
344
|
+
// SamplesOf returns the scored cases of a replay, in replay order. Unscored cases are excluded and counted as errors.
|
|
345
|
+
func SamplesOf[T any](results []ReplayResult[T]) (samples []ScoredSample, errors int) {
|
|
346
|
+
for _, r := range results {
|
|
347
|
+
if !r.Scored || math.IsNaN(r.Score) || math.IsInf(r.Score, 0) {
|
|
348
|
+
errors++
|
|
349
|
+
continue
|
|
350
|
+
}
|
|
351
|
+
samples = append(samples, ScoredSample{Label: r.Label, Score: r.Score, ID: r.ID})
|
|
352
|
+
}
|
|
353
|
+
return samples, errors
|
|
354
|
+
}
|