feat(backend): PolicyRunner ONNX inference wrapper

This commit is contained in:
meelstorm
2026-07-17 17:40:56 +00:00
committed by EugeneTes
parent 9f34314c2f
commit 0102f5f748
2 changed files with 107 additions and 0 deletions

View File

@@ -0,0 +1,55 @@
using System.IO;
using Backend;
using Xunit;
namespace Backend.Tests;
public class PolicyRunnerTests
{
private static string LocateOnnxModel()
{
var dir = AppContext.BaseDirectory;
while (dir is not null)
{
var candidate = Path.Combine(dir, "models", "ppo_lander.onnx");
if (File.Exists(candidate)) return candidate;
var parent = Directory.GetParent(dir);
if (parent is null) break;
dir = parent.FullName;
}
throw new FileNotFoundException(
"models/ppo_lander.onnx not found — run `python Training/export_onnx.py ...`");
}
[Fact]
public void LoadsWithoutError()
{
using var runner = new PolicyRunner(LocateOnnxModel());
}
[Fact]
public void SelectAction_ReturnsValidActionForZeroObs()
{
using var runner = new PolicyRunner(LocateOnnxModel());
var obs = new float[7]; // all zeros
int action = runner.SelectAction(obs);
Assert.InRange(action, 0, 3);
}
[Fact]
public void SelectAction_IsDeterministicForSameInput()
{
using var runner = new PolicyRunner(LocateOnnxModel());
var obs = new float[] { 0.1f, -0.2f, 0.05f, 0.0f, 0.0f, 1.0f, 0.0f };
int a1 = runner.SelectAction(obs);
int a2 = runner.SelectAction(obs);
Assert.Equal(a1, a2);
}
[Fact]
public void SelectAction_ThrowsIfObsLengthWrong()
{
using var runner = new PolicyRunner(LocateOnnxModel());
Assert.Throws<ArgumentException>(() => runner.SelectAction(new float[3]));
}
}

52
Backend/PolicyRunner.cs Normal file
View File

@@ -0,0 +1,52 @@
using Microsoft.ML.OnnxRuntime;
using Microsoft.ML.OnnxRuntime.Tensors;
namespace Backend;
/// <summary>
/// Loads a PPO policy ONNX model once and provides deterministic action selection.
/// Thread-safe: <see cref="InferenceSession"/> is safe for concurrent Run() calls.
/// </summary>
public sealed class PolicyRunner : IDisposable
{
public const int ObsDim = 7;
public const int NActions = 4;
private readonly InferenceSession _session;
private readonly string _inputName;
public PolicyRunner(string onnxPath)
{
if (!File.Exists(onnxPath))
throw new FileNotFoundException($"ONNX model not found at {onnxPath}");
_session = new InferenceSession(onnxPath);
_inputName = _session.InputMetadata.Keys.First();
}
/// <summary>Run one inference and return the argmax action index.</summary>
public int SelectAction(ReadOnlySpan<float> obs)
{
if (obs.Length != ObsDim)
throw new ArgumentException($"expected obs of length {ObsDim}, got {obs.Length}", nameof(obs));
var tensor = new DenseTensor<float>(new[] { 1, ObsDim });
for (int i = 0; i < ObsDim; i++) tensor[0, i] = obs[i];
using var results = _session.Run(new[]
{
NamedOnnxValue.CreateFromTensor(_inputName, tensor)
});
var logits = results.First().AsEnumerable<float>().ToArray();
// argmax
int best = 0;
float bestVal = logits[0];
for (int i = 1; i < logits.Length; i++)
{
if (logits[i] > bestVal) { best = i; bestVal = logits[i]; }
}
return best;
}
public void Dispose() => _session.Dispose();
}