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;