From 42435ad9150db459ee97a399c53050a7ec58012c Mon Sep 17 00:00:00 2001
From: Stevan Freeborn <65925598+StevanFreeborn@users.noreply.github.com>
Date: Thu, 13 Feb 2025 17:24:38 -0600
Subject: [PATCH] feat: refactor to proper app + write tests
---
.editorconfig | 2 +
.../BGR.Console.Tests.csproj | 9 +-
.../Integration/AppFactory.cs | 19 ++
.../Integration/RemovalCommandTests.cs | 24 ++
.../Unit/ImageSharpProcessorTests.cs | 241 ++++++++++++++++++
.../Unit/ModNetModelTests.cs | 66 +++++
.../Unit/ModelFactoryTests.cs | 67 +++++
src/BGR.Console.Tests/Unit/ModelTests.cs | 94 +++++++
.../Unit/OnnxInferenceRunnerTests.cs | 34 +++
src/BGR.Console.Tests/Unit/OnnxTensorTests.cs | 121 +++++++++
.../Unit/RemovalCommandTests.cs | 181 +++++++++++++
src/BGR.Console.Tests/Unit/RmbgModelTests.cs | 66 +++++
src/BGR.Console.Tests/Unit/SharpImageTests.cs | 86 +++++++
src/BGR.Console.Tests/Unit/U2NetModelTests.cs | 66 +++++
src/BGR.Console.Tests/Usings.cs | 12 +
src/BGR.Console/BGR.Console.csproj | 1 +
.../Common/HostBuilderExtensions.cs | 16 +-
src/BGR.Console/Logging/LoggerExtensions.cs | 63 +++++
src/BGR.Console/Program.cs | 201 +--------------
src/BGR.Console/Removal/IImage.cs | 3 +-
src/BGR.Console/Removal/IInferenceRunner.cs | 6 +
src/BGR.Console/Removal/IPixel.cs | 8 -
src/BGR.Console/Removal/ITensor.cs | 3 +-
src/BGR.Console/Removal/ImageProcessor.cs | 6 +-
.../Removal/ImageSharp/ImageSharpProcessor.cs | 32 ++-
.../Removal/ImageSharp/SharpImage.cs | 28 +-
.../Removal/ImageSharp/SharpPixel.cs | 10 -
.../Removal/Models/IModelFactory.cs | 6 +
src/BGR.Console/Removal/Models/ModNetModel.cs | 3 +-
src/BGR.Console/Removal/Models/Model.cs | 16 +-
.../Removal/Models/ModelFactory.cs | 21 ++
src/BGR.Console/Removal/Models/RmbgModel.cs | 4 +-
src/BGR.Console/Removal/Models/U2NetModel.cs | 5 +-
.../Removal/Onnx/OnnxInferenceRunner.cs | 18 ++
src/BGR.Console/Removal/Onnx/OnnxTensor.cs | 30 ++-
src/BGR.Console/Removal/RemovalCommand.cs | 139 ++++++++++
src/BGR.Console/Resources/ResourceManager.cs | 1 -
src/BGR.Console/Usings.cs | 7 +-
38 files changed, 1462 insertions(+), 253 deletions(-)
create mode 100644 src/BGR.Console.Tests/Integration/AppFactory.cs
create mode 100644 src/BGR.Console.Tests/Integration/RemovalCommandTests.cs
create mode 100644 src/BGR.Console.Tests/Unit/ImageSharpProcessorTests.cs
create mode 100644 src/BGR.Console.Tests/Unit/ModNetModelTests.cs
create mode 100644 src/BGR.Console.Tests/Unit/ModelFactoryTests.cs
create mode 100644 src/BGR.Console.Tests/Unit/ModelTests.cs
create mode 100644 src/BGR.Console.Tests/Unit/OnnxInferenceRunnerTests.cs
create mode 100644 src/BGR.Console.Tests/Unit/OnnxTensorTests.cs
create mode 100644 src/BGR.Console.Tests/Unit/RemovalCommandTests.cs
create mode 100644 src/BGR.Console.Tests/Unit/RmbgModelTests.cs
create mode 100644 src/BGR.Console.Tests/Unit/SharpImageTests.cs
create mode 100644 src/BGR.Console.Tests/Unit/U2NetModelTests.cs
create mode 100644 src/BGR.Console/Logging/LoggerExtensions.cs
delete mode 100644 src/BGR.Console/Removal/IPixel.cs
delete mode 100644 src/BGR.Console/Removal/ImageSharp/SharpPixel.cs
create mode 100644 src/BGR.Console/Removal/Models/IModelFactory.cs
create mode 100644 src/BGR.Console/Removal/Models/ModelFactory.cs
create mode 100644 src/BGR.Console/Removal/Onnx/OnnxInferenceRunner.cs
create mode 100644 src/BGR.Console/Removal/RemovalCommand.cs
diff --git a/.editorconfig b/.editorconfig
index 5d71be5..ee7b7fa 100644
--- a/.editorconfig
+++ b/.editorconfig
@@ -85,6 +85,8 @@ dotnet_diagnostic.CA1707.severity = none
dotnet_diagnostic.IDE0058.severity = none
dotnet_diagnostic.CA2007.severity = none
dotnet_diagnostic.CA1515.severity = none
+dotnet_diagnostic.IDE0100.severity = none
+dotnet_diagnostic.IDE0046.severity = none
# var preferences
csharp_style_var_elsewhere = true:suggestion
diff --git a/src/BGR.Console.Tests/BGR.Console.Tests.csproj b/src/BGR.Console.Tests/BGR.Console.Tests.csproj
index 633ac7a..34e37a0 100644
--- a/src/BGR.Console.Tests/BGR.Console.Tests.csproj
+++ b/src/BGR.Console.Tests/BGR.Console.Tests.csproj
@@ -11,9 +11,14 @@
runtime; build; native; contentfiles; analyzers; buildtransitive
all
+
+ runtime; build; native; contentfiles; analyzers; buildtransitive
+ all
+
+
runtime; build; native; contentfiles; analyzers; buildtransitive
@@ -25,14 +30,14 @@
true
- ./TestResults/coverage/
+ ./TestResults/Coverage/
cobertura
[BGR.Console]*
**/Program.cs
-
+
diff --git a/src/BGR.Console.Tests/Integration/AppFactory.cs b/src/BGR.Console.Tests/Integration/AppFactory.cs
new file mode 100644
index 0000000..91111d4
--- /dev/null
+++ b/src/BGR.Console.Tests/Integration/AppFactory.cs
@@ -0,0 +1,19 @@
+namespace BGR.Console.Tests.Integration;
+
+internal static class AppFactory
+{
+ public static CommandApp Create()
+ {
+ return Host.CreateDefaultBuilder()
+ .ConfigureLogging(static logging => logging.ClearProviders())
+ .ConfigureServices(static (_, services) =>
+ {
+ services.AddSingleton(new TestConsole());
+ services.AddSingleton();
+ services.AddSingleton();
+ services.AddSingleton();
+ services.AddSingleton();
+ })
+ .BuildApp();
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console.Tests/Integration/RemovalCommandTests.cs b/src/BGR.Console.Tests/Integration/RemovalCommandTests.cs
new file mode 100644
index 0000000..e14f24a
--- /dev/null
+++ b/src/BGR.Console.Tests/Integration/RemovalCommandTests.cs
@@ -0,0 +1,24 @@
+namespace BGR.Console.Tests.Integration;
+
+public class RemovalCommandTests
+{
+ private readonly CommandApp _app = AppFactory.Create();
+
+ [Fact]
+ public async Task RunAsync_WhenCalled_ItShouldRemoveImageBackground()
+ {
+ var imagePath = $"{Guid.NewGuid()}.png";
+ var outputPath = $"{Guid.NewGuid()}.png";
+
+ using var testImage = new Image(100, 100);
+ await testImage.SaveAsPngAsync(imagePath);
+
+ var result = await _app.RunAsync([imagePath, "--model", "u2net", "--output", outputPath]);
+
+ result.ShouldBe(0);
+ File.Exists(outputPath).ShouldBeTrue();
+
+ File.Delete(imagePath);
+ File.Delete(outputPath);
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console.Tests/Unit/ImageSharpProcessorTests.cs b/src/BGR.Console.Tests/Unit/ImageSharpProcessorTests.cs
new file mode 100644
index 0000000..7675f08
--- /dev/null
+++ b/src/BGR.Console.Tests/Unit/ImageSharpProcessorTests.cs
@@ -0,0 +1,241 @@
+namespace BGR.Console.Tests.Unit;
+
+public class ImageSharpProcessorTests : IDisposable
+{
+ private const string TestImagePath = "test.jpg";
+ private bool _isDisposed;
+ private readonly ImageSharpProcessor _sut = new();
+ private readonly Mock _modelMock = new();
+ private readonly Stream _testImageStream;
+
+ public ImageSharpProcessorTests()
+ {
+ if (File.Exists(TestImagePath) is false)
+ {
+ using var testImage = new Image(100, 100);
+ testImage.SaveAsJpeg(TestImagePath);
+ }
+
+ _modelMock.Setup(static x => x.InputWidth).Returns(320);
+ _modelMock.Setup(static x => x.InputHeight).Returns(320);
+
+ var stream = new MemoryStream();
+ using var image = new Image(100, 100);
+
+ for (var y = 0; y < image.Height; y++)
+ {
+ for (var x = 0; x < image.Width; x++)
+ {
+ image[x, y] = new Rgba32((byte)x, (byte)y, 128, 255);
+ }
+ }
+
+ image.SaveAsPng(stream);
+ stream.Position = 0;
+ _testImageStream = stream;
+ }
+
+ [Fact]
+ public async Task LoadImageAsync_WhenCalledWithValidPath_ItShouldReturnImage()
+ {
+ var result = await _sut.LoadImageAsync(TestImagePath);
+
+ result.ShouldBeOfType();
+ result.ShouldNotBeNull();
+ result.Width.ShouldBe(100);
+ result.Height.ShouldBe(100);
+ result.Data.Length.ShouldBeGreaterThan(0);
+ }
+
+ [Fact]
+ public async Task LoadImageAsync_WhenCalledWithValidPath_ItShouldReturnReusableStream()
+ {
+ var result = await _sut.LoadImageAsync(TestImagePath);
+
+ result.Data.Position.ShouldBe(0);
+ result.Data.CanRead.ShouldBeTrue();
+
+ var buffer = new byte[100];
+ await result.Data.ReadExactlyAsync(buffer);
+
+ result.Data.Position = 0;
+ await result.Data.ReadExactlyAsync(buffer);
+ }
+
+ [Fact]
+ public async Task CreateTensorInputAsync_WhenCalled_ItShouldResizeImageToModelDimensions()
+ {
+ const int modelWidth = 64;
+ const int modelHeight = 48;
+
+ _modelMock.Setup(static x => x.InputWidth).Returns(modelWidth);
+ _modelMock.Setup(static x => x.InputHeight).Returns(modelHeight);
+
+ var result = await _sut.CreateTensorInputAsync(_testImageStream, _modelMock.Object);
+
+ result.Width.ShouldBe(modelWidth);
+ result.Height.ShouldBe(modelHeight);
+ }
+
+ [Fact]
+ public async Task CreateTensorInputAsync_WhenCalled_ItShouldCreateTensorWithCorrectDimensions()
+ {
+ var result = await _sut.CreateTensorInputAsync(_testImageStream, _modelMock.Object);
+
+ result.ShouldBeOfType();
+ Should.NotThrow(() => result.GetValue(0, 2, 0, 0));
+ }
+
+ [Fact]
+ public async Task CreateTensorInputAsync_WhenCalled_ItShouldNormalizePixelValues()
+ {
+ var normalizedValue = 0.5f;
+ _modelMock.Setup(static x => x.NormalizeRed(It.IsAny())).Returns(normalizedValue);
+ _modelMock.Setup(static x => x.NormalizeGreen(It.IsAny())).Returns(normalizedValue);
+ _modelMock.Setup(static x => x.NormalizeBlue(It.IsAny())).Returns(normalizedValue);
+
+ var result = await _sut.CreateTensorInputAsync(_testImageStream, _modelMock.Object);
+
+ for (var y = 0; y < result.Height; y++)
+ {
+ for (var x = 0; x < result.Width; x++)
+ {
+ result.GetValue(0, 0, y, x).ShouldBe(normalizedValue); // Red
+ result.GetValue(0, 1, y, x).ShouldBe(normalizedValue); // Green
+ result.GetValue(0, 2, y, x).ShouldBe(normalizedValue); // Blue
+ }
+ }
+ }
+
+ [Fact]
+ public async Task CreateTensorInputAsync_WhenCalled_ItShouldCallNormalizeForEachChannel()
+ {
+ await _sut.CreateTensorInputAsync(_testImageStream, _modelMock.Object);
+
+ _modelMock.Verify(static x => x.NormalizeRed(It.IsAny()), Times.AtLeast(1));
+ _modelMock.Verify(static x => x.NormalizeGreen(It.IsAny()), Times.AtLeast(1));
+ _modelMock.Verify(static x => x.NormalizeBlue(It.IsAny()), Times.AtLeast(1));
+ }
+
+ [Fact]
+ public async Task GenerateMaskAsync_WhenCalled_ItShouldCreateMaskWithCorrectDimensions()
+ {
+ const int width = 64;
+ const int height = 48;
+
+ var tensor = new OnnxTensor(1, 1, height, width);
+ var stream = await _sut.GenerateMaskAsync(tensor, width, height);
+
+ using var mask = await Image.LoadAsync(stream);
+
+ mask.Width.ShouldBe(width);
+ mask.Height.ShouldBe(height);
+ }
+
+ [Fact]
+ public async Task GenerateMaskAsync_WhenCalled_ItShouldCreateGreyscaleMask()
+ {
+ var tensor = new OnnxTensor(1, 1, 100, 100);
+ var stream = await _sut.GenerateMaskAsync(tensor, 100, 100);
+
+ using var mask = await Image.LoadAsync(stream);
+
+ for (var y = 0; y < mask.Height; y++)
+ {
+ for (var x = 0; x < mask.Width; x++)
+ {
+ // we expect the mask to be greyscale
+ // so R, G, B should be equal
+ mask[x, y].R.ShouldBe(mask[x, y].G);
+ mask[x, y].G.ShouldBe(mask[x, y].B);
+ }
+ }
+ }
+
+ [Fact]
+ public async Task RemoveBackgroundAsync_WithValidImageAndMask_ShouldReturnProcessedStream()
+ {
+ var width = 2;
+ var height = 2;
+
+ using var imageStream = new MemoryStream();
+ using var image = new Image(width, height);
+
+ for (var x = 0; x < width; x++)
+ {
+ for (var y = 0; y < height; y++)
+ {
+ image[x, y] = new Rgba32(255, 0, 0, 255); // Red pixels
+ }
+ }
+
+ await image.SaveAsPngAsync(imageStream);
+ imageStream.Position = 0;
+
+ using var maskStream = new MemoryStream();
+ using var mask = new Image(width, height);
+
+ mask[0, 0] = new Rgba32(0, 0, 0, 255);
+ mask[0, 1] = new Rgba32(255, 255, 255, 255);
+ mask[1, 0] = new Rgba32(255, 255, 255, 255);
+ mask[1, 1] = new Rgba32(255, 255, 255, 255);
+
+ await mask.SaveAsPngAsync(maskStream);
+ maskStream.Position = 0;
+
+ var result = await _sut.RemoveBackgroundAsync(imageStream, maskStream);
+
+ result.ShouldNotBeNull();
+ result.Length.ShouldBeGreaterThan(0);
+
+ result.Position = 0;
+ using var resultImage = await Image.LoadAsync(result);
+ resultImage.Width.ShouldBe(width);
+ resultImage.Height.ShouldBe(height);
+
+ resultImage[0, 0].A.ShouldBe((byte)0);
+
+ resultImage[0, 1].R.ShouldBe((byte)255);
+ resultImage[0, 1].A.ShouldBe((byte)255);
+ resultImage[1, 0].R.ShouldBe((byte)255);
+ resultImage[1, 0].A.ShouldBe((byte)255);
+ resultImage[1, 1].R.ShouldBe((byte)255);
+ resultImage[1, 1].A.ShouldBe((byte)255);
+ }
+
+ [Fact]
+ public async Task SaveImageAsync_WhenCalled_ItShouldSaveImageToDiskAtProvidedPath()
+ {
+ using var image = new Image(100, 100);
+ var stream = new MemoryStream();
+ await image.SaveAsPngAsync(stream);
+
+ var path = $"{Guid.NewGuid()}.png";
+ await _sut.SaveImageAsync(stream, path);
+
+ File.Exists(path).ShouldBeTrue();
+ File.Delete(path);
+ }
+
+ public void Dispose()
+ {
+ Dispose(true);
+ GC.SuppressFinalize(this);
+ }
+
+ protected virtual void Dispose(bool disposing)
+ {
+ if (_isDisposed)
+ {
+ return;
+ }
+
+ if (disposing)
+ {
+ File.Delete(TestImagePath);
+ _testImageStream.Dispose();
+ }
+
+ _isDisposed = true;
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console.Tests/Unit/ModNetModelTests.cs b/src/BGR.Console.Tests/Unit/ModNetModelTests.cs
new file mode 100644
index 0000000..143caa7
--- /dev/null
+++ b/src/BGR.Console.Tests/Unit/ModNetModelTests.cs
@@ -0,0 +1,66 @@
+namespace BGR.Console.Tests.Unit;
+
+public class ModNetModelTests
+{
+ private readonly byte[] _sampleModelBytes = [0x01, 0x02, 0x03];
+ private readonly ModNetModel _sut;
+
+ public ModNetModelTests()
+ {
+ _sut = new ModNetModel(_sampleModelBytes);
+ }
+
+ [Fact]
+ public void Id_WhenCalled_ItShouldReturnModNetId()
+ {
+ ModNetModel.Id.ShouldBe("modnet");
+ }
+
+ [Fact]
+ public void InputWidth_WhenCalled_ItShouldReturn512()
+ {
+ _sut.InputWidth.ShouldBe(512);
+ }
+
+ [Fact]
+ public void InputHeight_WhenCalled_ItShouldReturn512()
+ {
+ _sut.InputHeight.ShouldBe(512);
+ }
+
+ [Fact]
+ public void RedNormalizationMean_WhenCalled_ItShouldReturnCorrectValue()
+ {
+ _sut.RedNormalizationMean.ShouldBe(0.485f);
+ }
+
+ [Fact]
+ public void GreenNormalizationMean_WhenCalled_ItShouldReturnCorrectValue()
+ {
+ _sut.GreenNormalizationMean.ShouldBe(0.456f);
+ }
+
+ [Fact]
+ public void BlueNormalizationMean_WhenCalled_ItShouldReturnCorrectValue()
+ {
+ _sut.BlueNormalizationMean.ShouldBe(0.406f);
+ }
+
+ [Fact]
+ public void RedNormalizationStd_WhenCalled_ItShouldReturnCorrectValue()
+ {
+ _sut.RedNormalizationStd.ShouldBe(0.229f);
+ }
+
+ [Fact]
+ public void GreenNormalizationStd_WhenCalled_ItShouldReturnCorrectValue()
+ {
+ _sut.GreenNormalizationStd.ShouldBe(0.224f);
+ }
+
+ [Fact]
+ public void BlueNormalizationStd_WhenCalled_ItShouldReturnCorrectValue()
+ {
+ _sut.BlueNormalizationStd.ShouldBe(0.225f);
+ }
+}
diff --git a/src/BGR.Console.Tests/Unit/ModelFactoryTests.cs b/src/BGR.Console.Tests/Unit/ModelFactoryTests.cs
new file mode 100644
index 0000000..209ad1d
--- /dev/null
+++ b/src/BGR.Console.Tests/Unit/ModelFactoryTests.cs
@@ -0,0 +1,67 @@
+namespace BGR.Console.Tests.Unit;
+
+public class ModelFactoryTests
+{
+ private readonly Mock _resourceManagerMock;
+ private readonly ModelFactory _sut;
+ private readonly byte[] _sampleModelBytes = [0x01, 0x02, 0x03];
+
+ public ModelFactoryTests()
+ {
+ _resourceManagerMock = new Mock();
+ _sut = new ModelFactory(_resourceManagerMock.Object);
+ }
+
+ [Theory]
+ [InlineData("u2net.onnx", typeof(U2NetModel))]
+ [InlineData("rmbg.onnx", typeof(RmbgModel))]
+ [InlineData("modnet.onnx", typeof(ModNetModel))]
+ public void Create_WhenCalledWithValidModelName_ItShouldReturnCorrectModelType(string resourceName, Type expectedType)
+ {
+ SetupResourceManagerMock(resourceName);
+
+ var result = _sut.Create(resourceName);
+
+ result.ShouldBeOfType(expectedType);
+ VerifyResourceManagerCalled(resourceName);
+ }
+
+ [Fact]
+ public void Create_WhenCalledWithUnknownModel_ItShouldThrowArgumentException()
+ {
+ var resourceName = "unknown.onnx";
+ SetupResourceManagerMock(resourceName);
+
+ var exception = Should.Throw(() => _sut.Create(resourceName));
+
+ exception.Message.ShouldBe($"Unknown model name: {resourceName}");
+ VerifyResourceManagerCalled(resourceName);
+ }
+
+ [Fact]
+ public void Create_WhenCalledWithU2NetModel_ItShouldReadTheResourceStreamToTheEnd()
+ {
+ var resourceName = $"{U2NetModel.Id}.onnx";
+ var memoryStream = new MemoryStream(_sampleModelBytes);
+
+ _resourceManagerMock.Setup(x => x.GetResource(resourceName))
+ .Returns(memoryStream);
+
+ _sut.Create(resourceName);
+
+ memoryStream.Position.ShouldBe(memoryStream.Length);
+ }
+
+ private void SetupResourceManagerMock(string resourceName)
+ {
+ var memoryStream = new MemoryStream(_sampleModelBytes);
+
+ _resourceManagerMock.Setup(x => x.GetResource(resourceName))
+ .Returns(memoryStream);
+ }
+
+ private void VerifyResourceManagerCalled(string resourceName)
+ {
+ _resourceManagerMock.Verify(x => x.GetResource(resourceName), Times.Once);
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console.Tests/Unit/ModelTests.cs b/src/BGR.Console.Tests/Unit/ModelTests.cs
new file mode 100644
index 0000000..ff6d839
--- /dev/null
+++ b/src/BGR.Console.Tests/Unit/ModelTests.cs
@@ -0,0 +1,94 @@
+namespace BGR.Console.Tests.Unit;
+
+public class ModelTests
+{
+ internal sealed class TestModel(byte[] modelBytes) : Model(modelBytes)
+ {
+ public override int InputWidth => 100;
+ public override int InputHeight => 100;
+ public override float RedNormalizationMean => 0.5f;
+ public override float GreenNormalizationMean => 0.5f;
+ public override float BlueNormalizationMean => 0.5f;
+ public override float RedNormalizationStd => 0.25f;
+ public override float GreenNormalizationStd => 0.25f;
+ public override float BlueNormalizationStd => 0.25f;
+ }
+
+ private readonly byte[] _sampleBytes = [0x01, 0x02, 0x03];
+ private readonly TestModel _sut;
+ private const float Delta = 0.00001f;
+
+ public ModelTests()
+ {
+ _sut = new TestModel(_sampleBytes);
+ }
+
+ [Theory]
+ [InlineData(0f)]
+ [InlineData(127.5f)]
+ [InlineData(255f)]
+ public void NormalizeRed_WhenCalled_ItShouldNormalizeCorrectly(float value)
+ {
+ var result = _sut.NormalizeRed(value);
+
+ var expected = ((value / 255f) - _sut.RedNormalizationMean) / _sut.RedNormalizationStd;
+ result.ShouldBe(expected, Delta);
+ }
+
+ [Theory]
+ [InlineData(0f)]
+ [InlineData(127.5f)]
+ [InlineData(255f)]
+ public void NormalizeGreen_WhenCalled_ItShouldNormalizeCorrectly(float value)
+ {
+ var result = _sut.NormalizeGreen(value);
+
+ var expected = ((value / 255f) - _sut.GreenNormalizationMean) / _sut.GreenNormalizationStd;
+ result.ShouldBe(expected, Delta);
+ }
+
+ [Theory]
+ [InlineData(0f)]
+ [InlineData(127.5f)]
+ [InlineData(255f)]
+ public void NormalizeBlue_WhenCalled_ItShouldNormalizeCorrectly(float value)
+ {
+ var result = _sut.NormalizeBlue(value);
+
+ var expected = ((value / 255f) - _sut.BlueNormalizationMean) / _sut.BlueNormalizationStd;
+ result.ShouldBe(expected, Delta);
+ }
+
+ [Theory]
+ [InlineData(-1f)]
+ [InlineData(256f)]
+ public void NormalizeRed_WhenCalledWithOutOfRangeValues_ItShouldStillNormalize(float value)
+ {
+ var result = _sut.NormalizeRed(value);
+
+ var expected = ((value / 255f) - _sut.RedNormalizationMean) / _sut.RedNormalizationStd;
+ result.ShouldBe(expected, Delta);
+ }
+
+ [Theory]
+ [InlineData(-1f)]
+ [InlineData(256f)]
+ public void NormalizeGreen_WhenCalledWithOutOfRangeValues_ItShouldStillNormalize(float value)
+ {
+ var result = _sut.NormalizeGreen(value);
+
+ var expected = ((value / 255f) - _sut.GreenNormalizationMean) / _sut.GreenNormalizationStd;
+ result.ShouldBe(expected, Delta);
+ }
+
+ [Theory]
+ [InlineData(-1f)]
+ [InlineData(256f)]
+ public void NormalizeBlue_WhenCalledWithOutOfRangeValues_ItShouldStillNormalize(float value)
+ {
+ var result = _sut.NormalizeBlue(value);
+
+ var expected = ((value / 255f) - _sut.BlueNormalizationMean) / _sut.BlueNormalizationStd;
+ result.ShouldBe(expected, Delta);
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console.Tests/Unit/OnnxInferenceRunnerTests.cs b/src/BGR.Console.Tests/Unit/OnnxInferenceRunnerTests.cs
new file mode 100644
index 0000000..ee9faae
--- /dev/null
+++ b/src/BGR.Console.Tests/Unit/OnnxInferenceRunnerTests.cs
@@ -0,0 +1,34 @@
+
+using Microsoft.ML.OnnxRuntime.Tensors;
+
+namespace BGR.Console.Tests.Unit;
+
+public class OnnxInferenceRunnerTests
+{
+ private const string TestModel = "u2net.onnx";
+ private readonly ResourceManager _resourceManager = new();
+ private readonly OnnxInferenceRunner _sut = new();
+ private readonly byte[] _sampleModelBytes;
+ private readonly Mock> _mockInputTensor = new();
+
+ public OnnxInferenceRunnerTests()
+ {
+ var modelStream = _resourceManager.GetResource(TestModel);
+ var modelBytes = new byte[modelStream.Length];
+ modelStream.ReadExactly(modelBytes);
+ _sampleModelBytes = modelBytes;
+
+ _mockInputTensor
+ .Setup(static x => x.ToTensor())
+ .Returns(new DenseTensor([1, 3, 320, 320]));
+ }
+
+ [Fact]
+ public void Run_WhenCalledWithValidInput_ItShouldReturnOutput()
+ {
+ var result = _sut.Run(_sampleModelBytes, _mockInputTensor.Object);
+
+ result.ShouldBeOfType();
+ result.ShouldNotBeNull();
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console.Tests/Unit/OnnxTensorTests.cs b/src/BGR.Console.Tests/Unit/OnnxTensorTests.cs
new file mode 100644
index 0000000..aaddc61
--- /dev/null
+++ b/src/BGR.Console.Tests/Unit/OnnxTensorTests.cs
@@ -0,0 +1,121 @@
+using BGR.Console.Removal.Onnx;
+
+using Microsoft.ML.OnnxRuntime.Tensors;
+
+namespace BGR.Console.Tests.Unit;
+
+public class OnnxTensorTests
+{
+ private const int DefaultBatchSize = 1;
+ private const int DefaultChannels = 3;
+ private const int DefaultHeight = 4;
+ private const int DefaultWidth = 5;
+
+ [Fact]
+ public void Constructor_WhenCalledWithDimensions_ItShouldCreateTensorWithCorrectDimensions()
+ {
+ var tensor = new OnnxTensor(DefaultBatchSize, DefaultChannels, DefaultHeight, DefaultWidth);
+
+ tensor.Height.ShouldBe(DefaultHeight);
+ tensor.Width.ShouldBe(DefaultWidth);
+ }
+
+ [Fact]
+ public void Constructor_WhenCalledWithExistingTensor_ItShouldCreateTensorWithSameDimensions()
+ {
+ var existingTensor = new DenseTensor([DefaultBatchSize, DefaultChannels, DefaultHeight, DefaultWidth]);
+
+ var tensor = new OnnxTensor(existingTensor);
+
+ tensor.Height.ShouldBe(DefaultHeight);
+ tensor.Width.ShouldBe(DefaultWidth);
+ }
+
+ [Theory]
+ [InlineData(0, 0, 0, 0, 1.0f)]
+ [InlineData(0, 1, 2, 3, 2.5f)]
+ [InlineData(0, 2, 3, 4, -1.0f)]
+ public void SetValue_WhenCalledWithValidValues_ItShouldSetCorrectValueAtPosition(int batch, int channel, int y, int x, float expectedValue)
+ {
+ var tensor = new OnnxTensor(DefaultBatchSize, DefaultChannels, DefaultHeight, DefaultWidth);
+
+ tensor.SetValue(batch, channel, y, x, expectedValue);
+ var actualValue = tensor.GetValue(batch, channel, y, x);
+
+ actualValue.ShouldBe(expectedValue);
+ }
+
+ [Theory]
+ [InlineData(0, 0, 0, 0, 1.0f)]
+ [InlineData(0, 1, 2, 3, 2.5f)]
+ [InlineData(0, 2, 3, 4, -1.0f)]
+ public void GetValue_WhenCalledWithValidValues_ItShouldReturnCorrectValue(int batch, int channel, int y, int x, float value)
+ {
+ var tensor = new OnnxTensor(DefaultBatchSize, DefaultChannels, DefaultHeight, DefaultWidth);
+ tensor.SetValue(batch, channel, y, x, value);
+
+ var result = tensor.GetValue(batch, channel, y, x);
+
+ result.ShouldBe(value);
+ }
+
+ [Fact]
+ public void ToTensor_WhenCalled_ItShouldReturnUnderlyingTensor()
+ {
+ var tensor = new OnnxTensor(DefaultBatchSize, DefaultChannels, DefaultHeight, DefaultWidth);
+ const float testValue = 42.0f;
+ tensor.SetValue(0, 0, 0, 0, testValue);
+
+ var result = tensor.ToTensor();
+
+ result.ShouldBeOfType>();
+ result[0, 0, 0, 0].ShouldBe(testValue);
+ }
+
+ [Theory]
+ [InlineData(-1, 0, 0, 0)]
+ [InlineData(1, 3, 0, 0)]
+ [InlineData(0, 0, 4, 0)]
+ [InlineData(0, 0, 0, 5)]
+ public void SetValue_WhenCalledWithInvalidIndices_ItShouldThrowIndexOutOfRangeException(int batch, int channel, int y, int x)
+ {
+ var tensor = new OnnxTensor(DefaultBatchSize, DefaultChannels, DefaultHeight, DefaultWidth);
+
+ try
+ {
+ tensor.SetValue(batch, channel, y, x, 0);
+ }
+ catch (IndexOutOfRangeException ex)
+ {
+ ex.Message.ShouldBe("Index was outside the bounds of the array.");
+ }
+ }
+
+ [Theory]
+ [InlineData(-1, 0, 0, 0)]
+ [InlineData(1, 3, 0, 0)]
+ [InlineData(0, 0, 4, 0)]
+ [InlineData(0, 0, 0, 5)]
+ public void GetValue_WhenCalledWithInvalidIndices_ItShouldThrowIndexOutOfRangeException(int batch, int channel, int y, int x)
+ {
+ var tensor = new OnnxTensor(DefaultBatchSize, DefaultChannels, DefaultHeight, DefaultWidth);
+
+ try
+ {
+ tensor.GetValue(batch, channel, y, x);
+ }
+ catch (IndexOutOfRangeException ex)
+ {
+ ex.Message.ShouldBe("Index was outside the bounds of the array.");
+ }
+ }
+
+ [Fact]
+ public void Constructor_WhenCalledWithZeroDimension_ItShouldCreateEmptyTensor()
+ {
+ var tensor = new OnnxTensor(0, 0, 0, 0);
+
+ tensor.Height.ShouldBe(0);
+ tensor.Width.ShouldBe(0);
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console.Tests/Unit/RemovalCommandTests.cs b/src/BGR.Console.Tests/Unit/RemovalCommandTests.cs
new file mode 100644
index 0000000..a65ca11
--- /dev/null
+++ b/src/BGR.Console.Tests/Unit/RemovalCommandTests.cs
@@ -0,0 +1,181 @@
+using Microsoft.Extensions.Logging;
+
+using Spectre.Console.Cli;
+
+namespace BGR.Console.Tests.Unit;
+
+public class RemovalCommandTests : IDisposable
+{
+ private bool _isDisposed;
+ private readonly Mock _modelFactoryMock = new();
+ private readonly Mock _imageProcessorMock = new();
+ private readonly Mock _inferenceRunnerMock = new();
+ private readonly TestConsole _console = new();
+ private readonly Mock> _loggerMock = new();
+ private readonly RemovalCommand _sut;
+
+ public RemovalCommandTests()
+ {
+ _sut = new RemovalCommand(
+ _modelFactoryMock.Object,
+ _imageProcessorMock.Object,
+ _inferenceRunnerMock.Object,
+ _console,
+ _loggerMock.Object
+ );
+ }
+
+ [Theory]
+ [InlineData("output.png", false)]
+ [InlineData("", true)]
+ public async Task ExecuteAsync_WhenCalled_ItShouldProcessImage(string outputPath, bool includeMask)
+ {
+ var imagePath = $"{Guid.NewGuid()}.png";
+
+ using var testImage = new Image(100, 100);
+ await testImage.SaveAsPngAsync(imagePath);
+
+ var resourceName = "u2net.onnx";
+
+ var settings = new RemovalCommand.Settings()
+ {
+ Image = imagePath,
+ Model = "u2net",
+ IncludeMask = includeMask,
+ Output = outputPath,
+ };
+
+ var model = new U2NetModel([1, 2, 3]);
+ var image = new SharpImage(100, 100, new MemoryStream());
+ var inputTensor = new OnnxTensor(1, 3, 320, 320);
+ var outputTensor = new OnnxTensor(1, 1, 320, 320);
+ var maskStream = new MemoryStream();
+ var outputStream = new MemoryStream();
+
+ _modelFactoryMock
+ .Setup(m => m.Create(resourceName))
+ .Returns(model);
+
+ _imageProcessorMock
+ .Setup(p => p.LoadImageAsync(imagePath))
+ .ReturnsAsync(image);
+
+ _imageProcessorMock
+ .Setup(p => p.CreateTensorInputAsync(image.Data, model))
+ .ReturnsAsync(inputTensor);
+
+ _inferenceRunnerMock
+ .Setup(r => r.Run(model.Bytes, inputTensor))
+ .Returns(outputTensor);
+
+ _imageProcessorMock
+ .Setup(p => p.GenerateMaskAsync(outputTensor, image.Width, image.Height))
+ .ReturnsAsync(maskStream);
+
+ _imageProcessorMock
+ .Setup(p => p.RemoveBackgroundAsync(image.Data, maskStream))
+ .ReturnsAsync(outputStream);
+
+ var commandContext = new CommandContext(
+ [],
+ new Mock().Object,
+ "test",
+ new object()
+ );
+
+ var result = await _sut.ExecuteAsync(commandContext, settings);
+
+ result.ShouldBe(0);
+
+ _modelFactoryMock.Verify(m => m.Create(resourceName), Times.Once);
+ _imageProcessorMock.Verify(p => p.LoadImageAsync(imagePath), Times.Once);
+ _imageProcessorMock.Verify(p => p.CreateTensorInputAsync(image.Data, model), Times.Once);
+ _inferenceRunnerMock.Verify(r => r.Run(model.Bytes, inputTensor), Times.Once);
+ _imageProcessorMock.Verify(p => p.GenerateMaskAsync(outputTensor, image.Width, image.Height), Times.Once);
+ _imageProcessorMock.Verify(p => p.RemoveBackgroundAsync(image.Data, maskStream), Times.Once);
+ _imageProcessorMock.Verify(p => p.SaveImageAsync(outputStream, It.IsAny()), Times.AtLeastOnce);
+
+ File.Delete(imagePath);
+ }
+
+ [Fact]
+ public void Validate_WhenCalledAndFileDoesNotExist_ItShouldReturnError()
+ {
+ var settings = new RemovalCommand.Settings()
+ {
+ Image = "invalid.png",
+ Model = "u2net",
+ IncludeMask = false,
+ Output = "output.png",
+ };
+
+ var result = settings.Validate();
+
+ result.Successful.ShouldBeFalse();
+ }
+
+ [Fact]
+ public void Validate_WhenCalledAndModelIsInvalid_ItShouldReturnError()
+ {
+ var imagePath = $"{Guid.NewGuid()}.png";
+
+ using var testImage = new Image(100, 100);
+ testImage.SaveAsPng(imagePath);
+
+ var settings = new RemovalCommand.Settings()
+ {
+ Image = imagePath,
+ Model = "invalid",
+ IncludeMask = false,
+ Output = "output.png",
+ };
+
+ var result = settings.Validate();
+
+ result.Successful.ShouldBeFalse();
+ }
+
+ [Fact]
+ public void Validate_WhenCalledAndSettingsValid_ItShouldReturnSuccess()
+ {
+ var imagePath = $"{Guid.NewGuid()}.png";
+
+ using var testImage = new Image(100, 100);
+ testImage.SaveAsPng(imagePath);
+
+ var settings = new RemovalCommand.Settings()
+ {
+ Image = imagePath,
+ Model = "u2net",
+ IncludeMask = false,
+ Output = "output.png",
+ };
+
+ var result = settings.Validate();
+
+ result.Successful.ShouldBeTrue();
+
+ File.Delete(imagePath);
+ }
+
+ public void Dispose()
+ {
+ Dispose(true);
+ GC.SuppressFinalize(this);
+ }
+
+ protected virtual void Dispose(bool disposing)
+ {
+ if (_isDisposed)
+ {
+ return;
+ }
+
+ if (disposing)
+ {
+ _console.Dispose();
+ }
+
+ _isDisposed = true;
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console.Tests/Unit/RmbgModelTests.cs b/src/BGR.Console.Tests/Unit/RmbgModelTests.cs
new file mode 100644
index 0000000..b60c662
--- /dev/null
+++ b/src/BGR.Console.Tests/Unit/RmbgModelTests.cs
@@ -0,0 +1,66 @@
+namespace BGR.Console.Tests.Unit;
+
+public class RmbgModelTests
+{
+ private readonly byte[] _sampleModelBytes = [0x01, 0x02, 0x03];
+ private readonly RmbgModel _sut;
+
+ public RmbgModelTests()
+ {
+ _sut = new RmbgModel(_sampleModelBytes);
+ }
+
+ [Fact]
+ public void Id_WhenCalled_ItShouldReturnRmbgId()
+ {
+ RmbgModel.Id.ShouldBe("rmbg");
+ }
+
+ [Fact]
+ public void InputWidth_WhenCalled_ItShouldReturn1024()
+ {
+ _sut.InputWidth.ShouldBe(1024);
+ }
+
+ [Fact]
+ public void InputHeight_WhenCalled_ItShouldReturn1024()
+ {
+ _sut.InputHeight.ShouldBe(1024);
+ }
+
+ [Fact]
+ public void RedNormalizationMean_WhenCalled_ItShouldHaveCorrectValue()
+ {
+ _sut.RedNormalizationMean.ShouldBe(0.485f);
+ }
+
+ [Fact]
+ public void GreenNormalizationMean_WhenCalled_ItShouldHaveCorrectValue()
+ {
+ _sut.GreenNormalizationMean.ShouldBe(0.456f);
+ }
+
+ [Fact]
+ public void BlueNormalizationMean_WhenCalled_ItShouldHaveCorrectValue()
+ {
+ _sut.BlueNormalizationMean.ShouldBe(0.406f);
+ }
+
+ [Fact]
+ public void RedNormalizationStd_WhenCalled_ItShouldHaveCorrectValue()
+ {
+ _sut.RedNormalizationStd.ShouldBe(0.229f);
+ }
+
+ [Fact]
+ public void GreenNormalizationStd_WhenCalled_ItShouldHaveCorrectValue()
+ {
+ _sut.GreenNormalizationStd.ShouldBe(0.224f);
+ }
+
+ [Fact]
+ public void BlueNormalizationStd_WhenCalled_ItShouldHaveCorrectValue()
+ {
+ _sut.BlueNormalizationStd.ShouldBe(0.225f);
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console.Tests/Unit/SharpImageTests.cs b/src/BGR.Console.Tests/Unit/SharpImageTests.cs
new file mode 100644
index 0000000..6c605d2
--- /dev/null
+++ b/src/BGR.Console.Tests/Unit/SharpImageTests.cs
@@ -0,0 +1,86 @@
+namespace BGR.Console.Tests.Unit;
+
+public class SharpImageTests : IDisposable
+{
+ private bool _isDisposed;
+ private const int DefaultWidth = 100;
+ private const int DefaultHeight = 200;
+ private readonly MemoryStream _sampleStream;
+
+ public SharpImageTests()
+ {
+ _sampleStream = new MemoryStream([0x01, 0x02, 0x03]);
+ }
+
+ [Fact]
+ public void Constructor_WhenCalled_ItShouldSetProperties()
+ {
+ var image = new SharpImage(DefaultWidth, DefaultHeight, _sampleStream);
+
+ image.Width.ShouldBe(DefaultWidth);
+ image.Height.ShouldBe(DefaultHeight);
+ image.Data.ShouldBe(_sampleStream);
+ }
+
+ [Theory]
+ [InlineData(1, 1)]
+ [InlineData(1920, 1080)]
+ [InlineData(int.MaxValue, int.MaxValue)]
+ public void Constructor_WhenCalledWithDifferentDimensions_ItShouldSetCorrectValues(int width, int height)
+ {
+ var image = new SharpImage(width, height, _sampleStream);
+
+ image.Width.ShouldBe(width);
+ image.Height.ShouldBe(height);
+ }
+
+ [Fact]
+ public void Constructor_WhenCalledWithNullStream_ItShouldThrowException()
+ {
+ var action = static () => new SharpImage(DefaultWidth, DefaultHeight, null!);
+
+ action.ShouldThrow();
+ }
+
+ [Theory]
+ [InlineData(0, 0)]
+ [InlineData(-1, -1)]
+ [InlineData(1, -1)]
+ [InlineData(-1, 1)]
+ public void Constructor_WhenCalledWithInvalidDimensions_ItShouldThrowException(int width, int height)
+ {
+ var action = () => new SharpImage(width, height, _sampleStream);
+
+ action.ShouldThrow();
+ }
+
+ [Fact]
+ public void Data_WhenCalled_ItShouldReturnSameStreamInstance()
+ {
+ var stream = new MemoryStream();
+ var image = new SharpImage(DefaultWidth, DefaultHeight, stream);
+
+ image.Data.ShouldBeSameAs(stream);
+ }
+
+ public void Dispose()
+ {
+ Dispose(true);
+ GC.SuppressFinalize(this);
+ }
+
+ protected virtual void Dispose(bool disposing)
+ {
+ if (_isDisposed)
+ {
+ return;
+ }
+
+ if (disposing)
+ {
+ _sampleStream.Dispose();
+ }
+
+ _isDisposed = true;
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console.Tests/Unit/U2NetModelTests.cs b/src/BGR.Console.Tests/Unit/U2NetModelTests.cs
new file mode 100644
index 0000000..ea09249
--- /dev/null
+++ b/src/BGR.Console.Tests/Unit/U2NetModelTests.cs
@@ -0,0 +1,66 @@
+namespace BGR.Console.Tests.Unit;
+
+public class U2NetModelTests
+{
+ private readonly byte[] _sampleModelBytes = [0x01, 0x02, 0x03];
+ private readonly U2NetModel _sut;
+
+ public U2NetModelTests()
+ {
+ _sut = new U2NetModel(_sampleModelBytes);
+ }
+
+ [Fact]
+ public void Id_WhenCalled_ItShouldReturnU2NetId()
+ {
+ U2NetModel.Id.ShouldBe("u2net");
+ }
+
+ [Fact]
+ public void InputWidth_WhenCalled_ItShouldReturn320()
+ {
+ _sut.InputWidth.ShouldBe(320);
+ }
+
+ [Fact]
+ public void InputHeight_WhenCalled_ItShouldReturn320()
+ {
+ _sut.InputHeight.ShouldBe(320);
+ }
+
+ [Fact]
+ public void RedNormalizationMean_WhenCalled_ItShouldHaveCorrectValue()
+ {
+ _sut.RedNormalizationMean.ShouldBe(0.485f);
+ }
+
+ [Fact]
+ public void GreenNormalizationMean_WhenCalled_ItShouldHaveCorrectValue()
+ {
+ _sut.GreenNormalizationMean.ShouldBe(0.456f);
+ }
+
+ [Fact]
+ public void BlueNormalizationMean_WhenCalled_ItShouldHaveCorrectValue()
+ {
+ _sut.BlueNormalizationMean.ShouldBe(0.406f);
+ }
+
+ [Fact]
+ public void RedNormalizationStd_WhenCalled_ItShouldHaveCorrectValue()
+ {
+ _sut.RedNormalizationStd.ShouldBe(0.229f);
+ }
+
+ [Fact]
+ public void GreenNormalizationStd_WhenCalled_ItShouldHaveCorrectValue()
+ {
+ _sut.GreenNormalizationStd.ShouldBe(0.224f);
+ }
+
+ [Fact]
+ public void BlueNormalizationStd_WhenCalled_ItShouldHaveCorrectValue()
+ {
+ _sut.BlueNormalizationStd.ShouldBe(0.225f);
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console.Tests/Usings.cs b/src/BGR.Console.Tests/Usings.cs
index 42140b5..94d985a 100644
--- a/src/BGR.Console.Tests/Usings.cs
+++ b/src/BGR.Console.Tests/Usings.cs
@@ -1,7 +1,19 @@
global using BGR.Console.Common;
+global using BGR.Console.Removal;
+global using BGR.Console.Removal.ImageSharp;
+global using BGR.Console.Removal.Models;
+global using BGR.Console.Removal.Onnx;
global using BGR.Console.Resources;
global using Microsoft.Extensions.DependencyInjection;
global using Microsoft.Extensions.Hosting;
+global using Microsoft.Extensions.Logging;
global using Moq;
+
+global using SixLabors.ImageSharp;
+global using SixLabors.ImageSharp.PixelFormats;
+
+global using Spectre.Console;
+global using Spectre.Console.Cli;
+global using Spectre.Console.Testing;
diff --git a/src/BGR.Console/BGR.Console.csproj b/src/BGR.Console/BGR.Console.csproj
index 771c83f..42c041e 100644
--- a/src/BGR.Console/BGR.Console.csproj
+++ b/src/BGR.Console/BGR.Console.csproj
@@ -6,6 +6,7 @@
+
diff --git a/src/BGR.Console/Common/HostBuilderExtensions.cs b/src/BGR.Console/Common/HostBuilderExtensions.cs
index b538825..dc69995 100644
--- a/src/BGR.Console/Common/HostBuilderExtensions.cs
+++ b/src/BGR.Console/Common/HostBuilderExtensions.cs
@@ -2,10 +2,22 @@ namespace BGR.Console.Common;
internal static class HostBuilderExtensions
{
- public static CommandApp BuildApp(this IHostBuilder builder)
+ public static CommandApp BuildApp(this IHostBuilder builder)
{
var registrar = new TypeRegistrar(builder);
- var app = new CommandApp(registrar);
+ var app = new CommandApp(registrar);
+
+ app.Configure(static c =>
+ c.SetExceptionHandler(static (ex, resolver) =>
+ {
+ var logger = resolver?.Resolve(typeof(ILogger)) as ILogger;
+ logger?.RemovalCommandFailed(ex);
+
+ var console = resolver?.Resolve(typeof(IAnsiConsole)) as IAnsiConsole;
+ console?.WriteLine($"[red]An error occurred while executing the command:[/]");
+ console?.WriteException(ex, ExceptionFormats.ShortenEverything);
+ })
+ );
return app;
}
diff --git a/src/BGR.Console/Logging/LoggerExtensions.cs b/src/BGR.Console/Logging/LoggerExtensions.cs
new file mode 100644
index 0000000..49a897c
--- /dev/null
+++ b/src/BGR.Console/Logging/LoggerExtensions.cs
@@ -0,0 +1,63 @@
+using System.Diagnostics;
+
+using ILogger = Microsoft.Extensions.Logging.ILogger;
+
+namespace BGR.Console.Logging;
+
+internal static class LoggerExtensions
+{
+ private static readonly Action RemovalCommandFailedMsg = LoggerMessage.Define(
+ LogLevel.Error,
+ new EventId(0, nameof(RemovalCommandFailed)),
+ "An error occurred while executing the command."
+ );
+
+ private static readonly Action TimeAndLogActionMsg = LoggerMessage.Define(
+ LogLevel.Information,
+ new EventId(0, nameof(TimeAndLogAction)),
+ "{Message} in {ElapsedMilliseconds}ms"
+ );
+
+ public static void RemovalCommandFailed(this ILogger logger, Exception ex)
+ {
+ RemovalCommandFailedMsg(logger, ex);
+ }
+
+ public static async Task TimeAndLogActionAsync(this ILogger logger, string message, Func action)
+ {
+ var sw = new Stopwatch();
+ sw.Start();
+ await action();
+ sw.Stop();
+ TimeAndLogActionMsg(logger, message, sw.ElapsedMilliseconds, default!);
+ }
+
+ public static async Task TimeAndLogActionAsync(this ILogger logger, string message, Func> action)
+ {
+ var sw = new Stopwatch();
+ sw.Start();
+ var result = await action();
+ sw.Stop();
+ TimeAndLogActionMsg(logger, message, sw.ElapsedMilliseconds, default!);
+ return result;
+ }
+
+ public static void TimeAndLogAction(this ILogger logger, string message, Action action)
+ {
+ var sw = new Stopwatch();
+ sw.Start();
+ action();
+ sw.Stop();
+ TimeAndLogActionMsg(logger, message, sw.ElapsedMilliseconds, default!);
+ }
+
+ public static T TimeAndLogAction(this ILogger logger, string message, Func action)
+ {
+ var sw = new Stopwatch();
+ sw.Start();
+ var result = action();
+ sw.Stop();
+ TimeAndLogActionMsg(logger, message, sw.ElapsedMilliseconds, default!);
+ return result;
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console/Program.cs b/src/BGR.Console/Program.cs
index f787117..74d6965 100644
--- a/src/BGR.Console/Program.cs
+++ b/src/BGR.Console/Program.cs
@@ -19,7 +19,11 @@ try
.ConfigureServices(static (_, services) =>
{
services.AddSerilog();
+ services.AddSingleton(AnsiConsole.Console);
services.AddSingleton();
+ services.AddSingleton();
+ services.AddSingleton();
+ services.AddSingleton();
})
.BuildApp()
.RunAsync(args);
@@ -35,200 +39,3 @@ finally
{
await Log.CloseAndFlushAsync();
}
-
-if (args.Length < 1)
-{
- Console.WriteLine("Usage: BackgroundRemover ");
- return;
-}
-
-var inputImagePath = args[0];
-var maskImagePath = Path.ChangeExtension(inputImagePath, null) + "_mask.png";
-var outputImagePath = Path.ChangeExtension(inputImagePath, null) + "_no_bg.png";
-
-try
-{
- var assembly = Assembly.GetExecutingAssembly();
- var resourceName = "BGR.Console.Resources.Files.rmbg.onnx";
-
- using var stream = assembly.GetManifestResourceStream(resourceName) ?? throw new FileNotFoundException("Model not found in embedded resources.");
- var modelBytes = new byte[stream.Length];
- stream.ReadExactly(modelBytes);
-
- using var image = await Image.LoadAsync(inputImagePath);
- var inputTensor = CreateTensorInput(image);
-
- using var options = new SessionOptions() { LogSeverityLevel = OrtLoggingLevel.ORT_LOGGING_LEVEL_ERROR };
- using InferenceSession session = new(modelBytes, options);
- var inputs = new List()
- {
- NamedOnnxValue.CreateFromTensor(session.InputNames[0], inputTensor),
- };
-
- using var results = session.Run(inputs);
- var outputTensor = results[0].AsTensor();
-
- using var mask = GenerateMask(outputTensor, image.Width, image.Height);
-
- using var bgRemoved = GetImageWithBackgroundRemoved(image, mask);
-
- var encoder = new PngEncoder { CompressionLevel = PngCompressionLevel.BestCompression };
-
- await mask.SaveAsync(maskImagePath, encoder);
- await bgRemoved.SaveAsync(outputImagePath, encoder);
-
- Console.WriteLine($"Background removed and saved to {outputImagePath}");
-}
-catch (Exception ex)
-{
- Console.WriteLine($"Error: {ex.Message}");
- throw;
-}
-
-static Tensor CreateTensorInput(Image image)
-{
- // U2Net expects input images to be 320x320. This is dependent on the model.
- const int targetWidth = 1024;
- const int targetHeight = 1024;
-
- // ImageNet normalization parameters
- // source:
- // - https://www.image-net.org/
- // - https://pytorch.org calculated these values from the ImageNet dataset
- // and they are commonly used for models trained on ImageNet so we use them here
- // to normalize the input image to better match the distribution of the data the model was trained on
- // NOTE: These values are not universal and may vary for different models
- const float rMean = 0.485f; // Mean value for Red channel
- const float gMean = 0.456f; // Mean value for Green channel
- const float bMean = 0.406f; // Mean value for Blue channel
- const float rStd = 0.229f; // Standard deviation for Red channel
- const float gStd = 0.224f; // Standard deviation for Green channel
- const float bStd = 0.225f; // Standard deviation for Blue channel
- const float pixelMax = 255f; // Maximum pixel intensity for normalization
-
- // Create a temporary image for preprocessing
- using var resized = image.Clone();
- resized.Mutate(x => x.Resize(targetWidth, targetHeight));
-
- // Create tensor of shape (1, 3, 320, 320)
- // 1 for batch size, 3 for RGB channels, 320x320 for image dimensions
- DenseTensor tensor = new([1, 3, targetHeight, targetWidth]);
-
- // Normalize pixel values and copy to tensor
- WalkImage(resized.Height, resized.Width, (x, y) =>
- {
- var pixel = resized[x, y];
-
- // u2net expects expect input images to be normalized using ImageNet mean and std
- // to better match the distribution of the data the model was trained on
- // Normalize to range [0, 1] and standardize using ImageNet mean/std
- // The tensor is filled with normalized pixel values
- tensor[0, 0, y, x] = ((pixel.R / pixelMax) - rMean) / rStd; // Red channel
- tensor[0, 1, y, x] = ((pixel.G / pixelMax) - gMean) / gStd; // Green channel
- tensor[0, 2, y, x] = ((pixel.B / pixelMax) - bMean) / bStd; // Blue channel
- });
-
- return tensor;
-}
-
-static Image GenerateMask(Tensor maskTensor, int width, int height)
-{
- var mask = new Image(width, height);
-
- var sourceHeight = maskTensor.Dimensions[2]; // Height of the original tensor mask
- var sourceWidth = maskTensor.Dimensions[3]; // Width of the original tensor mask
-
- using Image tempMask = new(sourceWidth, sourceHeight);
-
- // Sigmoid function parameters
- const float sigmoidScale = 1f; // Scaling factor for sigmoid activation
- const float sigmoidShift = 1f; // Shift factor in the denominator of the sigmoid function
- const float sigmoidDivisor = -1f; // Multiplier for the exponent in the sigmoid function
-
- static float CalculateSigmoid(float x)
- {
- return sigmoidScale / (sigmoidShift + MathF.Exp(sigmoidDivisor * x));
- }
-
- const float binarizationThreshold = 0.5f; // Threshold to determine foreground vs. background
- const float normalizationFactor = 2f; // Scales the thresholded value to enhance contrast
-
- // Pixel intensity values
- const byte maxIntensity = 255; // Maximum grayscale intensity
- const byte opaqueAlpha = 255; // Fully opaque alpha value
-
-
- WalkImage(sourceHeight, sourceWidth, (x, y) =>
- {
- // a sigmoid function is a function that produces an S-shaped curve
- // it is often used in machine learning and statistics to model probabilities
- // the sigmoid function is defined as:
- // f(x) = 1 / (1 + e^(-x))
- // where e is the base of the natural logarithm and x is the input value
-
- // the raw tensor values for our mask are going to be real unbounded numbers
- // i.e. -1.5, 0.5, 2.0, etc.
- // the sigmoid function will map these values to a range between 0 and 1
- // this allows us to say that value closer to 0 is background and value
- // closer to 1 is foreground
- var sigmoidValue = CalculateSigmoid(maskTensor[0, 0, y, x]);
-
- // now we want to threshold the sigmoid value to determine if it is foreground or background
- // we are arbitrarily choosing 0.5 as the threshold. so if the sigmoid value is greater than
- // 0.5 we will consider it foreground and if it is less than 0.5 we will consider it background
-
- // when a sigmoid value is greater than 0.5 we will subtract the threshold from it
- // and multiply it by 2 this way the intensity value will be larger for values closer to 1
- // and create more contrast in the mask
- var normalizedValue = sigmoidValue > binarizationThreshold
- ? (sigmoidValue - binarizationThreshold) * normalizationFactor
- : 0f;
-
- // Convert to an 8-bit grayscale intensity
- var intensity = (byte)(normalizedValue * maxIntensity);
-
- // Store the pixel with full opacity
- tempMask[x, y] = new Rgba32(intensity, intensity, intensity, opaqueAlpha);
- });
-
- // Resize the mask to match the target dimensions
- tempMask.Mutate(x => x.Resize(width, height));
-
- // Copy the resized mask to the final output image
- WalkImage(height, width, (x, y) => mask[x, y] = tempMask[x, y]);
-
- return mask;
-}
-
-static Image GetImageWithBackgroundRemoved(Image image, Image mask)
-{
- Image result = new(image.Width, image.Height);
-
- const byte alphaThreshold = 20;
- Rgba32 transparentPixel = new(0, 0, 0, 0);
-
- WalkImage(image.Height, image.Width, (x, y) =>
- {
- var sourcePixel = image[x, y];
- var maskPixel = mask[x, y];
-
- var alpha = maskPixel.R;
-
- result[x, y] = alpha > alphaThreshold
- ? new Rgba32(sourcePixel.R, sourcePixel.G, sourcePixel.B, sourcePixel.A)
- : transparentPixel;
- });
-
- return result;
-}
-
-static void WalkImage(int height, int width, Action action)
-{
- for (var y = 0; y < height; y++)
- {
- for (var x = 0; x < width; x++)
- {
- action(x, y);
- }
- }
-}
\ No newline at end of file
diff --git a/src/BGR.Console/Removal/IImage.cs b/src/BGR.Console/Removal/IImage.cs
index 1a66873..f25e08f 100644
--- a/src/BGR.Console/Removal/IImage.cs
+++ b/src/BGR.Console/Removal/IImage.cs
@@ -4,6 +4,5 @@ internal interface IImage
{
int Width { get; }
int Height { get; }
- void Resize(int width, int height);
- IPixel GetPixel(int x, int y);
+ Stream Data { get; }
}
\ No newline at end of file
diff --git a/src/BGR.Console/Removal/IInferenceRunner.cs b/src/BGR.Console/Removal/IInferenceRunner.cs
index e69de29..5966eb7 100644
--- a/src/BGR.Console/Removal/IInferenceRunner.cs
+++ b/src/BGR.Console/Removal/IInferenceRunner.cs
@@ -0,0 +1,6 @@
+namespace BGR.Console.Removal;
+
+internal interface IInferenceRunner
+{
+ ITensor Run(byte[] model, ITensor inputTensor);
+}
\ No newline at end of file
diff --git a/src/BGR.Console/Removal/IPixel.cs b/src/BGR.Console/Removal/IPixel.cs
deleted file mode 100644
index 4489608..0000000
--- a/src/BGR.Console/Removal/IPixel.cs
+++ /dev/null
@@ -1,8 +0,0 @@
-namespace BGR.Console.Removal;
-
-internal interface IPixel
-{
- float R { get; }
- float G { get; }
- float B { get; }
-}
\ No newline at end of file
diff --git a/src/BGR.Console/Removal/ITensor.cs b/src/BGR.Console/Removal/ITensor.cs
index 1f402a3..6a9550f 100644
--- a/src/BGR.Console/Removal/ITensor.cs
+++ b/src/BGR.Console/Removal/ITensor.cs
@@ -6,4 +6,5 @@ internal interface ITensor
int Width { get; }
void SetValue(int batch, int channel, int y, int x, T value);
float GetValue(int batch, int channel, int y, int x);
-}
\ No newline at end of file
+ Tensor ToTensor();
+}
diff --git a/src/BGR.Console/Removal/ImageProcessor.cs b/src/BGR.Console/Removal/ImageProcessor.cs
index 405f3de..e864234 100644
--- a/src/BGR.Console/Removal/ImageProcessor.cs
+++ b/src/BGR.Console/Removal/ImageProcessor.cs
@@ -2,12 +2,16 @@ namespace BGR.Console.Removal;
internal abstract class ImageProcessor
{
+ public abstract Task LoadImageAsync(string path);
+
public abstract Task> CreateTensorInputAsync(Stream image, Model model);
- public abstract Task GenerateMaskAsync(OnnxTensor maskTensor, int width, int height);
+ public abstract Task GenerateMaskAsync(ITensor maskTensor, int width, int height);
public abstract Task RemoveBackgroundAsync(Stream image, Stream mask);
+ public abstract Task SaveImageAsync(Stream image, string path);
+
protected static void WalkImage(int height, int width, Action action)
{
for (var y = 0; y < height; y++)
diff --git a/src/BGR.Console/Removal/ImageSharp/ImageSharpProcessor.cs b/src/BGR.Console/Removal/ImageSharp/ImageSharpProcessor.cs
index 1f5c594..dca520e 100644
--- a/src/BGR.Console/Removal/ImageSharp/ImageSharpProcessor.cs
+++ b/src/BGR.Console/Removal/ImageSharp/ImageSharpProcessor.cs
@@ -2,9 +2,26 @@ namespace BGR.Console.Removal.ImageSharp;
internal class ImageSharpProcessor : ImageProcessor
{
+ public override async Task LoadImageAsync(string path)
+ {
+ var image = await Image.LoadAsync(path);
+
+ if (image.Metadata.DecodedImageFormat is null)
+ {
+ throw new InvalidOperationException("Image format is not supported.");
+ }
+
+ var stream = new MemoryStream();
+ await image.SaveAsync(stream, image.Metadata.DecodedImageFormat);
+ stream.Position = 0;
+
+ return new SharpImage(image.Width, image.Height, stream);
+ }
+
public override async Task> CreateTensorInputAsync(Stream image, Model model)
{
using var resized = await Image.LoadAsync(image);
+
resized.Mutate(x => x.Resize(model.InputWidth, model.InputHeight));
const int batchSize = 1;
@@ -22,7 +39,7 @@ internal class ImageSharpProcessor : ImageProcessor
return tensor;
}
- public override async Task GenerateMaskAsync(OnnxTensor maskTensor, int width, int height)
+ public override async Task GenerateMaskAsync(ITensor maskTensor, int width, int height)
{
using var mask = new Image(width, height);
@@ -45,12 +62,16 @@ internal class ImageSharpProcessor : ImageProcessor
var stream = new MemoryStream();
await mask.SaveAsync(stream, new PngEncoder());
+ stream.Position = 0;
return stream;
}
public override async Task RemoveBackgroundAsync(Stream image, Stream mask)
{
+ image.Position = 0;
+ mask.Position = 0;
+
var imageWithBg = await Image.LoadAsync(image);
var maskImage = await Image.LoadAsync(mask);
using var imageWithBgRemoved = new Image(imageWithBg.Width, imageWithBg.Height);
@@ -72,9 +93,18 @@ internal class ImageSharpProcessor : ImageProcessor
var result = new MemoryStream();
await imageWithBgRemoved.SaveAsync(result, new PngEncoder());
+ result.Position = 0;
+
return result;
}
+ public override async Task SaveImageAsync(Stream image, string path)
+ {
+ image.Position = 0;
+ var img = await Image.LoadAsync(image);
+ await img.SaveAsync(path, new PngEncoder());
+ }
+
private static float Normalize(float value)
{
const float binarizationThreshold = 0.5f;
diff --git a/src/BGR.Console/Removal/ImageSharp/SharpImage.cs b/src/BGR.Console/Removal/ImageSharp/SharpImage.cs
index 577e1bd..dc0f952 100644
--- a/src/BGR.Console/Removal/ImageSharp/SharpImage.cs
+++ b/src/BGR.Console/Removal/ImageSharp/SharpImage.cs
@@ -1,19 +1,25 @@
namespace BGR.Console.Removal.ImageSharp;
-internal class SharpImage(Image image) : IImage
+internal class SharpImage : IImage
{
- private readonly Image _image = image;
+ public int Width { get; }
+ public int Height { get; }
+ public Stream Data { get; }
- public int Width => _image.Width;
- public int Height => _image.Height;
-
- public void Resize(int width, int height)
+ public SharpImage(int width, int height, Stream data)
{
- _image.Mutate(x => x.Resize(width, height));
- }
+ if (width <= 0)
+ {
+ throw new ArgumentOutOfRangeException(nameof(width), "must be greater than 0");
+ }
- public IPixel GetPixel(int x, int y)
- {
- return new SharpPixel(_image[x, y]);
+ if (height <= 0)
+ {
+ throw new ArgumentOutOfRangeException(nameof(height), "must be greater than 0");
+ }
+
+ Width = width;
+ Height = height;
+ Data = data ?? throw new ArgumentNullException(nameof(data));
}
}
\ No newline at end of file
diff --git a/src/BGR.Console/Removal/ImageSharp/SharpPixel.cs b/src/BGR.Console/Removal/ImageSharp/SharpPixel.cs
deleted file mode 100644
index be284b1..0000000
--- a/src/BGR.Console/Removal/ImageSharp/SharpPixel.cs
+++ /dev/null
@@ -1,10 +0,0 @@
-namespace BGR.Console.Removal.ImageSharp;
-
-internal class SharpPixel(Rgba32 pixel) : IPixel
-{
- private readonly Rgba32 _pixel = pixel;
-
- public float R => _pixel.R;
- public float G => _pixel.G;
- public float B => _pixel.B;
-}
\ No newline at end of file
diff --git a/src/BGR.Console/Removal/Models/IModelFactory.cs b/src/BGR.Console/Removal/Models/IModelFactory.cs
new file mode 100644
index 0000000..02e1389
--- /dev/null
+++ b/src/BGR.Console/Removal/Models/IModelFactory.cs
@@ -0,0 +1,6 @@
+namespace BGR.Console.Removal.Models;
+
+internal interface IModelFactory
+{
+ Model Create(string resourceName);
+}
\ No newline at end of file
diff --git a/src/BGR.Console/Removal/Models/ModNetModel.cs b/src/BGR.Console/Removal/Models/ModNetModel.cs
index 4c6612c..682d10f 100644
--- a/src/BGR.Console/Removal/Models/ModNetModel.cs
+++ b/src/BGR.Console/Removal/Models/ModNetModel.cs
@@ -1,7 +1,8 @@
namespace BGR.Console.Removal.Models;
-internal class ModNetModel : Model
+internal class ModNetModel(byte[] modelBytes) : Model(modelBytes)
{
+ public const string Id = "modnet";
public override int InputWidth => 512;
public override int InputHeight => 512;
public override float RedNormalizationMean => 0.485f;
diff --git a/src/BGR.Console/Removal/Models/Model.cs b/src/BGR.Console/Removal/Models/Model.cs
index b507c6f..b6f89b0 100644
--- a/src/BGR.Console/Removal/Models/Model.cs
+++ b/src/BGR.Console/Removal/Models/Model.cs
@@ -11,18 +11,28 @@ internal abstract class Model
public abstract float RedNormalizationStd { get; }
public abstract float GreenNormalizationStd { get; }
public abstract float BlueNormalizationStd { get; }
+ public byte[] Bytes { get; } = [];
- public float NormalizeRed(float value)
+ internal Model()
+ {
+ }
+
+ protected Model(byte[] modelBytes)
+ {
+ Bytes = modelBytes;
+ }
+
+ public virtual float NormalizeRed(float value)
{
return Normalize(value, RedNormalizationMean, RedNormalizationStd);
}
- public float NormalizeGreen(float value)
+ public virtual float NormalizeGreen(float value)
{
return Normalize(value, GreenNormalizationMean, GreenNormalizationStd);
}
- public float NormalizeBlue(float value)
+ public virtual float NormalizeBlue(float value)
{
return Normalize(value, BlueNormalizationMean, BlueNormalizationStd);
}
diff --git a/src/BGR.Console/Removal/Models/ModelFactory.cs b/src/BGR.Console/Removal/Models/ModelFactory.cs
new file mode 100644
index 0000000..de9087f
--- /dev/null
+++ b/src/BGR.Console/Removal/Models/ModelFactory.cs
@@ -0,0 +1,21 @@
+namespace BGR.Console.Removal.Models;
+
+internal class ModelFactory(IResourceManager resourceManager) : IModelFactory
+{
+ private readonly IResourceManager _resourceManager = resourceManager;
+
+ public Model Create(string resourceName)
+ {
+ var resource = _resourceManager.GetResource(resourceName);
+ var model = new byte[resource.Length];
+ resource.ReadExactly(model);
+
+ return resourceName switch
+ {
+ $"{U2NetModel.Id}.onnx" => new U2NetModel(model),
+ $"{RmbgModel.Id}.onnx" => new RmbgModel(model),
+ $"{ModNetModel.Id}.onnx" => new ModNetModel(model),
+ _ => throw new ArgumentException($"Unknown model name: {resourceName}")
+ };
+ }
+}
diff --git a/src/BGR.Console/Removal/Models/RmbgModel.cs b/src/BGR.Console/Removal/Models/RmbgModel.cs
index 09a9baf..0104b94 100644
--- a/src/BGR.Console/Removal/Models/RmbgModel.cs
+++ b/src/BGR.Console/Removal/Models/RmbgModel.cs
@@ -1,8 +1,8 @@
-
namespace BGR.Console.Removal.Models;
-internal class RmbgModel : Model
+internal class RmbgModel(byte[] modelBytes) : Model(modelBytes)
{
+ public const string Id = "rmbg";
public override int InputWidth => 1024;
public override int InputHeight => 1024;
public override float RedNormalizationMean => 0.485f;
diff --git a/src/BGR.Console/Removal/Models/U2NetModel.cs b/src/BGR.Console/Removal/Models/U2NetModel.cs
index b98863a..c0234b8 100644
--- a/src/BGR.Console/Removal/Models/U2NetModel.cs
+++ b/src/BGR.Console/Removal/Models/U2NetModel.cs
@@ -1,7 +1,8 @@
namespace BGR.Console.Removal.Models;
-internal class U2NetModel : Model
+internal class U2NetModel(byte[] modelBytes) : Model(modelBytes)
{
+ public const string Id = "u2net";
public override int InputWidth => 320;
public override int InputHeight => 320;
public override float RedNormalizationMean => 0.485f;
@@ -10,4 +11,4 @@ internal class U2NetModel : Model
public override float RedNormalizationStd => 0.229f;
public override float GreenNormalizationStd => 0.224f;
public override float BlueNormalizationStd => 0.225f;
-}
+}
\ No newline at end of file
diff --git a/src/BGR.Console/Removal/Onnx/OnnxInferenceRunner.cs b/src/BGR.Console/Removal/Onnx/OnnxInferenceRunner.cs
new file mode 100644
index 0000000..3abd2b3
--- /dev/null
+++ b/src/BGR.Console/Removal/Onnx/OnnxInferenceRunner.cs
@@ -0,0 +1,18 @@
+namespace BGR.Console.Removal.Onnx;
+
+internal class OnnxInferenceRunner : IInferenceRunner
+{
+ public ITensor Run(byte[] model, ITensor inputTensor)
+ {
+ using var options = new SessionOptions() { LogSeverityLevel = OrtLoggingLevel.ORT_LOGGING_LEVEL_ERROR };
+ using var session = new InferenceSession(model, options);
+ var inputs = new List()
+ {
+ NamedOnnxValue.CreateFromTensor(session.InputNames[0], inputTensor.ToTensor()),
+ };
+
+ var results = session.Run(inputs);
+ var outputTensor = results[0].AsTensor();
+ return new OnnxTensor(outputTensor);
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console/Removal/Onnx/OnnxTensor.cs b/src/BGR.Console/Removal/Onnx/OnnxTensor.cs
index 386e9e2..512704a 100644
--- a/src/BGR.Console/Removal/Onnx/OnnxTensor.cs
+++ b/src/BGR.Console/Removal/Onnx/OnnxTensor.cs
@@ -1,19 +1,28 @@
namespace BGR.Console.Removal.Onnx;
-public class OnnxTensor(
- int batchSize,
- int channels,
- int height,
- int width
-) : ITensor
+public class OnnxTensor : ITensor
{
- private readonly DenseTensor _tensor =
- new([batchSize, channels, height, width]);
+ private readonly Tensor _tensor;
public int Height => _tensor.Dimensions[2];
public int Width => _tensor.Dimensions[3];
+ public OnnxTensor(
+ int batchSize,
+ int channels,
+ int height,
+ int width
+ )
+ {
+ _tensor = new DenseTensor([batchSize, channels, height, width]);
+ }
+
+ public OnnxTensor(Tensor tensor)
+ {
+ _tensor = tensor;
+ }
+
public void SetValue(
int batch,
int channel,
@@ -34,4 +43,9 @@ public class OnnxTensor(
{
return _tensor[batch, channel, y, x];
}
+
+ public Tensor ToTensor()
+ {
+ return _tensor;
+ }
}
\ No newline at end of file
diff --git a/src/BGR.Console/Removal/RemovalCommand.cs b/src/BGR.Console/Removal/RemovalCommand.cs
new file mode 100644
index 0000000..c7c6bee
--- /dev/null
+++ b/src/BGR.Console/Removal/RemovalCommand.cs
@@ -0,0 +1,139 @@
+using System.Diagnostics;
+
+namespace BGR.Console.Removal;
+
+internal class RemovalCommand(
+ IModelFactory modelFactory,
+ ImageProcessor imageProcessor,
+ IInferenceRunner inferenceRunner,
+ IAnsiConsole console,
+ ILogger logger
+) : AsyncCommand
+{
+ private readonly IModelFactory _modelFactory = modelFactory;
+ private readonly ImageProcessor _imageProcessor = imageProcessor;
+ private readonly IInferenceRunner _inferenceRunner = inferenceRunner;
+ private readonly IAnsiConsole _console = console;
+ private readonly ILogger _logger = logger;
+
+ public override async Task ExecuteAsync(CommandContext context, Settings settings)
+ {
+ await _console.Status()
+ .Spinner(Spinner.Known.Dots)
+ .SpinnerStyle(Style.Parse("green"))
+ .StartAsync("Removing background...", async ctx =>
+ {
+ ctx.Status("Loading model...");
+ var model = _logger.TimeAndLogAction(
+ "Loading model",
+ () => _modelFactory.Create(settings.ResourceName)
+ );
+
+ ctx.Status("Loading image...");
+ var image = await _logger.TimeAndLogActionAsync(
+ "Loading image",
+ async () => await _imageProcessor.LoadImageAsync(settings.Image)
+ );
+
+ ctx.Status("Creating tensor input...");
+ var inputTensor = await _logger.TimeAndLogActionAsync(
+ "Creating tensor input",
+ async () => await _imageProcessor.CreateTensorInputAsync(image.Data, model)
+ );
+
+ ctx.Status("Running inference...");
+ var outputTensor = _logger.TimeAndLogAction(
+ "Running inference",
+ () => _inferenceRunner.Run(model.Bytes, inputTensor)
+ );
+
+ ctx.Status("Generating mask...");
+ var mask = await _logger.TimeAndLogActionAsync(
+ "Generating mask",
+ async () => await _imageProcessor.GenerateMaskAsync(outputTensor, image.Width, image.Height)
+ );
+
+ ctx.Status("Removing background...");
+ var output = await _logger.TimeAndLogActionAsync(
+ "Removing background",
+ async () => await _imageProcessor.RemoveBackgroundAsync(image.Data, mask)
+ );
+
+ if (settings.IncludeMask)
+ {
+ await _logger.TimeAndLogActionAsync(
+ "Saving mask",
+ async () => await _imageProcessor.SaveImageAsync(mask, settings.MaskPath)
+ );
+
+ _console.MarkupLine($"[bold]Mask saved to:[/] [blue]{settings.OutputPath}[/]");
+ }
+
+ await _logger.TimeAndLogActionAsync(
+ "Saving output",
+ async () => await _imageProcessor.SaveImageAsync(output, settings.OutputPath)
+ );
+
+ _console.MarkupLine($"[bold]Output saved to:[/] [green]{settings.OutputPath}[/]");
+ });
+
+ return 0;
+ }
+
+ internal class Settings : CommandSettings
+ {
+ private static readonly Dictionary Models = new()
+ {
+ { RmbgModel.Id, "rmbg.onnx" },
+ { ModNetModel.Id, "modnet.onnx" },
+ { U2NetModel.Id, "u2net.onnx" },
+ };
+
+ [CommandArgument(0, "")]
+ [Description("Path to the image file whose background you want to remove")]
+ public string Image { get; init; } = string.Empty;
+
+ [CommandOption("--model|-m")]
+ [Description("The model to use for background removal")]
+ public string Model { get; init; } = "rmbg";
+
+ [CommandOption("--include-mask|-i")]
+ [Description("Generate and output the mask used for background removal")]
+ public bool IncludeMask { get; init; } = false;
+
+ [CommandOption("--output|-o")]
+ [Description("Path to output image without background to. File extension will always be .png")]
+ public string Output { get; init; } = string.Empty;
+
+ public string ResourceName => Models[Model];
+
+ public string MaskPath => GetOutputPath("_mask");
+
+ public string OutputPath => GetOutputPath("_no_bg");
+
+ public override ValidationResult Validate()
+ {
+ if (File.Exists(Image) is false)
+ {
+ return ValidationResult.Error($"The image file '{Image}' does not exist.");
+ }
+
+ if (Models.ContainsKey(Model) is false)
+ {
+ return ValidationResult.Error($"The model '{Model}' is not supported.");
+ }
+
+ return ValidationResult.Success();
+ }
+
+ private string GetOutputPath(string modifier)
+ {
+ if (string.IsNullOrWhiteSpace(Output))
+ {
+ return Path.ChangeExtension(Image, null) + modifier + ".png";
+ }
+
+ return Path.ChangeExtension(Output, ".png");
+ }
+ }
+}
\ No newline at end of file
diff --git a/src/BGR.Console/Resources/ResourceManager.cs b/src/BGR.Console/Resources/ResourceManager.cs
index 845db3a..97d6df5 100644
--- a/src/BGR.Console/Resources/ResourceManager.cs
+++ b/src/BGR.Console/Resources/ResourceManager.cs
@@ -1,4 +1,3 @@
-
namespace BGR.Console.Resources;
internal sealed class ResourceManager : IResourceManager
diff --git a/src/BGR.Console/Usings.cs b/src/BGR.Console/Usings.cs
index 015bd70..8130861 100644
--- a/src/BGR.Console/Usings.cs
+++ b/src/BGR.Console/Usings.cs
@@ -1,9 +1,13 @@
+global using System.ComponentModel;
global using System.Reflection;
global using BGR.Console.Common;
+global using BGR.Console.Removal;
+global using BGR.Console.Removal.ImageSharp;
global using BGR.Console.Removal.Models;
global using BGR.Console.Removal.Onnx;
global using BGR.Console.Resources;
+global using BGR.Console.Logging;
global using Microsoft.Extensions.DependencyInjection;
global using Microsoft.Extensions.Hosting;
@@ -20,4 +24,5 @@ global using SixLabors.ImageSharp.Formats.Png;
global using SixLabors.ImageSharp.PixelFormats;
global using SixLabors.ImageSharp.Processing;
-global using Spectre.Console.Cli;
+global using Spectre.Console;
+global using Spectre.Console.Cli;
\ No newline at end of file