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