File size: 16,495 Bytes
b01bf09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b1e674a
 
 
 
b01bf09
 
 
 
0ae18df
 
b01bf09
 
 
 
 
 
b1e674a
 
b01bf09
 
b1e674a
b01bf09
 
 
 
 
 
 
 
 
b1e674a
 
b01bf09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b1e674a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b01bf09
 
 
 
 
 
 
 
 
b1e674a
 
 
 
 
 
 
 
 
b01bf09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
0ae18df
 
 
b01bf09
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b1e674a
 
 
 
617604d
 
 
 
 
 
 
 
 
b1e674a
 
 
 
 
 
 
 
0ae18df
b1e674a
 
 
 
 
b01bf09
 
b1e674a
 
b01bf09
 
b1e674a
 
 
 
 
0ae18df
b1e674a
0ae18df
b1e674a
 
 
b01bf09
b1e674a
 
 
 
 
 
 
 
 
 
 
 
b01bf09
b1e674a
 
 
 
 
 
 
 
 
 
 
 
 
 
b01bf09
b1e674a
b01bf09
 
 
 
 
 
0ae18df
 
b01bf09
 
0ae18df
 
b01bf09
 
 
 
 
 
0ae18df
b01bf09
0ae18df
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b1e674a
 
0ae18df
 
 
 
 
 
b1e674a
0ae18df
617604d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b01bf09
 
 
 
617604d
 
 
 
 
 
 
 
b01bf09
 
 
 
 
 
 
 
 
 
 
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
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
using System.Diagnostics;
using System.Text.Json;
using Microsoft.ML.OnnxRuntime;
using Microsoft.ML.OnnxRuntime.Tensors;
using Stateless;

namespace FindAJev.Bench;

public enum RunState { Running, Pending, ModelReady, SessionLoaded, WarmedUp, Measured, Scored, Failed }
public enum RunTrigger { Verify, Load, Warm, Measure, Score, Fail }

/// <summary>
/// The lifecycle of benchmarking one model, as a Stateless machine.
/// Every working state is a substate of Running, so a single Running.Permit(Fail) covers all of them.
/// </summary>
public sealed class RunMachine
{
    readonly ModelSpec _spec;
    readonly string _root;
    readonly int _threads, _warmup, _repeats;
    readonly List<Item> _items;
    readonly StateMachine<RunState, RunTrigger> _sm;
    readonly PolicyEngine? _pe;
    readonly int _cpus;
    string? _policyDenial;
    readonly List<TestMachine> _tests = new();

    InferenceSession? _session;
    List<Dictionary<string, OrtValue>>? _inputs;
    readonly List<double> _ms = new();
    readonly List<(int item, double ms)> _msItem = new();   // every timed call with the item it belongs to (all repeats)
    string[] _suiteOf = Array.Empty<string>();
    readonly List<float[]> _logits = new();
    readonly RunResult _r = new();

    public RunState State => _sm.State;
    public RunResult Result => _r;

    public RunMachine(ModelSpec spec, string root, List<Item> items, int threads, int warmup, int repeats,
                      PolicyEngine? pe = null, int cpus = 0)
    {
        (_spec, _root, _items, _threads, _warmup, _repeats) = (spec, root, items, threads, warmup, repeats);
        (_pe, _cpus) = (pe, cpus > 0 ? cpus : Environment.ProcessorCount);
        _r.Id = spec.Id; _r.Family = spec.Family; _r.Precision = spec.Precision; _r.Threads = threads;
        _r.Items = items.Count; _r.OrtVersion = OrtEnv.Instance().GetVersionString();
        _r.Cpu = CpuName();

        _sm = new StateMachine<RunState, RunTrigger>(RunState.Pending);

        _sm.Configure(RunState.Running).Permit(RunTrigger.Fail, RunState.Failed);

        _sm.Configure(RunState.Pending).SubstateOf(RunState.Running)
            .PermitIf(RunTrigger.Verify, RunState.ModelReady, () => File.Exists(OnnxPath) && _policyDenial is null,
                      "model file present and Cedar allows FetchModel + RunModel");

        _sm.Configure(RunState.ModelReady).SubstateOf(RunState.Running)
            .Permit(RunTrigger.Load, RunState.SessionLoaded);

        _sm.Configure(RunState.SessionLoaded).SubstateOf(RunState.Running)
            .OnEntry(CreateSession)
            .Permit(RunTrigger.Warm, RunState.WarmedUp);

        _sm.Configure(RunState.WarmedUp).SubstateOf(RunState.Running)
            .OnEntry(WarmUp)
            .Permit(RunTrigger.Measure, RunState.Measured);

        _sm.Configure(RunState.Measured).SubstateOf(RunState.Running)
            .OnEntry(Measure)
            .Permit(RunTrigger.Score, RunState.Scored);

        _sm.Configure(RunState.Scored)
            .OnEntry(Score);

        _sm.Configure(RunState.Failed)
            .OnEntry(t => { _r.State = "Failed"; });

        _sm.OnTransitioned(t => Events.Emit(new { e = "run", model = _spec.Id, from = t.Source.ToString(), to = t.Destination.ToString() }));
    }

    /// <summary>Ask Cedar whether this model may be fetched and run here. Sets _policyDenial when it may not.</summary>
    void EnforceRunPolicy()
    {
        if (_pe is null) return;
        foreach (var action in new[] { "FetchModel", "RunModel" })
        {
            var d = _pe.AuthorizeRun(action, _spec, _threads, _cpus);
            Events.Emit(new { e = "policy", scope = "run", model = _spec.Id, action, allow = d.Allow, by = d.Reasons });
            if (d.Error is not null) throw new InvalidOperationException($"Cedar error on {action}: {d.Error}");
            if (!d.Allow)
            {
                _policyDenial = $"Cedar denied {action} for {_spec.Id} ({(d.Reasons.Length > 0 ? string.Join(", ", d.Reasons) : "no permit policy matched")})";
                return;
            }
        }
    }

    string OnnxPath => Path.Combine(_root, _spec.Onnx);

    /// <summary>Drive the machine to a terminal state. Never throws; failures end in Failed with Result.Error set.</summary>
    public RunResult Run()
    {
        try
        {
            Events.Emit(Events.Graph(_sm, "run"));
            Events.Emit(TestMachine.Graph());
            if (_pe is not null)
            {
                var bad = _pe.Validate().Where(p => !p.Contains("warning:")).ToList();
                if (bad.Count > 0) throw new InvalidOperationException("policy validation failed: " + string.Join(" | ", bad).Replace("\n", " "));
            }
            EnforceRunPolicy();
            if (_policyDenial is not null) throw new UnauthorizedAccessException(_policyDenial);
            if (!_sm.CanFire(RunTrigger.Verify))
                throw new FileNotFoundException($"missing {OnnxPath} (run: python tools/fetch.py {_spec.Id})");
            _sm.Fire(RunTrigger.Verify);
            _sm.Fire(RunTrigger.Load);
            _sm.Fire(RunTrigger.Warm);
            _sm.Fire(RunTrigger.Measure);
            _sm.Fire(RunTrigger.Score);
        }
        catch (Exception e)
        {
            _r.Error = e.Message;
            if (_sm.CanFire(RunTrigger.Fail)) _sm.Fire(RunTrigger.Fail); else _r.State = "Failed";
        }
        finally { _session?.Dispose(); }
        return _r;
    }

    public string Dot() => Stateless.Graph.UmlDotGraph.Format(_sm.GetInfo());

    void CreateSession()
    {
        Progress("loading", 0, _items.Count, 0);
        var sw = Stopwatch.StartNew();
        var o = new SessionOptions
        {
            IntraOpNumThreads = _threads,
            InterOpNumThreads = 1,
            ExecutionMode = ExecutionMode.ORT_SEQUENTIAL,
            GraphOptimizationLevel = GraphOptimizationLevel.ORT_ENABLE_ALL,
        };
        _session = new InferenceSession(OnnxPath, o); // CPUExecutionProvider is the default; no other EP is appended
        _inputs = _items.Select(BuildInputs).ToList();
        _r.LoadSeconds = sw.Elapsed.TotalSeconds;
        // suite of each item: explicit (retrieval/tools files) or derived from the domain's policy pack (core files)
        _suiteOf = _items.Select(it => it.Suite.Length > 0 ? it.Suite
            : _pe is not null && _pe.Packs.TryGetValue(it.Domain, out var pk) ? pk.SuiteName : _pe is not null ? "classification" : "core").ToArray();
        Progress("loaded", 0, _items.Count, 0);
    }

    // Input names verified against the ONNX graphs with onnxruntime (Python): see README.
    Dictionary<string, OrtValue> BuildInputs(Item it)
    {
        static OrtValue L(long[] v, long[] shape) => OrtValue.CreateTensorValueFromMemory(v, shape);
        var n = it.Ids.Length;
        var ids = L(it.Ids, new long[] { 1, n });
        var att = L(Enumerable.Repeat(1L, n).ToArray(), new long[] { 1, n });
        var pos = L(it.Pos, new long[] { 1, it.Pos.Length });
        return _spec.Family switch
        {
            "gliner" => new() { ["input_ids"] = ids, ["attention_mask"] = att, ["label_positions"] = pos },
            "julia" or "laya" => new()
            {
                ["input_ids"] = ids, ["attention_mask"] = att, ["marker_pos"] = pos,
                ["marker_mask"] = OrtValue.CreateTensorValueFromMemory(Enumerable.Repeat(true, it.Pos.Length).ToArray(), new long[] { 1, it.Pos.Length }),
                ["qtype"] = L(new[] { it.QType }, new long[] { 1 }),
            },
            _ => throw new NotSupportedException($"family {_spec.Family}"),
        };
    }

    float[] Infer(int i)
    {
        using var run = new RunOptions();
        var inp = _inputs![i];
        using var outs = _session!.Run(run, inp.Keys.ToArray(), inp.Values.ToArray(), new[] { "logits" });
        return outs[0].GetTensorDataAsSpan<float>().ToArray();
    }

    void WarmUp()
    {
        Progress("warmup", 0, _warmup, 0);
        for (var k = 0; k < _warmup; k++) Infer(k % _items.Count);
    }

    static Prediction Predict(float[] logits, int n)
    {
        int arg = 0;
        for (var k = 1; k < n; k++) if (logits[k] > logits[arg]) arg = k;
        double sum = 0, second = 0;
        for (var k = 0; k < n; k++)
        {
            var e = Math.Exp(logits[k] - logits[arg]);
            sum += e;
            if (k != arg && e > second) second = e;
        }
        // softmax top probability and the gap to the runner-up (both uncalibrated), in percent
        return new Prediction(arg, (int)Math.Round(100.0 / sum), (int)Math.Round(100.0 * (1 - second) / sum));
    }

    double TimedInfer(int i)
    {
        var sw = Stopwatch.GetTimestamp();
        var lg = Infer(i);
        var ms = Stopwatch.GetElapsedTime(sw).TotalMilliseconds;
        _ms.Add(ms);
        _msItem.Add((i, ms));
        _lastLogits = lg;
        return ms;
    }
    float[] _lastLogits = Array.Empty<float>();

    void Measure()
    {
        // Per-decision latency at batch size 1, sequential, wall clock around session.Run only. Cedar and event emission
        // happen outside the stopwatch. Repeat 1 records logits and runs the per-test machines; every repeat feeds latency.
        var total = _repeats * _items.Count;
        Progress("measure", 0, total, 0);
        var rows = _pe is null ? null : Rows.Group(_items);
        if (rows is not null)
            Events.Emit(new
            {
                e = "plan", model = _spec.Id, tests = rows.Count,
                suites = rows.GroupBy(r => _suiteOf[r.Start]).ToDictionary(g => g.Key, g => g.Count()),
                domains = rows.GroupBy(r => r.Domain).ToDictionary(g => g.Key, g => g.Count()),
                packDomains = _pe!.Packs.ToDictionary(k => k.Key, k => k.Value.SuiteName),
                keys = rows.Select(r => r.Key + "|" + r.Domain),
            });

        for (var rep = 0; rep < _repeats; rep++)
        {
            if (rows is null || rep > 0)
            {
                for (var i = 0; i < _items.Count; i++)
                {
                    TimedInfer(i);
                    if (rep == 0) _logits.Add(_lastLogits);
                    if (_ms.Count % 25 == 0 || _ms.Count == total) Progress("measure", _ms.Count, total, _ms.Average());
                }
                continue;
            }
            for (var ri = 0; ri < rows.Count; ri++)
            {
                var row = rows[ri];
                var tm = new TestMachine(ri, row, _pe);
                tm.Start();
                var preds = new Prediction[row.Heads.Count];
                for (var h = 0; h < row.Heads.Count; h++)
                {
                    TimedInfer(row.Start + h);
                    _logits.Add(_lastLogits);
                    preds[h] = Predict(_lastLogits, row.Heads[h].N);
                    if (_ms.Count % 25 == 0 || _ms.Count == total) Progress("measure", _ms.Count, total, _ms.Average());
                }
                tm.Complete(preds);
                _tests.Add(tm);
                Events.Emit(tm.Describe());
            }
        }
    }

    /// <summary>Machine-readable progress line on stderr, consumed by space/server.py (which turns it into SSE events).</summary>
    void Progress(string phase, int done, int total, double meanMs) =>
        Console.Error.WriteLine("PROGRESS " + JsonSerializer.Serialize(new { model = _spec.Id, phase, done, total, meanMs = Math.Round(meanMs, 2) }));

    static double Pct(double[] sorted, double q) => sorted.Length == 0 ? 0 : sorted[Math.Min(sorted.Length - 1, (int)(sorted.Length * q))];

    void Score()
    {
        // Head-level argmax accuracy for every item.
        var hit = new bool[_items.Count];
        for (var i = 0; i < _items.Count; i++)
        {
            var lg = _logits[i];
            if (lg.Length < _items[i].N) throw new InvalidOperationException($"{_items[i].Id}: {lg.Length} logits for {_items[i].N} options");
            var arg = 0;
            for (var k = 1; k < _items[i].N; k++) if (lg[k] > lg[arg]) arg = k;
            hit[i] = arg == _items[i].Gold;
        }

        // Headline numbers stay on the original fast-decisions set (classification + automation, or "core" without policies) so they
        // remain comparable with earlier runs; retrieval and tools are reported per suite.
        bool Core(int i) => _suiteOf[i] is "classification" or "automation" or "core";
        var coreIdx = Enumerable.Range(0, _items.Count).Where(Core).ToArray();
        if (coreIdx.Length == 0) coreIdx = Enumerable.Range(0, _items.Count).ToArray();   // suite subset without fast-decisions: headline = everything run
        var inHeadline = coreIdx.ToHashSet();
        var coreMs = _msItem.Where(t => inHeadline.Contains(t.item)).Select(t => t.ms).OrderBy(x => x).ToArray();
        _r.MeanMs = coreMs.Average();
        _r.P50Ms = Pct(coreMs, 0.5);
        _r.P95Ms = Pct(coreMs, 0.95);
        _r.ItemsPerSec = 1000.0 / _r.MeanMs;
        _r.Accuracy = (double)coreIdx.Count(i => hit[i]) / coreIdx.Length;
        _r.AccuracyByDomain = coreIdx.GroupBy(i => _items[i].Domain).OrderBy(g => g.Key).ToDictionary(g => g.Key, g => (double)g.Count(i => hit[i]) / g.Count());

        foreach (var g in Enumerable.Range(0, _items.Count).GroupBy(i => _suiteOf[i]))
        {
            var ms = _msItem.Where(t => _suiteOf[t.item] == g.Key).Select(t => t.ms).OrderBy(x => x).ToArray();
            var tests = _tests.Where(t => t.Suite == g.Key).ToList();
            _r.Suites[g.Key] = new SuiteResult
            {
                Heads = g.Count(), Accuracy = (double)g.Count(i => hit[i]) / g.Count(),
                MeanMs = ms.Average(), P50Ms = Pct(ms, 0.5), P95Ms = Pct(ms, 0.95),
                Tests = tests.Count,
                Correct = tests.Count(t => t.State == TS.Correct), WrongButSafe = tests.Count(t => t.State == TS.WrongButSafe),
                Overblocked = tests.Count(t => t.State == TS.Overblocked), Unsafe = tests.Count(t => t.State == TS.Unsafe),
                Misclassified = tests.Count(t => t.State == TS.Misclassified), Errored = tests.Count(t => t.State == TS.Errored),
            };
        }
        foreach (var g in _tests.GroupBy(t => t.Suite))
            _r.Oracle[g.Key] = new OracleStats
            {
                Tests = g.Count(), Wrong = g.Count(t => t.LabelWrong), Flagged = g.Count(t => t.Flags.Length > 0),
                FlaggedWrong = g.Count(t => t.Flags.Length > 0 && t.LabelWrong),
                Unsafe = g.Count(t => t.State == TS.Unsafe), UnsafeFlagged = g.Count(t => t.State == TS.Unsafe && t.Flags.Length > 0),
            };
        if (_pe is not null)
            foreach (var rule in _pe.OraclePolicies.Keys)
            {
                var dom = _pe.OracleDomain(rule);
                var eligible = _tests.Where(t => dom is null || t.Row!.Domain == dom).ToList();
                if (eligible.Count == 0) continue;
                _r.OracleRules[rule] = new OracleRuleStat
                {
                    Eligible = eligible.Count,
                    EligibleWrong = eligible.Count(t => t.LabelWrong),
                    Flagged = eligible.Count(t => t.Flags.Contains(rule)),
                    FlaggedWrong = eligible.Count(t => t.Flags.Contains(rule) && t.LabelWrong),
                    FiredOnGold = eligible.Count(t => t.GoldFlags.Contains(rule)),
                };
            }
        _r.PeakRssMb = Process.GetCurrentProcess().PeakWorkingSet64 / 1048576.0;
        _r.State = "Scored";
    }

    /// <summary>One JSON line per test: predictions, gold, confidence/margin, oracle flags and outcome. Enables offline cross-model checks.</summary>
    public IEnumerable<string> PredictionLines() => _tests.Select(t => JsonSerializer.Serialize(new
    {
        k = t.Row!.Key, d = t.Row.Domain, suite = t.Suite, state = t.State.ToString(), wrong = t.LabelWrong, flags = t.Flags,
        heads = t.Row.Heads.Select((h, i) => new { t = h.Task, p = t.Predictions[i].Label, pl = h.Labels?[t.Predictions[i].Label], g = h.Gold,
                                                   c = t.Predictions[i].ConfPct, m = t.Predictions[i].MarginPct, n = h.N }),
    }));

    static string CpuName()
    {
        try
        {
            foreach (var l in File.ReadLines("/proc/cpuinfo"))
                if (l.StartsWith("model name")) return l.Split(':', 2)[1].Trim();
        }
        catch { }
        return "unknown";
    }
}