diff --git a/.env.example b/.env.example index 0e1138a..64570c6 100644 --- a/.env.example +++ b/.env.example @@ -166,6 +166,11 @@ DATABASE_URL=sqlite:///./data/ai_photo_edit.db # When false: Skips download, Smart Select uses Replicate API (requires REPLICATE_API_KEY) AUTO_DOWNLOAD_SAM=true +# Auto-download U2Net model on startup (true/false) +# When true (default): Downloads U2Net model (~176MB) on first startup for offline Remove Background +# When false: Skips download, Remove Background falls back to rembg (if installed) +AUTO_DOWNLOAD_U2NET=true + # Allow users to select model per-edit ALLOW_MODEL_OVERRIDE=true diff --git a/README.md b/README.md index 158e90c..41f2692 100644 --- a/README.md +++ b/README.md @@ -201,6 +201,22 @@ docker compose -f docker-compose.gpu.yml logs | grep -i sam If Docker created `./data/` as root and you can't write there without `sudo`, you can also use root's curl as above — the container reads the file regardless of owner. +**Remove Background fails ("Install u2net or rembg")** + +The U2Net model auto-downloads (~176MB) from GitHub on first use, same as SAM. If that download fails (DNS/firewall, see above) and `rembg` isn't installed either, you'll see this error. Fix it the same way — download directly on the host: + +```bash +mkdir -p ./data/models +sudo curl -L -o ./data/models/u2net.onnx \ + https://github.com/danielgatis/rembg/releases/download/v0.0.0/u2net.onnx +``` + +The file is ~176 MB. Once it exists at `./data/models/u2net.onnx`, the next "Remove Background" click picks it up — no rebuild or restart needed. Verify with: +```bash +docker compose logs -f | grep -i u2net +# Should show: "U2Net model loaded successfully with OpenCV DNN" +``` + **AI models not downloading (container DNS blocked)** If you ran `./install-local-gpu.sh`, this is already permanently fixed. Otherwise, the container's host firewall is blocking outbound DNS from the Docker bridge — apply the fix manually (does **not** affect container isolation): diff --git a/backend/app/services/u2net_model.py b/backend/app/services/u2net_model.py deleted file mode 100644 index fa269b1..0000000 --- a/backend/app/services/u2net_model.py +++ /dev/null @@ -1,500 +0,0 @@ -""" -U2Net Model Definition for Background Removal -Based on: https://github.com/xuebinqin/U-2-Net - -This is a simplified implementation that works with the standard U2Net weights. -""" - -import torch -import torch.nn as nn -import torch.nn.functional as F - - -class REBNCONV(nn.Module): - def __init__(self, in_ch=3, out_ch=3, dirate=1): - super(REBNCONV, self).__init__() - - self.conv_s1 = nn.Conv2d(in_ch, out_ch, 3, padding=1*dirate, dilation=1*dirate) - self.bn_s1 = nn.BatchNorm2d(out_ch) - self.relu_s1 = nn.ReLU(inplace=True) - - def forward(self, x): - hx = x - xout = self.relu_s1(self.bn_s1(self.conv_s1(hx))) - return xout - - -def _upsample_like(src, tar): - src = F.interpolate(src, size=tar.shape[2:], mode='bilinear', align_corners=False) - return src - - -class RSU7(nn.Module): - def __init__(self, in_ch=3, mid_ch=12, out_ch=3): - super(RSU7, self).__init__() - - self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1) - - self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1) - self.pool1 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=1) - self.pool2 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=1) - self.pool3 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=1) - self.pool4 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv5 = REBNCONV(mid_ch, mid_ch, dirate=1) - self.pool5 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv6 = REBNCONV(mid_ch, mid_ch, dirate=1) - - self.rebnconv7 = REBNCONV(mid_ch, mid_ch, dirate=2) - - self.rebnconv6d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv5d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv4d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv3d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv2d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv1d = REBNCONV(mid_ch*2, out_ch, dirate=1) - - def forward(self, x): - hx = x - hxin = self.rebnconvin(hx) - - hx1 = self.rebnconv1(hxin) - hx = self.pool1(hx1) - - hx2 = self.rebnconv2(hx) - hx = self.pool2(hx2) - - hx3 = self.rebnconv3(hx) - hx = self.pool3(hx3) - - hx4 = self.rebnconv4(hx) - hx = self.pool4(hx4) - - hx5 = self.rebnconv5(hx) - hx = self.pool5(hx5) - - hx6 = self.rebnconv6(hx) - - hx7 = self.rebnconv7(hx6) - - hx6d = self.rebnconv6d(torch.cat((hx7, hx6), 1)) - hx6dup = _upsample_like(hx6d, hx5) - - hx5d = self.rebnconv5d(torch.cat((hx6dup, hx5), 1)) - hx5dup = _upsample_like(hx5d, hx4) - - hx4d = self.rebnconv4d(torch.cat((hx5dup, hx4), 1)) - hx4dup = _upsample_like(hx4d, hx3) - - hx3d = self.rebnconv3d(torch.cat((hx4dup, hx3), 1)) - hx3dup = _upsample_like(hx3d, hx2) - - hx2d = self.rebnconv2d(torch.cat((hx3dup, hx2), 1)) - hx2dup = _upsample_like(hx2d, hx1) - - hx1d = self.rebnconv1d(torch.cat((hx2dup, hx1), 1)) - - return hx1d + hxin - - -class RSU6(nn.Module): - def __init__(self, in_ch=3, mid_ch=12, out_ch=3): - super(RSU6, self).__init__() - - self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1) - - self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1) - self.pool1 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=1) - self.pool2 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=1) - self.pool3 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=1) - self.pool4 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv5 = REBNCONV(mid_ch, mid_ch, dirate=1) - - self.rebnconv6 = REBNCONV(mid_ch, mid_ch, dirate=2) - - self.rebnconv5d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv4d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv3d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv2d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv1d = REBNCONV(mid_ch*2, out_ch, dirate=1) - - def forward(self, x): - hx = x - hxin = self.rebnconvin(hx) - - hx1 = self.rebnconv1(hxin) - hx = self.pool1(hx1) - - hx2 = self.rebnconv2(hx) - hx = self.pool2(hx2) - - hx3 = self.rebnconv3(hx) - hx = self.pool3(hx3) - - hx4 = self.rebnconv4(hx) - hx = self.pool4(hx4) - - hx5 = self.rebnconv5(hx) - - hx6 = self.rebnconv6(hx5) - - hx5d = self.rebnconv5d(torch.cat((hx6, hx5), 1)) - hx5dup = _upsample_like(hx5d, hx4) - - hx4d = self.rebnconv4d(torch.cat((hx5dup, hx4), 1)) - hx4dup = _upsample_like(hx4d, hx3) - - hx3d = self.rebnconv3d(torch.cat((hx4dup, hx3), 1)) - hx3dup = _upsample_like(hx3d, hx2) - - hx2d = self.rebnconv2d(torch.cat((hx3dup, hx2), 1)) - hx2dup = _upsample_like(hx2d, hx1) - - hx1d = self.rebnconv1d(torch.cat((hx2dup, hx1), 1)) - - return hx1d + hxin - - -class RSU5(nn.Module): - def __init__(self, in_ch=3, mid_ch=12, out_ch=3): - super(RSU5, self).__init__() - - self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1) - - self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1) - self.pool1 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=1) - self.pool2 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=1) - self.pool3 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=1) - - self.rebnconv5 = REBNCONV(mid_ch, mid_ch, dirate=2) - - self.rebnconv4d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv3d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv2d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv1d = REBNCONV(mid_ch*2, out_ch, dirate=1) - - def forward(self, x): - hx = x - hxin = self.rebnconvin(hx) - - hx1 = self.rebnconv1(hxin) - hx = self.pool1(hx1) - - hx2 = self.rebnconv2(hx) - hx = self.pool2(hx2) - - hx3 = self.rebnconv3(hx) - hx = self.pool3(hx3) - - hx4 = self.rebnconv4(hx) - - hx5 = self.rebnconv5(hx4) - - hx4d = self.rebnconv4d(torch.cat((hx5, hx4), 1)) - hx4dup = _upsample_like(hx4d, hx3) - - hx3d = self.rebnconv3d(torch.cat((hx4dup, hx3), 1)) - hx3dup = _upsample_like(hx3d, hx2) - - hx2d = self.rebnconv2d(torch.cat((hx3dup, hx2), 1)) - hx2dup = _upsample_like(hx2d, hx1) - - hx1d = self.rebnconv1d(torch.cat((hx2dup, hx1), 1)) - - return hx1d + hxin - - -class RSU4(nn.Module): - def __init__(self, in_ch=3, mid_ch=12, out_ch=3): - super(RSU4, self).__init__() - - self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1) - - self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1) - self.pool1 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=1) - self.pool2 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=1) - - self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=2) - - self.rebnconv3d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv2d = REBNCONV(mid_ch*2, mid_ch, dirate=1) - self.rebnconv1d = REBNCONV(mid_ch*2, out_ch, dirate=1) - - def forward(self, x): - hx = x - hxin = self.rebnconvin(hx) - - hx1 = self.rebnconv1(hxin) - hx = self.pool1(hx1) - - hx2 = self.rebnconv2(hx) - hx = self.pool2(hx2) - - hx3 = self.rebnconv3(hx) - - hx4 = self.rebnconv4(hx3) - - hx3d = self.rebnconv3d(torch.cat((hx4, hx3), 1)) - hx3dup = _upsample_like(hx3d, hx2) - - hx2d = self.rebnconv2d(torch.cat((hx3dup, hx2), 1)) - hx2dup = _upsample_like(hx2d, hx1) - - hx1d = self.rebnconv1d(torch.cat((hx2dup, hx1), 1)) - - return hx1d + hxin - - -class RSU4F(nn.Module): - def __init__(self, in_ch=3, mid_ch=12, out_ch=3): - super(RSU4F, self).__init__() - - self.rebnconvin = REBNCONV(in_ch, out_ch, dirate=1) - - self.rebnconv1 = REBNCONV(out_ch, mid_ch, dirate=1) - self.rebnconv2 = REBNCONV(mid_ch, mid_ch, dirate=2) - self.rebnconv3 = REBNCONV(mid_ch, mid_ch, dirate=4) - - self.rebnconv4 = REBNCONV(mid_ch, mid_ch, dirate=8) - - self.rebnconv3d = REBNCONV(mid_ch*2, mid_ch, dirate=4) - self.rebnconv2d = REBNCONV(mid_ch*2, mid_ch, dirate=2) - self.rebnconv1d = REBNCONV(mid_ch*2, out_ch, dirate=1) - - def forward(self, x): - hx = x - hxin = self.rebnconvin(hx) - - hx1 = self.rebnconv1(hxin) - hx2 = self.rebnconv2(hx1) - hx3 = self.rebnconv3(hx2) - - hx4 = self.rebnconv4(hx3) - - hx3d = self.rebnconv3d(torch.cat((hx4, hx3), 1)) - hx2d = self.rebnconv2d(torch.cat((hx3d, hx2), 1)) - hx1d = self.rebnconv1d(torch.cat((hx2d, hx1), 1)) - - return hx1d + hxin - - -class U2NET(nn.Module): - def __init__(self, in_ch=3, out_ch=1): - super(U2NET, self).__init__() - - self.stage1 = RSU7(in_ch, 32, 64) - self.pool12 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.stage2 = RSU6(64, 32, 128) - self.pool23 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.stage3 = RSU5(128, 64, 256) - self.pool34 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.stage4 = RSU4(256, 128, 512) - self.pool45 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.stage5 = RSU4F(512, 256, 512) - self.pool56 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.stage6 = RSU4F(512, 256, 512) - - # decoder - self.stage5d = RSU4F(1024, 256, 512) - self.stage4d = RSU4(1024, 128, 256) - self.stage3d = RSU5(512, 64, 128) - self.stage2d = RSU6(256, 32, 64) - self.stage1d = RSU7(128, 16, 64) - - self.side1 = nn.Conv2d(64, out_ch, 3, padding=1) - self.side2 = nn.Conv2d(64, out_ch, 3, padding=1) - self.side3 = nn.Conv2d(128, out_ch, 3, padding=1) - self.side4 = nn.Conv2d(256, out_ch, 3, padding=1) - self.side5 = nn.Conv2d(512, out_ch, 3, padding=1) - self.side6 = nn.Conv2d(512, out_ch, 3, padding=1) - - self.outconv = nn.Conv2d(6*out_ch, out_ch, 1) - - def forward(self, x): - hx = x - - # stage 1 - hx1 = self.stage1(hx) - hx = self.pool12(hx1) - - # stage 2 - hx2 = self.stage2(hx) - hx = self.pool23(hx2) - - # stage 3 - hx3 = self.stage3(hx) - hx = self.pool34(hx3) - - # stage 4 - hx4 = self.stage4(hx) - hx = self.pool45(hx4) - - # stage 5 - hx5 = self.stage5(hx) - hx = self.pool56(hx5) - - # stage 6 - hx6 = self.stage6(hx) - hx6up = _upsample_like(hx6, hx5) - - # decoder - hx5d = self.stage5d(torch.cat((hx6up, hx5), 1)) - hx5dup = _upsample_like(hx5d, hx4) - - hx4d = self.stage4d(torch.cat((hx5dup, hx4), 1)) - hx4dup = _upsample_like(hx4d, hx3) - - hx3d = self.stage3d(torch.cat((hx4dup, hx3), 1)) - hx3dup = _upsample_like(hx3d, hx2) - - hx2d = self.stage2d(torch.cat((hx3dup, hx2), 1)) - hx2dup = _upsample_like(hx2d, hx1) - - hx1d = self.stage1d(torch.cat((hx2dup, hx1), 1)) - - # side output - d1 = self.side1(hx1d) - - d2 = self.side2(hx2d) - d2 = _upsample_like(d2, d1) - - d3 = self.side3(hx3d) - d3 = _upsample_like(d3, d1) - - d4 = self.side4(hx4d) - d4 = _upsample_like(d4, d1) - - d5 = self.side5(hx5d) - d5 = _upsample_like(d5, d1) - - d6 = self.side6(hx6) - d6 = _upsample_like(d6, d1) - - d0 = self.outconv(torch.cat((d1, d2, d3, d4, d5, d6), 1)) - - return torch.sigmoid(d0), torch.sigmoid(d1), torch.sigmoid(d2), torch.sigmoid(d3), torch.sigmoid(d4), torch.sigmoid(d5), torch.sigmoid(d6) - - -class U2NETP(nn.Module): - """Smaller/faster U2Net variant (u2netp)""" - - def __init__(self, in_ch=3, out_ch=1): - super(U2NETP, self).__init__() - - self.stage1 = RSU7(in_ch, 16, 64) - self.pool12 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.stage2 = RSU6(64, 16, 64) - self.pool23 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.stage3 = RSU5(64, 16, 64) - self.pool34 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.stage4 = RSU4(64, 16, 64) - self.pool45 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.stage5 = RSU4F(64, 16, 64) - self.pool56 = nn.MaxPool2d(2, stride=2, ceil_mode=True) - - self.stage6 = RSU4F(64, 16, 64) - - # decoder - self.stage5d = RSU4F(128, 16, 64) - self.stage4d = RSU4(128, 16, 64) - self.stage3d = RSU5(128, 16, 64) - self.stage2d = RSU6(128, 16, 64) - self.stage1d = RSU7(128, 16, 64) - - self.side1 = nn.Conv2d(64, out_ch, 3, padding=1) - self.side2 = nn.Conv2d(64, out_ch, 3, padding=1) - self.side3 = nn.Conv2d(64, out_ch, 3, padding=1) - self.side4 = nn.Conv2d(64, out_ch, 3, padding=1) - self.side5 = nn.Conv2d(64, out_ch, 3, padding=1) - self.side6 = nn.Conv2d(64, out_ch, 3, padding=1) - - self.outconv = nn.Conv2d(6*out_ch, out_ch, 1) - - def forward(self, x): - hx = x - - hx1 = self.stage1(hx) - hx = self.pool12(hx1) - - hx2 = self.stage2(hx) - hx = self.pool23(hx2) - - hx3 = self.stage3(hx) - hx = self.pool34(hx3) - - hx4 = self.stage4(hx) - hx = self.pool45(hx4) - - hx5 = self.stage5(hx) - hx = self.pool56(hx5) - - hx6 = self.stage6(hx) - hx6up = _upsample_like(hx6, hx5) - - hx5d = self.stage5d(torch.cat((hx6up, hx5), 1)) - hx5dup = _upsample_like(hx5d, hx4) - - hx4d = self.stage4d(torch.cat((hx5dup, hx4), 1)) - hx4dup = _upsample_like(hx4d, hx3) - - hx3d = self.stage3d(torch.cat((hx4dup, hx3), 1)) - hx3dup = _upsample_like(hx3d, hx2) - - hx2d = self.stage2d(torch.cat((hx3dup, hx2), 1)) - hx2dup = _upsample_like(hx2d, hx1) - - hx1d = self.stage1d(torch.cat((hx2dup, hx1), 1)) - - d1 = self.side1(hx1d) - - d2 = self.side2(hx2d) - d2 = _upsample_like(d2, d1) - - d3 = self.side3(hx3d) - d3 = _upsample_like(d3, d1) - - d4 = self.side4(hx4d) - d4 = _upsample_like(d4, d1) - - d5 = self.side5(hx5d) - d5 = _upsample_like(d5, d1) - - d6 = self.side6(hx6) - d6 = _upsample_like(d6, d1) - - d0 = self.outconv(torch.cat((d1, d2, d3, d4, d5, d6), 1)) - - return torch.sigmoid(d0), torch.sigmoid(d1), torch.sigmoid(d2), torch.sigmoid(d3), torch.sigmoid(d4), torch.sigmoid(d5), torch.sigmoid(d6) diff --git a/backend/entrypoint.sh b/backend/entrypoint.sh index 65f3d26..f51828e 100644 --- a/backend/entrypoint.sh +++ b/backend/entrypoint.sh @@ -5,8 +5,9 @@ # This script runs when the container starts. It: # 1. Initializes the database # 2. Downloads SAM model automatically (can be disabled with AUTO_DOWNLOAD_SAM=false) -# 3. Downloads sample eye images if the catalog is empty -# 4. Starts the FastAPI server +# 3. Downloads U2Net model automatically (can be disabled with AUTO_DOWNLOAD_U2NET=false) +# 4. Downloads sample eye images if the catalog is empty +# 5. Starts the FastAPI server # ============================================================================= set -e @@ -61,6 +62,36 @@ else fi fi +echo "" +echo "Checking U2Net model (Remove Background)..." +echo "------------------------------------------" +if [ -f "/app/data/models/u2net.onnx" ] || [ -f "/app/data/models/u2netp.onnx" ]; then + echo "✓ U2Net model found - Remove Background will use local AI (free, offline)" +else + # Auto-download U2Net unless explicitly disabled + AUTO_DOWNLOAD_U2NET="${AUTO_DOWNLOAD_U2NET:-true}" + if [ "$AUTO_DOWNLOAD_U2NET" = "true" ]; then + echo "U2Net model not found. Downloading automatically..." + echo "(This is a one-time ~176MB download that persists across rebuilds)" + echo "" + python /scripts/download_u2net_model.py u2net || { + echo "" + echo "⚠ U2Net download failed (non-fatal)" + echo " Remove Background will fall back to rembg (if installed)" + echo " To retry later: docker exec -it ai-photo-edit-backend python /scripts/download_u2net_model.py u2net" + } + else + echo "" + echo "⚠ U2Net model not found (AUTO_DOWNLOAD_U2NET=false)" + echo "" + echo " Remove Background will fall back to rembg (if installed)" + echo "" + echo " To enable FREE offline Remove Background, run:" + echo " docker exec -it ai-photo-edit-backend python /scripts/download_u2net_model.py u2net" + echo "" + fi +fi + echo "" echo "Checking GPU capabilities..." echo "------------------------------------------" diff --git a/docker-compose.gpu.yml b/docker-compose.gpu.yml index 02dd9ff..eda3690 100644 --- a/docker-compose.gpu.yml +++ b/docker-compose.gpu.yml @@ -120,6 +120,7 @@ services: - SECRET_KEY=${SECRET_KEY:-change-this-secret-key-in-production} - CORS_ORIGINS=* - AUTO_DOWNLOAD_SAM=${AUTO_DOWNLOAD_SAM:-true} + - AUTO_DOWNLOAD_U2NET=${AUTO_DOWNLOAD_U2NET:-true} # ── NVIDIA GPU passthrough ──────────────────────────────────────────────── # Requires nvidia-container-toolkit; see prerequisites at top of this file.