diff --git a/Directory.Packages.props b/Directory.Packages.props index 72c7935..7166ec9 100644 --- a/Directory.Packages.props +++ b/Directory.Packages.props @@ -8,6 +8,7 @@ + diff --git a/bench/PrivacyFilter.Net.Benchmarks/Program.cs b/bench/PrivacyFilter.Net.Benchmarks/Program.cs index d662e66..7faf3af 100644 --- a/bench/PrivacyFilter.Net.Benchmarks/Program.cs +++ b/bench/PrivacyFilter.Net.Benchmarks/Program.cs @@ -97,7 +97,39 @@ public int[] ViterbiDense() => public int[] ViterbiSparse() => _decoder.Decode(_logProbabilities, TokenCount); [Benchmark] - public int[] ArgMax() => PrivacyFilter.ArgMax(_logProbabilities, TokenCount, _classCount); + public int[] ArgMaxScalar() => + ArgMaxScalar(_logProbabilities, TokenCount, _classCount); + + [Benchmark] + public int[] ArgMaxTensorPrimitives() => + PrivacyFilter.ArgMax(_logProbabilities, TokenCount, _classCount); + + private static int[] ArgMaxScalar( + ReadOnlySpan scores, + int tokenCount, + int classCount) + { + var labels = new int[tokenCount]; + for (int token = 0; token < tokenCount; token++) + { + int offset = token * classCount; + int bestLabel = 0; + float bestScore = scores[offset]; + for (int label = 1; label < classCount; label++) + { + float score = scores[offset + label]; + if (score > bestScore) + { + bestScore = score; + bestLabel = label; + } + } + + labels[token] = bestLabel; + } + + return labels; + } private static int[] DecodeDense( float[] emissions, diff --git a/src/PrivacyFilter.Net/PrivacyFilter.Net.csproj b/src/PrivacyFilter.Net/PrivacyFilter.Net.csproj index 128d749..a9dbc25 100644 --- a/src/PrivacyFilter.Net/PrivacyFilter.Net.csproj +++ b/src/PrivacyFilter.Net/PrivacyFilter.Net.csproj @@ -47,6 +47,7 @@ + diff --git a/src/PrivacyFilter.Net/PrivacyFilter.cs b/src/PrivacyFilter.Net/PrivacyFilter.cs index a134d07..78f04a9 100644 --- a/src/PrivacyFilter.Net/PrivacyFilter.cs +++ b/src/PrivacyFilter.Net/PrivacyFilter.cs @@ -1,6 +1,7 @@ using System.Buffers; using Microsoft.ML.OnnxRuntime; using Microsoft.ML.OnnxRuntime.Tensors; +using TensorPrimitives = System.Numerics.Tensors.TensorPrimitives; namespace PrivacyFilterNet; @@ -247,20 +248,8 @@ internal static int[] ArgMax(ReadOnlySpan scores, int tokenCount, int cla var labels = new int[tokenCount]; for (int token = 0; token < tokenCount; token++) { - int offset = token * classCount; - int bestLabel = 0; - float bestScore = scores[offset]; - for (int label = 1; label < classCount; label++) - { - float score = scores[offset + label]; - if (score > bestScore) - { - bestScore = score; - bestLabel = label; - } - } - - labels[token] = bestLabel; + labels[token] = TensorPrimitives.IndexOfMax( + scores.Slice(token * classCount, classCount)); } return labels; diff --git a/tests/PrivacyFilter.Net.Tests/DecoderTests.cs b/tests/PrivacyFilter.Net.Tests/DecoderTests.cs index 6576db0..24b435f 100644 --- a/tests/PrivacyFilter.Net.Tests/DecoderTests.cs +++ b/tests/PrivacyFilter.Net.Tests/DecoderTests.cs @@ -84,6 +84,28 @@ public void SparseViterbiMatchesDenseReference() decoder.Decode(tiedEmissions, 8)); } + [Fact] + public void TensorArgMaxMatchesScalarForFiniteScores() + { + const int tokenCount = 128; + int classCount = ClassNames.Length; + var scores = new float[tokenCount * classCount]; + var random = new Random(42); + for (int index = 0; index < scores.Length; index++) + { + scores[index] = (random.NextSingle() * 20) - 10; + } + + Assert.Equal( + ArgMaxScalar(scores, tokenCount, classCount), + PrivacyFilter.ArgMax(scores, tokenCount, classCount)); + + Array.Clear(scores); + Assert.Equal( + ArgMaxScalar(scores, tokenCount, classCount), + PrivacyFilter.ArgMax(scores, tokenCount, classCount)); + } + [Fact] public void SpanDecoderTrimsAndRedacts() { @@ -113,6 +135,33 @@ public void SpanDecoderTrimsAndRedacts() private static int ClassIndex(string name) => Array.IndexOf(ClassNames, name); + private static int[] ArgMaxScalar( + ReadOnlySpan scores, + int tokenCount, + int classCount) + { + var labels = new int[tokenCount]; + for (int token = 0; token < tokenCount; token++) + { + int offset = token * classCount; + int bestLabel = 0; + float bestScore = scores[offset]; + for (int label = 1; label < classCount; label++) + { + float score = scores[offset + label]; + if (score > bestScore) + { + bestScore = score; + bestLabel = label; + } + } + + labels[token] = bestLabel; + } + + return labels; + } + private static int[] DecodeDenseReference(float[] emissions, int tokenCount, int classCount) { const float invalid = -1e9f;