feat: initial working POC

This commit is contained in:
Stevan Freeborn
2025-02-06 23:26:58 -06:00
commit 651869c27f
7 changed files with 722 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
*.onnx filter=lfs diff=lfs merge=lfs -text
+484
View File
@@ -0,0 +1,484 @@
## Ignore Visual Studio temporary files, build results, and
## files generated by popular Visual Studio add-ons.
##
## Get latest from `dotnet new gitignore`
# dotenv files
.env
# User-specific files
*.rsuser
*.suo
*.user
*.userosscache
*.sln.docstates
# User-specific files (MonoDevelop/Xamarin Studio)
*.userprefs
# Mono auto generated files
mono_crash.*
# Build results
[Dd]ebug/
[Dd]ebugPublic/
[Rr]elease/
[Rr]eleases/
x64/
x86/
[Ww][Ii][Nn]32/
[Aa][Rr][Mm]/
[Aa][Rr][Mm]64/
bld/
[Bb]in/
[Oo]bj/
[Ll]og/
[Ll]ogs/
# Visual Studio 2015/2017 cache/options directory
.vs/
# Uncomment if you have tasks that create the project's static files in wwwroot
#wwwroot/
# Visual Studio 2017 auto generated files
Generated\ Files/
# MSTest test Results
[Tt]est[Rr]esult*/
[Bb]uild[Ll]og.*
# NUnit
*.VisualState.xml
TestResult.xml
nunit-*.xml
# Build Results of an ATL Project
[Dd]ebugPS/
[Rr]eleasePS/
dlldata.c
# Benchmark Results
BenchmarkDotNet.Artifacts/
# .NET
project.lock.json
project.fragment.lock.json
artifacts/
# Tye
.tye/
# ASP.NET Scaffolding
ScaffoldingReadMe.txt
# StyleCop
StyleCopReport.xml
# Files built by Visual Studio
*_i.c
*_p.c
*_h.h
*.ilk
*.meta
*.obj
*.iobj
*.pch
*.pdb
*.ipdb
*.pgc
*.pgd
*.rsp
*.sbr
*.tlb
*.tli
*.tlh
*.tmp
*.tmp_proj
*_wpftmp.csproj
*.log
*.tlog
*.vspscc
*.vssscc
.builds
*.pidb
*.svclog
*.scc
# Chutzpah Test files
_Chutzpah*
# Visual C++ cache files
ipch/
*.aps
*.ncb
*.opendb
*.opensdf
*.sdf
*.cachefile
*.VC.db
*.VC.VC.opendb
# Visual Studio profiler
*.psess
*.vsp
*.vspx
*.sap
# Visual Studio Trace Files
*.e2e
# TFS 2012 Local Workspace
$tf/
# Guidance Automation Toolkit
*.gpState
# ReSharper is a .NET coding add-in
_ReSharper*/
*.[Rr]e[Ss]harper
*.DotSettings.user
# TeamCity is a build add-in
_TeamCity*
# DotCover is a Code Coverage Tool
*.dotCover
# AxoCover is a Code Coverage Tool
.axoCover/*
!.axoCover/settings.json
# Coverlet is a free, cross platform Code Coverage Tool
coverage*.json
coverage*.xml
coverage*.info
# Visual Studio code coverage results
*.coverage
*.coveragexml
# NCrunch
_NCrunch_*
.*crunch*.local.xml
nCrunchTemp_*
# MightyMoose
*.mm.*
AutoTest.Net/
# Web workbench (sass)
.sass-cache/
# Installshield output folder
[Ee]xpress/
# DocProject is a documentation generator add-in
DocProject/buildhelp/
DocProject/Help/*.HxT
DocProject/Help/*.HxC
DocProject/Help/*.hhc
DocProject/Help/*.hhk
DocProject/Help/*.hhp
DocProject/Help/Html2
DocProject/Help/html
# Click-Once directory
publish/
# Publish Web Output
*.[Pp]ublish.xml
*.azurePubxml
# Note: Comment the next line if you want to checkin your web deploy settings,
# but database connection strings (with potential passwords) will be unencrypted
*.pubxml
*.publishproj
# Microsoft Azure Web App publish settings. Comment the next line if you want to
# checkin your Azure Web App publish settings, but sensitive information contained
# in these scripts will be unencrypted
PublishScripts/
# NuGet Packages
*.nupkg
# NuGet Symbol Packages
*.snupkg
# The packages folder can be ignored because of Package Restore
**/[Pp]ackages/*
# except build/, which is used as an MSBuild target.
!**/[Pp]ackages/build/
# Uncomment if necessary however generally it will be regenerated when needed
#!**/[Pp]ackages/repositories.config
# NuGet v3's project.json files produces more ignorable files
*.nuget.props
*.nuget.targets
# Microsoft Azure Build Output
csx/
*.build.csdef
# Microsoft Azure Emulator
ecf/
rcf/
# Windows Store app package directories and files
AppPackages/
BundleArtifacts/
Package.StoreAssociation.xml
_pkginfo.txt
*.appx
*.appxbundle
*.appxupload
# Visual Studio cache files
# files ending in .cache can be ignored
*.[Cc]ache
# but keep track of directories ending in .cache
!?*.[Cc]ache/
# Others
ClientBin/
~$*
*~
*.dbmdl
*.dbproj.schemaview
*.jfm
*.pfx
*.publishsettings
orleans.codegen.cs
# Including strong name files can present a security risk
# (https://github.com/github/gitignore/pull/2483#issue-259490424)
#*.snk
# Since there are multiple workflows, uncomment next line to ignore bower_components
# (https://github.com/github/gitignore/pull/1529#issuecomment-104372622)
#bower_components/
# RIA/Silverlight projects
Generated_Code/
# Backup & report files from converting an old project file
# to a newer Visual Studio version. Backup files are not needed,
# because we have git ;-)
_UpgradeReport_Files/
Backup*/
UpgradeLog*.XML
UpgradeLog*.htm
ServiceFabricBackup/
*.rptproj.bak
# SQL Server files
*.mdf
*.ldf
*.ndf
# Business Intelligence projects
*.rdl.data
*.bim.layout
*.bim_*.settings
*.rptproj.rsuser
*- [Bb]ackup.rdl
*- [Bb]ackup ([0-9]).rdl
*- [Bb]ackup ([0-9][0-9]).rdl
# Microsoft Fakes
FakesAssemblies/
# GhostDoc plugin setting file
*.GhostDoc.xml
# Node.js Tools for Visual Studio
.ntvs_analysis.dat
node_modules/
# Visual Studio 6 build log
*.plg
# Visual Studio 6 workspace options file
*.opt
# Visual Studio 6 auto-generated workspace file (contains which files were open etc.)
*.vbw
# Visual Studio 6 auto-generated project file (contains which files were open etc.)
*.vbp
# Visual Studio 6 workspace and project file (working project files containing files to include in project)
*.dsw
*.dsp
# Visual Studio 6 technical files
*.ncb
*.aps
# Visual Studio LightSwitch build output
**/*.HTMLClient/GeneratedArtifacts
**/*.DesktopClient/GeneratedArtifacts
**/*.DesktopClient/ModelManifest.xml
**/*.Server/GeneratedArtifacts
**/*.Server/ModelManifest.xml
_Pvt_Extensions
# Paket dependency manager
.paket/paket.exe
paket-files/
# FAKE - F# Make
.fake/
# CodeRush personal settings
.cr/personal
# Python Tools for Visual Studio (PTVS)
__pycache__/
*.pyc
# Cake - Uncomment if you are using it
# tools/**
# !tools/packages.config
# Tabs Studio
*.tss
# Telerik's JustMock configuration file
*.jmconfig
# BizTalk build output
*.btp.cs
*.btm.cs
*.odx.cs
*.xsd.cs
# OpenCover UI analysis results
OpenCover/
# Azure Stream Analytics local run output
ASALocalRun/
# MSBuild Binary and Structured Log
*.binlog
# NVidia Nsight GPU debugger configuration file
*.nvuser
# MFractors (Xamarin productivity tool) working folder
.mfractor/
# Local History for Visual Studio
.localhistory/
# Visual Studio History (VSHistory) files
.vshistory/
# BeatPulse healthcheck temp database
healthchecksdb
# Backup folder for Package Reference Convert tool in Visual Studio 2017
MigrationBackup/
# Ionide (cross platform F# VS Code tools) working folder
.ionide/
# Fody - auto-generated XML schema
FodyWeavers.xsd
# VS Code files for those working on multiple tools
.vscode/*
!.vscode/settings.json
!.vscode/tasks.json
!.vscode/launch.json
!.vscode/extensions.json
*.code-workspace
# Local History for Visual Studio Code
.history/
# Windows Installer files from build outputs
*.cab
*.msi
*.msix
*.msm
*.msp
# JetBrains Rider
*.sln.iml
.idea/
##
## Visual studio for Mac
##
# globs
Makefile.in
*.userprefs
*.usertasks
config.make
config.status
aclocal.m4
install-sh
autom4te.cache/
*.tar.gz
tarballs/
test-results/
# Mac bundle stuff
*.dmg
*.app
# content below from: https://github.com/github/gitignore/blob/main/Global/macOS.gitignore
# General
.DS_Store
.AppleDouble
.LSOverride
# Icon must end with two \r
Icon
# Thumbnails
._*
# Files that might appear in the root of a volume
.DocumentRevisions-V100
.fseventsd
.Spotlight-V100
.TemporaryItems
.Trashes
.VolumeIcon.icns
.com.apple.timemachine.donotpresent
# Directories potentially created on remote AFP share
.AppleDB
.AppleDesktop
Network Trash Folder
Temporary Items
.apdisk
# content below from: https://github.com/github/gitignore/blob/main/Global/Windows.gitignore
# Windows thumbnail cache files
Thumbs.db
ehthumbs.db
ehthumbs_vista.db
# Dump file
*.stackdump
# Folder config file
[Dd]esktop.ini
# Recycle Bin used on file shares
$RECYCLE.BIN/
# Windows Installer files
*.cab
*.msi
*.msix
*.msm
*.msp
# Windows shortcuts
*.lnk
# Vim temporary swap files
*.swp
+21
View File
@@ -0,0 +1,21 @@
<Project Sdk="Microsoft.NET.Sdk">
<PropertyGroup>
<OutputType>Exe</OutputType>
<TargetFramework>net9.0</TargetFramework>
<ImplicitUsings>enable</ImplicitUsings>
<Nullable>enable</Nullable>
</PropertyGroup>
<ItemGroup>
<PackageReference Include="Microsoft.ML.OnnxRuntime" Version="1.20.1" />
<PackageReference Include="SixLabors.ImageSharp" Version="3.1.6" />
</ItemGroup>
<ItemGroup>
<EmbeddedResource Include="modnet.onnx" />
<EmbeddedResource Include="rmbg.onnx" />
<EmbeddedResource Include="u2net.onnx" />
</ItemGroup>
</Project>
+207
View File
@@ -0,0 +1,207 @@
using System.Reflection;
using Microsoft.ML.OnnxRuntime;
using Microsoft.ML.OnnxRuntime.Tensors;
using SixLabors.ImageSharp;
using SixLabors.ImageSharp.Formats.Png;
using SixLabors.ImageSharp.PixelFormats;
using SixLabors.ImageSharp.Processing;
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.u2net.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 = Image.Load<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 };
mask.Save(maskImagePath, encoder);
bgRemoved.Save(outputImagePath, encoder);
Console.WriteLine($"Background removed and saved to {outputImagePath}");
}
catch (Exception ex)
{
Console.WriteLine($"Error: {ex.Message}");
throw;
}
static DenseTensor<float> CreateTensorInput(Image<Rgba32> image)
{
// U2Net expects input images to be 320x320. This is dependent on the model.
const int targetWidth = 320;
const int targetHeight = 320;
// 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);
}
}
}
+3
View File
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:07c308cf0fc7e6e8b2065a12ed7fc07e1de8febb7dc7839d7b7f15dd66584df9
size 25888640
+3
View File
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:fcea23951a378f92634834888896cc1eec54655366ae6e949282646ce17c5420
size 366087549
+3
View File
@@ -0,0 +1,3 @@
version https://git-lfs.github.com/spec/v1
oid sha256:8d10d2f3bb75ae3b6d527c77944fc5e7dcd94b29809d47a739a7a728a912b491
size 175997641