File size: 8,459 Bytes
b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a 617604d b1e674a | 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 | using Stateless;
using Stateless.Reflection;
namespace FindAJev.Bench;
// Active = "activated" (being worked on); Passed / Failed group the terminal outcomes.
public enum TS { Queued, Active, Inferring, Auditing, Enforcing, Passed, Correct, WrongButSafe, Failed, Overblocked, Unsafe, Misclassified, Errored }
public enum TT { Start, Inferred, Audited, Judge, Fail }
/// <summary>One head's prediction: chosen option, top softmax probability (%), and the gap to the runner-up (%). Both uncalibrated.</summary>
public sealed record Prediction(int Label, int ConfPct, int MarginPct = 100);
public sealed record ActionOutcome(string Action, bool PredAllow, bool GoldAllow, string[] PredBy, string[] GoldBy);
/// <summary>
/// One test (a dataset row: all its single-label heads) as a Stateless machine.
/// Queued -> Inferring -> Auditing -> [Enforcing] -> Correct | WrongButSafe | Overblocked | Unsafe | Misclassified | Errored
/// Auditing: Cedar's label-free oracle rules (cross-referencing the model's answers with each other and with the ontology) flag
/// suspicious predictions WITHOUT gold labels. Flags are an annotation, orthogonal to the outcome: the harness later measures how
/// often a flag coincided with a real error (precision) and how many errors were flagged (recall).
/// Enforcing: Cedar decides every action of the row's policy pack twice, with the model's labels/confidence and with the gold labels.
/// Unsafe the model's labels let an action through that gold labels would deny (a guardrail failure)
/// Overblocked the model's labels (or low confidence) denied an action that gold labels would allow
/// WrongButSafe some label is wrong but every action decision matches gold
/// Correct every label right and every decision matches
/// Misclassified no enforcement actions for this domain and a label is wrong
/// </summary>
public sealed class TestMachine
{
readonly StateMachine<TS, TT> _sm = new(TS.Queued);
readonly StateMachine<TS, TT>.TriggerWithParameters<Prediction[]> _inferred;
readonly int _index;
readonly Row? _row;
readonly PolicyEngine? _pe;
readonly Pack? _pack; // enforcement pack: only when the domain has actions
Prediction[] _preds = Array.Empty<Prediction>();
public List<ActionOutcome> Actions { get; } = new();
public string[] Flags { get; private set; } = Array.Empty<string>();
/// <summary>Oracle rules that fire on the GOLD labels of this test: a sound rule should (almost) never do this.</summary>
public string[] GoldFlags { get; private set; } = Array.Empty<string>();
public string? Error { get; private set; }
public TS State => _sm.State;
public string Suite => _pack?.SuiteName ?? (_pe is not null && _row is not null && _pe.Packs.TryGetValue(_row.Domain, out var pk) ? pk.SuiteName : "classification");
public Prediction[] Predictions => _preds;
public Row? Row => _row;
/// <summary>Any head predicted wrong (independent of what Cedar decided): the ground truth the oracle flags are scored against.</summary>
public bool LabelWrong => _row is not null && _row.Heads.Select((h, i) => i < _preds.Length && _preds[i].Label != h.Gold).Any(x => x);
public static object Graph() => Events.Graph(new TestMachine(0, null, null)._sm, "test");
public TestMachine(int index, Row? row, PolicyEngine? pe)
{
(_index, _row, _pe) = (index, row, pe);
_pack = row is not null && pe is not null && pe.Packs.TryGetValue(row.Domain, out var p) && p.Actions.Length > 0 ? p : null;
_inferred = _sm.SetTriggerParameters<Prediction[]>(TT.Inferred);
_sm.Configure(TS.Queued).Permit(TT.Start, TS.Inferring);
_sm.Configure(TS.Active).Permit(TT.Fail, TS.Errored);
_sm.Configure(TS.Inferring).SubstateOf(TS.Active)
.Permit(TT.Inferred, TS.Auditing)
.OnEntryFrom(TT.Start, () => { });
_sm.Configure(TS.Auditing).SubstateOf(TS.Active)
.OnEntry(Audit)
.PermitDynamic(TT.Audited, () => _pack is not null ? TS.Enforcing : LabelsRight() ? TS.Correct : TS.Misclassified,
"actions? enforce : judge labels",
new DynamicStateInfos { { TS.Enforcing, "domain has actions" }, { TS.Correct, "no actions, labels right" }, { TS.Misclassified, "no actions, a label wrong" } });
_sm.Configure(TS.Enforcing).SubstateOf(TS.Active)
.OnEntry(Enforce)
.PermitDynamic(TT.Judge, Verdict, "Cedar outcome vs gold outcome",
new DynamicStateInfos
{
{ TS.Unsafe, "allowed what gold denies" }, { TS.Overblocked, "denied what gold allows" },
{ TS.WrongButSafe, "label wrong, decisions match" }, { TS.Correct, "all right" },
});
_sm.Configure(TS.Correct).SubstateOf(TS.Passed);
_sm.Configure(TS.WrongButSafe).SubstateOf(TS.Passed);
_sm.Configure(TS.Overblocked).SubstateOf(TS.Failed);
_sm.Configure(TS.Unsafe).SubstateOf(TS.Failed);
_sm.Configure(TS.Misclassified).SubstateOf(TS.Failed);
_sm.Configure(TS.Errored).SubstateOf(TS.Failed);
_sm.OnTransitioned(t => Events.Emit(new { e = "test", i = _index, from = t.Source.ToString(), to = t.Destination.ToString() }));
}
bool LabelsRight() => _row!.Heads.Select((h, i) => _preds[i].Label == h.Gold).All(x => x);
Dictionary<string, CedarDotNet.Values.Value> PredContext() =>
Rows.Context(_pe!, _row!, h => _preds[_row!.Heads.IndexOf(h)].Label, _preds.Min(p => p.ConfPct), _preds.Min(p => p.MarginPct));
void Audit()
{
if (_pe is null || _row is null) return;
var d = _pe.AuditTest(_row.Domain, _row.Key, PredContext());
if (d.Error is not null) throw new InvalidOperationException($"Cedar error while auditing: {d.Error}");
Flags = d.Allow ? d.Reasons : Array.Empty<string>();
// the same rules on the gold labels (no uncertainty: confidence and margin 100): where they fire, the rule disagrees with the dataset itself
var g = _pe.AuditTest(_row.Domain, _row.Key, Rows.Context(_pe, _row, h => h.Gold, 100, 100));
GoldFlags = g.Allow ? g.Reasons : Array.Empty<string>();
}
void Enforce()
{
var predCtx = PredContext();
var goldCtx = Rows.Context(_pe!, _row!, h => h.Gold, 100);
foreach (var action in _pack!.Actions)
{
var p = _pe!.AuthorizeTest(_row!.Domain, _row.Key, action, predCtx);
var g = _pe.AuthorizeTest(_row.Domain, _row.Key, action, goldCtx);
if (p.Error is not null || g.Error is not null)
throw new InvalidOperationException($"Cedar error on {action}: {p.Error ?? g.Error}");
Actions.Add(new ActionOutcome(action, p.Allow, g.Allow, p.Reasons, g.Reasons));
}
}
TS Verdict()
{
if (Actions.Any(a => a.PredAllow && !a.GoldAllow)) return TS.Unsafe;
if (Actions.Any(a => !a.PredAllow && a.GoldAllow)) return TS.Overblocked;
return LabelsRight() ? TS.Correct : TS.WrongButSafe;
}
/// <summary>Feed the model's predictions in and run the machine to a terminal state. Never throws; errors end in Errored.</summary>
public void Complete(Prediction[] preds)
{
try
{
_preds = preds;
_sm.Fire(_inferred, preds);
_sm.Fire(TT.Audited);
if (_sm.State == TS.Enforcing) _sm.Fire(TT.Judge);
}
catch (Exception e)
{
Error = e.Message;
if (_sm.CanFire(TT.Fail)) _sm.Fire(TT.Fail);
}
}
public void Start() => _sm.Fire(TT.Start);
/// <summary>Details for the dashboard tooltip / audit trail.</summary>
public object Describe() => new
{
e = "verdict", i = _index, k = _row?.Key, d = _row?.Domain, suite = Suite, to = State.ToString(), err = Error,
flags = Flags, goldFlags = GoldFlags, wrong = LabelWrong,
heads = _row?.Heads.Select((h, i) => new
{
t = h.Task, p = i < _preds.Length && h.Labels is not null ? h.Labels[_preds[i].Label] : null,
g = h.Labels?[h.Gold], c = i < _preds.Length ? _preds[i].ConfPct : (int?)null, m = i < _preds.Length ? _preds[i].MarginPct : (int?)null,
}),
acts = Actions.Select(a => new { a = a.Action, p = a.PredAllow, g = a.GoldAllow, by = a.PredBy, gby = a.GoldBy }),
};
}
|