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 } /// /// 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. /// public sealed class RunMachine { readonly ModelSpec _spec; readonly string _root; readonly int _threads, _warmup, _repeats; readonly List _items; readonly StateMachine _sm; readonly PolicyEngine? _pe; readonly int _cpus; string? _policyDenial; readonly List _tests = new(); InferenceSession? _session; List>? _inputs; readonly List _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(); readonly List _logits = new(); readonly RunResult _r = new(); public RunState State => _sm.State; public RunResult Result => _r; public RunMachine(ModelSpec spec, string root, List 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.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() })); } /// Ask Cedar whether this model may be fetched and run here. Sets _policyDenial when it may not. 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); /// Drive the machine to a terminal state. Never throws; failures end in Failed with Result.Error set. 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 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().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(); 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()); } } } /// Machine-readable progress line on stderr, consumed by space/server.py (which turns it into SSE events). 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"; } /// One JSON line per test: predictions, gold, confidence/margin, oracle flags and outcome. Enables offline cross-model checks. public IEnumerable 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"; } }