feat: refactor to proper app + write tests
This commit is contained in:
@@ -85,6 +85,8 @@ dotnet_diagnostic.CA1707.severity = none
|
|||||||
dotnet_diagnostic.IDE0058.severity = none
|
dotnet_diagnostic.IDE0058.severity = none
|
||||||
dotnet_diagnostic.CA2007.severity = none
|
dotnet_diagnostic.CA2007.severity = none
|
||||||
dotnet_diagnostic.CA1515.severity = none
|
dotnet_diagnostic.CA1515.severity = none
|
||||||
|
dotnet_diagnostic.IDE0100.severity = none
|
||||||
|
dotnet_diagnostic.IDE0046.severity = none
|
||||||
|
|
||||||
# var preferences
|
# var preferences
|
||||||
csharp_style_var_elsewhere = true:suggestion
|
csharp_style_var_elsewhere = true:suggestion
|
||||||
|
|||||||
@@ -11,9 +11,14 @@
|
|||||||
<IncludeAssets>runtime; build; native; contentfiles; analyzers; buildtransitive</IncludeAssets>
|
<IncludeAssets>runtime; build; native; contentfiles; analyzers; buildtransitive</IncludeAssets>
|
||||||
<PrivateAssets>all</PrivateAssets>
|
<PrivateAssets>all</PrivateAssets>
|
||||||
</PackageReference>
|
</PackageReference>
|
||||||
|
<PackageReference Include="coverlet.msbuild" Version="6.0.4">
|
||||||
|
<IncludeAssets>runtime; build; native; contentfiles; analyzers; buildtransitive</IncludeAssets>
|
||||||
|
<PrivateAssets>all</PrivateAssets>
|
||||||
|
</PackageReference>
|
||||||
<PackageReference Include="Moq" Version="4.20.72" />
|
<PackageReference Include="Moq" Version="4.20.72" />
|
||||||
<PackageReference Include="Shouldly" Version="4.3.0" />
|
<PackageReference Include="Shouldly" Version="4.3.0" />
|
||||||
<PackageReference Include="Microsoft.NET.Test.Sdk" Version="17.12.0" />
|
<PackageReference Include="Microsoft.NET.Test.Sdk" Version="17.12.0" />
|
||||||
|
<PackageReference Include="Spectre.Console.Testing" Version="0.49.1" />
|
||||||
<PackageReference Include="xunit" Version="2.9.3" />
|
<PackageReference Include="xunit" Version="2.9.3" />
|
||||||
<PackageReference Include="xunit.runner.visualstudio" Version="3.0.1">
|
<PackageReference Include="xunit.runner.visualstudio" Version="3.0.1">
|
||||||
<IncludeAssets>runtime; build; native; contentfiles; analyzers; buildtransitive</IncludeAssets>
|
<IncludeAssets>runtime; build; native; contentfiles; analyzers; buildtransitive</IncludeAssets>
|
||||||
@@ -25,14 +30,14 @@
|
|||||||
|
|
||||||
<PropertyGroup>
|
<PropertyGroup>
|
||||||
<CollectCoverage>true</CollectCoverage>
|
<CollectCoverage>true</CollectCoverage>
|
||||||
<CoverletOutput>./TestResults/coverage/</CoverletOutput>
|
<CoverletOutput>./TestResults/Coverage/</CoverletOutput>
|
||||||
<CoverletOutputFormat>cobertura</CoverletOutputFormat>
|
<CoverletOutputFormat>cobertura</CoverletOutputFormat>
|
||||||
<Include>[BGR.Console]*</Include>
|
<Include>[BGR.Console]*</Include>
|
||||||
<ExcludeByFile>**/Program.cs</ExcludeByFile>
|
<ExcludeByFile>**/Program.cs</ExcludeByFile>
|
||||||
</PropertyGroup>
|
</PropertyGroup>
|
||||||
|
|
||||||
<Target Name="GenerateHtmlCoverageReport" AfterTargets="GenerateCoverageResultAfterTest">
|
<Target Name="GenerateHtmlCoverageReport" AfterTargets="GenerateCoverageResultAfterTest">
|
||||||
<Exec Command="reportgenerator -reports:./TestResults/coverage/*.xml -targetdir:./TestResults/coverage/report/ -reporttypes:Html_Dark" />
|
<Exec Command="reportgenerator -reports:./TestResults/Coverage/*.xml -targetdir:./TestResults/Coverage/Report/ -reporttypes:Html_Dark" />
|
||||||
</Target>
|
</Target>
|
||||||
|
|
||||||
<ItemGroup>
|
<ItemGroup>
|
||||||
|
|||||||
@@ -0,0 +1,19 @@
|
|||||||
|
namespace BGR.Console.Tests.Integration;
|
||||||
|
|
||||||
|
internal static class AppFactory
|
||||||
|
{
|
||||||
|
public static CommandApp<RemovalCommand> Create()
|
||||||
|
{
|
||||||
|
return Host.CreateDefaultBuilder()
|
||||||
|
.ConfigureLogging(static logging => logging.ClearProviders())
|
||||||
|
.ConfigureServices(static (_, services) =>
|
||||||
|
{
|
||||||
|
services.AddSingleton<IAnsiConsole>(new TestConsole());
|
||||||
|
services.AddSingleton<IResourceManager, ResourceManager>();
|
||||||
|
services.AddSingleton<ImageProcessor, ImageSharpProcessor>();
|
||||||
|
services.AddSingleton<IInferenceRunner, OnnxInferenceRunner>();
|
||||||
|
services.AddSingleton<IModelFactory, ModelFactory>();
|
||||||
|
})
|
||||||
|
.BuildApp();
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,24 @@
|
|||||||
|
namespace BGR.Console.Tests.Integration;
|
||||||
|
|
||||||
|
public class RemovalCommandTests
|
||||||
|
{
|
||||||
|
private readonly CommandApp<RemovalCommand> _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<Rgba32>(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);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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<Model> _modelMock = new();
|
||||||
|
private readonly Stream _testImageStream;
|
||||||
|
|
||||||
|
public ImageSharpProcessorTests()
|
||||||
|
{
|
||||||
|
if (File.Exists(TestImagePath) is false)
|
||||||
|
{
|
||||||
|
using var testImage = new Image<Rgba32>(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<Rgba32>(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<SharpImage>();
|
||||||
|
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<OnnxTensor>();
|
||||||
|
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<float>())).Returns(normalizedValue);
|
||||||
|
_modelMock.Setup(static x => x.NormalizeGreen(It.IsAny<float>())).Returns(normalizedValue);
|
||||||
|
_modelMock.Setup(static x => x.NormalizeBlue(It.IsAny<float>())).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<float>()), Times.AtLeast(1));
|
||||||
|
_modelMock.Verify(static x => x.NormalizeGreen(It.IsAny<float>()), Times.AtLeast(1));
|
||||||
|
_modelMock.Verify(static x => x.NormalizeBlue(It.IsAny<float>()), 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<Rgba32>(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<Rgba32>(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<Rgba32>(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<Rgba32>(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<Rgba32>(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<Rgba32>(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;
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,67 @@
|
|||||||
|
namespace BGR.Console.Tests.Unit;
|
||||||
|
|
||||||
|
public class ModelFactoryTests
|
||||||
|
{
|
||||||
|
private readonly Mock<IResourceManager> _resourceManagerMock;
|
||||||
|
private readonly ModelFactory _sut;
|
||||||
|
private readonly byte[] _sampleModelBytes = [0x01, 0x02, 0x03];
|
||||||
|
|
||||||
|
public ModelFactoryTests()
|
||||||
|
{
|
||||||
|
_resourceManagerMock = new Mock<IResourceManager>();
|
||||||
|
_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<ArgumentException>(() => _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);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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<ITensor<float>> _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<float>([1, 3, 320, 320]));
|
||||||
|
}
|
||||||
|
|
||||||
|
[Fact]
|
||||||
|
public void Run_WhenCalledWithValidInput_ItShouldReturnOutput()
|
||||||
|
{
|
||||||
|
var result = _sut.Run(_sampleModelBytes, _mockInputTensor.Object);
|
||||||
|
|
||||||
|
result.ShouldBeOfType<OnnxTensor>();
|
||||||
|
result.ShouldNotBeNull();
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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<float>([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<DenseTensor<float>>();
|
||||||
|
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);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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<IModelFactory> _modelFactoryMock = new();
|
||||||
|
private readonly Mock<ImageProcessor> _imageProcessorMock = new();
|
||||||
|
private readonly Mock<IInferenceRunner> _inferenceRunnerMock = new();
|
||||||
|
private readonly TestConsole _console = new();
|
||||||
|
private readonly Mock<ILogger<RemovalCommand>> _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<Rgba32>(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<IRemainingArguments>().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<string>()), 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<Rgba32>(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<Rgba32>(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;
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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<ArgumentNullException>();
|
||||||
|
}
|
||||||
|
|
||||||
|
[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<ArgumentOutOfRangeException>();
|
||||||
|
}
|
||||||
|
|
||||||
|
[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;
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,7 +1,19 @@
|
|||||||
global using BGR.Console.Common;
|
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.Resources;
|
||||||
|
|
||||||
global using Microsoft.Extensions.DependencyInjection;
|
global using Microsoft.Extensions.DependencyInjection;
|
||||||
global using Microsoft.Extensions.Hosting;
|
global using Microsoft.Extensions.Hosting;
|
||||||
|
global using Microsoft.Extensions.Logging;
|
||||||
|
|
||||||
global using Moq;
|
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;
|
||||||
|
|||||||
@@ -6,6 +6,7 @@
|
|||||||
|
|
||||||
<ItemGroup>
|
<ItemGroup>
|
||||||
<InternalsVisibleTo Include="$(AssemblyName).Tests" />
|
<InternalsVisibleTo Include="$(AssemblyName).Tests" />
|
||||||
|
<InternalsVisibleTo Include="DynamicProxyGenAssembly2" />
|
||||||
</ItemGroup>
|
</ItemGroup>
|
||||||
|
|
||||||
<ItemGroup>
|
<ItemGroup>
|
||||||
|
|||||||
@@ -2,10 +2,22 @@ namespace BGR.Console.Common;
|
|||||||
|
|
||||||
internal static class HostBuilderExtensions
|
internal static class HostBuilderExtensions
|
||||||
{
|
{
|
||||||
public static CommandApp BuildApp(this IHostBuilder builder)
|
public static CommandApp<RemovalCommand> BuildApp(this IHostBuilder builder)
|
||||||
{
|
{
|
||||||
var registrar = new TypeRegistrar(builder);
|
var registrar = new TypeRegistrar(builder);
|
||||||
var app = new CommandApp(registrar);
|
var app = new CommandApp<RemovalCommand>(registrar);
|
||||||
|
|
||||||
|
app.Configure(static c =>
|
||||||
|
c.SetExceptionHandler(static (ex, resolver) =>
|
||||||
|
{
|
||||||
|
var logger = resolver?.Resolve(typeof(ILogger<RemovalCommand>)) as ILogger<RemovalCommand>;
|
||||||
|
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;
|
return app;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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<ILogger, Exception> RemovalCommandFailedMsg = LoggerMessage.Define(
|
||||||
|
LogLevel.Error,
|
||||||
|
new EventId(0, nameof(RemovalCommandFailed)),
|
||||||
|
"An error occurred while executing the command."
|
||||||
|
);
|
||||||
|
|
||||||
|
private static readonly Action<ILogger, string, long, Exception> TimeAndLogActionMsg = LoggerMessage.Define<string, long>(
|
||||||
|
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<Task> action)
|
||||||
|
{
|
||||||
|
var sw = new Stopwatch();
|
||||||
|
sw.Start();
|
||||||
|
await action();
|
||||||
|
sw.Stop();
|
||||||
|
TimeAndLogActionMsg(logger, message, sw.ElapsedMilliseconds, default!);
|
||||||
|
}
|
||||||
|
|
||||||
|
public static async Task<T> TimeAndLogActionAsync<T>(this ILogger logger, string message, Func<Task<T>> 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<T>(this ILogger logger, string message, Func<T> action)
|
||||||
|
{
|
||||||
|
var sw = new Stopwatch();
|
||||||
|
sw.Start();
|
||||||
|
var result = action();
|
||||||
|
sw.Stop();
|
||||||
|
TimeAndLogActionMsg(logger, message, sw.ElapsedMilliseconds, default!);
|
||||||
|
return result;
|
||||||
|
}
|
||||||
|
}
|
||||||
+4
-197
@@ -19,7 +19,11 @@ try
|
|||||||
.ConfigureServices(static (_, services) =>
|
.ConfigureServices(static (_, services) =>
|
||||||
{
|
{
|
||||||
services.AddSerilog();
|
services.AddSerilog();
|
||||||
|
services.AddSingleton(AnsiConsole.Console);
|
||||||
services.AddSingleton<IResourceManager, ResourceManager>();
|
services.AddSingleton<IResourceManager, ResourceManager>();
|
||||||
|
services.AddSingleton<ImageProcessor, ImageSharpProcessor>();
|
||||||
|
services.AddSingleton<IInferenceRunner, OnnxInferenceRunner>();
|
||||||
|
services.AddSingleton<IModelFactory, ModelFactory>();
|
||||||
})
|
})
|
||||||
.BuildApp()
|
.BuildApp()
|
||||||
.RunAsync(args);
|
.RunAsync(args);
|
||||||
@@ -35,200 +39,3 @@ finally
|
|||||||
{
|
{
|
||||||
await Log.CloseAndFlushAsync();
|
await Log.CloseAndFlushAsync();
|
||||||
}
|
}
|
||||||
|
|
||||||
if (args.Length < 1)
|
|
||||||
{
|
|
||||||
Console.WriteLine("Usage: BackgroundRemover <input_image_path>");
|
|
||||||
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<Rgba32>(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>()
|
|
||||||
{
|
|
||||||
NamedOnnxValue.CreateFromTensor(session.InputNames[0], inputTensor),
|
|
||||||
};
|
|
||||||
|
|
||||||
using var results = session.Run(inputs);
|
|
||||||
var outputTensor = results[0].AsTensor<float>();
|
|
||||||
|
|
||||||
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<float> CreateTensorInput(Image<Rgba32> 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<float> 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<Rgba32> GenerateMask(Tensor<float> maskTensor, int width, int height)
|
|
||||||
{
|
|
||||||
var mask = new Image<Rgba32>(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<Rgba32> 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<Rgba32> GetImageWithBackgroundRemoved(Image<Rgba32> image, Image<Rgba32> mask)
|
|
||||||
{
|
|
||||||
Image<Rgba32> 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<int, int> action)
|
|
||||||
{
|
|
||||||
for (var y = 0; y < height; y++)
|
|
||||||
{
|
|
||||||
for (var x = 0; x < width; x++)
|
|
||||||
{
|
|
||||||
action(x, y);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -4,6 +4,5 @@ internal interface IImage
|
|||||||
{
|
{
|
||||||
int Width { get; }
|
int Width { get; }
|
||||||
int Height { get; }
|
int Height { get; }
|
||||||
void Resize(int width, int height);
|
Stream Data { get; }
|
||||||
IPixel GetPixel(int x, int y);
|
|
||||||
}
|
}
|
||||||
@@ -0,0 +1,6 @@
|
|||||||
|
namespace BGR.Console.Removal;
|
||||||
|
|
||||||
|
internal interface IInferenceRunner
|
||||||
|
{
|
||||||
|
ITensor<float> Run(byte[] model, ITensor<float> inputTensor);
|
||||||
|
}
|
||||||
@@ -1,8 +0,0 @@
|
|||||||
namespace BGR.Console.Removal;
|
|
||||||
|
|
||||||
internal interface IPixel
|
|
||||||
{
|
|
||||||
float R { get; }
|
|
||||||
float G { get; }
|
|
||||||
float B { get; }
|
|
||||||
}
|
|
||||||
@@ -6,4 +6,5 @@ internal interface ITensor<T>
|
|||||||
int Width { get; }
|
int Width { get; }
|
||||||
void SetValue(int batch, int channel, int y, int x, T value);
|
void SetValue(int batch, int channel, int y, int x, T value);
|
||||||
float GetValue(int batch, int channel, int y, int x);
|
float GetValue(int batch, int channel, int y, int x);
|
||||||
|
Tensor<T> ToTensor();
|
||||||
}
|
}
|
||||||
@@ -2,12 +2,16 @@ namespace BGR.Console.Removal;
|
|||||||
|
|
||||||
internal abstract class ImageProcessor
|
internal abstract class ImageProcessor
|
||||||
{
|
{
|
||||||
|
public abstract Task<IImage> LoadImageAsync(string path);
|
||||||
|
|
||||||
public abstract Task<ITensor<float>> CreateTensorInputAsync(Stream image, Model model);
|
public abstract Task<ITensor<float>> CreateTensorInputAsync(Stream image, Model model);
|
||||||
|
|
||||||
public abstract Task<Stream> GenerateMaskAsync(OnnxTensor maskTensor, int width, int height);
|
public abstract Task<Stream> GenerateMaskAsync(ITensor<float> maskTensor, int width, int height);
|
||||||
|
|
||||||
public abstract Task<Stream> RemoveBackgroundAsync(Stream image, Stream mask);
|
public abstract Task<Stream> RemoveBackgroundAsync(Stream image, Stream mask);
|
||||||
|
|
||||||
|
public abstract Task SaveImageAsync(Stream image, string path);
|
||||||
|
|
||||||
protected static void WalkImage(int height, int width, Action<int, int> action)
|
protected static void WalkImage(int height, int width, Action<int, int> action)
|
||||||
{
|
{
|
||||||
for (var y = 0; y < height; y++)
|
for (var y = 0; y < height; y++)
|
||||||
|
|||||||
@@ -2,9 +2,26 @@ namespace BGR.Console.Removal.ImageSharp;
|
|||||||
|
|
||||||
internal class ImageSharpProcessor : ImageProcessor
|
internal class ImageSharpProcessor : ImageProcessor
|
||||||
{
|
{
|
||||||
|
public override async Task<IImage> LoadImageAsync(string path)
|
||||||
|
{
|
||||||
|
var image = await Image.LoadAsync<Rgba32>(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<ITensor<float>> CreateTensorInputAsync(Stream image, Model model)
|
public override async Task<ITensor<float>> CreateTensorInputAsync(Stream image, Model model)
|
||||||
{
|
{
|
||||||
using var resized = await Image.LoadAsync<Rgba32>(image);
|
using var resized = await Image.LoadAsync<Rgba32>(image);
|
||||||
|
|
||||||
resized.Mutate(x => x.Resize(model.InputWidth, model.InputHeight));
|
resized.Mutate(x => x.Resize(model.InputWidth, model.InputHeight));
|
||||||
|
|
||||||
const int batchSize = 1;
|
const int batchSize = 1;
|
||||||
@@ -22,7 +39,7 @@ internal class ImageSharpProcessor : ImageProcessor
|
|||||||
return tensor;
|
return tensor;
|
||||||
}
|
}
|
||||||
|
|
||||||
public override async Task<Stream> GenerateMaskAsync(OnnxTensor maskTensor, int width, int height)
|
public override async Task<Stream> GenerateMaskAsync(ITensor<float> maskTensor, int width, int height)
|
||||||
{
|
{
|
||||||
using var mask = new Image<Rgba32>(width, height);
|
using var mask = new Image<Rgba32>(width, height);
|
||||||
|
|
||||||
@@ -45,12 +62,16 @@ internal class ImageSharpProcessor : ImageProcessor
|
|||||||
|
|
||||||
var stream = new MemoryStream();
|
var stream = new MemoryStream();
|
||||||
await mask.SaveAsync(stream, new PngEncoder());
|
await mask.SaveAsync(stream, new PngEncoder());
|
||||||
|
stream.Position = 0;
|
||||||
|
|
||||||
return stream;
|
return stream;
|
||||||
}
|
}
|
||||||
|
|
||||||
public override async Task<Stream> RemoveBackgroundAsync(Stream image, Stream mask)
|
public override async Task<Stream> RemoveBackgroundAsync(Stream image, Stream mask)
|
||||||
{
|
{
|
||||||
|
image.Position = 0;
|
||||||
|
mask.Position = 0;
|
||||||
|
|
||||||
var imageWithBg = await Image.LoadAsync<Rgba32>(image);
|
var imageWithBg = await Image.LoadAsync<Rgba32>(image);
|
||||||
var maskImage = await Image.LoadAsync<Rgba32>(mask);
|
var maskImage = await Image.LoadAsync<Rgba32>(mask);
|
||||||
using var imageWithBgRemoved = new Image<Rgba32>(imageWithBg.Width, imageWithBg.Height);
|
using var imageWithBgRemoved = new Image<Rgba32>(imageWithBg.Width, imageWithBg.Height);
|
||||||
@@ -72,9 +93,18 @@ internal class ImageSharpProcessor : ImageProcessor
|
|||||||
|
|
||||||
var result = new MemoryStream();
|
var result = new MemoryStream();
|
||||||
await imageWithBgRemoved.SaveAsync(result, new PngEncoder());
|
await imageWithBgRemoved.SaveAsync(result, new PngEncoder());
|
||||||
|
result.Position = 0;
|
||||||
|
|
||||||
return result;
|
return result;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
public override async Task SaveImageAsync(Stream image, string path)
|
||||||
|
{
|
||||||
|
image.Position = 0;
|
||||||
|
var img = await Image.LoadAsync<Rgba32>(image);
|
||||||
|
await img.SaveAsync(path, new PngEncoder());
|
||||||
|
}
|
||||||
|
|
||||||
private static float Normalize(float value)
|
private static float Normalize(float value)
|
||||||
{
|
{
|
||||||
const float binarizationThreshold = 0.5f;
|
const float binarizationThreshold = 0.5f;
|
||||||
|
|||||||
@@ -1,19 +1,25 @@
|
|||||||
namespace BGR.Console.Removal.ImageSharp;
|
namespace BGR.Console.Removal.ImageSharp;
|
||||||
|
|
||||||
internal class SharpImage(Image<Rgba32> image) : IImage
|
internal class SharpImage : IImage
|
||||||
{
|
{
|
||||||
private readonly Image<Rgba32> _image = image;
|
public int Width { get; }
|
||||||
|
public int Height { get; }
|
||||||
|
public Stream Data { get; }
|
||||||
|
|
||||||
public int Width => _image.Width;
|
public SharpImage(int width, int height, Stream data)
|
||||||
public int Height => _image.Height;
|
|
||||||
|
|
||||||
public void Resize(int width, int height)
|
|
||||||
{
|
{
|
||||||
_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)
|
if (height <= 0)
|
||||||
{
|
{
|
||||||
return new SharpPixel(_image[x, y]);
|
throw new ArgumentOutOfRangeException(nameof(height), "must be greater than 0");
|
||||||
|
}
|
||||||
|
|
||||||
|
Width = width;
|
||||||
|
Height = height;
|
||||||
|
Data = data ?? throw new ArgumentNullException(nameof(data));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -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;
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,6 @@
|
|||||||
|
namespace BGR.Console.Removal.Models;
|
||||||
|
|
||||||
|
internal interface IModelFactory
|
||||||
|
{
|
||||||
|
Model Create(string resourceName);
|
||||||
|
}
|
||||||
@@ -1,7 +1,8 @@
|
|||||||
namespace BGR.Console.Removal.Models;
|
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 InputWidth => 512;
|
||||||
public override int InputHeight => 512;
|
public override int InputHeight => 512;
|
||||||
public override float RedNormalizationMean => 0.485f;
|
public override float RedNormalizationMean => 0.485f;
|
||||||
|
|||||||
@@ -11,18 +11,28 @@ internal abstract class Model
|
|||||||
public abstract float RedNormalizationStd { get; }
|
public abstract float RedNormalizationStd { get; }
|
||||||
public abstract float GreenNormalizationStd { get; }
|
public abstract float GreenNormalizationStd { get; }
|
||||||
public abstract float BlueNormalizationStd { 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);
|
return Normalize(value, RedNormalizationMean, RedNormalizationStd);
|
||||||
}
|
}
|
||||||
|
|
||||||
public float NormalizeGreen(float value)
|
public virtual float NormalizeGreen(float value)
|
||||||
{
|
{
|
||||||
return Normalize(value, GreenNormalizationMean, GreenNormalizationStd);
|
return Normalize(value, GreenNormalizationMean, GreenNormalizationStd);
|
||||||
}
|
}
|
||||||
|
|
||||||
public float NormalizeBlue(float value)
|
public virtual float NormalizeBlue(float value)
|
||||||
{
|
{
|
||||||
return Normalize(value, BlueNormalizationMean, BlueNormalizationStd);
|
return Normalize(value, BlueNormalizationMean, BlueNormalizationStd);
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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}")
|
||||||
|
};
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,8 +1,8 @@
|
|||||||
|
|
||||||
namespace BGR.Console.Removal.Models;
|
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 InputWidth => 1024;
|
||||||
public override int InputHeight => 1024;
|
public override int InputHeight => 1024;
|
||||||
public override float RedNormalizationMean => 0.485f;
|
public override float RedNormalizationMean => 0.485f;
|
||||||
|
|||||||
@@ -1,7 +1,8 @@
|
|||||||
namespace BGR.Console.Removal.Models;
|
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 InputWidth => 320;
|
||||||
public override int InputHeight => 320;
|
public override int InputHeight => 320;
|
||||||
public override float RedNormalizationMean => 0.485f;
|
public override float RedNormalizationMean => 0.485f;
|
||||||
|
|||||||
@@ -0,0 +1,18 @@
|
|||||||
|
namespace BGR.Console.Removal.Onnx;
|
||||||
|
|
||||||
|
internal class OnnxInferenceRunner : IInferenceRunner
|
||||||
|
{
|
||||||
|
public ITensor<float> Run(byte[] model, ITensor<float> inputTensor)
|
||||||
|
{
|
||||||
|
using var options = new SessionOptions() { LogSeverityLevel = OrtLoggingLevel.ORT_LOGGING_LEVEL_ERROR };
|
||||||
|
using var session = new InferenceSession(model, options);
|
||||||
|
var inputs = new List<NamedOnnxValue>()
|
||||||
|
{
|
||||||
|
NamedOnnxValue.CreateFromTensor(session.InputNames[0], inputTensor.ToTensor()),
|
||||||
|
};
|
||||||
|
|
||||||
|
var results = session.Run(inputs);
|
||||||
|
var outputTensor = results[0].AsTensor<float>();
|
||||||
|
return new OnnxTensor(outputTensor);
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,19 +1,28 @@
|
|||||||
namespace BGR.Console.Removal.Onnx;
|
namespace BGR.Console.Removal.Onnx;
|
||||||
|
|
||||||
public class OnnxTensor(
|
public class OnnxTensor : ITensor<float>
|
||||||
int batchSize,
|
|
||||||
int channels,
|
|
||||||
int height,
|
|
||||||
int width
|
|
||||||
) : ITensor<float>
|
|
||||||
{
|
{
|
||||||
private readonly DenseTensor<float> _tensor =
|
private readonly Tensor<float> _tensor;
|
||||||
new([batchSize, channels, height, width]);
|
|
||||||
|
|
||||||
public int Height => _tensor.Dimensions[2];
|
public int Height => _tensor.Dimensions[2];
|
||||||
|
|
||||||
public int Width => _tensor.Dimensions[3];
|
public int Width => _tensor.Dimensions[3];
|
||||||
|
|
||||||
|
public OnnxTensor(
|
||||||
|
int batchSize,
|
||||||
|
int channels,
|
||||||
|
int height,
|
||||||
|
int width
|
||||||
|
)
|
||||||
|
{
|
||||||
|
_tensor = new DenseTensor<float>([batchSize, channels, height, width]);
|
||||||
|
}
|
||||||
|
|
||||||
|
public OnnxTensor(Tensor<float> tensor)
|
||||||
|
{
|
||||||
|
_tensor = tensor;
|
||||||
|
}
|
||||||
|
|
||||||
public void SetValue(
|
public void SetValue(
|
||||||
int batch,
|
int batch,
|
||||||
int channel,
|
int channel,
|
||||||
@@ -34,4 +43,9 @@ public class OnnxTensor(
|
|||||||
{
|
{
|
||||||
return _tensor[batch, channel, y, x];
|
return _tensor[batch, channel, y, x];
|
||||||
}
|
}
|
||||||
|
|
||||||
|
public Tensor<float> ToTensor()
|
||||||
|
{
|
||||||
|
return _tensor;
|
||||||
|
}
|
||||||
}
|
}
|
||||||
@@ -0,0 +1,139 @@
|
|||||||
|
using System.Diagnostics;
|
||||||
|
|
||||||
|
namespace BGR.Console.Removal;
|
||||||
|
|
||||||
|
internal class RemovalCommand(
|
||||||
|
IModelFactory modelFactory,
|
||||||
|
ImageProcessor imageProcessor,
|
||||||
|
IInferenceRunner inferenceRunner,
|
||||||
|
IAnsiConsole console,
|
||||||
|
ILogger<RemovalCommand> logger
|
||||||
|
) : AsyncCommand<RemovalCommand.Settings>
|
||||||
|
{
|
||||||
|
private readonly IModelFactory _modelFactory = modelFactory;
|
||||||
|
private readonly ImageProcessor _imageProcessor = imageProcessor;
|
||||||
|
private readonly IInferenceRunner _inferenceRunner = inferenceRunner;
|
||||||
|
private readonly IAnsiConsole _console = console;
|
||||||
|
private readonly ILogger<RemovalCommand> _logger = logger;
|
||||||
|
|
||||||
|
public override async Task<int> 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<string, string> Models = new()
|
||||||
|
{
|
||||||
|
{ RmbgModel.Id, "rmbg.onnx" },
|
||||||
|
{ ModNetModel.Id, "modnet.onnx" },
|
||||||
|
{ U2NetModel.Id, "u2net.onnx" },
|
||||||
|
};
|
||||||
|
|
||||||
|
[CommandArgument(0, "<image>")]
|
||||||
|
[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");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,4 +1,3 @@
|
|||||||
|
|
||||||
namespace BGR.Console.Resources;
|
namespace BGR.Console.Resources;
|
||||||
|
|
||||||
internal sealed class ResourceManager : IResourceManager
|
internal sealed class ResourceManager : IResourceManager
|
||||||
|
|||||||
@@ -1,9 +1,13 @@
|
|||||||
|
global using System.ComponentModel;
|
||||||
global using System.Reflection;
|
global using System.Reflection;
|
||||||
|
|
||||||
global using BGR.Console.Common;
|
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.Models;
|
||||||
global using BGR.Console.Removal.Onnx;
|
global using BGR.Console.Removal.Onnx;
|
||||||
global using BGR.Console.Resources;
|
global using BGR.Console.Resources;
|
||||||
|
global using BGR.Console.Logging;
|
||||||
|
|
||||||
global using Microsoft.Extensions.DependencyInjection;
|
global using Microsoft.Extensions.DependencyInjection;
|
||||||
global using Microsoft.Extensions.Hosting;
|
global using Microsoft.Extensions.Hosting;
|
||||||
@@ -20,4 +24,5 @@ global using SixLabors.ImageSharp.Formats.Png;
|
|||||||
global using SixLabors.ImageSharp.PixelFormats;
|
global using SixLabors.ImageSharp.PixelFormats;
|
||||||
global using SixLabors.ImageSharp.Processing;
|
global using SixLabors.ImageSharp.Processing;
|
||||||
|
|
||||||
|
global using Spectre.Console;
|
||||||
global using Spectre.Console.Cli;
|
global using Spectre.Console.Cli;
|
||||||
Reference in New Issue
Block a user