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";
}
}
|