feat: initial project setup
This commit is contained in:
@@ -0,0 +1,6 @@
|
||||
/python
|
||||
*.pyc
|
||||
/.workbuddy
|
||||
/outputs
|
||||
/models
|
||||
run.bat
|
||||
@@ -0,0 +1,213 @@
|
||||
# BiRefNet WebUI
|
||||
|
||||
**自包含**的网页版智能抠图工具:模型代码、权重、Python 运行时全部在项目目录内,
|
||||
不依赖 ComfyUI,整个文件夹拷到别的机器(有 NVIDIA 显卡)也能直接跑。
|
||||
|
||||
浏览器里拖图 → 一键抠图 → 对比 → 下载;也可以直接填两个目录,整目录批量抠图。
|
||||
|
||||
- **零 Web 框架依赖**:标准库 `http.server` 实现,不装 Gradio / FastAPI / Flask
|
||||
- **自带运行时**:`python/` 内置 Python 3.12 + torch(CUDA),无需配置环境
|
||||
- **GPU 串行队列**:同一时刻只跑一个模型,8G 显存也能稳定使用
|
||||
- **目录批量**:给一个图片目录 + 一个输出目录,结果按 `RMBG_原名_时间戳.后缀` 落盘
|
||||
- 数值行为与 ComfyUI 节点 `comfyui_birefnet_ll` 对齐:同款预处理、`fast-foreground-estimation` 前景精修
|
||||
|
||||
## 快速开始
|
||||
|
||||
```
|
||||
方式一:双击 run.bat
|
||||
方式二:python\python.exe Webui.py
|
||||
```
|
||||
|
||||
启动后自动打开 `http://127.0.0.1:7861/`。
|
||||
|
||||
> 把整个项目文件夹移动或拷贝到其它路径/机器后仍可直接运行;
|
||||
> 若目标机器没有 NVIDIA 显卡,会自动退回 CPU(较慢)。
|
||||
|
||||
## 目录结构
|
||||
|
||||
```
|
||||
BiRefNet_WebUI/
|
||||
├─ Webui.py 启动入口
|
||||
├─ run.bat 双击启动(优先用自带 python\ 运行时)
|
||||
├─ selftest.py 端到端自检(python\python.exe selftest.py)
|
||||
├─ selftest_batch.py 批量抠图自检(纯逻辑 + 假引擎流水线 + HTTP 契约,秒级)
|
||||
├─ python/ ★ 自带 Python 3.12 运行时(torch 2.13 + CUDA)
|
||||
├─ models/ ★ BiRefNet 权重(birefnet / Portrait)
|
||||
├─ vendor/comfyui_birefnet_ll/ ★ 模型代码(来自 ComfyUI 节点,含 LICENSE)
|
||||
├─ birefnet_web/ 服务端
|
||||
│ ├─ compat.py 模型代码定位与 folder_paths 垫片
|
||||
│ ├─ fsutil.py 图片后缀集合、文件名清洗(不依赖 torch)
|
||||
│ ├─ imageops.py 预处理 / 后处理
|
||||
│ ├─ batch.py 批量:路径校验、目录扫描、输出命名
|
||||
│ ├─ engine.py 模型扫描、加载缓存与推理
|
||||
│ └─ server.py HTTP 服务、REST API、任务队列
|
||||
├─ web/ 前端(原生 HTML/CSS/JS,无构建步骤)
|
||||
│ ├─ zip.js 纯 JS 的 ZIP 打包器(跨任务打包结果用)
|
||||
│ └─ test_app.js 前端回归测试(node web/test_app.js,无需浏览器)
|
||||
├─ outputs/ 结果输出(按任务分目录)
|
||||
└─ config.json 首次启动自动生成
|
||||
```
|
||||
|
||||
★ = 随项目分发的核心资源。
|
||||
|
||||
## 使用
|
||||
|
||||
1. **左侧「图片」**:拖入 / 点击选择 / `Ctrl+V` 粘贴。**一次只放一张**:新图会替换上一张,
|
||||
重复拖入同一张会被去重提示;一次拖入多张时只保留最后一张
|
||||
2. **「模型」**:选择权重(默认扫描项目 `models/`,也可加别的目录)
|
||||
3. **「参数」**:按需调整(都有与节点一致的默认值)
|
||||
4. 点击 **开始抠图**,右侧实时显示进度;**新任务的结果追加到列表末尾,
|
||||
已有记录不会被清空**(换参数、传新图、重新处理都不影响)
|
||||
5. 结果卡片上**左右拖动**分割线对比原图与结果,支持 抠图/遮罩/原图 三种视图、单图下载、ZIP 打包
|
||||
6. **结果列表交互**:
|
||||
- 每条记录的文件名下方有一行**参数摘要**(该次处理使用的模型/尺寸/背景等,
|
||||
完整参数在悬停提示里),跨任务累积成历史
|
||||
- **点击文件名或「重新处理」** → 把该记录的原图放回待处理槽,**并把它当时使用的
|
||||
参数重新赋值到左侧面板**,可直接重跑(同一记录重复点击会去重)
|
||||
- **参数记忆**:每张图最近一次使用的参数按文件名自动保存(localStorage),
|
||||
之后再把同名图片拖进来时自动套用;同一张图多次处理时总是记录最新的参数
|
||||
- **点击卡片本体** → 选中该记录(主题色描边),再点一次取消,`Esc` 也可取消
|
||||
- **点卡片右上角 ×** → 从结果列表移除该条记录(仅移除列表项,`outputs/` 下的文件保留)
|
||||
- **打包下载 ZIP** 覆盖列表里的全部记录(可能来自多个任务,同名自动加序号)
|
||||
- 刷新页面后,服务端还在内存里的任务历史会自动恢复显示(含参数)
|
||||
|
||||
> 需要按目录整批处理时,用左侧面板 **4 · 批量处理**,详见下一节。
|
||||
|
||||
## 批量处理(整目录抠图)
|
||||
|
||||
不想一张张拖图时,用左侧面板 **4 · 批量处理**:
|
||||
|
||||
1. **图片目录(模板路径)**:待处理图片所在目录,例 `F:\photos\待抠图`
|
||||
2. **输出目录**:结果写到哪;**留空=项目 `outputs/` 目录**(输入框默认已填好绝对路径)
|
||||
3. 勾选 **包含子目录** 则递归扫描子目录
|
||||
4. 点 **开始批量抠图** → 结果区顶部出现逐文件清单:`源文件名 → 输出文件名 → 状态 / 耗时`
|
||||
|
||||
抠图参数(模型、尺寸、背景、精修…)**沿用左侧「模型 / 参数」面板当前设置**,与单图模式完全一致。
|
||||
|
||||
### 目录与落盘规则
|
||||
|
||||
| 行为 | 说明 |
|
||||
| --- | --- |
|
||||
| 路径必须已存在 | 输入目录、输出目录都必须真实存在;**本工具不会创建任何目录**,路径不对直接在界面上报错 |
|
||||
| 默认输出目录 | 项目内 `outputs/`(前端按服务端返回的绝对路径预填) |
|
||||
| 命名规则 | `RMBG_<原文件名主干>_<Unix 时间戳>.<后缀>`;主干超过 **20 字符**按字符截断(中文同样按字符算) |
|
||||
| 后缀 | 默认沿用原后缀;但透明输出需要 alpha 通道,遇到 `jpg/jpeg/bmp` 会退化成 `.png`(否则会被压成白底) |
|
||||
| 遮罩文件 | 默认**不产出**;勾选「同时输出遮罩」后才额外写一份 `<主结果同名>_mask.png`(灰度,白=前景) |
|
||||
| 重名不覆盖 | 同一秒内的同名结果自动加 `_1`、`_2` … 递增避让 |
|
||||
| 扫描范围 | 只认图片后缀;**跳过 `RMBG_` 开头的文件**(避免把上次产物又抠一遍);输出目录在输入目录内时会被排除 |
|
||||
| 输入 = 输出 | 允许:结果与源图同目录,靠 `RMBG_` 前缀区分 |
|
||||
| 原图 | 批量模式**不拷贝、不修改**源图,就地读取 |
|
||||
| 结果展示 | 批量结果不铺成单图卡片(几百张会拖慢页面),只在结果区顶部列清单;文件本体已在输出目录,无需再下载 |
|
||||
|
||||
> 批量任务与单图任务共用同一个串行队列(同一时刻只跑一张),不会抢显存;关掉页面不影响服务端继续跑完。
|
||||
|
||||
## 参数说明
|
||||
|
||||
「模型 / 设备」与「参数」两个面板中每一项的含义(均有与 ComfyUI 节点一致的默认值,改了会记入参数摘要):
|
||||
|
||||
### 设备与精度
|
||||
|
||||
| 参数 | 说明 | 默认 |
|
||||
| -- | ------------------------------------------------------------------------ | ---- |
|
||||
| 设备 | 自动(GPU)/ CPU | auto |
|
||||
| 精度 | auto = fp32 权重 + fp16 autocast(官方推荐做法);float16 最省显存;出现 NaN 会自动回退 fp32 重算 | auto |
|
||||
| 架构 | v1(新版 safetensors)/ old(旧版 .pth)/ 自动识别 | auto |
|
||||
|
||||
### 预处理尺寸(送入模型前的缩放方式)
|
||||
|
||||
BiRefNet 推理前会把原图缩放到固定输入尺寸,这一项决定怎么缩:
|
||||
|
||||
| 模式 | 说明 |
|
||||
| -------------------- | ------------------------------------ |
|
||||
| 固定 1024×1024(square) | 官方推荐做法:不论原图比例,直接拉伸到 1024×1024,精度最稳定 |
|
||||
| 长边自适应(longest) | 保持宽高比,把长边缩放到指定值(边长对齐 32 的倍数),另一边按比例算 |
|
||||
| 自定义宽高(custom) | 手动指定「宽 / 高」两个输入框(默认各 1024),不保持原图比例 |
|
||||
|
||||
> 输入尺寸越小推理越快、越省显存(如 512);越大细节越丰富但更慢。
|
||||
|
||||
### 其余参数
|
||||
|
||||
| 参数 | 说明 | 默认 |
|
||||
| ------- | ------------------------------------------------------------------------------------- | -------- |
|
||||
| 插值方式 | 预处理缩图与遮罩还原回原尺寸时使用的插值算法(bilinear / nearest / bicubic 等) | bilinear |
|
||||
| 遮罩阈值 | 大于 0 时,把低于该阈值的 alpha 概率直接置零(不做二值化,只清零弱值),可清理弱边缘杂讯;0 = 不处理 | 0 |
|
||||
| 背景 | 抠图结果的合成方式:**透明(PNG Alpha)** 输出带 alpha 通道的 PNG;**纯色** 则把前景合成到指定颜色上(含绿幕、色板快选) | 透明 |
|
||||
| 前景精修 | fast-foreground-estimation:按遮罩估计前景色,去除半透明边缘残留的原背景色(白底抠图边缘发白就是它治的)。超大图(超过像素上限)会自动跳过并提示 | 开 |
|
||||
| — 大核 r1 | 精修第一轮 box blur 的核尺寸,用于估计整体前景色 | 90 |
|
||||
| — 小核 r2 | 精修第二轮 box blur 的核尺寸,用于收细边缘 | 6 |
|
||||
| 同时输出遮罩 | 额外输出一张 8bit 灰度 PNG(黑白遮罩,白色 = 前景),可用于后期合成。**默认关闭**,需要时才勾 | 关 |
|
||||
| 最终尺寸 · 最长边 | 按原图比例等比缩放结果,使**最长边等于**该值;0 = 保持原图尺寸。例:2000×3000 填 1440 → 960×1440;3000×2000 填 1440 → 1440×960 | 0 |
|
||||
|
||||
> 「最终尺寸 · 最长边」精确定义:设原图(抠图结果原尺寸)为 `W×H`,填充值为 `L`
|
||||
>
|
||||
> - `L = 0`:不缩放,原样输出(默认)
|
||||
> - `L > 0`:`k = L / max(W, H)`,输出 `round(W·k) × round(H·k)`,最长边恰好等于 `L`;短边按比例取整,可能与理论值差 1 像素
|
||||
> - 原图最长边**小于** `L` 时会**放大**(这点与旧的「输出长边上限」不同,旧参数只缩不放,已被本参数取代)
|
||||
> - 缩放同时作用于抠图结果与(勾选时的)遮罩文件,两者始终保持同尺寸
|
||||
|
||||
## 命令行参数
|
||||
|
||||
```
|
||||
python\python.exe Webui.py [--host 0.0.0.0] [--port 7861] [--device auto|cpu|cuda]
|
||||
[--dtype auto|float32|float16|bfloat16]
|
||||
[--model-dir 路径]... 追加模型目录
|
||||
[--node-dir 路径] 指定模型代码目录(默认用项目内 vendor/)
|
||||
[--output-dir 路径] 结果输出目录(默认 ./outputs)
|
||||
[--max-cached-models N] 显存中驻留的模型数(默认 1)
|
||||
[--no-browser] [--save-config] [--check]
|
||||
```
|
||||
|
||||
`--check` 只做环境自检(依赖、GPU、模型扫描)后退出,适合排障。
|
||||
|
||||
## 模型代码与权重
|
||||
|
||||
模型代码内联在 `vendor/comfyui_birefnet_ll/`(来源:ComfyUI 节点 comfyui_birefnet_ll,
|
||||
遵循其 LICENSE)。`birefnet_web/compat.py` 会在 import 阶段注入一个最小 `folder_paths`
|
||||
垫片(模型包只在「加载骨干预训练权重」时用到它,本工具始终以 `bb_pretrained=False`
|
||||
构建,因此垫片返回 None 即可),从而脱离 ComfyUI 独立运行。
|
||||
|
||||
权重文件(放入 `models/` 即可,扫描时自动识别):
|
||||
|
||||
- 新版 `*.safetensors`:[ZhengPeng7 的 HuggingFace 仓库](https://huggingface.co/ZhengPeng7),如 `General.safetensors`、`Portrait.safetensors`
|
||||
- 旧版 `BiRefNet-DIS_ep580.pth` / `BiRefNet-ep480.pth`
|
||||
- 骨干权重(`swin_*` / `pvt_*`)会被自动忽略,不会被当成抠图模型
|
||||
|
||||
也可用 `--model-dir` 或 `config.json` 的 `model_dirs` 追加其它目录
|
||||
(例如继续共用 ComfyUI 的 `models/BiRefNet`,两处权重会合并去重显示)。
|
||||
|
||||
## HTTP API
|
||||
|
||||
| 方法 | 路径 | 说明 |
|
||||
| ---- | ----------------------------------------- | ----------------------------------------------- |
|
||||
| GET | `/api/state` | 环境、模型列表、默认参数 |
|
||||
| GET | `/api/models` · POST `/api/models/reload` | 模型列表 / 重新扫描 |
|
||||
| POST | `/api/tasks` | multipart 提交:`files`(可多份)+ `options`(JSON,字段同上) |
|
||||
| POST | `/api/batch` | JSON 提交目录批量:`input_dir` / `output_dir`(可空=项目 outputs)/ `recursive` / `options`;两个目录都必须已存在 |
|
||||
| GET | `/api/tasks` | 任务历史列表(含各自参数与 `mode`,前端据此恢复结果历史) |
|
||||
| GET | `/api/tasks/<id>` | 任务状态与进度(前端 400ms 轮询);`?tail=N` 只回传最后 N 张,批量任务靠它压载荷 |
|
||||
| POST | `/api/tasks/<id>/cancel` | 取消(阶段粒度) |
|
||||
| GET | `/api/tasks/<id>/file/<index>/<kind>` | 取结果,kind = cutout / mask / original |
|
||||
| GET | `/api/tasks/<id>/zip?kind=cutout` | 打包下载 |
|
||||
| POST | `/api/engine/unload` | 释放显存 |
|
||||
| POST | `/api/shutdown` | 关闭服务 |
|
||||
|
||||
## 常见问题
|
||||
|
||||
- **双击 run.bat 一闪而过**:确认 `run.bat` 是 CRLF 换行(纯 LF 的批处理会被 cmd 静默终止)
|
||||
- **显存不足(OOM)**:精度选 `float16`,或预处理尺寸改小(如 512);默认同时只驻留 1 个模型
|
||||
- **没有可用模型**:确认 `*.safetensors` 在 `models/` 下,或用 `--model-dir` 指定
|
||||
- **结果边缘有原背景色残留**:确认「前景精修」开启
|
||||
- **批量时报「输出目录不存在」**:本工具不会自动建目录,请先在资源管理器里把输出目录建好再填
|
||||
- **批量产物多出 `_mask.png`**:说明勾了「同时输出遮罩」(默认不勾);取消勾选即可只留抠图结果
|
||||
- **输出尺寸不对**:检查「最终尺寸 · 最长边」——填了非 0 值就会把结果等比缩放到该最长边(小图会被放大),想原样输出请填 0
|
||||
- **旧版本里设过「输出长边上限」**:该参数已改名为「最终尺寸 · 最长边」,旧值会自动带到新输入框;浏览器里存的老参数还会自动修正一次「同时输出遮罩」的默认值
|
||||
- **批量结果名被截断**:设计如此——原文件名主干超过 20 字符会截断,避免超长文件名;时间戳用来区分多次处理
|
||||
- **CPU 很慢**:正常现象,建议 512 分辨率 + float32
|
||||
- **想换 Python**:设环境变量 `BIREFNET_PYTHON=<解释器路径>`(需已装 `requirements.txt` 依赖);
|
||||
缺依赖时入口会自动切换到自带的 `python\` 运行时
|
||||
|
||||
## 环境要求
|
||||
|
||||
自带运行时已就绪,无需安装。若用自己的解释器:Python 3.9+,
|
||||
`pip install -r requirements.txt`(torch / torchvision / numpy / pillow / safetensors /
|
||||
timm / einops / kornia / huggingface_hub / tqdm,可选 opencv-python)。
|
||||
@@ -0,0 +1,357 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""BiRefNet WebUI 启动入口。
|
||||
|
||||
本项目**自包含**:模型代码、权重、Python 运行时都在项目目录内,不再需要 ComfyUI。
|
||||
|
||||
F:\\BiRefNet_WebUI\\
|
||||
├─ Webui.py
|
||||
├─ python\\ 自带 Python 3.12 运行时(torch + CUDA)
|
||||
├─ models\\ BiRefNet 权重(*.safetensors)
|
||||
├─ vendor\\comfyui_birefnet_ll\\ 模型代码
|
||||
└─ birefnet_web\\ web\\ outputs\\
|
||||
|
||||
启动方式(任选其一):
|
||||
|
||||
run.bat 双击即可
|
||||
python\\python.exe Webui.py 用自带运行时
|
||||
python Webui.py 用系统 Python(缺依赖时会自动切到自带运行时)
|
||||
|
||||
常用参数:
|
||||
--port 7861 监听端口
|
||||
--host 0.0.0.0 允许局域网访问
|
||||
--device cpu 强制 CPU 推理
|
||||
--dtype float32 强制单精度(默认 GPU 上 fp32 权重 + fp16 autocast)
|
||||
--model-dir D:\\models 追加模型目录(可重复)
|
||||
--node-dir <path> 指定模型代码目录(默认用项目内 vendor/)
|
||||
--no-browser 启动后不自动打开浏览器
|
||||
--check 仅自检环境后退出
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
import webbrowser
|
||||
from pathlib import Path
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parent
|
||||
CONFIG_PATH = PROJECT_ROOT / "config.json"
|
||||
|
||||
# 内置 Python 带有 python3xx._pth(隔离模式),脚本所在目录不会自动进入 sys.path,
|
||||
# 这里显式补上,保证无论用哪个解释器都能 import 到本项目包。
|
||||
if str(PROJECT_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
|
||||
#: 项目自带的独立 Python 运行时
|
||||
BUNDLED_PYTHON = PROJECT_ROOT / "python" / "python.exe"
|
||||
|
||||
#: 防止自动切换解释器时无限递归
|
||||
_RELAUNCH_FLAG = "BIREFNET_RELAUNCHED"
|
||||
|
||||
#: 运行必需的三方包 -> pip 包名
|
||||
_REQUIRED = (
|
||||
("torch", "torch"),
|
||||
("torchvision", "torchvision"),
|
||||
("numpy", "numpy"),
|
||||
("PIL", "pillow"),
|
||||
("safetensors", "safetensors"),
|
||||
("timm", "timm"),
|
||||
("einops", "einops"),
|
||||
("kornia", "kornia"),
|
||||
)
|
||||
|
||||
CONFIG_DEFAULTS: Dict[str, object] = {
|
||||
"host": "127.0.0.1",
|
||||
"port": 7861,
|
||||
"node_dir": "",
|
||||
"model_dirs": [],
|
||||
"output_dir": "outputs",
|
||||
"device": "auto",
|
||||
"dtype": "auto",
|
||||
"max_cached_models": 1,
|
||||
"max_tasks": 50,
|
||||
"open_browser": True,
|
||||
}
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 依赖自检 / 解释器自动切换
|
||||
# --------------------------------------------------------------------------- #
|
||||
def missing_packages() -> List[str]:
|
||||
"""返回缺失的 pip 包名列表。"""
|
||||
missing: List[str] = []
|
||||
for mod, pkg in _REQUIRED:
|
||||
try:
|
||||
__import__(mod)
|
||||
except Exception:
|
||||
missing.append(pkg)
|
||||
return missing
|
||||
|
||||
|
||||
def bundled_python() -> Optional[Path]:
|
||||
"""项目自带运行时是否可用。"""
|
||||
return BUNDLED_PYTHON if BUNDLED_PYTHON.is_file() else None
|
||||
|
||||
|
||||
def relaunch_with_bundled_python(missing: List[str]) -> None:
|
||||
"""当前解释器缺依赖、而项目自带运行时可用时,直接把进程换成自带运行时。"""
|
||||
if not missing or os.environ.get(_RELAUNCH_FLAG):
|
||||
return
|
||||
exe = bundled_python()
|
||||
if exe is None:
|
||||
return
|
||||
try:
|
||||
if Path(sys.executable).resolve() == exe.resolve():
|
||||
return
|
||||
except OSError:
|
||||
pass
|
||||
os.environ[_RELAUNCH_FLAG] = "1"
|
||||
print(f"[环境] 当前解释器缺少 {', '.join(missing)},自动切换到项目自带运行时:{exe}", flush=True)
|
||||
try:
|
||||
os.execv(str(exe), [str(exe), str(Path(__file__).resolve()), *sys.argv[1:]])
|
||||
except OSError as exc: # 切换失败就继续走下面的报错分支
|
||||
print(f"[环境] 切换失败:{exc}", file=sys.stderr, flush=True)
|
||||
|
||||
|
||||
def preflight() -> None:
|
||||
"""尽早给出人话的依赖缺失提示,而不是一堆 ImportError 堆栈。"""
|
||||
missing = missing_packages()
|
||||
if missing:
|
||||
relaunch_with_bundled_python(missing)
|
||||
if missing:
|
||||
print("缺少依赖:" + ", ".join(missing), file=sys.stderr)
|
||||
if bundled_python() is None:
|
||||
print(
|
||||
f"项目自带运行时不存在({BUNDLED_PYTHON})。\n"
|
||||
"请从 ComfyUI 便携版的 python_embeded 目录复制一份到项目的 python\\ 下,"
|
||||
"或用你自己的解释器安装依赖:",
|
||||
file=sys.stderr,
|
||||
)
|
||||
print(f' "{sys.executable}" -m pip install -r requirements.txt', file=sys.stderr)
|
||||
sys.exit(2)
|
||||
if not _has_cv2():
|
||||
print("[提示] 未检测到 opencv-python,前景精修将退化为高斯模糊实现(结果略有差异)")
|
||||
|
||||
|
||||
def _has_cv2() -> bool:
|
||||
try:
|
||||
import cv2 # noqa: F401
|
||||
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 配置
|
||||
# --------------------------------------------------------------------------- #
|
||||
def load_config(path: Path) -> Dict[str, object]:
|
||||
config = dict(CONFIG_DEFAULTS)
|
||||
if path.is_file():
|
||||
try:
|
||||
data = json.loads(path.read_text(encoding="utf-8"))
|
||||
if isinstance(data, dict):
|
||||
config.update(data)
|
||||
except (json.JSONDecodeError, OSError) as exc:
|
||||
print(f"[警告] 读取 {path} 失败,使用默认配置:{exc}")
|
||||
return config
|
||||
|
||||
|
||||
def save_config(path: Path, config: Dict[str, object]) -> None:
|
||||
try:
|
||||
path.write_text(json.dumps(config, ensure_ascii=False, indent=2), encoding="utf-8")
|
||||
except OSError as exc: # pragma: no cover
|
||||
print(f"[警告] 写入 {path} 失败:{exc}")
|
||||
|
||||
|
||||
def parse_args(argv: Optional[List[str]] = None) -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
prog="Webui.py",
|
||||
description="BiRefNet WebUI —— 自包含的网页版智能抠图服务",
|
||||
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
|
||||
)
|
||||
parser.add_argument("--host", default=None, help="监听地址,0.0.0.0 可局域网访问")
|
||||
parser.add_argument("--port", type=int, default=None, help="监听端口")
|
||||
parser.add_argument("--device", default=None, choices=["auto", "cpu", "cuda"], help="推理设备")
|
||||
parser.add_argument("--dtype", default=None, choices=["auto", "float32", "float16", "bfloat16"], help="推理精度")
|
||||
parser.add_argument("--model-dir", action="append", default=None, help="追加模型目录,可重复指定")
|
||||
parser.add_argument("--node-dir", default=None, help="模型代码目录(默认用项目内 vendor/)")
|
||||
parser.add_argument("--output-dir", default=None, help="结果输出目录")
|
||||
parser.add_argument("--max-cached-models", type=int, default=None, help="显存中最多缓存的模型数量")
|
||||
parser.add_argument("--max-tasks", type=int, default=None, help="内存中保留的任务记录数量")
|
||||
parser.add_argument("--config", default=str(CONFIG_PATH), help="配置文件路径")
|
||||
parser.add_argument("--no-browser", action="store_true", help="启动后不自动打开浏览器")
|
||||
parser.add_argument("--save-config", action="store_true", help="把本次参数写回配置文件")
|
||||
parser.add_argument("--check", action="store_true", help="仅做环境自检并打印可用模型,然后退出")
|
||||
return parser.parse_args(argv)
|
||||
|
||||
|
||||
def merge_config(args: argparse.Namespace) -> Dict[str, object]:
|
||||
config_path = Path(args.config).expanduser().resolve()
|
||||
config = load_config(config_path)
|
||||
|
||||
if args.host is not None:
|
||||
config["host"] = args.host
|
||||
if args.port is not None:
|
||||
config["port"] = args.port
|
||||
if args.device is not None:
|
||||
config["device"] = args.device
|
||||
if args.dtype is not None:
|
||||
config["dtype"] = args.dtype
|
||||
if args.output_dir is not None:
|
||||
config["output_dir"] = args.output_dir
|
||||
if args.max_cached_models is not None:
|
||||
config["max_cached_models"] = args.max_cached_models
|
||||
if args.max_tasks is not None:
|
||||
config["max_tasks"] = args.max_tasks
|
||||
if args.model_dir:
|
||||
config["model_dirs"] = list(dict.fromkeys(list(config.get("model_dirs") or []) + args.model_dir))
|
||||
if args.no_browser:
|
||||
config["open_browser"] = False
|
||||
|
||||
# 命令行显式指定的代码目录优先级最高,单独存一个键,避免被 config.json 里的旧值盖住
|
||||
config["_cli_node_dir"] = args.node_dir
|
||||
config["_path"] = str(config_path)
|
||||
config["_first_run"] = not config_path.is_file()
|
||||
return config
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 主流程
|
||||
# --------------------------------------------------------------------------- #
|
||||
def build_app(config: Dict[str, object]):
|
||||
from birefnet_web import compat
|
||||
from birefnet_web.engine import InferenceEngine, ModelRegistry
|
||||
from birefnet_web.server import TaskManager, WebUIServer
|
||||
|
||||
# 1) 模型代码目录:命令行 > 环境变量 > 项目内 vendor/ > 旧配置 > 外部 ComfyUI
|
||||
node_dir = compat.find_node_dir(config.get("_cli_node_dir"), config.get("node_dir") or None)
|
||||
|
||||
# 2) 权重目录:项目内 models/ 优先
|
||||
model_dirs: List[Path] = []
|
||||
for d in compat.default_model_dirs(node_dir):
|
||||
if d not in model_dirs:
|
||||
model_dirs.append(d)
|
||||
for item in config.get("model_dirs") or []:
|
||||
if item:
|
||||
p = Path(str(item)).expanduser()
|
||||
if p not in model_dirs:
|
||||
model_dirs.append(p)
|
||||
|
||||
registry = ModelRegistry(model_dirs)
|
||||
models = registry.scan()
|
||||
|
||||
output_root = Path(str(config.get("output_dir") or "outputs")).expanduser()
|
||||
if not output_root.is_absolute():
|
||||
output_root = PROJECT_ROOT / output_root
|
||||
|
||||
engine = InferenceEngine(registry, node_dir=node_dir, max_cached=int(config.get("max_cached_models") or 1))
|
||||
logger = lambda msg: print(f"[{time.strftime('%H:%M:%S')}] {msg}", flush=True) # noqa: E731
|
||||
tasks = TaskManager(engine, output_root, max_tasks=int(config.get("max_tasks") or 50), logger=logger)
|
||||
|
||||
server = WebUIServer(
|
||||
(str(config.get("host") or "127.0.0.1"), int(config.get("port") or 7861)),
|
||||
registry=registry,
|
||||
engine=engine,
|
||||
tasks=tasks,
|
||||
web_dir=PROJECT_ROOT / "web",
|
||||
output_root=output_root,
|
||||
node_dir=node_dir,
|
||||
logger=logger,
|
||||
project_root=PROJECT_ROOT,
|
||||
)
|
||||
return server, registry, models, node_dir, model_dirs, output_root
|
||||
|
||||
|
||||
def _rel_or_abs(path: Path) -> str:
|
||||
"""项目内的路径显示成相对路径,读起来更舒服。"""
|
||||
try:
|
||||
return str(path.resolve().relative_to(PROJECT_ROOT))
|
||||
except (ValueError, OSError):
|
||||
return str(path)
|
||||
|
||||
|
||||
def print_banner(server, models, node_dir, model_dirs, output_root, config) -> None:
|
||||
host, port = server.server_address[0], server.server_address[1]
|
||||
shown = "127.0.0.1" if host in ("0.0.0.0", "::") else host
|
||||
url = f"http://{shown}:{port}/"
|
||||
from birefnet_web.engine import describe_environment
|
||||
|
||||
env = describe_environment(server.engine)
|
||||
line = "─" * 66
|
||||
print(line)
|
||||
print(" BiRefNet WebUI · 智能抠图")
|
||||
print(line)
|
||||
print(f" 访问地址 {url}")
|
||||
if host in ("0.0.0.0", "::"):
|
||||
print(f" 局域网 同一网络下用 http://<本机IP>:{port}/ 访问")
|
||||
print(f" 运行时 {_rel_or_abs(Path(sys.executable))} (python {env['python']})")
|
||||
print(f" 推理库 torch {env['torch']} / torchvision {env['torchvision']}")
|
||||
if env["cuda"]:
|
||||
for dev in env["devices"]:
|
||||
print(f" GPU [{dev['index']}] {dev['name']} {dev['free_mem']}G 空闲 / {dev['total_mem']}G")
|
||||
else:
|
||||
print(" GPU 未检测到 CUDA,将使用 CPU(较慢)")
|
||||
print(f" 模型代码 {_rel_or_abs(node_dir)}")
|
||||
for d in model_dirs:
|
||||
mark = "✔" if Path(d).is_dir() else "✘ 不存在"
|
||||
print(f" 模型目录 {_rel_or_abs(Path(d))} {mark}")
|
||||
print(f" 可用模型 {len(models)} 个" + (f":{', '.join(m.name for m in models)}" if models else "(未找到权重文件)"))
|
||||
if not models:
|
||||
print(" ↳ 请把 *.safetensors 放到上面的模型目录")
|
||||
print(f" 输出目录 {_rel_or_abs(output_root)}")
|
||||
print(f" 配置文件 {config.get('_path')}")
|
||||
print(line)
|
||||
print(" 按 Ctrl+C 退出")
|
||||
print(line, flush=True)
|
||||
|
||||
|
||||
def main(argv: Optional[List[str]] = None) -> int:
|
||||
args = parse_args(argv)
|
||||
preflight()
|
||||
config = merge_config(args)
|
||||
|
||||
try:
|
||||
server, registry, models, node_dir, model_dirs, output_root = build_app(config)
|
||||
except Exception as exc:
|
||||
print(f"[启动失败] {type(exc).__name__}: {exc}", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
# 首启或显式要求时落盘配置(补全探测到的路径,便于下次直接改文件)
|
||||
config_to_save = {k: v for k, v in config.items() if not k.startswith("_")}
|
||||
config_to_save["node_dir"] = str(node_dir)
|
||||
if config.get("_first_run") or args.save_config:
|
||||
save_config(Path(str(config["_path"])), config_to_save)
|
||||
|
||||
if args.check:
|
||||
print_banner(server, models, node_dir, model_dirs, output_root, config)
|
||||
print("自检完成(--check 模式,不启动服务)")
|
||||
return 0
|
||||
|
||||
print_banner(server, models, node_dir, model_dirs, output_root, config)
|
||||
|
||||
if config.get("open_browser"):
|
||||
host, port = server.server_address[0], server.server_address[1]
|
||||
shown = "127.0.0.1" if host in ("0.0.0.0", "::") else host
|
||||
threading.Thread(
|
||||
target=lambda: (time.sleep(1.0), webbrowser.open(f"http://{shown}:{port}/")),
|
||||
daemon=True,
|
||||
).start()
|
||||
|
||||
try:
|
||||
server.serve_forever(poll_interval=0.4)
|
||||
except KeyboardInterrupt:
|
||||
print("\n收到 Ctrl+C,正在退出…")
|
||||
finally:
|
||||
server.engine.unload()
|
||||
server.server_close()
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,14 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""BiRefNet WebUI —— 基于 ComfyUI 节点 comfyui_birefnet_ll 的网页版抠图工具。
|
||||
|
||||
模块划分:
|
||||
compat.py 定位 ComfyUI 节点目录、注入 folder_paths 垫片、导入模型类
|
||||
fsutil.py 与推理无关的小工具(图片后缀、文件名清洗、最终尺寸换算)
|
||||
imageops.py 图像预处理与后处理(纯 numpy / PIL / torch)
|
||||
batch.py 目录批量抠图:路径校验、目录扫描、输出命名
|
||||
engine.py 模型扫描、加载缓存与推理
|
||||
server.py HTTP 服务、REST API、任务队列
|
||||
"""
|
||||
|
||||
__version__ = "1.0.0"
|
||||
APP_NAME = "BiRefNet WebUI"
|
||||
@@ -0,0 +1,222 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""目录批量抠图:路径校验 / 目录扫描 / 输出文件命名。
|
||||
|
||||
三条硬约束(对应需求):
|
||||
|
||||
1. **校验阶段只读不写**:输入目录、输出目录都必须**已经存在**,本模块
|
||||
绝不 ``mkdir``。路径写错时直接报错,而不是顺手把目录树建出来 ——
|
||||
这样一次误输入(或网页上的恶意路径)不会在磁盘上留下任何东西。
|
||||
2. **命名固定**:``RMBG_<原主干>_<Unix 时间戳>.<后缀>``,原主干超过
|
||||
:data:`STEM_MAXLEN` 个字符时截断。主干统一走
|
||||
:func:`birefnet_web.imageops.safe_stem` 清洗,杜绝 ``..`` 之类的穿越。
|
||||
3. **后缀自适应**:透明输出需要 alpha 通道,``.jpg`` / ``.bmp`` 装不下,
|
||||
此时退化为 ``.png``(强行按原后缀保存会被 PIL 静默压成白底图)。
|
||||
|
||||
本模块不依赖 torch,可以脱离推理环境单独测试。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Iterable, List, Optional, Sequence
|
||||
|
||||
from .fsutil import ALPHA_SUFFIXES, IMAGE_SUFFIXES, safe_stem
|
||||
|
||||
#: 输出文件名统一前缀(扫描输入目录时也用来自我过滤,避免把产物再抠一遍)
|
||||
BATCH_PREFIX = "RMBG_"
|
||||
|
||||
#: 原文件名主干的最大长度(超过即截断)
|
||||
STEM_MAXLEN = 20
|
||||
|
||||
#: 重名避让时最多尝试的序号
|
||||
_MAX_DEDUP = 999
|
||||
|
||||
|
||||
class BatchPathError(ValueError):
|
||||
"""路径不合规:不存在 / 不是目录 / 不可写 / 命中系统目录黑名单。"""
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 路径解析与校验
|
||||
# --------------------------------------------------------------------------- #
|
||||
def _critical_dirs() -> List[Path]:
|
||||
"""返回当前系统上「不该往里写文件」的目录(仅用于提示级别的高危拦截)。"""
|
||||
roots: List[Path] = []
|
||||
names = ("SystemRoot", "windir", "ProgramFiles", "ProgramFiles(x86)", "ProgramData")
|
||||
for name in names:
|
||||
raw = os.environ.get(name)
|
||||
if not raw:
|
||||
continue
|
||||
p = Path(raw)
|
||||
if p.is_absolute():
|
||||
roots.append(p)
|
||||
if os.name != "nt":
|
||||
roots.extend(Path(p) for p in ("/bin", "/sbin", "/etc", "/usr", "/boot", "/dev", "/proc", "/sys"))
|
||||
return roots
|
||||
|
||||
|
||||
def resolve_user_path(raw: object, base: Path) -> Path:
|
||||
"""把用户在界面上敲的一行路径规范化成绝对路径。
|
||||
|
||||
Args:
|
||||
raw: 原始输入(可能带首尾空格、成对引号、``~`` 或 ``%VAR%``)。
|
||||
base: 相对路径的基准目录(本项目的 ``PROJECT_ROOT``)。
|
||||
|
||||
Returns:
|
||||
规范化后的绝对路径(**不校验存在性**,也不做任何创建)。
|
||||
|
||||
Raises:
|
||||
BatchPathError: 输入为空。
|
||||
"""
|
||||
text = str(raw or "").strip().strip('"').strip("'").strip()
|
||||
if not text:
|
||||
raise BatchPathError("路径不能为空")
|
||||
text = os.path.expandvars(os.path.expanduser(text))
|
||||
path = Path(text)
|
||||
if not path.is_absolute():
|
||||
path = base / path
|
||||
return Path(os.path.normpath(str(path)))
|
||||
|
||||
|
||||
def check_input_dir(raw: object, base: Path) -> Path:
|
||||
"""校验「图片所在目录」,必须已存在且是目录。
|
||||
|
||||
Raises:
|
||||
BatchPathError: 路径为空 / 不存在 / 不是目录。
|
||||
"""
|
||||
path = resolve_user_path(raw, base)
|
||||
if not path.exists():
|
||||
raise BatchPathError(f"输入目录不存在:{path}(本工具不会自动创建目录,请先确认路径)")
|
||||
if not path.is_dir():
|
||||
raise BatchPathError(f"输入路径不是目录:{path}")
|
||||
return path
|
||||
|
||||
|
||||
def check_output_dir(raw: object, base: Path, *, allow_system_dir: bool = False) -> Path:
|
||||
"""校验「结果输出目录」,必须已存在、是目录且可写。
|
||||
|
||||
Args:
|
||||
raw: 用户输入或默认输出目录。
|
||||
base: 相对路径基准。
|
||||
allow_system_dir: 为 True 时跳过系统目录黑名单(测试用)。
|
||||
|
||||
Raises:
|
||||
BatchPathError: 路径为空 / 不存在 / 不是目录 / 命中黑名单 / 不可写。
|
||||
"""
|
||||
path = resolve_user_path(raw, base)
|
||||
if not path.exists():
|
||||
raise BatchPathError(f"输出目录不存在:{path}(本工具不会自动创建目录,请先手动建好)")
|
||||
if not path.is_dir():
|
||||
raise BatchPathError(f"输出路径不是目录:{path}")
|
||||
|
||||
if not allow_system_dir:
|
||||
if path == Path(path.anchor): # C:\ 或 /
|
||||
raise BatchPathError(f"拒绝把结果写进磁盘根目录:{path}")
|
||||
for critical in _critical_dirs():
|
||||
if path == critical or critical in path.parents:
|
||||
raise BatchPathError(f"拒绝把结果写进系统目录:{path}")
|
||||
|
||||
if not os.access(path, os.W_OK):
|
||||
raise BatchPathError(f"输出目录不可写(权限不足):{path}")
|
||||
return path
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 目录扫描
|
||||
# --------------------------------------------------------------------------- #
|
||||
def scan_images(
|
||||
root: Path,
|
||||
*,
|
||||
recursive: bool = False,
|
||||
skip_dirs: Sequence[Path] = (),
|
||||
) -> List[Path]:
|
||||
"""扫描目录下的图片文件。
|
||||
|
||||
Args:
|
||||
root: 已校验过的图片目录。
|
||||
recursive: 是否包含子目录。
|
||||
skip_dirs: 需要跳过的目录(通常是输出目录,避免读到自己的产物)。
|
||||
|
||||
Returns:
|
||||
排序后的图片路径列表(先按目录、再按文件名,保证同一批次的处理顺序稳定)。
|
||||
"""
|
||||
skipped = [Path(d).resolve() for d in skip_dirs]
|
||||
walker: Iterable[Path] = root.rglob("*") if recursive else root.glob("*")
|
||||
|
||||
found: List[Path] = []
|
||||
for path in walker:
|
||||
try:
|
||||
if not path.is_file():
|
||||
continue
|
||||
except OSError: # 权限/坏链接
|
||||
continue
|
||||
if path.suffix.lower() not in IMAGE_SUFFIXES:
|
||||
continue
|
||||
# 本工具自己的产物(RMBG_ 前缀)不参与批量,避免反复抠同一张图
|
||||
if path.name.startswith(BATCH_PREFIX):
|
||||
continue
|
||||
if skipped:
|
||||
parent = path.parent.resolve()
|
||||
if any(parent == s or s in parent.parents or parent in s.parents for s in skipped):
|
||||
continue
|
||||
found.append(path)
|
||||
|
||||
return sorted(found, key=lambda p: (str(p.parent).lower(), p.name.lower()))
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 输出命名
|
||||
# --------------------------------------------------------------------------- #
|
||||
def choose_suffix(orig_suffix: str, background: str) -> str:
|
||||
"""决定输出文件后缀:透明模式下如果原后缀装不下 alpha,就退化为 ``.png``。
|
||||
|
||||
Args:
|
||||
orig_suffix: 原文件后缀(含点,如 ``.jpg``)。
|
||||
background: 背景模式,``transparent`` 或其它(纯色)。
|
||||
|
||||
Returns:
|
||||
合法的输出后缀字符串。
|
||||
"""
|
||||
suffix = (orig_suffix or "").lower()
|
||||
if background == "transparent":
|
||||
return suffix if suffix in ALPHA_SUFFIXES else ".png"
|
||||
return suffix if suffix in IMAGE_SUFFIXES else ".png"
|
||||
|
||||
|
||||
def build_output_name(source: Path, timestamp: int, background: str = "transparent") -> str:
|
||||
"""按 ``RMBG_<主干(≤20)>_<Unix 时间戳>.<后缀>`` 组装输出文件名。
|
||||
|
||||
Args:
|
||||
source: 原始图片路径。
|
||||
timestamp: Unix 时间戳(秒),由调用方在**处理时**取,便于追溯。
|
||||
background: 背景模式,决定后缀是否允许退化为 ``.png``。
|
||||
|
||||
Returns:
|
||||
输出文件名(不含目录)。
|
||||
"""
|
||||
stem = safe_stem(source.name, fallback="image", maxlen=STEM_MAXLEN)
|
||||
return f"{BATCH_PREFIX}{stem}_{int(timestamp)}{choose_suffix(source.suffix, background)}"
|
||||
|
||||
|
||||
def unique_path(directory: Path, name: str) -> Path:
|
||||
"""在同一秒内出现同名输出时,用 ``_1`` / ``_2`` 递增避让,绝不覆盖已有文件。
|
||||
|
||||
Args:
|
||||
directory: 输出目录(已存在)。
|
||||
name: 目标文件名。
|
||||
|
||||
Returns:
|
||||
可安全写入的完整路径。
|
||||
|
||||
Raises:
|
||||
BatchPathError: 避让序号耗尽(同名文件超过 :data:`_MAX_DEDUP` 个)。
|
||||
"""
|
||||
target = directory / name
|
||||
if not target.exists():
|
||||
return target
|
||||
for i in range(1, _MAX_DEDUP + 1):
|
||||
candidate = directory / f"{target.stem}_{i}{target.suffix}"
|
||||
if not candidate.exists():
|
||||
return candidate
|
||||
raise BatchPathError(f"输出目录中同名文件过多,请检查:{name}")
|
||||
@@ -0,0 +1,262 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""模型代码定位与 ComfyUI 解耦层。
|
||||
|
||||
本项目把原来 ComfyUI 节点 ``comfyui_birefnet_ll`` 的模型代码**内联**到了
|
||||
``vendor/comfyui_birefnet_ll/``,因此默认情况下完全自包含,不再依赖外部 ComfyUI 安装。
|
||||
|
||||
定位优先级(``find_node_dir``):
|
||||
1. 命令行 ``--node-dir``
|
||||
2. 环境变量 ``BIREFNET_NODE_DIR``
|
||||
3. **项目内 ``vendor/comfyui_birefnet_ll``(默认,开箱即用)**
|
||||
4. ``config.json`` 中的 ``node_dir`` 字段(兼容旧配置)
|
||||
5. 自动扫描本机 ComfyUI 安装(便于跟随上游节点升级)
|
||||
|
||||
模型包在 import 阶段依赖 ComfyUI 的 ``folder_paths``,但只用于「按文件名查权重路径」
|
||||
(即骨干网络预训练权重)。本项目始终以 ``bb_pretrained=False`` 构建骨干网络,不会真正
|
||||
读取这类权重,因此这里注入一个最小实现的 ``folder_paths`` 垫片即可脱离 ComfyUI 运行。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
from pathlib import Path
|
||||
from typing import Dict, Iterable, List, Optional, Tuple
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 路径常量
|
||||
# --------------------------------------------------------------------------- #
|
||||
PROJECT_ROOT = Path(__file__).resolve().parent.parent
|
||||
|
||||
#: 项目内自带的模型代码(默认来源,随项目一起分发)
|
||||
VENDOR_DIR = PROJECT_ROOT / "vendor" / "comfyui_birefnet_ll"
|
||||
|
||||
#: 项目内自带的权重目录(默认来源)
|
||||
LOCAL_MODEL_DIR = PROJECT_ROOT / "models"
|
||||
|
||||
#: 环境变量:手动指定模型代码目录
|
||||
NODE_DIR_ENV = "BIREFNET_NODE_DIR"
|
||||
|
||||
#: 节点仓库名,用于兜底扫描外部 ComfyUI 安装
|
||||
NODE_REPO_NAME = "comfyui_birefnet_ll"
|
||||
|
||||
#: 旧版架构权重文件名(与 birefnetNode.py 保持一致)
|
||||
OLD_MODEL_FILES: Tuple[str, ...] = ("BiRefNet-DIS_ep580.pth", "BiRefNet-ep480.pth")
|
||||
|
||||
#: 判定「是有效模型代码目录」的标志文件
|
||||
_MARKER = ("birefnet", "models", "birefnet.py")
|
||||
|
||||
_BUNDLED_PYTHON = PROJECT_ROOT / "python" / "python.exe"
|
||||
|
||||
_INSTALL_HINT = (
|
||||
"缺少运行依赖。本项目自带独立 Python 运行时,请使用:\n"
|
||||
f' "{_BUNDLED_PYTHON}" Webui.py\n'
|
||||
"或双击 run.bat。\n"
|
||||
"若想用自己的解释器,请先安装依赖:pip install -r requirements.txt"
|
||||
)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 模型代码目录定位
|
||||
# --------------------------------------------------------------------------- #
|
||||
def _is_code_dir(path: Path) -> bool:
|
||||
"""目录内是否含 ``birefnet/models/birefnet.py``。"""
|
||||
try:
|
||||
return path.joinpath(*_MARKER).is_file()
|
||||
except OSError:
|
||||
return False
|
||||
|
||||
|
||||
def _scan_external_node_dirs() -> List[Path]:
|
||||
"""扫描本机常见的 ComfyUI 安装位置(仅在需要跟随上游升级时才会命中)。"""
|
||||
found: List[Path] = []
|
||||
home = Path.home()
|
||||
roots: List[Path] = [home]
|
||||
for letter in "CDEFGHIJ":
|
||||
roots.append(Path(f"{letter}:/"))
|
||||
patterns = (
|
||||
"ComfyUI/custom_nodes/" + NODE_REPO_NAME,
|
||||
"ComfyUI*/ComfyUI/custom_nodes/" + NODE_REPO_NAME,
|
||||
"ComfyUI*/custom_nodes/" + NODE_REPO_NAME,
|
||||
"ComfyUI*/ComfyUI/custom_nodes/ComfyUI_BiRefNet_ll",
|
||||
"ComfyUI*/custom_nodes/ComfyUI_BiRefNet_ll",
|
||||
"*/*/ComfyUI/custom_nodes/" + NODE_REPO_NAME,
|
||||
)
|
||||
for root in roots:
|
||||
if not root.exists():
|
||||
continue
|
||||
for pat in patterns:
|
||||
try:
|
||||
found.extend(sorted(root.glob(pat)))
|
||||
except (OSError, ValueError):
|
||||
continue
|
||||
return found
|
||||
|
||||
|
||||
def _dedup(paths: Iterable[Path]) -> List[Path]:
|
||||
seen, out = set(), []
|
||||
for p in paths:
|
||||
try:
|
||||
key = str(p.resolve()).lower() if p.exists() else str(p).lower()
|
||||
except OSError:
|
||||
key = str(p).lower()
|
||||
if key in seen:
|
||||
continue
|
||||
seen.add(key)
|
||||
out.append(p)
|
||||
return out
|
||||
|
||||
|
||||
def candidate_code_dirs(config_hint: Optional[str] = None) -> List[Path]:
|
||||
"""按优先级生成候选模型代码目录。"""
|
||||
cands: List[Path] = []
|
||||
|
||||
env = os.environ.get(NODE_DIR_ENV)
|
||||
if env:
|
||||
cands.append(Path(env))
|
||||
|
||||
# 项目自带(默认命中)
|
||||
cands.append(VENDOR_DIR)
|
||||
|
||||
# 旧配置中记录的路径(兼容早期版本写死的 ComfyUI 节点路径)
|
||||
if config_hint:
|
||||
cands.append(Path(config_hint))
|
||||
|
||||
# 外部 ComfyUI 安装(兜底)
|
||||
cands.append(PROJECT_ROOT.parent / NODE_REPO_NAME)
|
||||
cands.extend(_scan_external_node_dirs())
|
||||
|
||||
return _dedup(cands)
|
||||
|
||||
|
||||
def find_node_dir(explicit: Optional[str] = None, config_hint: Optional[str] = None) -> Path:
|
||||
"""定位模型代码目录(需含 ``birefnet/models/birefnet.py``)。
|
||||
|
||||
正常情况下会直接命中项目内的 ``vendor/comfyui_birefnet_ll``。
|
||||
"""
|
||||
cands: List[Path] = []
|
||||
if explicit:
|
||||
cands.append(Path(explicit))
|
||||
cands.extend(candidate_code_dirs(config_hint))
|
||||
|
||||
for cand in cands:
|
||||
if _is_code_dir(cand):
|
||||
try:
|
||||
return cand.resolve()
|
||||
except OSError:
|
||||
return cand
|
||||
|
||||
raise FileNotFoundError(
|
||||
"未找到 BiRefNet 模型代码目录(应包含 birefnet/models/birefnet.py)。\n"
|
||||
f"默认位置:{VENDOR_DIR}\n"
|
||||
"若该目录缺失或被误删,可从 ComfyUI 节点 comfyui_birefnet_ll 重新拷贝,或用以下方式指定:\n"
|
||||
" 1) 启动参数 --node-dir <路径>\n"
|
||||
f" 2) 环境变量 {NODE_DIR_ENV}=<路径>\n"
|
||||
" 3) config.json 中的 node_dir 字段\n"
|
||||
"已尝试的候选目录:\n - " + "\n - ".join(str(c) for c in cands[:8])
|
||||
)
|
||||
|
||||
|
||||
def default_model_dirs(code_dir: Optional[Path] = None) -> List[Path]:
|
||||
"""返回默认权重搜索目录:项目内 ``models/`` 优先,其次外部目录。"""
|
||||
dirs: List[Path] = [LOCAL_MODEL_DIR]
|
||||
|
||||
# 若用户显式指向了外部 ComfyUI 节点,则顺带搜索对应的 models/BiRefNet
|
||||
if code_dir is not None:
|
||||
p = Path(code_dir)
|
||||
try:
|
||||
if p.resolve() != VENDOR_DIR.resolve():
|
||||
dirs.append(p.parent.parent / "models" / "BiRefNet")
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
return _dedup(dirs)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# folder_paths 垫片
|
||||
# --------------------------------------------------------------------------- #
|
||||
class _FolderPathsShim(types.ModuleType):
|
||||
"""仅实现 birefnet/config.py 用到的接口,一律返回 None(不查权重)。"""
|
||||
|
||||
def __init__(self, model_dirs: Iterable[Path]) -> None:
|
||||
super().__init__("folder_paths")
|
||||
self.models_dir = str(next(iter(model_dirs), Path("models")))
|
||||
self.folder_names_and_paths: Dict[str, Tuple[List[str], set]] = {
|
||||
"birefnet": ([self.models_dir], {".pt", ".pth", ".safetensors", ".ckpt"})
|
||||
}
|
||||
self.supported_pt_extensions = {".pt", ".pth", ".safetensors", ".ckpt"}
|
||||
|
||||
def get_folder_paths(self, folder_name: str) -> List[str]:
|
||||
return list(self.folder_names_and_paths.get(folder_name, ([], None))[0])
|
||||
|
||||
def get_full_path(self, folder_name: str, filename: str) -> Optional[str]:
|
||||
return None
|
||||
|
||||
def get_filename_list(self, folder_name: str) -> List[str]:
|
||||
return []
|
||||
|
||||
def add_model_folder_path(self, folder_name: str, full_folder_path: str) -> None:
|
||||
if folder_name not in self.folder_names_and_paths:
|
||||
self.folder_names_and_paths[folder_name] = ([], set())
|
||||
self.folder_names_and_paths[folder_name][0].append(str(full_folder_path))
|
||||
|
||||
def filter_files_extensions(self, files, extensions): # pragma: no cover
|
||||
return [f for f in files if os.path.splitext(f)[1] in extensions]
|
||||
|
||||
|
||||
def ensure_folder_paths(model_dirs: Iterable[Path]) -> bool:
|
||||
"""确保 ``folder_paths`` 可用。返回 True 表示用了真实实现(即运行在 ComfyUI 内)。"""
|
||||
if "folder_paths" in sys.modules:
|
||||
return not isinstance(sys.modules["folder_paths"], _FolderPathsShim)
|
||||
try: # 若真的在 ComfyUI 环境内运行,优先用真实实现
|
||||
importlib.import_module("folder_paths")
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
sys.modules["folder_paths"] = _FolderPathsShim(model_dirs)
|
||||
return False
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 模型类导入
|
||||
# --------------------------------------------------------------------------- #
|
||||
_cached: Dict[str, object] = {}
|
||||
|
||||
|
||||
def load_model_classes(node_dir: Path, model_dirs: Iterable[Path]) -> Dict[str, object]:
|
||||
"""导入并返回 ``{"BiRefNet":..., "OldBiRefNet":..., "check_state_dict":..., "shim": bool}``。
|
||||
|
||||
结果会被缓存,重复调用不会重复导入。
|
||||
"""
|
||||
if _cached:
|
||||
return _cached
|
||||
|
||||
node_dir = Path(node_dir).resolve()
|
||||
shim = ensure_folder_paths(model_dirs)
|
||||
|
||||
p = str(node_dir)
|
||||
if p not in sys.path:
|
||||
sys.path.insert(0, p)
|
||||
|
||||
try:
|
||||
from birefnet.models.birefnet import BiRefNet
|
||||
from birefnet.utils import check_state_dict
|
||||
except ImportError as exc: # pragma: no cover - 环境问题
|
||||
raise ImportError(f"导入 birefnet 模型包失败({node_dir}):{exc}\n{_INSTALL_HINT}") from exc
|
||||
|
||||
try:
|
||||
from birefnet_old.models.birefnet import BiRefNet as OldBiRefNet
|
||||
except Exception: # 旧包缺失不影响新模型使用
|
||||
OldBiRefNet = None # type: ignore[assignment]
|
||||
|
||||
_cached.update(
|
||||
BiRefNet=BiRefNet,
|
||||
OldBiRefNet=OldBiRefNet,
|
||||
check_state_dict=check_state_dict,
|
||||
node_dir=node_dir,
|
||||
used_folder_paths_shim=shim,
|
||||
)
|
||||
return _cached
|
||||
@@ -0,0 +1,499 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""模型注册表与推理引擎。
|
||||
|
||||
设计要点:
|
||||
* 模型代码直接复用 ComfyUI 节点(见 compat.py),不复制、不魔改;
|
||||
* 模型按 (路径, mtime, 设备, 精度, 架构) 做 LRU 缓存,默认只驻留 1 个,避免显存爆炸;
|
||||
* float16 / bfloat16 推理若出现 NaN/Inf 会自动回退 float32 重算一次并记录告警。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import os
|
||||
import threading
|
||||
import time
|
||||
from collections import OrderedDict
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Callable, Dict, List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from . import compat, imageops
|
||||
|
||||
ProgressFn = Optional[Callable[[str, float], None]]
|
||||
|
||||
#: 不是抠图模型,而是骨干网络权重,扫描时跳过
|
||||
_BACKBONE_PREFIXES = ("swin_", "pvt_v2_")
|
||||
|
||||
#: 支持加载的权重后缀
|
||||
_MODEL_SUFFIXES = (".safetensors", ".pth", ".pt", ".ckpt")
|
||||
|
||||
#: 精度选项 -> torch dtype;auto 由设备决定
|
||||
DTYPES: Dict[str, Optional[torch.dtype]] = {
|
||||
"auto": None,
|
||||
"float32": torch.float32,
|
||||
"float16": torch.float16,
|
||||
"bfloat16": torch.bfloat16,
|
||||
}
|
||||
|
||||
#: 前景精修超过该像素量时自动跳过(避免内存与耗时失控)
|
||||
_REFINE_PIXEL_LIMIT = 40_000_000
|
||||
|
||||
|
||||
@dataclass
|
||||
class ModelInfo:
|
||||
"""一个可用的模型权重文件。"""
|
||||
|
||||
key: str
|
||||
name: str
|
||||
file: str
|
||||
path: str
|
||||
size_mb: float
|
||||
mtime: float
|
||||
arch: str # "v1" | "old"
|
||||
bb_index: int
|
||||
backbone: str
|
||||
directory: str
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"key": self.key,
|
||||
"name": self.name,
|
||||
"file": self.file,
|
||||
"path": self.path,
|
||||
"size_mb": round(self.size_mb, 1),
|
||||
"arch": self.arch,
|
||||
"bb_index": self.bb_index,
|
||||
"backbone": self.backbone,
|
||||
"directory": self.directory,
|
||||
}
|
||||
|
||||
|
||||
def guess_bb_index(filename: str) -> int:
|
||||
"""根据文件名猜骨干网络:lite 系列用 swin_v1_t(3),其余用 swin_v1_l(6)。"""
|
||||
return 3 if "lite" in filename.lower() else 6
|
||||
|
||||
|
||||
def guess_arch(filename: str) -> str:
|
||||
return "old" if os.path.basename(filename) in compat.OLD_MODEL_FILES else "v1"
|
||||
|
||||
|
||||
class ModelRegistry:
|
||||
"""扫描并维护可用模型列表。"""
|
||||
|
||||
def __init__(self, model_dirs: List[Path]) -> None:
|
||||
self.model_dirs = [Path(d) for d in model_dirs]
|
||||
self._models: "OrderedDict[str, ModelInfo]" = OrderedDict()
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def scan(self) -> List[ModelInfo]:
|
||||
found: "OrderedDict[str, ModelInfo]" = OrderedDict()
|
||||
for idx, d in enumerate(self.model_dirs):
|
||||
if not d or not d.is_dir():
|
||||
continue
|
||||
try:
|
||||
entries = sorted(d.iterdir(), key=lambda p: p.name.lower())
|
||||
except OSError:
|
||||
continue
|
||||
for entry in entries:
|
||||
if not entry.is_file():
|
||||
continue
|
||||
if entry.suffix.lower() not in _MODEL_SUFFIXES:
|
||||
continue
|
||||
if entry.name.lower().startswith(_BACKBONE_PREFIXES):
|
||||
continue
|
||||
try:
|
||||
stat = entry.stat()
|
||||
except OSError:
|
||||
continue
|
||||
stem = entry.stem
|
||||
bb_index = guess_bb_index(stem)
|
||||
bb = {3: "swin_v1_t", 6: "swin_v1_l"}[bb_index]
|
||||
key = f"{stem}@{idx}" if stem in found else stem
|
||||
if key in found: # 极少数重名情况
|
||||
key = f"{stem}@{idx}"
|
||||
found[key] = ModelInfo(
|
||||
key=key,
|
||||
name=stem,
|
||||
file=entry.name,
|
||||
path=str(entry.resolve()),
|
||||
size_mb=stat.st_size / (1024 * 1024),
|
||||
mtime=stat.st_mtime,
|
||||
arch=guess_arch(entry.name),
|
||||
bb_index=bb_index,
|
||||
backbone=bb,
|
||||
directory=str(d),
|
||||
)
|
||||
with self._lock:
|
||||
self._models = found
|
||||
return list(found.values())
|
||||
|
||||
def list(self) -> List[ModelInfo]:
|
||||
with self._lock:
|
||||
if not self._models:
|
||||
return self.scan()
|
||||
return list(self._models.values())
|
||||
|
||||
def get(self, key: str) -> ModelInfo:
|
||||
for info in self.list():
|
||||
if key in (info.key, info.name, info.file, info.path):
|
||||
return info
|
||||
raise KeyError(f"未找到模型 {key!r},可用模型:{[m.key for m in self.list()]}")
|
||||
|
||||
|
||||
class InferenceEngine:
|
||||
"""串行执行抠图推理(GPU 上同一时刻只跑一个模型)。"""
|
||||
|
||||
def __init__(self, registry: ModelRegistry, node_dir: Path, max_cached: int = 1) -> None:
|
||||
self.registry = registry
|
||||
self.node_dir = Path(node_dir)
|
||||
self.max_cached = max(1, int(max_cached))
|
||||
self._cache: "OrderedDict[tuple, Tuple[object, str, torch.dtype]]" = OrderedDict()
|
||||
self._lock = threading.RLock()
|
||||
self._classes: Optional[dict] = None
|
||||
self.last_warning: Optional[str] = None
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# 模型加载
|
||||
# ------------------------------------------------------------------ #
|
||||
def _model_classes(self) -> dict:
|
||||
if self._classes is None:
|
||||
self._classes = compat.load_model_classes(self.node_dir, self.registry.model_dirs)
|
||||
return self._classes
|
||||
|
||||
@staticmethod
|
||||
def resolve_device(device: str) -> str:
|
||||
device = (device or "auto").strip().lower()
|
||||
if device in ("auto", "", "gpu"):
|
||||
return "cuda" if torch.cuda.is_available() else "cpu"
|
||||
if device.startswith("cuda"):
|
||||
if not torch.cuda.is_available():
|
||||
raise RuntimeError("请求使用 CUDA,但当前 torch 未检测到可用 GPU")
|
||||
return device
|
||||
return "cpu"
|
||||
|
||||
@staticmethod
|
||||
def resolve_dtype(dtype: str, device: str) -> torch.dtype:
|
||||
dt = DTYPES.get((dtype or "auto").lower(), None)
|
||||
if dt is not None:
|
||||
return dt
|
||||
# auto:CUDA 下用 fp32 权重 + fp16 autocast(官方推荐),CPU 下 fp32
|
||||
return torch.float32
|
||||
|
||||
@staticmethod
|
||||
def _autocast_dtype(dtype: str, device: str) -> Optional[torch.dtype]:
|
||||
if not device.startswith("cuda"):
|
||||
return None
|
||||
key = (dtype or "auto").lower()
|
||||
if key in ("float16", "auto"):
|
||||
return torch.float16
|
||||
if key == "bfloat16":
|
||||
return torch.bfloat16
|
||||
return None
|
||||
|
||||
def _read_state_dict(self, info: ModelInfo) -> dict:
|
||||
if info.path.lower().endswith(".safetensors"):
|
||||
import safetensors.torch
|
||||
|
||||
return safetensors.torch.load_file(info.path, device="cpu")
|
||||
try:
|
||||
sd = torch.load(info.path, map_location="cpu", weights_only=True)
|
||||
except Exception:
|
||||
sd = torch.load(info.path, map_location="cpu", weights_only=False)
|
||||
check_state_dict = self._model_classes()["check_state_dict"]
|
||||
return check_state_dict(sd)
|
||||
|
||||
def _build_model(self, info: ModelInfo, state_dict: dict, dtype: torch.dtype):
|
||||
"""构建网络并载入权重;骨干索引猜错时自动换一个重试。"""
|
||||
classes = self._model_classes()
|
||||
BiRefNet = classes["BiRefNet"]
|
||||
OldBiRefNet = classes["OldBiRefNet"]
|
||||
|
||||
if info.arch == "old":
|
||||
if OldBiRefNet is None:
|
||||
raise RuntimeError("该权重属于旧版架构,但节点内缺少 birefnet_old 包")
|
||||
candidates = [-1]
|
||||
else:
|
||||
candidates = [info.bb_index] + [i for i in (6, 3) if i != info.bb_index]
|
||||
|
||||
last_err: Optional[Exception] = None
|
||||
for bb_index in candidates:
|
||||
try:
|
||||
if bb_index < 0:
|
||||
model = OldBiRefNet(bb_pretrained=False)
|
||||
version = "old"
|
||||
else:
|
||||
model = BiRefNet(bb_pretrained=False, bb_index=bb_index)
|
||||
version = "v1"
|
||||
if dtype != torch.float32:
|
||||
model = model.to(dtype=dtype)
|
||||
model.load_state_dict(state_dict)
|
||||
if bb_index >= 0 and bb_index != info.bb_index:
|
||||
self.last_warning = (
|
||||
f"模型 {info.name} 的骨干索引自动修正为 {bb_index}(文件名推断为 {info.bb_index})"
|
||||
)
|
||||
return model, version, bb_index
|
||||
except Exception as exc: # 尺寸不匹配等
|
||||
last_err = exc
|
||||
continue
|
||||
raise RuntimeError(f"加载模型 {info.name} 失败:{last_err}") from last_err
|
||||
|
||||
def get_model(self, key: str, device: str = "auto", dtype: str = "auto", arch: str = "auto"):
|
||||
info = self.registry.get(key)
|
||||
if arch in ("v1", "old"):
|
||||
info = ModelInfo(**{**info.__dict__, "arch": arch})
|
||||
device = self.resolve_device(device)
|
||||
torch_dtype = self.resolve_dtype(dtype, device)
|
||||
cache_key = (info.path, info.mtime, device, str(torch_dtype), info.arch)
|
||||
|
||||
with self._lock:
|
||||
if cache_key in self._cache:
|
||||
self._cache.move_to_end(cache_key)
|
||||
model, version, _ = self._cache[cache_key]
|
||||
return model, version, info
|
||||
state_dict = self._read_state_dict(info)
|
||||
model, version, _ = self._build_model(info, state_dict, torch_dtype)
|
||||
del state_dict
|
||||
model = model.to(device)
|
||||
model.eval()
|
||||
for p in model.parameters():
|
||||
p.requires_grad_(False)
|
||||
self._cache[cache_key] = (model, version, torch_dtype)
|
||||
self._cache.move_to_end(cache_key)
|
||||
self._evict_locked()
|
||||
return model, version, info
|
||||
|
||||
def _evict_locked(self) -> None:
|
||||
while len(self._cache) > self.max_cached:
|
||||
_, (model, _, _) = self._cache.popitem(last=False)
|
||||
del model
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
def unload(self) -> None:
|
||||
with self._lock:
|
||||
self._cache.clear()
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# ------------------------------------------------------------------ #
|
||||
# 推理
|
||||
# ------------------------------------------------------------------ #
|
||||
def remove_background(
|
||||
self,
|
||||
image: Image.Image,
|
||||
options: Dict[str, object],
|
||||
progress: ProgressFn = None,
|
||||
) -> Dict[str, object]:
|
||||
"""对单张图片执行抠图。
|
||||
|
||||
options 支持(均有默认值):
|
||||
model / device / dtype / arch
|
||||
resolution_mode: square | longest | custom
|
||||
width / height / longest_side
|
||||
upscale_method / mask_threshold
|
||||
refine_foreground / blur_size / blur_size_two
|
||||
background: transparent | color
|
||||
bg_color: "#ffffff"
|
||||
output_mask: bool(默认 False:不额外产出遮罩文件)
|
||||
final_longest_side: int (0 = 保持原图尺寸;>0 = 等比缩放到该最长边)
|
||||
"""
|
||||
t0 = time.perf_counter()
|
||||
self.last_warning = None
|
||||
|
||||
def report(stage: str, frac: float) -> None:
|
||||
if progress:
|
||||
try:
|
||||
progress(stage, max(0.0, min(1.0, frac)))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
model_key = str(options.get("model") or "")
|
||||
if not model_key:
|
||||
raise ValueError("未指定模型")
|
||||
device = str(options.get("device") or "auto")
|
||||
dtype = str(options.get("dtype") or "auto")
|
||||
arch = str(options.get("arch") or "auto")
|
||||
|
||||
report("加载模型", 0.05)
|
||||
model, version, info = self.get_model(model_key, device, dtype, arch)
|
||||
model_dtype = next(model.parameters()).dtype
|
||||
model_device = next(model.parameters()).device
|
||||
|
||||
report("预处理", 0.2)
|
||||
rgb = imageops.pil_to_rgb_array(image)
|
||||
src_h, src_w = rgb.shape[:2]
|
||||
in_h, in_w = imageops.build_input_size(
|
||||
src_h,
|
||||
src_w,
|
||||
mode=str(options.get("resolution_mode") or "square"),
|
||||
width=int(options.get("width") or 1024),
|
||||
height=int(options.get("height") or 1024),
|
||||
longest_side=int(options.get("longest_side") or 1024),
|
||||
)
|
||||
upscale_method = str(options.get("upscale_method") or "bilinear")
|
||||
tensor = imageops.preprocess(rgb, (in_h, in_w), upscale_method)
|
||||
|
||||
report("推理", 0.35)
|
||||
x = tensor.to(model_device)
|
||||
if model_dtype != torch.float32:
|
||||
x = x.to(model_dtype)
|
||||
|
||||
autocast_dtype = self._autocast_dtype(dtype, device)
|
||||
alpha = self._forward(model, x, model_device, model_dtype, autocast_dtype)
|
||||
del x
|
||||
|
||||
# 精度兜底:半精度偶发 NaN 时用 fp32 重算一次
|
||||
if not bool(torch.isfinite(alpha).all()):
|
||||
self.last_warning = "半精度推理出现 NaN/Inf,已自动回退 float32 重算"
|
||||
self.unload()
|
||||
model, version, info = self.get_model(model_key, device, "float32", arch)
|
||||
model_device = next(model.parameters()).device
|
||||
x = tensor.to(model_device)
|
||||
alpha = self._forward(model, x, model_device, torch.float32, None)
|
||||
del x
|
||||
del tensor
|
||||
|
||||
report("后处理", 0.75)
|
||||
alpha = imageops.upscale_mask(alpha, src_h, src_w, upscale_method)
|
||||
alpha = imageops.filter_mask(alpha, float(options.get("mask_threshold") or 0.0))
|
||||
alpha_np = alpha.squeeze(0).squeeze(0).to(torch.float32).cpu().numpy()
|
||||
del alpha
|
||||
alpha_np = np.clip(alpha_np, 0.0, 1.0)
|
||||
|
||||
want_mask = bool(options.get("output_mask", False))
|
||||
background = str(options.get("background") or "transparent")
|
||||
rgb_out = rgb
|
||||
refine = bool(options.get("refine_foreground", True))
|
||||
if refine and src_h * src_w > _REFINE_PIXEL_LIMIT:
|
||||
refine = False
|
||||
self.last_warning = (
|
||||
f"图像像素 {src_h * src_w / 1e6:.1f}MP 超过前景精修上限,已自动跳过该步骤"
|
||||
)
|
||||
|
||||
if refine:
|
||||
report("前景精修", 0.85)
|
||||
rgb_out = np.clip(
|
||||
np.rint(
|
||||
imageops.refine_foreground(
|
||||
rgb.astype(np.float32) / 255.0,
|
||||
alpha_np,
|
||||
int(options.get("blur_size") or 90),
|
||||
int(options.get("blur_size_two") or 6),
|
||||
)
|
||||
* 255.0
|
||||
),
|
||||
0,
|
||||
255,
|
||||
).astype(np.uint8)
|
||||
|
||||
if background == "color":
|
||||
cutout = imageops.compose_on_color(
|
||||
rgb_out, alpha_np, imageops.parse_color(options.get("bg_color"))
|
||||
)
|
||||
else:
|
||||
cutout = imageops.to_pil_rgba(rgb_out, alpha_np)
|
||||
|
||||
mask_pil = imageops.to_pil_mask(alpha_np) if want_mask else None
|
||||
|
||||
# 最终尺寸:按原图比例把最长边缩放到指定像素(0 = 保持原尺寸)
|
||||
final_side = int(options.get("final_longest_side") or 0)
|
||||
new_size = imageops.fit_longest_side(cutout.size, final_side)
|
||||
if new_size != cutout.size:
|
||||
cutout = cutout.resize(new_size, Image.LANCZOS)
|
||||
if mask_pil is not None:
|
||||
mask_pil = mask_pil.resize(new_size, Image.LANCZOS)
|
||||
|
||||
report("保存结果", 0.95)
|
||||
# 前景占比:便于前端提示「模型可能没抠到东西」
|
||||
coverage = float((alpha_np > 0.5).mean()) if alpha_np.size else 0.0
|
||||
return {
|
||||
"cutout": cutout,
|
||||
"mask": mask_pil,
|
||||
"width": cutout.width,
|
||||
"height": cutout.height,
|
||||
"source_width": src_w,
|
||||
"source_height": src_h,
|
||||
"input_size": [in_w, in_h],
|
||||
"coverage": round(coverage, 4),
|
||||
"model": info.name,
|
||||
"model_path": info.path,
|
||||
"backbone": info.backbone,
|
||||
"arch": version,
|
||||
"device": str(model_device),
|
||||
"dtype": str(model_dtype),
|
||||
"elapsed": round(time.perf_counter() - t0, 2),
|
||||
"warning": self.last_warning,
|
||||
"notes": self._notes(background, refine),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _notes(background: str, refine: bool) -> List[str]:
|
||||
notes = []
|
||||
if refine:
|
||||
notes.append("已启用前景精修(去除边缘残留背景色)")
|
||||
if background == "color":
|
||||
notes.append("已合成到自定义背景色")
|
||||
return notes
|
||||
|
||||
@staticmethod
|
||||
def _forward(model, x, device: torch.device, model_dtype: torch.dtype, autocast_dtype):
|
||||
"""前向一次,返回 sigmoid 后的低分辨率遮罩 (1,1,h,w)。"""
|
||||
use_autocast = autocast_dtype is not None and str(device).startswith("cuda")
|
||||
ctx = (
|
||||
torch.autocast(device_type="cuda", dtype=autocast_dtype, enabled=True)
|
||||
if use_autocast
|
||||
else torch.autocast(device_type="cpu", enabled=False)
|
||||
)
|
||||
with torch.inference_mode(), ctx:
|
||||
out = model(x)
|
||||
pred = out[-1] if isinstance(out, (list, tuple)) else out
|
||||
alpha = pred.sigmoid().float()
|
||||
return alpha
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 环境信息
|
||||
# --------------------------------------------------------------------------- #
|
||||
def describe_environment(engine: InferenceEngine) -> dict:
|
||||
"""给前端展示的设备/运行环境信息。"""
|
||||
try:
|
||||
import torch as _torch # noqa: F401
|
||||
import torchvision as _tv
|
||||
|
||||
tv_version = getattr(_tv, "__version__", "?")
|
||||
except Exception: # pragma: no cover
|
||||
tv_version = "?"
|
||||
|
||||
devices = []
|
||||
if torch.cuda.is_available():
|
||||
for i in range(torch.cuda.device_count()):
|
||||
try:
|
||||
props = torch.cuda.get_device_properties(i)
|
||||
free, total = torch.cuda.mem_get_info(i)
|
||||
devices.append(
|
||||
{
|
||||
"index": i,
|
||||
"name": props.name,
|
||||
"total_mem": round(total / 1024**3, 1),
|
||||
"free_mem": round(free / 1024**3, 1),
|
||||
"cc": f"{props.major}.{props.minor}",
|
||||
}
|
||||
)
|
||||
except Exception:
|
||||
devices.append({"index": i, "name": "CUDA 设备", "total_mem": 0, "free_mem": 0})
|
||||
return {
|
||||
"python": os.sys.version.split()[0],
|
||||
"torch": torch.__version__,
|
||||
"torchvision": tv_version,
|
||||
"cuda": torch.cuda.is_available(),
|
||||
"cuda_version": getattr(torch.version, "cuda", None),
|
||||
"devices": devices,
|
||||
"default_device": "cuda" if torch.cuda.is_available() else "cpu",
|
||||
}
|
||||
@@ -0,0 +1,67 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""与推理无关的通用小工具:图片后缀集合、文件名清洗与尺寸换算。
|
||||
|
||||
刻意不 import torch / numpy —— 让路径校验、命名、尺寸计算这类纯逻辑可以脱离
|
||||
推理环境被单独测试(见 :mod:`birefnet_web.batch`)。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from typing import Set, Tuple
|
||||
|
||||
#: 允许上传 / 扫描的图片后缀
|
||||
IMAGE_SUFFIXES: Set[str] = {".png", ".jpg", ".jpeg", ".webp", ".bmp", ".tif", ".tiff", ".gif"}
|
||||
|
||||
#: 能装下 alpha 通道的后缀(透明输出只能用这些;jpg/bmp 装不下)
|
||||
ALPHA_SUFFIXES: Set[str] = {".png", ".webp", ".tif", ".tiff"}
|
||||
|
||||
#: 文件名中不允许出现的字符(Windows 硬限制 + 控制字符)
|
||||
BAD_NAME_CHARS = '<>:"/\\|?*\x00-\x1f'
|
||||
|
||||
|
||||
def safe_stem(name: str, fallback: str = "image", maxlen: int = 80) -> str:
|
||||
"""把外部来源的文件名清洗成安全的文件名主干(防路径穿越 + 去掉非法字符)。
|
||||
|
||||
Args:
|
||||
name: 原始文件名或路径。
|
||||
fallback: 清洗后为空时的兜底名字。
|
||||
maxlen: 主干最大长度,超出截断。
|
||||
|
||||
Returns:
|
||||
仅含安全字符的文件名主干(不含扩展名)。
|
||||
"""
|
||||
stem = os.path.splitext(os.path.basename(name or ""))[0]
|
||||
stem = "".join("_" if ch in BAD_NAME_CHARS else ch for ch in stem).strip(" .")
|
||||
stem = stem[:maxlen] or fallback
|
||||
return stem
|
||||
|
||||
|
||||
def fit_longest_side(size: Tuple[int, int], longest: int) -> Tuple[int, int]:
|
||||
"""把 ``(宽, 高)`` 按原比例缩放到「最长边 == longest」。
|
||||
|
||||
与推理无关的纯几何换算,因此放在这里(而不是 imageops):结果尺寸只由
|
||||
原尺寸和用户输入决定,可以脱离 torch 单测。
|
||||
|
||||
与 ``imageops.build_input_size(mode="longest")`` 的区别:那个算的是**网络
|
||||
输入**尺寸,会对齐到 32 的倍数;这个算的是**最终成品**尺寸,不做对齐,
|
||||
最长边精确等于用户填的值。
|
||||
|
||||
例:``(2000, 3000)`` + ``1440`` → ``(960, 1440)``;
|
||||
``(3000, 2000)`` + ``1440`` → ``(1440, 960)``。
|
||||
|
||||
Args:
|
||||
size: 原尺寸 ``(宽, 高)``,与 ``PIL.Image.size`` 一致。
|
||||
longest: 目标最长边像素;``<= 0`` 表示保持原尺寸(用户默认值)。
|
||||
|
||||
Returns:
|
||||
新尺寸 ``(宽, 高)``;无需缩放时原样返回。比例很小/很大时也至少 1 像素。
|
||||
"""
|
||||
w, h = int(size[0]), int(size[1])
|
||||
if longest <= 0 or w <= 0 or h <= 0:
|
||||
return w, h
|
||||
current = max(w, h)
|
||||
if current == longest:
|
||||
return w, h
|
||||
ratio = float(longest) / float(current)
|
||||
return max(1, int(round(w * ratio))), max(1, int(round(h * ratio)))
|
||||
@@ -0,0 +1,254 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""图像预处理 / 后处理。
|
||||
|
||||
所有数值行为对齐 ComfyUI 节点 ``comfyui_birefnet_ll``:
|
||||
* 预处理 —— 与 birefnetNode.ImagePreprocessor 相同的 Resize + ImageNet Normalize
|
||||
* 遮罩还原 —— 等价 comfy.utils.common_upscale(F.interpolate,无抗锯齿)
|
||||
* 前景精修 —— fast-foreground-estimation(Photoroom)的 box-blur 版本
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from functools import lru_cache
|
||||
from typing import Dict, Iterable, Optional, Sequence, Tuple
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image, ImageOps
|
||||
|
||||
# 与推理无关的通用工具(后缀集合、文件名清洗、尺寸换算)放在 fsutil 里,
|
||||
# 这里重新导出,保持 `imageops.safe_stem` 这类既有调用点不变。
|
||||
from .fsutil import ( # noqa: F401
|
||||
ALPHA_SUFFIXES,
|
||||
IMAGE_SUFFIXES,
|
||||
fit_longest_side,
|
||||
safe_stem,
|
||||
)
|
||||
|
||||
try: # cv2 的 box blur 速度更快、背景更纯(与节点一致)
|
||||
import cv2
|
||||
|
||||
_HAS_CV2 = True
|
||||
except Exception: # pragma: no cover
|
||||
cv2 = None # type: ignore[assignment]
|
||||
_HAS_CV2 = False
|
||||
|
||||
IMAGENET_MEAN = [0.485, 0.456, 0.406]
|
||||
IMAGENET_STD = [0.229, 0.224, 0.225]
|
||||
|
||||
#: 与节点 interpolation_modes_mapping 对齐(torchvision InterpolationMode 的数值)
|
||||
INTERPOLATIONS: Dict[str, int] = {
|
||||
"nearest": 0,
|
||||
"bilinear": 2,
|
||||
"bicubic": 3,
|
||||
"nearest-exact": 0,
|
||||
}
|
||||
|
||||
#: 输出遮罩还原时允许的 F.interpolate 模式
|
||||
UPSCALE_METHODS = ("bilinear", "nearest", "nearest-exact", "bicubic")
|
||||
|
||||
#: 预处理尺寸必须是 32 的倍数(Swin 骨干下采样 32 倍)
|
||||
SIZE_ALIGN = 32
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 基础转换
|
||||
# --------------------------------------------------------------------------- #
|
||||
def pil_to_rgb_array(image: Image.Image) -> np.ndarray:
|
||||
"""PIL -> uint8 RGB 数组,同时按 EXIF 方向自动旋转。"""
|
||||
if image.mode == "RGBA":
|
||||
# 有透明通道时先合成到白底,避免 alpha 变黑
|
||||
bg = Image.new("RGBA", image.size, (255, 255, 255, 255))
|
||||
image = Image.alpha_composite(bg, image)
|
||||
elif image.mode not in ("RGB", "L"):
|
||||
image = image.convert("RGB")
|
||||
image = ImageOps.exif_transpose(image)
|
||||
if image.mode != "RGB":
|
||||
image = image.convert("RGB")
|
||||
return np.asarray(image, dtype=np.uint8)
|
||||
|
||||
|
||||
def to_pil_rgba(rgb: np.ndarray, alpha: Optional[np.ndarray] = None) -> Image.Image:
|
||||
"""uint8 RGB (+ float alpha[0,1]) -> PIL 图片。"""
|
||||
if alpha is None:
|
||||
return Image.fromarray(np.ascontiguousarray(rgb), "RGB")
|
||||
a8 = np.clip(np.rint(alpha * 255.0), 0, 255).astype(np.uint8)
|
||||
rgba = np.dstack([rgb, a8])
|
||||
return Image.fromarray(np.ascontiguousarray(rgba), "RGBA")
|
||||
|
||||
|
||||
def to_pil_mask(alpha: np.ndarray) -> Image.Image:
|
||||
"""float mask[0,1] -> 8bit 灰度图。"""
|
||||
m8 = np.clip(np.rint(alpha * 255.0), 0, 255).astype(np.uint8)
|
||||
return Image.fromarray(np.ascontiguousarray(m8), "L")
|
||||
|
||||
|
||||
def parse_color(value: object, default: Tuple[int, int, int] = (255, 255, 255)) -> Tuple[int, int, int]:
|
||||
"""解析 ``#rgb`` / ``#rrggbb`` / ``(r,g,b)`` / int 为 RGB 三元组。"""
|
||||
if value is None:
|
||||
return default
|
||||
if isinstance(value, (list, tuple)) and len(value) >= 3:
|
||||
return tuple(int(np.clip(v, 0, 255)) for v in value[:3]) # type: ignore[return-value]
|
||||
if isinstance(value, int):
|
||||
return ((value >> 16) & 0xFF, (value >> 8) & 0xFF, value & 0xFF)
|
||||
if isinstance(value, str):
|
||||
s = value.strip().lstrip("#")
|
||||
if len(s) == 3:
|
||||
s = "".join(ch * 2 for ch in s)
|
||||
if len(s) == 6:
|
||||
try:
|
||||
v = int(s, 16)
|
||||
return ((v >> 16) & 0xFF, (v >> 8) & 0xFF, v & 0xFF)
|
||||
except ValueError:
|
||||
pass
|
||||
return default
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 输入尺寸
|
||||
# --------------------------------------------------------------------------- #
|
||||
def _align(v: float, align: int = SIZE_ALIGN, minimum: int = SIZE_ALIGN) -> int:
|
||||
v = int(round(v / align) * align)
|
||||
return max(minimum, v)
|
||||
|
||||
|
||||
def build_input_size(
|
||||
src_h: int,
|
||||
src_w: int,
|
||||
mode: str = "square",
|
||||
width: int = 1024,
|
||||
height: int = 1024,
|
||||
longest_side: int = 1024,
|
||||
) -> Tuple[int, int]:
|
||||
"""计算网络输入尺寸 (h, w)。
|
||||
|
||||
mode:
|
||||
square 固定 1024x1024(与节点默认一致,也是官方推荐分辨率)
|
||||
longest 保持宽高比,长边 = longest_side,边长对齐到 32 的倍数
|
||||
custom 使用自定义 width / height
|
||||
"""
|
||||
if src_h <= 0 or src_w <= 0:
|
||||
return max(SIZE_ALIGN, int(height)), max(SIZE_ALIGN, int(width))
|
||||
|
||||
if mode == "custom":
|
||||
h, w = _align(int(height)), _align(int(width))
|
||||
elif mode == "longest":
|
||||
target = max(SIZE_ALIGN, int(longest_side))
|
||||
scale = target / float(max(src_h, src_w))
|
||||
h, w = _align(src_h * scale), _align(src_w * scale)
|
||||
else: # square
|
||||
h, w = _align(int(height)), _align(int(width))
|
||||
return h, w
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 预处理
|
||||
# --------------------------------------------------------------------------- #
|
||||
@lru_cache(maxsize=32)
|
||||
def _transform(size_hw: Tuple[int, int], method: str):
|
||||
"""缓存 torchvision 变换(与节点 ImagePreprocessor 完全一致)。"""
|
||||
from torchvision import transforms
|
||||
|
||||
interp = INTERPOLATIONS.get(method, 2)
|
||||
return transforms.Compose(
|
||||
[
|
||||
transforms.Resize(size_hw, interpolation=interp),
|
||||
transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def preprocess(rgb: np.ndarray, size_hw: Tuple[int, int], method: str = "bilinear") -> torch.Tensor:
|
||||
"""uint8 RGB (h,w,3) -> 归一化 float32 张量 (1,3,H,W)。"""
|
||||
arr = np.asarray(rgb, dtype=np.float32) / 255.0 # 新建可写数组,避免 from_numpy 告警
|
||||
tensor = torch.from_numpy(np.ascontiguousarray(arr)).permute(2, 0, 1).unsqueeze(0)
|
||||
return _transform(tuple(size_hw), method)(tensor)
|
||||
|
||||
|
||||
def upscale_mask(mask_bchw: torch.Tensor, height: int, width: int, method: str = "bilinear") -> torch.Tensor:
|
||||
"""把网络输出的低分辨率遮罩还原到原图尺寸(等价 comfy.utils.common_upscale)。"""
|
||||
mode = method if method in UPSCALE_METHODS else "bilinear"
|
||||
if tuple(mask_bchw.shape[-2:]) == (height, width):
|
||||
return mask_bchw
|
||||
return torch.nn.functional.interpolate(mask_bchw, size=(height, width), mode=mode)
|
||||
|
||||
|
||||
def filter_mask(mask: torch.Tensor, threshold: float) -> torch.Tensor:
|
||||
"""低于阈值的概率直接置零(与节点 util.filter_mask 一致)。"""
|
||||
if threshold <= 0:
|
||||
return mask
|
||||
return mask * (mask > threshold).to(mask.dtype)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 前景精修(fast-foreground-estimation)
|
||||
# --------------------------------------------------------------------------- #
|
||||
def _box_blur(arr: np.ndarray, r: int) -> np.ndarray:
|
||||
"""box blur;r 为核尺寸。cv2 可用时优先使用,保证与节点结果一致。
|
||||
|
||||
注意:cv2 对 (h, w, 1) 的输入会返回 (h, w),这里统一补齐通道维,
|
||||
否则后续与 (h, w, 3) 广播会报错(原节点也是靠 `[:, :, None]` 兜住这一点)。
|
||||
"""
|
||||
r = int(max(1, r))
|
||||
limit = 2 * min(arr.shape[0], arr.shape[1]) + 1
|
||||
r = int(min(r, limit))
|
||||
if _HAS_CV2:
|
||||
out = cv2.blur(np.ascontiguousarray(arr, dtype=np.float32), (r, r))
|
||||
return out[..., None] if (out.ndim == 2 and arr.ndim == 3) else out
|
||||
# 兜底:torchvision 高斯模糊
|
||||
from torchvision.transforms import functional as TF
|
||||
|
||||
if r % 2 == 0:
|
||||
r += 1
|
||||
t = torch.from_numpy(np.ascontiguousarray(arr)).permute(2, 0, 1).unsqueeze(0).float()
|
||||
out = TF.gaussian_blur(t, r) if r > 1 else t
|
||||
return out.squeeze(0).permute(1, 2, 0).numpy()
|
||||
|
||||
|
||||
def _fb_step(image: np.ndarray, F: np.ndarray, B: np.ndarray, alpha: np.ndarray, r: int):
|
||||
a = alpha[:, :, None] if alpha.ndim == 2 else alpha
|
||||
blurred_alpha = _box_blur(a, r)
|
||||
blurred_F = _box_blur(F * a, r) / (blurred_alpha + 1e-5)
|
||||
blurred_B = _box_blur(B * (1.0 - a), r) / ((1.0 - blurred_alpha) + 1e-5)
|
||||
out = blurred_F + a * (image - a * blurred_F - (1.0 - a) * blurred_B)
|
||||
return np.clip(out, 0.0, 1.0), blurred_B
|
||||
|
||||
|
||||
def refine_foreground(rgb01: np.ndarray, alpha: np.ndarray, r1: int = 90, r2: int = 6) -> np.ndarray:
|
||||
"""估计前景色,消除半透明边缘残留的原背景色。
|
||||
|
||||
参考 https://github.com/Photoroom/fast-foreground-estimation
|
||||
"""
|
||||
image = np.ascontiguousarray(rgb01, dtype=np.float32)
|
||||
a = alpha.astype(np.float32)
|
||||
F, blur_B = _fb_step(image, image, image, a, r1)
|
||||
F2, _ = _fb_step(image, F, blur_B, a, r2)
|
||||
return np.clip(F2, 0.0, 1.0)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 合成输出
|
||||
# --------------------------------------------------------------------------- #
|
||||
def compose_on_color(rgb: np.ndarray, alpha: np.ndarray, color: Sequence[int]) -> Image.Image:
|
||||
"""按遮罩把前景合成到纯色背景上,返回 RGB 图。"""
|
||||
a = alpha[:, :, None].astype(np.float32)
|
||||
bg = np.array(color, dtype=np.float32).reshape(1, 1, 3)
|
||||
out = rgb.astype(np.float32) * a + bg * (1.0 - a)
|
||||
return Image.fromarray(np.clip(np.rint(out), 0, 255).astype(np.uint8), "RGB")
|
||||
|
||||
|
||||
def save_image(image: Image.Image, path: str, quality: int = 95) -> str:
|
||||
"""保存图片(自动建目录)。PNG/WebP 支持透明,JPEG 不支持。"""
|
||||
os.makedirs(os.path.dirname(os.path.abspath(path)), exist_ok=True)
|
||||
ext = os.path.splitext(path)[1].lower()
|
||||
if ext in (".jpg", ".jpeg"):
|
||||
if image.mode == "RGBA":
|
||||
bg = Image.new("RGBA", image.size, (255, 255, 255, 255))
|
||||
image = Image.alpha_composite(bg, image).convert("RGB")
|
||||
image.save(path, quality=quality, subsampling=0, optimize=True)
|
||||
elif ext == ".webp":
|
||||
image.save(path, quality=quality, method=4)
|
||||
else:
|
||||
image.save(path, optimize=False)
|
||||
return path
|
||||
@@ -0,0 +1,937 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""HTTP 服务:静态页面 + REST API + 串行任务队列。
|
||||
|
||||
刻意不依赖任何 Web 框架(Gradio / FastAPI / Flask 都不需要):
|
||||
* 传输层 —— 标准库 http.server.ThreadingHTTPServer
|
||||
* 表单解析 —— 自研 multipart/form-data 解析(`cgi` 模块在 Python 3.13 已移除)
|
||||
* 并发模型 —— 单 worker 线程串行消费任务,避免 GPU 显存被打爆
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
import traceback
|
||||
import uuid
|
||||
import zipfile
|
||||
from dataclasses import dataclass, field
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from pathlib import Path
|
||||
from typing import Dict, Iterable, List, Optional, Sequence, Tuple
|
||||
from urllib.parse import parse_qs, unquote, urlparse
|
||||
|
||||
from PIL import Image
|
||||
|
||||
from . import APP_NAME, __version__, batch, imageops
|
||||
from .engine import InferenceEngine, ModelRegistry, describe_environment
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 常量
|
||||
# --------------------------------------------------------------------------- #
|
||||
MAX_UPLOAD_BYTES = 512 * 1024 * 1024 # 单次请求体上限
|
||||
MIME_TYPES = {
|
||||
".html": "text/html; charset=utf-8",
|
||||
".css": "text/css; charset=utf-8",
|
||||
".js": "application/javascript; charset=utf-8",
|
||||
".json": "application/json; charset=utf-8",
|
||||
".png": "image/png",
|
||||
".jpg": "image/jpeg",
|
||||
".jpeg": "image/jpeg",
|
||||
".webp": "image/webp",
|
||||
".bmp": "image/bmp",
|
||||
".tif": "image/tiff",
|
||||
".tiff": "image/tiff",
|
||||
".gif": "image/gif",
|
||||
".svg": "image/svg+xml",
|
||||
".ico": "image/x-icon",
|
||||
".woff2": "font/woff2",
|
||||
}
|
||||
#: 允许上传 / 扫描的图片后缀(与批量模块共用同一份定义,避免两处不一致)
|
||||
IMAGE_SUFFIXES = imageops.IMAGE_SUFFIXES
|
||||
|
||||
DEFAULT_OPTIONS: Dict[str, object] = {
|
||||
"model": "",
|
||||
"device": "auto",
|
||||
"dtype": "auto",
|
||||
"arch": "auto",
|
||||
"resolution_mode": "square",
|
||||
"width": 1024,
|
||||
"height": 1024,
|
||||
"longest_side": 1024,
|
||||
"upscale_method": "bilinear",
|
||||
"mask_threshold": 0.0,
|
||||
"refine_foreground": True,
|
||||
"blur_size": 90,
|
||||
"blur_size_two": 6,
|
||||
"background": "transparent",
|
||||
"bg_color": "#ffffff",
|
||||
"output_mask": False,
|
||||
"final_longest_side": 0,
|
||||
}
|
||||
|
||||
|
||||
class TaskCanceled(BaseException):
|
||||
"""取消信号。
|
||||
|
||||
故意继承 BaseException:这样它会穿过引擎里 `except Exception` 的进度回调保护,
|
||||
直达任务循环,实现阶段级取消。
|
||||
"""
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# multipart/form-data 解析
|
||||
# --------------------------------------------------------------------------- #
|
||||
@dataclass
|
||||
class FormPart:
|
||||
name: str
|
||||
filename: Optional[str]
|
||||
content_type: str
|
||||
data: bytes
|
||||
|
||||
|
||||
def _parse_disposition(value: str) -> Tuple[str, Optional[str]]:
|
||||
name, filename = "", None
|
||||
for chunk in value.split(";"):
|
||||
chunk = chunk.strip()
|
||||
low = chunk.lower()
|
||||
if low.startswith("name=") and name == "":
|
||||
name = chunk[5:].strip().strip('"')
|
||||
elif low.startswith("filename="):
|
||||
raw = chunk[9:].strip().strip('"')
|
||||
if raw:
|
||||
filename = raw
|
||||
return name, filename
|
||||
|
||||
|
||||
def parse_multipart(body: bytes, boundary: bytes) -> List[FormPart]:
|
||||
"""把一个 multipart/form-data 请求体拆成若干 FormPart。"""
|
||||
parts: List[FormPart] = []
|
||||
delim = b"--" + boundary
|
||||
for raw in body.split(delim):
|
||||
if not raw or raw in (b"--", b"--\r\n", b"\r\n"):
|
||||
continue
|
||||
if raw.startswith(b"--"): # 结束标记
|
||||
break
|
||||
if raw.startswith(b"\r\n"):
|
||||
raw = raw[2:]
|
||||
elif raw.startswith(b"\n"):
|
||||
raw = raw[1:]
|
||||
head, sep, data = raw.partition(b"\r\n\r\n")
|
||||
if not sep:
|
||||
head, sep, data = raw.partition(b"\n\n")
|
||||
if not sep:
|
||||
continue
|
||||
if data.endswith(b"\r\n"):
|
||||
data = data[:-2]
|
||||
elif data.endswith(b"\n"):
|
||||
data = data[:-1]
|
||||
|
||||
headers: Dict[str, str] = {}
|
||||
for line in head.split(b"\r\n"):
|
||||
if b":" not in line:
|
||||
continue
|
||||
key, _, val = line.partition(b":")
|
||||
headers[key.strip().lower().decode("latin-1")] = (
|
||||
val.strip().decode("utf-8", "replace")
|
||||
)
|
||||
name, filename = _parse_disposition(headers.get("content-disposition", ""))
|
||||
if not name:
|
||||
continue
|
||||
parts.append(
|
||||
FormPart(
|
||||
name=name,
|
||||
filename=filename,
|
||||
content_type=headers.get("content-type", "application/octet-stream"),
|
||||
data=data,
|
||||
)
|
||||
)
|
||||
return parts
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 任务模型
|
||||
# --------------------------------------------------------------------------- #
|
||||
@dataclass
|
||||
class ImageJob:
|
||||
index: int
|
||||
filename: str
|
||||
label: str
|
||||
input_path: str
|
||||
state: str = "pending" # pending | running | done | failed | canceled
|
||||
stage: str = "等待中"
|
||||
progress: float = 0.0
|
||||
error: Optional[str] = None
|
||||
width: int = 0
|
||||
height: int = 0
|
||||
source_width: int = 0
|
||||
source_height: int = 0
|
||||
coverage: float = 0.0
|
||||
elapsed: float = 0.0
|
||||
input_size: Tuple[int, int] = (0, 0)
|
||||
warning: Optional[str] = None
|
||||
outputs: Dict[str, str] = field(default_factory=dict)
|
||||
#: 批量模式:写进用户输出目录的文件名(RMBG_<原主干>_<时间戳>.<后缀>)
|
||||
output_name: Optional[str] = None
|
||||
|
||||
def to_dict(self, task_id: str) -> dict:
|
||||
urls = {
|
||||
kind: f"/api/tasks/{task_id}/file/{self.index}/{kind}" for kind in self.outputs
|
||||
}
|
||||
return {
|
||||
"index": self.index,
|
||||
"name": self.filename,
|
||||
"label": self.label,
|
||||
"state": self.state,
|
||||
"stage": self.stage,
|
||||
"progress": round(self.progress, 3),
|
||||
"error": self.error,
|
||||
"width": self.width,
|
||||
"height": self.height,
|
||||
"source_width": self.source_width,
|
||||
"source_height": self.source_height,
|
||||
"input_size": list(self.input_size),
|
||||
"coverage": self.coverage,
|
||||
"elapsed": self.elapsed,
|
||||
"warning": self.warning,
|
||||
"output_name": self.output_name,
|
||||
"urls": urls,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class Task:
|
||||
id: str
|
||||
created: float
|
||||
options: dict
|
||||
images: List[ImageJob]
|
||||
output_dir: str
|
||||
#: upload = 网页上传单图(结果落在 outputs/<task_id>);batch = 目录批量(结果落在用户输出目录)
|
||||
mode: str = "upload"
|
||||
input_dir: Optional[str] = None
|
||||
recursive: bool = False
|
||||
state: str = "queued" # queued | running | done | partial | failed | canceled
|
||||
stage: str = "排队中"
|
||||
progress: float = 0.0
|
||||
error: Optional[str] = None
|
||||
finished_at: Optional[float] = None
|
||||
cancel_requested: bool = False
|
||||
|
||||
def to_dict(self, include_options: bool = False, tail: int = 0) -> dict:
|
||||
"""序列化任务。
|
||||
|
||||
Args:
|
||||
include_options: 是否带上参数快照。
|
||||
tail: 大于 0 时只回传最后 N 张图片的状态。
|
||||
批量任务动辄几百张,前端轮询必须靠它把载荷压住
|
||||
(处理是按序串行的,所以窗口外的一定已经是终态,不会漏更新)。
|
||||
"""
|
||||
total = len(self.images)
|
||||
start = max(0, total - int(tail)) if tail and tail > 0 else 0
|
||||
payload = {
|
||||
"id": self.id,
|
||||
"created": self.created,
|
||||
"mode": self.mode,
|
||||
"input_dir": self.input_dir,
|
||||
"output_dir": self.output_dir,
|
||||
"recursive": self.recursive,
|
||||
"state": self.state,
|
||||
"stage": self.stage,
|
||||
"progress": round(self.progress, 3),
|
||||
"error": self.error,
|
||||
"elapsed": round((self.finished_at or time.time()) - self.created, 2),
|
||||
"counts": self._counts(),
|
||||
"images_total": total,
|
||||
"images_from": start,
|
||||
"images": [img.to_dict(self.id) for img in self.images[start:]],
|
||||
}
|
||||
if include_options:
|
||||
payload["options"] = self.options
|
||||
return payload
|
||||
|
||||
def _counts(self) -> dict:
|
||||
out = {"total": len(self.images), "done": 0, "failed": 0, "pending": 0}
|
||||
for img in self.images:
|
||||
if img.state == "done":
|
||||
out["done"] += 1
|
||||
elif img.state == "failed":
|
||||
out["failed"] += 1
|
||||
elif img.state != "canceled":
|
||||
out["pending"] += 1
|
||||
return out
|
||||
|
||||
def refresh_progress(self) -> None:
|
||||
if not self.images:
|
||||
self.progress = 0.0
|
||||
else:
|
||||
self.progress = sum(i.progress for i in self.images) / len(self.images)
|
||||
|
||||
|
||||
class TaskManager:
|
||||
"""单 worker 串行任务队列。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
engine: InferenceEngine,
|
||||
output_root: Path,
|
||||
max_tasks: int = 50,
|
||||
logger=None,
|
||||
) -> None:
|
||||
self.engine = engine
|
||||
self.output_root = Path(output_root)
|
||||
self.output_root.mkdir(parents=True, exist_ok=True)
|
||||
self.max_tasks = max(1, int(max_tasks))
|
||||
self.log = logger or (lambda msg: None)
|
||||
self._tasks: Dict[str, Task] = {}
|
||||
self._lock = threading.RLock()
|
||||
self._queue: "queue.Queue[Optional[Task]]" = queue.Queue()
|
||||
self._worker = threading.Thread(target=self._work_loop, name="birefnet-worker", daemon=True)
|
||||
self._worker.start()
|
||||
|
||||
# -------------------------- 对外接口 -------------------------- #
|
||||
def submit(self, uploads: Sequence[Tuple[str, bytes]], options: dict) -> Task:
|
||||
if not uploads:
|
||||
raise ValueError("没有收到任何图片")
|
||||
task_id = time.strftime("%Y%m%d-%H%M%S-") + uuid.uuid4().hex[:4]
|
||||
output_dir = self.output_root / task_id
|
||||
input_dir = output_dir / "input"
|
||||
input_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
jobs: List[ImageJob] = []
|
||||
for idx, (filename, data) in enumerate(uploads):
|
||||
label = imageops.safe_stem(filename, fallback=f"image_{idx + 1}")
|
||||
suffix = Path(filename or "").suffix.lower()
|
||||
if suffix not in IMAGE_SUFFIXES:
|
||||
suffix = ".png"
|
||||
stored = input_dir / f"{idx + 1:03d}_{label}{suffix}"
|
||||
stored.write_bytes(data)
|
||||
jobs.append(ImageJob(index=idx, filename=filename or stored.name, label=label, input_path=str(stored)))
|
||||
|
||||
task = Task(id=task_id, created=time.time(), options=options, images=jobs, output_dir=str(output_dir))
|
||||
with self._lock:
|
||||
self._tasks[task_id] = task
|
||||
self._prune_locked()
|
||||
self._queue.put(task)
|
||||
self.log(f"[任务 {task_id}] 已提交,共 {len(jobs)} 张图片")
|
||||
return task
|
||||
|
||||
def submit_batch(
|
||||
self,
|
||||
input_dir: Path,
|
||||
output_dir: Path,
|
||||
options: dict,
|
||||
*,
|
||||
recursive: bool = False,
|
||||
) -> Task:
|
||||
"""按目录批量建任务:原图不拷贝,结果直接写进用户指定的输出目录。
|
||||
|
||||
Args:
|
||||
input_dir: 已校验存在的图片目录。
|
||||
output_dir: 已校验存在的输出目录(本方法不创建任何目录)。
|
||||
options: 与单图模式完全相同的参数集合。
|
||||
recursive: 是否包含子目录。
|
||||
|
||||
Returns:
|
||||
已入队的新任务。
|
||||
|
||||
Raises:
|
||||
ValueError: 目录下没有可处理的图片。
|
||||
"""
|
||||
# 输入 == 输出时不能把输出目录当跳过项,否则一张都扫不到
|
||||
skip = [] if Path(input_dir) == Path(output_dir) else [output_dir]
|
||||
files = batch.scan_images(input_dir, recursive=recursive, skip_dirs=skip)
|
||||
if not files:
|
||||
raise ValueError(f"目录下没有找到可处理的图片:{input_dir}")
|
||||
|
||||
task_id = time.strftime("%Y%m%d-%H%M%S-") + uuid.uuid4().hex[:4]
|
||||
jobs = [
|
||||
ImageJob(
|
||||
index=idx,
|
||||
filename=path.name,
|
||||
label=imageops.safe_stem(path.name, fallback=f"image_{idx + 1}"),
|
||||
input_path=str(path),
|
||||
)
|
||||
for idx, path in enumerate(files)
|
||||
]
|
||||
task = Task(
|
||||
id=task_id,
|
||||
created=time.time(),
|
||||
options=options,
|
||||
images=jobs,
|
||||
output_dir=str(output_dir),
|
||||
mode="batch",
|
||||
input_dir=str(input_dir),
|
||||
recursive=recursive,
|
||||
)
|
||||
with self._lock:
|
||||
self._tasks[task_id] = task
|
||||
self._prune_locked()
|
||||
self._queue.put(task)
|
||||
self.log(f"[批量 {task_id}] 已提交,{len(jobs)} 张图片 → 输出目录 {output_dir}")
|
||||
return task
|
||||
|
||||
def get(self, task_id: str) -> Optional[Task]:
|
||||
with self._lock:
|
||||
return self._tasks.get(task_id)
|
||||
|
||||
def list_tasks(self, limit: int = 20) -> List[dict]:
|
||||
"""任务列表(含各自使用的参数,供前端恢复结果历史时直接显示参数)。"""
|
||||
with self._lock:
|
||||
tasks = sorted(self._tasks.values(), key=lambda t: t.created, reverse=True)[:limit]
|
||||
return [t.to_dict(include_options=True) for t in tasks]
|
||||
|
||||
def cancel(self, task_id: str) -> bool:
|
||||
task = self.get(task_id)
|
||||
if task is None or task.finished_at is not None:
|
||||
return False
|
||||
task.cancel_requested = True
|
||||
if task.state == "queued":
|
||||
task.state = "canceled"
|
||||
task.stage = "已取消"
|
||||
task.finished_at = time.time()
|
||||
self.log(f"[任务 {task_id}] 收到取消请求")
|
||||
return True
|
||||
|
||||
def _prune_locked(self) -> None:
|
||||
if len(self._tasks) <= self.max_tasks:
|
||||
return
|
||||
ordered = sorted(self._tasks.values(), key=lambda t: t.created)
|
||||
for stale in ordered[: len(self._tasks) - self.max_tasks]:
|
||||
if stale.finished_at is not None:
|
||||
self._tasks.pop(stale.id, None)
|
||||
|
||||
# -------------------------- worker -------------------------- #
|
||||
def _work_loop(self) -> None:
|
||||
while True:
|
||||
task = self._queue.get()
|
||||
if task is None:
|
||||
return
|
||||
if task.cancel_requested:
|
||||
continue
|
||||
try:
|
||||
self._process(task)
|
||||
except TaskCanceled: # 正常取消,不需要堆栈
|
||||
task.state = "canceled"
|
||||
task.stage = "已取消"
|
||||
task.finished_at = time.time()
|
||||
except BaseException as exc: # noqa: BLE001 - worker 必须兜住一切
|
||||
self.log(f"[任务 {task.id}] 异常终止:{exc}")
|
||||
traceback.print_exc()
|
||||
task.state = "failed"
|
||||
task.error = str(exc)
|
||||
task.stage = "失败"
|
||||
task.finished_at = time.time()
|
||||
|
||||
def _process(self, task: Task) -> None:
|
||||
task.state = "running"
|
||||
task.stage = "开始处理"
|
||||
self.log(f"[任务 {task.id}] 开始处理,参数:{json.dumps(task.options, ensure_ascii=False)}")
|
||||
for job in task.images:
|
||||
if task.cancel_requested:
|
||||
if job.state in ("pending", "running"):
|
||||
job.state = "canceled"
|
||||
job.stage = "已取消"
|
||||
job.progress = 1.0
|
||||
continue
|
||||
try:
|
||||
self._process_one(task, job)
|
||||
except TaskCanceled:
|
||||
# 当前图片已被标记为 canceled;剩余图片由上面的分支收尾
|
||||
continue
|
||||
|
||||
counts = task._counts()
|
||||
task.refresh_progress()
|
||||
task.progress = 1.0
|
||||
task.finished_at = time.time()
|
||||
if task.cancel_requested:
|
||||
task.state = "canceled"
|
||||
task.stage = "已取消"
|
||||
elif counts["failed"] and counts["done"]:
|
||||
task.state = "partial"
|
||||
task.stage = "部分完成"
|
||||
elif counts["failed"]:
|
||||
task.state = "failed"
|
||||
task.stage = "失败"
|
||||
else:
|
||||
task.state = "done"
|
||||
task.stage = "完成"
|
||||
self.log(
|
||||
f"[任务 {task.id}] 结束:{task.state},成功 {counts['done']} / 失败 {counts['failed']},"
|
||||
f"耗时 {task.finished_at - task.created:.1f}s"
|
||||
)
|
||||
|
||||
def _process_one(self, task: Task, job: ImageJob) -> None:
|
||||
job.state = "running"
|
||||
job.progress = 0.02
|
||||
task.refresh_progress()
|
||||
|
||||
def on_progress(stage: str, frac: float) -> None:
|
||||
if task.cancel_requested:
|
||||
raise TaskCanceled()
|
||||
job.stage = stage
|
||||
job.progress = 0.05 + 0.9 * float(frac)
|
||||
task.stage = f"{job.label} · {stage}"
|
||||
task.refresh_progress()
|
||||
|
||||
try:
|
||||
with Image.open(job.input_path) as im:
|
||||
im.load()
|
||||
pil = im.copy()
|
||||
job.source_width, job.source_height = pil.size
|
||||
result = self.engine.remove_background(pil, task.options, on_progress)
|
||||
|
||||
job.outputs = self._store_outputs(task, job, result)
|
||||
job.width = result["width"]
|
||||
job.height = result["height"]
|
||||
job.coverage = result["coverage"]
|
||||
job.elapsed = result["elapsed"]
|
||||
job.input_size = tuple(result["input_size"]) # type: ignore[assignment]
|
||||
job.warning = result["warning"]
|
||||
job.state = "done"
|
||||
job.stage = "完成"
|
||||
job.progress = 1.0
|
||||
target = f" → {job.output_name}" if job.output_name else ""
|
||||
self.log(
|
||||
f"[{'批量' if task.mode == 'batch' else '任务'} {task.id}] {job.label}{target} 完成:"
|
||||
f"{job.width}x{job.height},前景占比 {job.coverage:.1%},耗时 {job.elapsed}s"
|
||||
)
|
||||
except TaskCanceled:
|
||||
job.state = "canceled"
|
||||
job.stage = "已取消"
|
||||
job.progress = 1.0
|
||||
raise
|
||||
except Exception as exc: # 单张失败不影响其余图片
|
||||
job.state = "failed"
|
||||
job.stage = "失败"
|
||||
job.progress = 1.0
|
||||
job.error = f"{type(exc).__name__}: {exc}"
|
||||
self.log(f"[任务 {task.id}] {job.label} 处理失败:{job.error}")
|
||||
traceback.print_exc()
|
||||
finally:
|
||||
task.refresh_progress()
|
||||
|
||||
def _store_outputs(self, task: Task, job: ImageJob, result: dict) -> Dict[str, str]:
|
||||
"""把单张结果落盘,返回 ``kind -> 路径``。
|
||||
|
||||
* 单图模式:沿用 ``outputs/<task_id>/<序号>_<主干>_cutout.png``。
|
||||
* 批量模式:写进用户指定的输出目录,文件名按需求固定为
|
||||
``RMBG_<原主干(≤20)>_<Unix 时间戳>.<原后缀>``;同名时递增 ``_1``
|
||||
避让而不是覆盖(见 :mod:`birefnet_web.batch`)。
|
||||
"""
|
||||
outputs: Dict[str, str] = {}
|
||||
if task.mode == "batch":
|
||||
out_dir = Path(task.output_dir)
|
||||
name = batch.build_output_name(
|
||||
Path(job.input_path),
|
||||
int(time.time()),
|
||||
str(task.options.get("background") or "transparent"),
|
||||
)
|
||||
target = batch.unique_path(out_dir, name)
|
||||
imageops.save_image(result["cutout"], str(target))
|
||||
job.output_name = target.name
|
||||
outputs["cutout"] = str(target)
|
||||
if result["mask"] is not None:
|
||||
mask_target = batch.unique_path(out_dir, f"{target.stem}_mask.png")
|
||||
imageops.save_image(result["mask"], str(mask_target))
|
||||
outputs["mask"] = str(mask_target)
|
||||
else:
|
||||
cutout_png = os.path.join(task.output_dir, f"{job.index + 1:03d}_{job.label}_cutout.png")
|
||||
imageops.save_image(result["cutout"], cutout_png)
|
||||
outputs["cutout"] = cutout_png
|
||||
if result["mask"] is not None:
|
||||
mask_png = os.path.join(task.output_dir, f"{job.index + 1:03d}_{job.label}_mask.png")
|
||||
imageops.save_image(result["mask"], mask_png)
|
||||
outputs["mask"] = mask_png
|
||||
outputs["original"] = job.input_path
|
||||
return outputs
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# HTTP 处理器
|
||||
# --------------------------------------------------------------------------- #
|
||||
class WebUIHandler(BaseHTTPRequestHandler):
|
||||
server_version = f"BiRefNetWebUI/{__version__}"
|
||||
protocol_version = "HTTP/1.1"
|
||||
|
||||
# ------------------------- 基础工具 ------------------------- #
|
||||
@property
|
||||
def app(self) -> "WebUIServer": # type: ignore[override]
|
||||
return self.server # type: ignore[return-value]
|
||||
|
||||
def log_message(self, fmt: str, *args) -> None: # noqa: A003
|
||||
path = getattr(self, "path", "")
|
||||
# 轮询与静态资源太吵,不打印
|
||||
if path.startswith("/api/tasks/") or path.startswith("/static/") or path == "/favicon.ico":
|
||||
return
|
||||
super().log_message(fmt, *args)
|
||||
|
||||
def _send(self, status: int, body: bytes, content_type: str, extra: Optional[dict] = None) -> None:
|
||||
try:
|
||||
self.send_response(status)
|
||||
self.send_header("Content-Type", content_type)
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.send_header("Cache-Control", "no-store")
|
||||
for k, v in (extra or {}).items():
|
||||
self.send_header(k, v)
|
||||
self.end_headers()
|
||||
if self.command != "HEAD":
|
||||
self.wfile.write(body)
|
||||
except (BrokenPipeError, ConnectionResetError, ConnectionAbortedError):
|
||||
pass
|
||||
|
||||
def _send_json(self, payload: object, status: int = 200) -> None:
|
||||
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
|
||||
self._send(status, body, "application/json; charset=utf-8")
|
||||
|
||||
def _error(self, status: int, message: str) -> None:
|
||||
self._send_json({"ok": False, "error": message}, status=status)
|
||||
|
||||
def _read_body(self) -> bytes:
|
||||
try:
|
||||
length = int(self.headers.get("Content-Length") or 0)
|
||||
except ValueError:
|
||||
length = 0
|
||||
if length <= 0:
|
||||
return b""
|
||||
if length > MAX_UPLOAD_BYTES:
|
||||
raise ValueError(
|
||||
f"请求体过大({length / 1024 / 1024:.1f} MB),上限 {MAX_UPLOAD_BYTES / 1024 / 1024:.0f} MB"
|
||||
)
|
||||
return self.rfile.read(length)
|
||||
|
||||
# ------------------------- 路由 ------------------------- #
|
||||
def do_GET(self) -> None: # noqa: N802
|
||||
try:
|
||||
parsed = urlparse(self.path)
|
||||
path = unquote(parsed.path)
|
||||
query = parse_qs(parsed.query)
|
||||
|
||||
if path in ("/", "/index.html"):
|
||||
return self._serve_static("index.html")
|
||||
if path.startswith("/static/"):
|
||||
return self._serve_static(path[len("/static/"):])
|
||||
if path == "/favicon.ico":
|
||||
return self._send(204, b"", "image/x-icon")
|
||||
if path == "/api/health":
|
||||
return self._send_json({"ok": True, "app": APP_NAME, "version": __version__})
|
||||
if path == "/api/state":
|
||||
return self._api_state()
|
||||
if path == "/api/models":
|
||||
return self._send_json({"ok": True, "models": [m.to_dict() for m in self.app.registry.list()]})
|
||||
if path == "/api/tasks":
|
||||
return self._send_json({"ok": True, "tasks": self.app.tasks.list_tasks()})
|
||||
|
||||
parts = path.strip("/").split("/")
|
||||
# /api/tasks/<id>[/...]
|
||||
if len(parts) >= 3 and parts[0] == "api" and parts[1] == "tasks":
|
||||
task_id = parts[2]
|
||||
if len(parts) == 3:
|
||||
task = self.app.tasks.get(task_id)
|
||||
if task is None:
|
||||
return self._error(404, "任务不存在")
|
||||
return self._send_json({"ok": True, "task": self._task_payload(task, query)})
|
||||
if len(parts) == 6 and parts[3] == "file":
|
||||
return self._serve_result(task_id, parts[4], parts[5])
|
||||
if len(parts) == 4 and parts[3] == "zip":
|
||||
return self._serve_zip(task_id, query)
|
||||
return self._error(404, f"未知路径 {path}")
|
||||
except (BrokenPipeError, ConnectionResetError, ConnectionAbortedError):
|
||||
pass
|
||||
except Exception as exc: # pragma: no cover
|
||||
traceback.print_exc()
|
||||
self._error(500, f"{type(exc).__name__}: {exc}")
|
||||
|
||||
def do_HEAD(self) -> None: # noqa: N802
|
||||
self.do_GET()
|
||||
|
||||
def do_POST(self) -> None: # noqa: N802
|
||||
try:
|
||||
parsed = urlparse(self.path)
|
||||
path = unquote(parsed.path)
|
||||
|
||||
if path == "/api/models/reload":
|
||||
models = self.app.registry.scan()
|
||||
return self._send_json({"ok": True, "models": [m.to_dict() for m in models]})
|
||||
if path == "/api/engine/unload":
|
||||
self.app.engine.unload()
|
||||
return self._send_json({"ok": True})
|
||||
if path == "/api/shutdown":
|
||||
self._send_json({"ok": True, "message": "服务正在关闭"})
|
||||
threading.Thread(target=self.app.stop_soon, daemon=True).start()
|
||||
return
|
||||
if path == "/api/tasks":
|
||||
return self._create_task()
|
||||
if path == "/api/batch":
|
||||
return self._create_batch_task()
|
||||
|
||||
parts = path.strip("/").split("/")
|
||||
if len(parts) == 4 and parts[0] == "api" and parts[1] == "tasks" and parts[3] == "cancel":
|
||||
ok = self.app.tasks.cancel(parts[2])
|
||||
return self._send_json({"ok": ok})
|
||||
return self._error(404, f"未知路径 {path}")
|
||||
except (BrokenPipeError, ConnectionResetError, ConnectionAbortedError):
|
||||
pass
|
||||
except Exception as exc: # pragma: no cover
|
||||
traceback.print_exc()
|
||||
self._error(500, f"{type(exc).__name__}: {exc}")
|
||||
|
||||
# ------------------------- 具体处理 ------------------------- #
|
||||
@staticmethod
|
||||
def _int_param(query: dict, key: str, default: int) -> int:
|
||||
"""从 query(parse_qs 的结果)里安全地取一个整数参数。"""
|
||||
try:
|
||||
return int(str((query.get(key) or [default])[0]))
|
||||
except (TypeError, ValueError):
|
||||
return default
|
||||
|
||||
def _task_payload(self, task: Task, query: dict) -> dict:
|
||||
"""任务详情;``?tail=N`` 只回传最后 N 张(批量任务轮询用,压住载荷)。"""
|
||||
return task.to_dict(include_options=True, tail=max(0, self._int_param(query, "tail", 0)))
|
||||
|
||||
def _merge_options(self, incoming: object) -> dict:
|
||||
"""把前端传来的参数合并到默认参数上(只接受白名单键)。"""
|
||||
options = dict(DEFAULT_OPTIONS)
|
||||
if isinstance(incoming, dict):
|
||||
options.update({k: v for k, v in incoming.items() if k in DEFAULT_OPTIONS})
|
||||
return options
|
||||
|
||||
def _pin_model(self, options: dict) -> Optional[str]:
|
||||
"""确保 ``options['model']`` 指向一个真实存在的权重。
|
||||
|
||||
Returns:
|
||||
出错时返回给人看的错误信息;成功返回 None。
|
||||
"""
|
||||
models = self.app.registry.list()
|
||||
if not models:
|
||||
return "模型目录中没有找到任何权重文件(*.safetensors / *.pth)"
|
||||
options["model"] = options.get("model") or self.app.default_model_key(
|
||||
[m.to_dict() for m in models]
|
||||
)
|
||||
try:
|
||||
self.app.registry.get(str(options["model"]))
|
||||
except KeyError:
|
||||
options["model"] = models[0].key
|
||||
return None
|
||||
|
||||
def _api_state(self) -> None:
|
||||
app = self.app
|
||||
models = [m.to_dict() for m in app.registry.list()]
|
||||
default_model = app.default_model_key(models)
|
||||
self._send_json(
|
||||
{
|
||||
"ok": True,
|
||||
"app": APP_NAME,
|
||||
"version": __version__,
|
||||
"environment": describe_environment(app.engine),
|
||||
"node_dir": str(app.node_dir),
|
||||
"model_dirs": [str(d) for d in app.registry.model_dirs],
|
||||
"output_dir": str(app.output_root),
|
||||
"project_root": str(app.project_root),
|
||||
"models": models,
|
||||
"defaults": {**DEFAULT_OPTIONS, "model": default_model},
|
||||
}
|
||||
)
|
||||
|
||||
def _create_task(self) -> None:
|
||||
content_type = self.headers.get("Content-Type", "")
|
||||
if "multipart/form-data" not in content_type:
|
||||
return self._error(400, "Content-Type 必须是 multipart/form-data")
|
||||
boundary = None
|
||||
for chunk in content_type.split(";"):
|
||||
chunk = chunk.strip()
|
||||
if chunk.lower().startswith("boundary="):
|
||||
boundary = chunk[9:].strip().strip('"').encode("latin-1")
|
||||
if not boundary:
|
||||
return self._error(400, "缺少 multipart boundary")
|
||||
|
||||
try:
|
||||
body = self._read_body()
|
||||
except ValueError as exc:
|
||||
return self._error(413, str(exc))
|
||||
|
||||
options = dict(DEFAULT_OPTIONS)
|
||||
uploads: List[Tuple[str, bytes]] = []
|
||||
for part in parse_multipart(body, boundary):
|
||||
if part.name == "options":
|
||||
try:
|
||||
options = self._merge_options(json.loads(part.data.decode("utf-8") or "{}"))
|
||||
except json.JSONDecodeError:
|
||||
return self._error(400, "options 字段不是合法 JSON")
|
||||
elif part.name in ("files", "images", "file"):
|
||||
if part.data:
|
||||
uploads.append((part.filename or f"upload_{len(uploads) + 1}.png", part.data))
|
||||
|
||||
if not uploads:
|
||||
return self._error(400, "没有收到任何图片,请选择 PNG / JPG / WebP 文件")
|
||||
|
||||
error = self._pin_model(options)
|
||||
if error:
|
||||
return self._error(400, error)
|
||||
|
||||
try:
|
||||
task = self.app.tasks.submit(uploads, options)
|
||||
except ValueError as exc:
|
||||
return self._error(400, str(exc))
|
||||
self._send_json({"ok": True, "task_id": task.id, "count": len(uploads)})
|
||||
|
||||
def _create_batch_task(self) -> None:
|
||||
"""POST /api/batch —— 目录批量抠图(JSON 传路径,不上传文件)。
|
||||
|
||||
请求体:
|
||||
{
|
||||
"input_dir": "F:\\\\photos\\\\待抠图", # 必填,图片所在目录
|
||||
"output_dir": "F:\\\\photos\\\\results", # 可空 → 用项目 outputs/
|
||||
"recursive": false, # 是否含子目录
|
||||
"options": { …与单图模式相同的参数… }
|
||||
}
|
||||
|
||||
两个目录都必须**已经存在**:路径不对就报错,本接口绝不会创建目录。
|
||||
"""
|
||||
try:
|
||||
payload = json.loads(self._read_body().decode("utf-8") or "{}")
|
||||
except (ValueError, UnicodeDecodeError):
|
||||
return self._error(400, "请求体不是合法 JSON")
|
||||
if not isinstance(payload, dict):
|
||||
return self._error(400, "请求体必须是 JSON 对象")
|
||||
|
||||
base = self.app.project_root
|
||||
raw_out = str(payload.get("output_dir") or "").strip() or str(self.app.output_root)
|
||||
try:
|
||||
input_dir = batch.check_input_dir(payload.get("input_dir"), base)
|
||||
output_dir = batch.check_output_dir(raw_out, base)
|
||||
except batch.BatchPathError as exc:
|
||||
return self._error(400, str(exc))
|
||||
|
||||
options = self._merge_options(payload.get("options"))
|
||||
error = self._pin_model(options)
|
||||
if error:
|
||||
return self._error(400, error)
|
||||
|
||||
try:
|
||||
task = self.app.tasks.submit_batch(
|
||||
input_dir, output_dir, options, recursive=bool(payload.get("recursive"))
|
||||
)
|
||||
except ValueError as exc:
|
||||
return self._error(400, str(exc))
|
||||
self._send_json(
|
||||
{
|
||||
"ok": True,
|
||||
"task_id": task.id,
|
||||
"mode": "batch",
|
||||
"count": len(task.images),
|
||||
"input_dir": str(input_dir),
|
||||
"output_dir": str(output_dir),
|
||||
}
|
||||
)
|
||||
|
||||
def _task_file(self, task_id: str, index: int, kind: str) -> Optional[str]:
|
||||
task = self.app.tasks.get(task_id)
|
||||
if task is None:
|
||||
return None
|
||||
for job in task.images:
|
||||
if job.index == index:
|
||||
return job.outputs.get(kind)
|
||||
return None
|
||||
|
||||
def _serve_result(self, task_id: str, idx_part: str, kind: str) -> None:
|
||||
try:
|
||||
index = int(idx_part)
|
||||
except ValueError:
|
||||
return self._error(400, "图片序号非法")
|
||||
if kind not in ("cutout", "mask", "original"):
|
||||
return self._error(400, f"未知的结果类型 {kind!r}")
|
||||
path = self._task_file(task_id, index, kind)
|
||||
if not path or not os.path.isfile(path):
|
||||
return self._error(404, "结果文件不存在")
|
||||
suffix = Path(path).suffix.lower()
|
||||
with open(path, "rb") as fh:
|
||||
data = fh.read()
|
||||
self._send(200, data, MIME_TYPES.get(suffix, "application/octet-stream"))
|
||||
|
||||
def _serve_zip(self, task_id: str, query: dict) -> None:
|
||||
task = self.app.tasks.get(task_id)
|
||||
if task is None:
|
||||
return self._error(404, "任务不存在")
|
||||
kinds = (query.get("kind") or ["cutout"])[0].split(",")
|
||||
want = [k for k in kinds if k in ("cutout", "mask", "original")] or ["cutout"]
|
||||
buffer = io.BytesIO()
|
||||
added = 0
|
||||
with zipfile.ZipFile(buffer, "w", zipfile.ZIP_DEFLATED) as zf:
|
||||
for job in task.images:
|
||||
for kind in want:
|
||||
path = job.outputs.get(kind)
|
||||
if path and os.path.isfile(path):
|
||||
zf.write(path, arcname=os.path.basename(path))
|
||||
added += 1
|
||||
if added == 0:
|
||||
return self._error(404, "没有可打包的结果文件")
|
||||
self._send(
|
||||
200,
|
||||
buffer.getvalue(),
|
||||
"application/zip",
|
||||
extra={"Content-Disposition": f'attachment; filename="birefnet_{task_id}.zip"'},
|
||||
)
|
||||
|
||||
def _serve_static(self, rel_path: str) -> None:
|
||||
web_dir = self.app.web_dir.resolve()
|
||||
target = (web_dir / rel_path.replace("\\", "/")).resolve()
|
||||
try:
|
||||
target.relative_to(web_dir) # 防路径穿越
|
||||
except ValueError:
|
||||
return self._error(403, "非法路径")
|
||||
if not target.is_file():
|
||||
return self._error(404, f"资源不存在:{rel_path}")
|
||||
suffix = target.suffix.lower()
|
||||
with open(target, "rb") as fh:
|
||||
data = fh.read()
|
||||
cache = "no-cache" if suffix in (".html", ".js", ".css") else "public, max-age=3600"
|
||||
self._send(200, data, MIME_TYPES.get(suffix, "application/octet-stream"), {"Cache-Control": cache})
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 服务容器
|
||||
# --------------------------------------------------------------------------- #
|
||||
class WebUIServer(ThreadingHTTPServer):
|
||||
daemon_threads = True
|
||||
allow_reuse_address = True
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
address: Tuple[str, int],
|
||||
registry: ModelRegistry,
|
||||
engine: InferenceEngine,
|
||||
tasks: TaskManager,
|
||||
web_dir: Path,
|
||||
output_root: Path,
|
||||
node_dir: Path,
|
||||
logger=None,
|
||||
project_root: Optional[Path] = None,
|
||||
) -> None:
|
||||
super().__init__(address, WebUIHandler)
|
||||
self.registry = registry
|
||||
self.engine = engine
|
||||
self.tasks = tasks
|
||||
self.web_dir = Path(web_dir)
|
||||
self.output_root = Path(output_root)
|
||||
self.node_dir = Path(node_dir)
|
||||
#: 相对路径型用户输入(批量处理的目录)以此为准
|
||||
self.project_root = Path(project_root) if project_root else self.web_dir.parent
|
||||
self.log = logger or (lambda msg: None)
|
||||
|
||||
def default_model_key(self, models: Optional[Iterable[dict]] = None) -> str:
|
||||
models = list(models if models is not None else (m.to_dict() for m in self.registry.list()))
|
||||
if not models:
|
||||
return ""
|
||||
for preferred in ("Portrait", "General", "General-HR"):
|
||||
for m in models:
|
||||
if m["name"] == preferred:
|
||||
return m["key"]
|
||||
return models[0]["key"]
|
||||
|
||||
def stop_soon(self) -> None:
|
||||
time.sleep(0.4)
|
||||
self.log("收到关闭指令,服务即将退出")
|
||||
self.shutdown()
|
||||
+12
@@ -0,0 +1,12 @@
|
||||
{
|
||||
"host": "127.0.0.1",
|
||||
"port": 7861,
|
||||
"node_dir": "",
|
||||
"model_dirs": [],
|
||||
"output_dir": "outputs",
|
||||
"device": "auto",
|
||||
"dtype": "auto",
|
||||
"max_cached_models": 1,
|
||||
"max_tasks": 50,
|
||||
"open_browser": true
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
# BiRefNet WebUI 依赖清单
|
||||
#
|
||||
# 说明:如果你直接用 ComfyUI 便携版自带的嵌入式 Python 运行(推荐),
|
||||
# 这些包通常已经全部存在,无需再装。只有用独立 venv 时才需要:
|
||||
# python -m pip install -r requirements.txt
|
||||
|
||||
torch>=2.1
|
||||
torchvision>=0.16
|
||||
numpy>=1.24
|
||||
pillow>=9.5
|
||||
safetensors>=0.4
|
||||
timm>=0.9
|
||||
einops>=0.7
|
||||
kornia>=0.7
|
||||
huggingface_hub>=0.20
|
||||
tqdm>=4.65
|
||||
opencv-python>=4.8
|
||||
+179
@@ -0,0 +1,179 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""端到端自检脚本:验证依赖、模型加载与抠图结果。
|
||||
|
||||
用法(本项目自包含,直接用自带的 Python 运行时):
|
||||
python\\python.exe selftest.py # 用内置合成图测试第一个可用模型
|
||||
python\\python.exe selftest.py --image photo.jpg # 用真实图片测试
|
||||
python\\python.exe selftest.py --device cpu # 强制 CPU
|
||||
python\\python.exe selftest.py --thorough # 把所有模型都跑一遍(慢)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import sys
|
||||
import time
|
||||
from pathlib import Path
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parent
|
||||
if str(PROJECT_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
|
||||
import numpy as np # noqa: E402
|
||||
from PIL import Image, ImageDraw # noqa: E402
|
||||
|
||||
|
||||
def make_test_image(width: int = 900, height: int = 640) -> Image.Image:
|
||||
"""合成一张「纯色背景 + 高对比主体」的测试图,便于断言遮罩是否合理。"""
|
||||
y, x = np.mgrid[0:height, 0:width].astype(np.float32)
|
||||
bg = np.stack(
|
||||
[120 + 90 * x / width, 130 + 60 * y / height, 200 - 90 * x / width], axis=-1
|
||||
).astype(np.uint8)
|
||||
image = Image.fromarray(bg, "RGB")
|
||||
draw = ImageDraw.Draw(image)
|
||||
# 主体:一个亮色椭圆(占画面约 20% 面积)
|
||||
box = (width * 0.25, height * 0.18, width * 0.75, height * 0.82)
|
||||
draw.ellipse(box, fill=(245, 240, 235))
|
||||
draw.ellipse((width * 0.36, height * 0.30, width * 0.42, height * 0.40), fill=(40, 40, 45))
|
||||
draw.ellipse((width * 0.58, height * 0.30, width * 0.64, height * 0.40), fill=(40, 40, 45))
|
||||
return image
|
||||
|
||||
|
||||
def stats(arr: np.ndarray, name: str) -> None:
|
||||
print(
|
||||
f" {name:<12} 形状={arr.shape} 最小值={arr.min():.3f} "
|
||||
f"最大值={arr.max():.3f} 均值={arr.mean():.3f}"
|
||||
)
|
||||
|
||||
|
||||
def main(argv=None) -> int:
|
||||
parser = argparse.ArgumentParser(description="BiRefNet WebUI 自检")
|
||||
parser.add_argument("--image", default=None, help="测试图片路径(默认使用内置合成图)")
|
||||
parser.add_argument("--model", default=None, help="模型 key / 名称(默认取第一个可用模型)")
|
||||
parser.add_argument("--device", default="auto", choices=["auto", "cpu", "cuda"])
|
||||
parser.add_argument("--dtype", default="auto", choices=["auto", "float32", "float16", "bfloat16"])
|
||||
parser.add_argument("--thorough", action="store_true", help="遍历所有模型")
|
||||
parser.add_argument("--final-side", type=int, default=0,
|
||||
help="最终尺寸·最长边(0 = 保持原图尺寸;>0 = 等比缩放到该最长边)")
|
||||
parser.add_argument("--out", default=str(PROJECT_ROOT / "outputs" / "selftest"), help="结果输出目录")
|
||||
args = parser.parse_args(argv)
|
||||
|
||||
from birefnet_web import compat
|
||||
from birefnet_web.engine import InferenceEngine, ModelRegistry, describe_environment
|
||||
from birefnet_web import imageops
|
||||
|
||||
print("=" * 68)
|
||||
print("BiRefNet WebUI 自检")
|
||||
print("=" * 68)
|
||||
|
||||
node_dir = compat.find_node_dir(None)
|
||||
print(f"模型代码 {node_dir}")
|
||||
|
||||
model_dirs = compat.default_model_dirs(node_dir)
|
||||
registry = ModelRegistry(model_dirs)
|
||||
models = registry.scan()
|
||||
print(f"模型目录 {', '.join(str(d) for d in model_dirs)}")
|
||||
for m in models:
|
||||
print(f" · {m.name:<22} {m.size_mb:>7.1f} MB arch={m.arch} bb={m.backbone} {m.path}")
|
||||
if not models:
|
||||
print("[失败] 没有找到任何模型权重,请先准备 *.safetensors")
|
||||
return 1
|
||||
|
||||
env = describe_environment(InferenceEngine(registry, node_dir))
|
||||
print(f"运行环境 python {env['python']} · torch {env['torch']} · cuda={env['cuda']}")
|
||||
for dev in env["devices"]:
|
||||
print(f" [{dev['index']}] {dev['name']} {dev['free_mem']}G/{dev['total_mem']}G")
|
||||
|
||||
if args.image:
|
||||
src = Image.open(args.image)
|
||||
src.load()
|
||||
print(f"测试图片 {args.image} {src.size}")
|
||||
else:
|
||||
src = make_test_image()
|
||||
print(f"测试图片 内置合成图 {src.size}")
|
||||
|
||||
engine = InferenceEngine(registry, node_dir=node_dir, max_cached=1)
|
||||
targets = models if args.thorough else [registry.get(args.model) if args.model else models[0]]
|
||||
out_dir = Path(args.out)
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
failures = 0
|
||||
for info in targets:
|
||||
print("-" * 68)
|
||||
print(f"模型 {info.name}(arch={info.arch}, bb={info.backbone})")
|
||||
stages = []
|
||||
t0 = time.perf_counter()
|
||||
try:
|
||||
result = engine.remove_background(
|
||||
src,
|
||||
{
|
||||
"model": info.key,
|
||||
"device": args.device,
|
||||
"dtype": args.dtype,
|
||||
"resolution_mode": "square",
|
||||
"width": 1024,
|
||||
"height": 1024,
|
||||
"upscale_method": "bilinear",
|
||||
"mask_threshold": 0.0,
|
||||
"refine_foreground": True,
|
||||
"background": "transparent",
|
||||
"output_mask": True,
|
||||
"final_longest_side": args.final_side,
|
||||
},
|
||||
progress=lambda stage, frac: stages.append(f"{stage}({frac:.0%})"),
|
||||
)
|
||||
except Exception as exc:
|
||||
failures += 1
|
||||
print(f" [失败] {type(exc).__name__}: {exc}")
|
||||
continue
|
||||
|
||||
cutout, mask = result["cutout"], result["mask"]
|
||||
alpha = np.asarray(cutout.split()[-1]) if cutout.mode == "RGBA" else None
|
||||
|
||||
print(f" 设备/精度 {result['device']} / {result['dtype']} 输入 {result['input_size']}")
|
||||
print(f" 输出 {cutout.mode} {cutout.size},耗时 {result['elapsed']}s")
|
||||
print(f" 前景占比 {result['coverage']:.1%}")
|
||||
print(f" 阶段 {' → '.join(stages)}")
|
||||
if result.get("warning"):
|
||||
print(f" [告警] {result['warning']}")
|
||||
if alpha is not None:
|
||||
stats(alpha.astype(np.float32) / 255.0, "alpha")
|
||||
|
||||
# 最终尺寸:结果尺寸必须等于「原图按比例缩放到最长边 = final_side」
|
||||
expect_size = imageops.fit_longest_side((src.width, src.height), args.final_side)
|
||||
if cutout.size != expect_size:
|
||||
failures += 1
|
||||
print(f" [失败] 最终尺寸不符:期望 {expect_size},实际 {cutout.size}")
|
||||
else:
|
||||
print(f" [通过] 最终尺寸 {cutout.size}(最长边参数 {args.final_side})")
|
||||
|
||||
cutout_path = out_dir / f"{info.name}_cutout.png"
|
||||
imageops.save_image(cutout, str(cutout_path))
|
||||
if mask is not None:
|
||||
imageops.save_image(mask, str(out_dir / f"{info.name}_mask.png"))
|
||||
print(f" 已保存 {cutout_path}")
|
||||
|
||||
if args.image is None:
|
||||
coverage = result["coverage"]
|
||||
if not 0.05 < coverage < 0.95:
|
||||
failures += 1
|
||||
print(f" [失败] 合成图期望前景占比在 5%~95% 之间,实际 {coverage:.1%}")
|
||||
elif alpha is not None and not (
|
||||
alpha.min() < 40 and alpha.max() > 215
|
||||
):
|
||||
failures += 1
|
||||
print(" [失败] alpha 动态范围不足,遮罩可能异常")
|
||||
else:
|
||||
print(" [通过] 遮罩分布合理")
|
||||
print(f" 总耗时 {time.perf_counter() - t0:.1f}s")
|
||||
|
||||
print("=" * 68)
|
||||
if failures:
|
||||
print(f"自检未通过:{failures} 项失败")
|
||||
return 1
|
||||
print("自检全部通过 ✔")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,386 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""批量抠图自检:路径校验 / 命名规则 / 任务流水线 / HTTP 接口契约。
|
||||
|
||||
**不加载任何真实权重**(用假引擎直接产出结果图),所以跑起来只要几秒。
|
||||
|
||||
用法:
|
||||
python\\python.exe selftest_batch.py # 全量
|
||||
python selftest_batch.py --keep # 保留临时目录,便于看产物
|
||||
python selftest_batch.py --logic-only # 只跑不依赖 torch 的纯逻辑部分
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import shutil
|
||||
import sys
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
import urllib.error
|
||||
import urllib.request
|
||||
from pathlib import Path
|
||||
from typing import List, Optional
|
||||
|
||||
PROJECT_ROOT = Path(__file__).resolve().parent
|
||||
if str(PROJECT_ROOT) not in sys.path:
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
|
||||
_FAILS: List[str] = []
|
||||
|
||||
|
||||
def check(name: str, cond: bool, extra: object = "") -> None:
|
||||
print(f" {'[通过]' if cond else '[失败]'} {name}" + ("" if cond else f" ← {extra}"))
|
||||
if not cond:
|
||||
_FAILS.append(name)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 1) 纯逻辑:路径校验 + 命名规则(不需要 torch)
|
||||
# --------------------------------------------------------------------------- #
|
||||
def part_logic(tmp: Path) -> None:
|
||||
from birefnet_web import batch
|
||||
|
||||
src, out = tmp / "src", tmp / "out"
|
||||
out.mkdir(parents=True)
|
||||
(src / "sub").mkdir(parents=True)
|
||||
(out / "RMBG_already_1.png").write_bytes(b"x")
|
||||
(src / "photo.jpg").write_bytes(b"x")
|
||||
(src / "a_pretty_long_original_filename_here.png").write_bytes(b"x")
|
||||
(src / "note.txt").write_bytes(b"x")
|
||||
(src / "RMBG_old_1.png").write_bytes(b"x")
|
||||
(src / "sub" / "deep.webp").write_bytes(b"x")
|
||||
|
||||
print("\n[1] 路径校验")
|
||||
missing = tmp / "nope"
|
||||
for label, fn in (("输入", batch.check_input_dir), ("输出", batch.check_output_dir)):
|
||||
try:
|
||||
fn(str(missing), PROJECT_ROOT)
|
||||
check(f"{label}目录不存在时抛错", False, "没有抛 BatchPathError")
|
||||
except batch.BatchPathError as exc:
|
||||
check(f"{label}目录不存在时抛错", True)
|
||||
check(f"{label}报错时未创建目录", not missing.exists())
|
||||
print(f" → {exc}")
|
||||
try:
|
||||
batch.check_input_dir(str(src / "photo.jpg"), PROJECT_ROOT)
|
||||
check("输入是文件时抛错", False, "没有抛 BatchPathError")
|
||||
except batch.BatchPathError as exc:
|
||||
check("输入是文件时抛错", True)
|
||||
print(f" → {exc}")
|
||||
try:
|
||||
batch.check_output_dir(str(PROJECT_ROOT), PROJECT_ROOT, allow_system_dir=False)
|
||||
check("项目目录可写时通过", True)
|
||||
except batch.BatchPathError as exc:
|
||||
check("项目目录可写时通过", False, exc)
|
||||
check("相对路径基于项目根解析", batch.resolve_user_path("outputs", PROJECT_ROOT) == PROJECT_ROOT / "outputs")
|
||||
check("成对引号被去掉", batch.resolve_user_path(' "outputs" ', PROJECT_ROOT) == PROJECT_ROOT / "outputs")
|
||||
|
||||
print("\n[2] 目录扫描")
|
||||
names = sorted(p.name for p in batch.scan_images(src))
|
||||
check("非递归只扫顶层图片", names == ["a_pretty_long_original_filename_here.png", "photo.jpg"], names)
|
||||
deep = batch.scan_images(src, recursive=True)
|
||||
check("递归包含子目录", len(deep) == 3, [p.name for p in deep])
|
||||
check("跳过非图片文件", all(p.suffix != ".txt" for p in deep))
|
||||
check("跳过本工具产物 RMBG_*", all(not p.name.startswith("RMBG_") for p in deep))
|
||||
check("跳过输出目录", all(out not in p.parents for p in batch.scan_images(src, recursive=True, skip_dirs=[out])))
|
||||
|
||||
print("\n[3] 命名规则")
|
||||
long_src = src / "a_pretty_long_original_filename_here.png"
|
||||
check("RMBG_ + 原主干 + _ + 时间戳 + 原后缀",
|
||||
batch.build_output_name(src / "photo.jpg", 1700000000, "color") == "RMBG_photo_1700000000.jpg")
|
||||
check("主干超 20 字符截断",
|
||||
batch.build_output_name(long_src, 1700000000, "transparent") == "RMBG_a_pretty_long_origin_1700000000.png",
|
||||
batch.build_output_name(long_src, 1700000000, "transparent"))
|
||||
check("中文按字符截断(不是按字节)",
|
||||
batch.build_output_name(Path("一二三四五六七八九十一二三四五六七八九十一二三.png"), 7, "transparent")
|
||||
== "RMBG_一二三四五六七八九十一二三四五六七八九十_7.png")
|
||||
check("透明 + jpg 退化 .png(jpg 装不下 alpha)",
|
||||
batch.build_output_name(src / "photo.jpg", 1, "transparent") == "RMBG_photo_1.png")
|
||||
check("纯色 + jpg 保持 .jpg",
|
||||
batch.build_output_name(src / "photo.jpg", 1, "color") == "RMBG_photo_1.jpg")
|
||||
check("非法字符被清洗", batch.build_output_name(Path("a?c*d.jpg"), 1, "color") == "RMBG_a_c_d_1.jpg")
|
||||
check("同名文件避让为 _1", batch.unique_path(out, "RMBG_already_1.png").name == "RMBG_already_1_1.png")
|
||||
check("无同名时原样返回", batch.unique_path(out, "RMBG_new_1.png").name == "RMBG_new_1.png")
|
||||
|
||||
print("\n[4] 最终尺寸 · 最长边换算")
|
||||
from birefnet_web import fsutil
|
||||
|
||||
fit = fsutil.fit_longest_side
|
||||
check("0 = 保持原图尺寸", fit((2000, 3000), 0) == (2000, 3000))
|
||||
check("负数同样按不缩放处理", fit((2000, 3000), -1) == (2000, 3000))
|
||||
check("需求例 2:竖图 2000×3000 → 1440 得 960×1440", fit((2000, 3000), 1440) == (960, 1440), fit((2000, 3000), 1440))
|
||||
check("需求例 3:横图 3000×2000 → 1440 得 1440×960", fit((3000, 2000), 1440) == (1440, 960), fit((3000, 2000), 1440))
|
||||
check("最长边已等于目标值时不动", fit((1440, 960), 1440) == (1440, 960))
|
||||
check("小图按比例放大(800×600 → 1600 得 1600×1200)", fit((800, 600), 1600) == (1600, 1200), fit((800, 600), 1600))
|
||||
check("正方形等比缩放", fit((1000, 1000), 500) == (500, 500))
|
||||
check("极端比例也不会算出 0 像素", fit((10000, 3), 10) == (10, 1), fit((10000, 3), 10))
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 2) 任务流水线(假引擎,不加载权重)
|
||||
# --------------------------------------------------------------------------- #
|
||||
class FakeModel:
|
||||
key = "fake:safetensors"
|
||||
name = "Fake"
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {"key": self.key, "name": self.name, "file": "fake.safetensors", "arch": "v1",
|
||||
"size_mb": 1.0, "backbone": "swin_v1_tiny", "path": "fake"}
|
||||
|
||||
|
||||
class FakeRegistry:
|
||||
model_dirs: List[Path] = []
|
||||
|
||||
def list(self) -> List[FakeModel]:
|
||||
return [FakeModel()]
|
||||
|
||||
def get(self, key: str) -> FakeModel:
|
||||
if key != FakeModel.key:
|
||||
raise KeyError(key)
|
||||
return FakeModel()
|
||||
|
||||
|
||||
class FakeEngine:
|
||||
"""只按参数产出一张图,用来验证落盘与命名,不碰 torch。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.calls = 0
|
||||
|
||||
def unload(self) -> None:
|
||||
pass
|
||||
|
||||
def remove_background(self, image, options, progress=None): # noqa: ANN001
|
||||
self.calls += 1
|
||||
if progress:
|
||||
progress("推理", 0.5)
|
||||
size = image.size
|
||||
transparent = options.get("background") == "transparent"
|
||||
rgba = image.convert("RGBA")
|
||||
rgba.putalpha(128)
|
||||
cutout = rgba if transparent else rgba.convert("RGB")
|
||||
return {
|
||||
"cutout": cutout,
|
||||
"mask": image.convert("L") if options.get("output_mask") else None,
|
||||
"width": size[0], "height": size[1], "coverage": 0.5, "elapsed": 0.01,
|
||||
"input_size": (1024, 1024), "warning": None, "device": "cpu", "dtype": "float32",
|
||||
}
|
||||
|
||||
|
||||
def _seed_images(root: Path) -> None:
|
||||
from PIL import Image
|
||||
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
Image.new("RGB", (64, 48), (200, 30, 30)).save(root / "photo.jpg")
|
||||
Image.new("RGB", (32, 32), (30, 200, 30)).save(root / "a_pretty_long_original_filename_here.png")
|
||||
(root / "note.txt").write_text("not an image", encoding="utf-8")
|
||||
|
||||
|
||||
def _wait(task, timeout: float = 30.0) -> bool:
|
||||
deadline = time.time() + timeout
|
||||
while time.time() < deadline:
|
||||
if task.finished_at is not None:
|
||||
return True
|
||||
time.sleep(0.05)
|
||||
return False
|
||||
|
||||
|
||||
def part_pipeline(tmp: Path) -> None:
|
||||
from birefnet_web.server import TaskManager
|
||||
|
||||
print("\n[4] 批量任务流水线(假引擎)")
|
||||
src, out = tmp / "pipe_src", tmp / "pipe_out"
|
||||
_seed_images(src)
|
||||
# 两个不同子目录里的同名文件 → 检验同一秒内的重名避让
|
||||
(src / "a").mkdir()
|
||||
(src / "b").mkdir()
|
||||
from PIL import Image
|
||||
|
||||
Image.new("RGB", (16, 16), (0, 0, 255)).save(src / "a" / "pic.png")
|
||||
Image.new("RGB", (16, 16), (255, 0, 255)).save(src / "b" / "pic.png")
|
||||
out.mkdir()
|
||||
|
||||
engine = FakeEngine()
|
||||
manager = TaskManager(engine, tmp / "default_out", max_tasks=5)
|
||||
options = {"background": "transparent", "output_mask": False}
|
||||
|
||||
task = manager.submit_batch(src, out, options, recursive=True)
|
||||
check("任务模式标记为 batch", task.mode == "batch")
|
||||
check("原图未被拷贝进任务目录(就地读取)",
|
||||
all(Path(j.input_path).parent != tmp / "default_out" for j in task.images))
|
||||
check("递归扫描命中 4 张", len(task.images) == 4, len(task.images))
|
||||
check("等待处理结束", _wait(task))
|
||||
check("任务状态为 done", task.state == "done", task.state)
|
||||
|
||||
produced = sorted(p.name for p in out.iterdir())
|
||||
print(f" 产物:{produced}")
|
||||
check("产物数量 = 图片数量", len(produced) == 4, produced)
|
||||
check("全部带 RMBG_ 前缀", all(n.startswith("RMBG_") for n in produced))
|
||||
check("jpg + 透明 → .png", any(n.startswith("RMBG_photo_") and n.endswith(".png") for n in produced))
|
||||
check("长主干截断到 20 字符",
|
||||
any(n.startswith("RMBG_a_pretty_long_origin_") for n in produced))
|
||||
check("同名文件避让(无覆盖)",
|
||||
sum(1 for n in produced if n.startswith("RMBG_pic_")) == 2, produced)
|
||||
check("每张都记录了 output_name",
|
||||
all(j.output_name and (out / j.output_name).is_file() for j in task.images))
|
||||
|
||||
# 纯色背景 + 同时输出遮罩
|
||||
out2 = tmp / "pipe_out_color"
|
||||
out2.mkdir()
|
||||
task2 = manager.submit_batch(src / "a", out2, {"background": "color", "output_mask": True})
|
||||
check("子目录只有 1 张", len(task2.images) == 1)
|
||||
check("等待第二个任务结束", _wait(task2))
|
||||
names2 = sorted(p.name for p in out2.iterdir())
|
||||
print(f" 产物:{names2}")
|
||||
check("纯色 + png 保持 .png", any(n.startswith("RMBG_pic_") and n.endswith(".png") and "_mask" not in n for n in names2))
|
||||
check("同时输出遮罩 _mask.png", any(n.endswith("_mask.png") for n in names2), names2)
|
||||
|
||||
# 输入目录不存在 → 由 HTTP 层拦下,这里只验目录空的情况
|
||||
try:
|
||||
manager.submit_batch(tmp / "empty_in", out, options)
|
||||
check("空目录抛 ValueError", False, "没有抛")
|
||||
except ValueError as exc:
|
||||
check("空目录抛 ValueError", True)
|
||||
print(f" → {exc}")
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# 3) HTTP 契约(前端就是照这个调的)
|
||||
# --------------------------------------------------------------------------- #
|
||||
def _post(url: str, payload: dict) -> tuple:
|
||||
body = json.dumps(payload).encode("utf-8")
|
||||
req = urllib.request.Request(url, data=body, headers={"Content-Type": "application/json"}, method="POST")
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=10) as res: # noqa: S310
|
||||
return res.status, json.loads(res.read().decode("utf-8"))
|
||||
except urllib.error.HTTPError as exc:
|
||||
return exc.code, json.loads(exc.read().decode("utf-8"))
|
||||
|
||||
|
||||
def _get(url: str) -> tuple:
|
||||
try:
|
||||
with urllib.request.urlopen(url, timeout=10) as res: # noqa: S310
|
||||
return res.status, json.loads(res.read().decode("utf-8"))
|
||||
except urllib.error.HTTPError as exc:
|
||||
return exc.code, json.loads(exc.read().decode("utf-8"))
|
||||
|
||||
|
||||
def part_http(tmp: Path) -> None:
|
||||
from birefnet_web.server import TaskManager, WebUIServer
|
||||
|
||||
print("\n[5] HTTP 接口 /api/batch")
|
||||
engine = FakeEngine()
|
||||
out_root = tmp / "http_default_out"
|
||||
tasks = TaskManager(engine, out_root, max_tasks=5)
|
||||
server = WebUIServer(
|
||||
("127.0.0.1", 0),
|
||||
registry=FakeRegistry(), # type: ignore[arg-type]
|
||||
engine=engine, # type: ignore[arg-type]
|
||||
tasks=tasks,
|
||||
web_dir=PROJECT_ROOT / "web",
|
||||
output_root=out_root,
|
||||
node_dir=PROJECT_ROOT / "vendor",
|
||||
logger=lambda msg: None,
|
||||
project_root=PROJECT_ROOT,
|
||||
)
|
||||
port = server.server_address[1]
|
||||
threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.1}, daemon=True).start()
|
||||
base = f"http://127.0.0.1:{port}"
|
||||
try:
|
||||
src = tmp / "http_src"
|
||||
_seed_images(src)
|
||||
out = tmp / "http_out"
|
||||
out.mkdir()
|
||||
|
||||
status, data = _post(f"{base}/api/batch", {"input_dir": str(tmp / "missing")})
|
||||
check("输入目录不存在 → 400", status == 400, (status, data))
|
||||
check("错误信息直接点出不存在", "不存在" in str(data.get("error")), data)
|
||||
check("仍然没有创建该目录", not (tmp / "missing").exists())
|
||||
|
||||
status, data = _post(f"{base}/api/batch", {"input_dir": str(src), "output_dir": str(tmp / "missing_out")})
|
||||
check("输出目录不存在 → 400", status == 400, (status, data))
|
||||
check("仍然没有创建输出目录", not (tmp / "missing_out").exists())
|
||||
|
||||
status, data = _post(f"{base}/api/batch", {"input_dir": str(src), "output_dir": str(src)})
|
||||
check("输入 == 输出时不会「一张都扫不到」", status == 200 and data.get("count") == 2, (status, data))
|
||||
if status == 200:
|
||||
_wait(tasks.get(data["task_id"]))
|
||||
for f in sorted(src.glob("RMBG_*")):
|
||||
f.unlink()
|
||||
|
||||
status, data = _post(f"{base}/api/batch",
|
||||
{"input_dir": str(src), "output_dir": str(out), "options": {"background": "transparent"}})
|
||||
check("正常提交 → 200", status == 200 and data.get("count") == 2, (status, data))
|
||||
task_id = data.get("task_id")
|
||||
check("回传输出目录绝对路径", data.get("output_dir") == str(out), data.get("output_dir"))
|
||||
|
||||
if task_id:
|
||||
_wait(tasks.get(task_id))
|
||||
status, detail = _get(f"{base}/api/tasks/{task_id}?tail=1")
|
||||
task = detail.get("task", {})
|
||||
check("任务详情带 mode/input_dir/output_dir",
|
||||
task.get("mode") == "batch" and task.get("input_dir") == str(src) and task.get("output_dir") == str(out),
|
||||
task.get("mode"))
|
||||
check("tail=1 只回传 1 张(载荷可控)", len(task.get("images", [])) == 1, len(task.get("images", [])))
|
||||
check("images_total = 2", task.get("images_total") == 2, task.get("images_total"))
|
||||
check("images_from 指向窗口起点", task.get("images_from") == 1, task.get("images_from"))
|
||||
check("批量任务未被塞进单图结果网格的判据(mode=batch)", task.get("mode") == "batch")
|
||||
opts = task.get("options") or {}
|
||||
check("未指定的参数落到新默认值(不输出遮罩 / 最终尺寸不限)",
|
||||
opts.get("output_mask") is False and opts.get("final_longest_side") == 0,
|
||||
{k: opts.get(k) for k in ("output_mask", "final_longest_side")})
|
||||
|
||||
status, listing = _get(f"{base}/api/tasks")
|
||||
modes = [t.get("mode") for t in listing.get("tasks", [])]
|
||||
check("任务列表里能区分两种模式", set(modes) <= {"batch", "upload"} and "batch" in modes, modes)
|
||||
|
||||
status, bad = _post(f"{base}/api/batch", {"input_dir": str(src), "recursive": True})
|
||||
check("留空输出目录 → 落到项目 outputs", status == 200 and bad.get("output_dir") == str(out_root),
|
||||
bad.get("output_dir"))
|
||||
if status == 200:
|
||||
_wait(tasks.get(bad["task_id"]))
|
||||
finally:
|
||||
server.shutdown()
|
||||
server.server_close()
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
def main(argv: Optional[List[str]] = None) -> int:
|
||||
parser = argparse.ArgumentParser(description="批量抠图自检")
|
||||
parser.add_argument("--keep", action="store_true", help="保留临时目录")
|
||||
parser.add_argument("--logic-only", action="store_true", help="只跑不依赖 torch 的纯逻辑部分")
|
||||
args = parser.parse_args(argv)
|
||||
|
||||
tmp = Path(tempfile.mkdtemp(prefix="birefnet_batch_"))
|
||||
print("=" * 68)
|
||||
print("批量抠图自检")
|
||||
print(f"临时目录 {tmp}")
|
||||
print("=" * 68)
|
||||
try:
|
||||
part_logic(tmp)
|
||||
if args.logic_only:
|
||||
print("\n[跳过] 任务流水线与 HTTP 契约(--logic-only)")
|
||||
else:
|
||||
try:
|
||||
part_pipeline(tmp)
|
||||
part_http(tmp)
|
||||
except ImportError as exc:
|
||||
print(f"\n[跳过] 任务流水线与 HTTP 契约(缺少依赖:{exc})")
|
||||
finally:
|
||||
if args.keep:
|
||||
print(f"\n临时目录已保留:{tmp}")
|
||||
else:
|
||||
shutil.rmtree(tmp, ignore_errors=True)
|
||||
|
||||
print("=" * 68)
|
||||
if _FAILS:
|
||||
print(f"自检未通过:{len(_FAILS)} 项失败")
|
||||
for name in _FAILS:
|
||||
print(f" · {name}")
|
||||
return 1
|
||||
print("自检全部通过 ✔")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
Vendored
+25
@@ -0,0 +1,25 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2024 lldacing
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
|
||||
---
|
||||
|
||||
The code and models of BiRefNet are released under the MIT License.
|
||||
+203
@@ -0,0 +1,203 @@
|
||||
import os
|
||||
import math
|
||||
|
||||
import folder_paths
|
||||
|
||||
|
||||
class Config:
|
||||
def __init__(self, bb_index: int = 6) -> None:
|
||||
# PATH settings
|
||||
# Make up your file system as: SYS_HOME_DIR/codes/dis/BiRefNet, SYS_HOME_DIR/datasets/dis/xx, SYS_HOME_DIR/weights/xx
|
||||
# self.sys_home_dir = [os.path.expanduser('~'), '/mnt/data'][0] # Default, custom
|
||||
# self.data_root_dir = os.path.join(self.sys_home_dir, 'datasets/dis')
|
||||
|
||||
# TASK settings
|
||||
self.task = ['DIS5K', 'COD', 'HRSOD', 'General', 'General-2K', 'Matting'][0]
|
||||
self.testsets = {
|
||||
# Benchmarks
|
||||
'DIS5K': ','.join(['DIS-VD', 'DIS-TE1', 'DIS-TE2', 'DIS-TE3', 'DIS-TE4']),
|
||||
'COD': ','.join(['CHAMELEON', 'NC4K', 'TE-CAMO', 'TE-COD10K']),
|
||||
'HRSOD': ','.join(['DAVIS-S', 'TE-HRSOD', 'TE-UHRSD', 'DUT-OMRON', 'TE-DUTS']),
|
||||
# Practical use
|
||||
'General': ','.join(['DIS-VD', 'TE-P3M-500-NP']),
|
||||
'General-2K': ','.join(['DIS-VD', 'TE-P3M-500-NP']),
|
||||
'Matting': ','.join(['TE-P3M-500-NP', 'TE-AM-2k']),
|
||||
}[self.task]
|
||||
# datasets_all = '+'.join([ds for ds in (os.listdir(os.path.join(self.data_root_dir, self.task)) if os.path.isdir(os.path.join(self.data_root_dir, self.task)) else []) if ds not in self.testsets.split(',')])
|
||||
self.training_set = {
|
||||
'DIS5K': ['DIS-TR', 'DIS-TR+DIS-TE1+DIS-TE2+DIS-TE3+DIS-TE4'][0],
|
||||
'COD': 'TR-COD10K+TR-CAMO',
|
||||
'HRSOD': ['TR-DUTS', 'TR-HRSOD', 'TR-UHRSD', 'TR-DUTS+TR-HRSOD', 'TR-DUTS+TR-UHRSD', 'TR-HRSOD+TR-UHRSD', 'TR-DUTS+TR-HRSOD+TR-UHRSD'][5],
|
||||
'General': 'DIS-TE1+DIS-TE2+DIS-TE3+DIS-TE4+DIS-TR+TR-HRSOD+TE-HRSOD+TR-HRS10K+TE-HRS10K+TR-UHRSD+TE-UHRSD+TR-P3M-10k+TE-P3M-500-P+TR-humans+DIS-VD-ori', # datasets_all
|
||||
'General-2K': 'DIS-TE1+DIS-TE2+DIS-TE3+DIS-TE4+DIS-TR+TR-HRSOD+TE-HRSOD+TR-HRS10K+TE-HRS10K+TR-UHRSD+TE-UHRSD+TR-P3M-10k+TE-P3M-500-P+TR-humans+DIS-VD-ori', # datasets_all
|
||||
'Matting': 'TR-P3M-10k+TE-P3M-500-NP+TR-humans+TR-Distrinctions-646', # datasets_all
|
||||
}[self.task]
|
||||
|
||||
# Data settings
|
||||
self.size = (1024, 1024) if self.task not in ['General-2K'] else (2560, 1440) # wid, hei. Can be overwritten by dynamic_size in training.
|
||||
self.dynamic_size = [None, ((512-256, 2048+256), (512-256, 2048+256))][0] # wid, hei. It might cause errors in using compile.
|
||||
self.background_color_synthesis = False # whether to use pure bg color to replace the original backgrounds.
|
||||
|
||||
# Faster-Training settings
|
||||
self.load_all = False and self.dynamic_size is None # Turn it on/off by your case. It may consume a lot of CPU memory. And for multi-GPU (N), it would cost N times the CPU memory to load the data.
|
||||
self.compile = True # 1. Trigger CPU memory leak in some extend, which is an inherent problem of PyTorch.
|
||||
# Machines with > 70GB CPU memory can run the whole training on DIS5K with default setting.
|
||||
# 2. Higher PyTorch version may fix it: https://github.com/pytorch/pytorch/issues/119607.
|
||||
# 3. But compile in 2.0.1 < Pytorch < 2.5.0 seems to bring no acceleration for training.
|
||||
self.precisionHigh = True
|
||||
|
||||
# MODEL settings
|
||||
self.ms_supervision = True
|
||||
self.out_ref = self.ms_supervision and True
|
||||
self.dec_ipt = True
|
||||
self.dec_ipt_split = True
|
||||
self.cxt_num = [0, 3][1] # multi-scale skip connections from encoder
|
||||
self.mul_scl_ipt = ['', 'add', 'cat'][2]
|
||||
self.dec_att = ['', 'ASPP', 'ASPPDeformable'][2]
|
||||
self.squeeze_block = ['', 'BasicDecBlk_x1', 'ResBlk_x4', 'ASPP_x3', 'ASPPDeformable_x3'][1]
|
||||
self.dec_blk = ['BasicDecBlk', 'ResBlk'][0]
|
||||
|
||||
# TRAINING settings
|
||||
self.batch_size = 4
|
||||
self.finetune_last_epochs = [
|
||||
0,
|
||||
{
|
||||
'DIS5K': -40,
|
||||
'COD': -20,
|
||||
'HRSOD': -20,
|
||||
'General': -20,
|
||||
'General-2K': -20,
|
||||
'Matting': -20,
|
||||
}[self.task]
|
||||
][1] # choose 0 to skip
|
||||
self.lr = (1e-4 if 'DIS5K' in self.task else 1e-5) * math.sqrt(self.batch_size / 4) # DIS needs high lr to converge faster. Adapt the lr linearly
|
||||
self.num_workers = max(4, self.batch_size) # will be decrease to min(it, batch_size) at the initialization of the data_loader
|
||||
|
||||
# Backbone settings
|
||||
self.bb = [
|
||||
'vgg16', 'vgg16bn', 'resnet50', # 0, 1, 2
|
||||
'swin_v1_t', 'swin_v1_s', # 3, 4
|
||||
'swin_v1_b', 'swin_v1_l', # 5-bs9, 6-bs4
|
||||
'pvt_v2_b0', 'pvt_v2_b1', # 7, 8
|
||||
'pvt_v2_b2', 'pvt_v2_b5', # 9-bs10, 10-bs5
|
||||
][bb_index]
|
||||
self.lateral_channels_in_collection = {
|
||||
'vgg16': [512, 256, 128, 64], 'vgg16bn': [512, 256, 128, 64], 'resnet50': [1024, 512, 256, 64],
|
||||
'pvt_v2_b2': [512, 320, 128, 64], 'pvt_v2_b5': [512, 320, 128, 64],
|
||||
'swin_v1_b': [1024, 512, 256, 128], 'swin_v1_l': [1536, 768, 384, 192],
|
||||
'swin_v1_t': [768, 384, 192, 96], 'swin_v1_s': [768, 384, 192, 96],
|
||||
'pvt_v2_b0': [256, 160, 64, 32], 'pvt_v2_b1': [512, 320, 128, 64],
|
||||
}[self.bb]
|
||||
if self.mul_scl_ipt == 'cat':
|
||||
self.lateral_channels_in_collection = [channel * 2 for channel in self.lateral_channels_in_collection]
|
||||
self.cxt = self.lateral_channels_in_collection[1:][::-1][-self.cxt_num:] if self.cxt_num else []
|
||||
|
||||
# MODEL settings - inactive
|
||||
self.lat_blk = ['BasicLatBlk'][0]
|
||||
self.dec_channels_inter = ['fixed', 'adap'][0]
|
||||
self.refine = ['', 'itself', 'RefUNet', 'Refiner', 'RefinerPVTInChannels4'][0]
|
||||
self.progressive_ref = self.refine and True
|
||||
self.ender = self.progressive_ref and False
|
||||
self.scale = self.progressive_ref and 2
|
||||
self.auxiliary_classification = False # Only for DIS5K, where class labels are saved in `dataset.py`.
|
||||
self.refine_iteration = 1
|
||||
self.freeze_bb = False
|
||||
self.model = [
|
||||
'BiRefNet',
|
||||
'BiRefNetC2F',
|
||||
][0]
|
||||
|
||||
# TRAINING settings - inactive
|
||||
self.preproc_methods = ['flip', 'enhance', 'rotate', 'pepper', 'crop'][:4 if not self.background_color_synthesis else 1]
|
||||
self.optimizer = ['Adam', 'AdamW'][1]
|
||||
self.lr_decay_epochs = [1e5] # Set to negative N to decay the lr in the last N-th epoch.
|
||||
self.lr_decay_rate = 0.5
|
||||
# Loss
|
||||
if self.task in ['Matting']:
|
||||
self.lambdas_pix_last = {
|
||||
'bce': 30 * 1,
|
||||
'iou': 0.5 * 0,
|
||||
'iou_patch': 0.5 * 0,
|
||||
'mae': 100 * 1,
|
||||
'mse': 30 * 0,
|
||||
'triplet': 3 * 0,
|
||||
'reg': 100 * 0,
|
||||
'ssim': 10 * 1,
|
||||
'cnt': 5 * 0,
|
||||
'structure': 5 * 0,
|
||||
}
|
||||
elif self.task in ['General', 'General-2K']:
|
||||
self.lambdas_pix_last = {
|
||||
'bce': 30 * 1,
|
||||
'iou': 0.5 * 1,
|
||||
'iou_patch': 0.5 * 0,
|
||||
'mae': 100 * 1,
|
||||
'mse': 30 * 0,
|
||||
'triplet': 3 * 0,
|
||||
'reg': 100 * 0,
|
||||
'ssim': 10 * 1,
|
||||
'cnt': 5 * 0,
|
||||
'structure': 5 * 0,
|
||||
}
|
||||
else:
|
||||
self.lambdas_pix_last = {
|
||||
# not 0 means opening this loss
|
||||
# original rate -- 1 : 30 : 1.5 : 0.2, bce x 30
|
||||
'bce': 30 * 1, # high performance
|
||||
'iou': 0.5 * 1, # 0 / 255
|
||||
'iou_patch': 0.5 * 0, # 0 / 255, win_size = (64, 64)
|
||||
'mae': 30 * 0,
|
||||
'mse': 30 * 0, # can smooth the saliency map
|
||||
'triplet': 3 * 0,
|
||||
'reg': 100 * 0,
|
||||
'ssim': 10 * 1, # help contours,
|
||||
'cnt': 5 * 0, # help contours
|
||||
'structure': 5 * 0, # structure loss from codes of MVANet. A little improvement on DIS-TE[1,2,3], a bit more decrease on DIS-TE4.
|
||||
}
|
||||
self.lambdas_cls = {
|
||||
'ce': 5.0
|
||||
}
|
||||
|
||||
# PATH settings - inactive
|
||||
# https://drive.google.com/drive/folders/1cmce_emsS8A5ha5XT2c_CZiJzlLM81ms
|
||||
# self.weights_root_dir = os.path.join(self.sys_home_dir, 'weights/cv')
|
||||
# self.weights = {
|
||||
# 'pvt_v2_b2': os.path.join(self.weights_root_dir, 'pvt_v2_b2.pth'),
|
||||
# 'pvt_v2_b5': os.path.join(self.weights_root_dir, ['pvt_v2_b5.pth', 'pvt_v2_b5_22k.pth'][0]),
|
||||
# 'swin_v1_b': os.path.join(self.weights_root_dir, ['swin_base_patch4_window12_384_22kto1k.pth', 'swin_base_patch4_window12_384_22k.pth'][0]),
|
||||
# 'swin_v1_l': os.path.join(self.weights_root_dir, ['swin_large_patch4_window12_384_22kto1k.pth', 'swin_large_patch4_window12_384_22k.pth'][0]),
|
||||
# 'swin_v1_t': os.path.join(self.weights_root_dir, ['swin_tiny_patch4_window7_224_22kto1k_finetune.pth'][0]),
|
||||
# 'swin_v1_s': os.path.join(self.weights_root_dir, ['swin_small_patch4_window7_224_22kto1k_finetune.pth'][0]),
|
||||
# 'pvt_v2_b0': os.path.join(self.weights_root_dir, ['pvt_v2_b0.pth'][0]),
|
||||
# 'pvt_v2_b1': os.path.join(self.weights_root_dir, ['pvt_v2_b1.pth'][0]),
|
||||
# }
|
||||
weight_paths_name = "birefnet"
|
||||
self.weights = {
|
||||
'pvt_v2_b2': folder_paths.get_full_path(weight_paths_name, 'pvt_v2_b2.pth'),
|
||||
'pvt_v2_b5': folder_paths.get_full_path(weight_paths_name, ['pvt_v2_b5.pth', 'pvt_v2_b5_22k.pth'][0]),
|
||||
'swin_v1_b': folder_paths.get_full_path(weight_paths_name, ['swin_base_patch4_window12_384_22kto1k.pth', 'swin_base_patch4_window12_384_22k.pth'][0]),
|
||||
'swin_v1_l': folder_paths.get_full_path(weight_paths_name, ['swin_large_patch4_window12_384_22kto1k.pth', 'swin_large_patch4_window12_384_22k.pth'][0]),
|
||||
'swin_v1_t': folder_paths.get_full_path(weight_paths_name, ['swin_tiny_patch4_window7_224_22kto1k_finetune.pth'][0]),
|
||||
'swin_v1_s': folder_paths.get_full_path(weight_paths_name, ['swin_small_patch4_window7_224_22kto1k_finetune.pth'][0]),
|
||||
'pvt_v2_b0': folder_paths.get_full_path(weight_paths_name, ['pvt_v2_b0.pth'][0]),
|
||||
'pvt_v2_b1': folder_paths.get_full_path(weight_paths_name, ['pvt_v2_b1.pth'][0]),
|
||||
}
|
||||
|
||||
# Callbacks - inactive
|
||||
self.verbose_eval = True
|
||||
self.only_S_MAE = False
|
||||
self.SDPA_enabled = False # Bugs. Slower and errors occur in multi-GPUs
|
||||
|
||||
# others
|
||||
self.device = [0, 'cpu'][0] # .to(0) == .to('cuda:0')
|
||||
|
||||
self.batch_size_valid = 1
|
||||
self.rand_seed = 7
|
||||
# run_sh_file = [f for f in os.listdir('.') if 'train.sh' == f] + [os.path.join('..', f) for f in os.listdir('..') if 'train.sh' == f]
|
||||
# if run_sh_file:
|
||||
# with open(run_sh_file[0], 'r') as f:
|
||||
# lines = f.readlines()
|
||||
# self.save_last = int([l.strip() for l in lines if '"{}")'.format(self.task) in l and 'val_last=' in l][0].split('val_last=')[-1].split()[0])
|
||||
# self.save_step = int([l.strip() for l in lines if '"{}")'.format(self.task) in l and 'step=' in l][0].split('step=')[-1].split()[0])
|
||||
|
||||
|
||||
+173
@@ -0,0 +1,173 @@
|
||||
import os
|
||||
import random
|
||||
import numpy as np
|
||||
import cv2
|
||||
from tqdm import tqdm
|
||||
from PIL import Image
|
||||
from torch.utils import data
|
||||
from torchvision import transforms
|
||||
|
||||
from .image_proc import preproc
|
||||
from .config import Config
|
||||
from .utils import path_to_image
|
||||
|
||||
|
||||
Image.MAX_IMAGE_PIXELS = None # remove DecompressionBombWarning
|
||||
config = Config()
|
||||
_class_labels_TR_sorted = (
|
||||
'Airplane, Ant, Antenna, Archery, Axe, BabyCarriage, Bag, BalanceBeam, Balcony, Balloon, Basket, BasketballHoop, Beatle, Bed, Bee, Bench, Bicycle, '
|
||||
'BicycleFrame, BicycleStand, Boat, Bonsai, BoomLift, Bridge, BunkBed, Butterfly, Button, Cable, CableLift, Cage, Camcorder, Cannon, Canoe, Car, '
|
||||
'CarParkDropArm, Carriage, Cart, Caterpillar, CeilingLamp, Centipede, Chair, Clip, Clock, Clothes, CoatHanger, Comb, ConcretePumpTruck, Crack, Crane, '
|
||||
'Cup, DentalChair, Desk, DeskChair, Diagram, DishRack, DoorHandle, Dragonfish, Dragonfly, Drum, Earphone, Easel, ElectricIron, Excavator, Eyeglasses, '
|
||||
'Fan, Fence, Fencing, FerrisWheel, FireExtinguisher, Fishing, Flag, FloorLamp, Forklift, GasStation, Gate, Gear, Goal, Golf, GymEquipment, Hammock, '
|
||||
'Handcart, Handcraft, Handrail, HangGlider, Harp, Harvester, Headset, Helicopter, Helmet, Hook, HorizontalBar, Hydrovalve, IroningTable, Jewelry, Key, '
|
||||
'KidsPlayground, Kitchenware, Kite, Knife, Ladder, LaundryRack, Lightning, Lobster, Locust, Machine, MachineGun, MagazineRack, Mantis, Medal, MemorialArchway, '
|
||||
'Microphone, Missile, MobileHolder, Monitor, Mosquito, Motorcycle, MovingTrolley, Mower, MusicPlayer, MusicStand, ObservationTower, Octopus, OilWell, '
|
||||
'OlympicLogo, OperatingTable, OutdoorFitnessEquipment, Parachute, Pavilion, Piano, Pipe, PlowHarrow, PoleVault, Punchbag, Rack, Racket, Rifle, Ring, Robot, '
|
||||
'RockClimbing, Rope, Sailboat, Satellite, Scaffold, Scale, Scissor, Scooter, Sculpture, Seadragon, Seahorse, Seal, SewingMachine, Ship, Shoe, ShoppingCart, '
|
||||
'ShoppingTrolley, Shower, Shrimp, Signboard, Skateboarding, Skeleton, Skiing, Spade, SpeedBoat, Spider, Spoon, Stair, Stand, Stationary, SteeringWheel, '
|
||||
'Stethoscope, Stool, Stove, StreetLamp, SweetStand, Swing, Sword, TV, Table, TableChair, TableLamp, TableTennis, Tank, Tapeline, Teapot, Telescope, Tent, '
|
||||
'TobaccoPipe, Toy, Tractor, TrafficLight, TrafficSign, Trampoline, TransmissionTower, Tree, Tricycle, TrimmerCover, Tripod, Trombone, Truck, Trumpet, Tuba, '
|
||||
'UAV, Umbrella, UnevenBars, UtilityPole, VacuumCleaner, Violin, Wakesurfing, Watch, WaterTower, WateringPot, Well, WellLid, Wheel, Wheelchair, WindTurbine, Windmill, WineGlass, WireWhisk, Yacht'
|
||||
)
|
||||
class_labels_TR_sorted = _class_labels_TR_sorted.split(', ')
|
||||
|
||||
|
||||
class MyData(data.Dataset):
|
||||
def __init__(self, datasets, data_size, is_train=True):
|
||||
# data_size is None when using dynamic_size or data_size is manually set to None (for inference in the original size).
|
||||
self.is_train = is_train
|
||||
self.data_size = data_size
|
||||
self.load_all = config.load_all
|
||||
self.device = config.device
|
||||
valid_extensions = ['.png', '.jpg', '.PNG', '.JPG', '.JPEG']
|
||||
|
||||
if self.is_train and config.auxiliary_classification:
|
||||
self.cls_name2id = {_name: _id for _id, _name in enumerate(class_labels_TR_sorted)}
|
||||
self.transform_image = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
||||
])
|
||||
self.transform_label = transforms.Compose([
|
||||
transforms.ToTensor(),
|
||||
])
|
||||
dataset_root = os.path.join(config.data_root_dir, config.task)
|
||||
# datasets can be a list of different datasets for training on combined sets.
|
||||
self.image_paths = []
|
||||
for dataset in datasets.split('+'):
|
||||
image_root = os.path.join(dataset_root, dataset, 'im')
|
||||
self.image_paths += [os.path.join(image_root, p) for p in os.listdir(image_root) if any(p.endswith(ext) for ext in valid_extensions)]
|
||||
self.label_paths = []
|
||||
for p in self.image_paths:
|
||||
for ext in valid_extensions:
|
||||
## 'im' and 'gt' may need modifying
|
||||
p_gt = p.replace('/im/', '/gt/')[:-(len(p.split('.')[-1])+1)] + ext
|
||||
file_exists = False
|
||||
if os.path.exists(p_gt):
|
||||
self.label_paths.append(p_gt)
|
||||
file_exists = True
|
||||
break
|
||||
if not file_exists:
|
||||
print('Not exists:', p_gt)
|
||||
|
||||
if len(self.label_paths) != len(self.image_paths):
|
||||
set_image_paths = set([os.path.splitext(p.split(os.sep)[-1])[0] for p in self.image_paths])
|
||||
set_label_paths = set([os.path.splitext(p.split(os.sep)[-1])[0] for p in self.label_paths])
|
||||
print('Path diff:', set_image_paths - set_label_paths)
|
||||
raise ValueError(f"There are different numbers of images ({len(self.label_paths)}) and labels ({len(self.image_paths)})")
|
||||
|
||||
if self.load_all:
|
||||
self.images_loaded, self.labels_loaded = [], []
|
||||
self.class_labels_loaded = []
|
||||
# for image_path, label_path in zip(self.image_paths, self.label_paths):
|
||||
for image_path, label_path in tqdm(zip(self.image_paths, self.label_paths), total=len(self.image_paths)):
|
||||
_image = path_to_image(image_path, size=self.data_size, color_type='rgb')
|
||||
_label = path_to_image(label_path, size=self.data_size, color_type='gray')
|
||||
self.images_loaded.append(_image)
|
||||
self.labels_loaded.append(_label)
|
||||
self.class_labels_loaded.append(
|
||||
self.cls_name2id[label_path.split('/')[-1].split('#')[3]] if self.is_train and config.auxiliary_classification else -1
|
||||
)
|
||||
|
||||
def __getitem__(self, index):
|
||||
if self.load_all:
|
||||
image = self.images_loaded[index]
|
||||
label = self.labels_loaded[index]
|
||||
class_label = self.class_labels_loaded[index] if self.is_train and config.auxiliary_classification else -1
|
||||
else:
|
||||
image = path_to_image(self.image_paths[index], size=self.data_size, color_type='rgb')
|
||||
label = path_to_image(self.label_paths[index], size=self.data_size, color_type='gray')
|
||||
class_label = self.cls_name2id[self.label_paths[index].split('/')[-1].split('#')[3]] if self.is_train and config.auxiliary_classification else -1
|
||||
|
||||
# loading image and label
|
||||
if self.is_train:
|
||||
if config.background_color_synthesis:
|
||||
image.putalpha(label)
|
||||
array_image = np.array(image)
|
||||
array_foreground = array_image[:, :, :3].astype(np.float32)
|
||||
array_mask = (array_image[:, :, 3:] / 255).astype(np.float32)
|
||||
array_background = np.zeros_like(array_foreground)
|
||||
choice = random.random()
|
||||
if choice < 0.4:
|
||||
# Black/Gray/White backgrounds
|
||||
array_background[:, :, :] = random.randint(0, 255)
|
||||
elif choice < 0.8:
|
||||
# Background color that similar to the foreground object. Hard negative samples.
|
||||
foreground_pixel_number = np.sum(array_mask > 0)
|
||||
color_foreground_mean = np.mean(array_foreground * array_mask, axis=(0, 1)) * (np.prod(array_foreground.shape[:2]) / foreground_pixel_number)
|
||||
color_up_or_down = random.choice((-1, 1))
|
||||
# Up or down for 20% range from 255 or 0, respectively.
|
||||
color_foreground_mean += (255 - color_foreground_mean if color_up_or_down == 1 else color_foreground_mean) * (random.random() * 0.2) * color_up_or_down
|
||||
array_background[:, :, :] = color_foreground_mean
|
||||
else:
|
||||
# Any color
|
||||
for idx_channel in range(3):
|
||||
array_background[:, :, idx_channel] = random.randint(0, 255)
|
||||
array_foreground_background = array_foreground * array_mask + array_background * (1 - array_mask)
|
||||
image = Image.fromarray(array_foreground_background.astype(np.uint8))
|
||||
image, label = preproc(image, label, preproc_methods=config.preproc_methods)
|
||||
# else:
|
||||
# if _label.shape[0] > 2048 or _label.shape[1] > 2048:
|
||||
# _image = cv2.resize(_image, (2048, 2048), interpolation=cv2.INTER_LINEAR)
|
||||
# _label = cv2.resize(_label, (2048, 2048), interpolation=cv2.INTER_LINEAR)
|
||||
|
||||
# At present, we use fixed sizes in inference, instead of consistent dynamic size with training.
|
||||
if self.is_train:
|
||||
if config.dynamic_size is None:
|
||||
image, label = self.transform_image(image), self.transform_label(label)
|
||||
else:
|
||||
size_div_32 = (int(image.size[0] // 32 * 32), int(image.size[1] // 32 * 32))
|
||||
if image.size != size_div_32:
|
||||
image = image.resize(size_div_32)
|
||||
label = label.resize(size_div_32)
|
||||
image, label = self.transform_image(image), self.transform_label(label)
|
||||
|
||||
if self.is_train:
|
||||
return image, label, class_label
|
||||
else:
|
||||
return image, label, self.label_paths[index]
|
||||
|
||||
def __len__(self):
|
||||
return len(self.image_paths)
|
||||
|
||||
|
||||
def custom_collate_fn(batch):
|
||||
if config.dynamic_size:
|
||||
dynamic_size = tuple(sorted(config.dynamic_size))
|
||||
dynamic_size_batch = (random.randint(dynamic_size[0][0], dynamic_size[0][1]) // 32 * 32, random.randint(dynamic_size[1][0], dynamic_size[1][1]) // 32 * 32) # select a value randomly in the range of [dynamic_size[0/1][0], dynamic_size[0/1][1]].
|
||||
data_size = dynamic_size_batch
|
||||
else:
|
||||
data_size = config.size
|
||||
new_batch = []
|
||||
transform_image = transforms.Compose([
|
||||
transforms.Resize(data_size[::-1]),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
||||
])
|
||||
transform_label = transforms.Compose([
|
||||
transforms.Resize(data_size[::-1]),
|
||||
transforms.ToTensor(),
|
||||
])
|
||||
for image, label, class_label in batch:
|
||||
new_batch.append((transform_image(image), transform_label(label), class_label))
|
||||
return data._utils.collate.default_collate(new_batch)
|
||||
+116
@@ -0,0 +1,116 @@
|
||||
import random
|
||||
from PIL import Image, ImageEnhance
|
||||
import numpy as np
|
||||
import cv2
|
||||
|
||||
|
||||
def refine_foreground(image, mask, r=90):
|
||||
if mask.size != image.size:
|
||||
mask = mask.resize(image.size)
|
||||
image = np.array(image) / 255.0
|
||||
mask = np.array(mask) / 255.0
|
||||
estimated_foreground = FB_blur_fusion_foreground_estimator_2(image, mask, r=r)
|
||||
image_masked = Image.fromarray((estimated_foreground * 255.0).astype(np.uint8))
|
||||
return image_masked
|
||||
|
||||
|
||||
def FB_blur_fusion_foreground_estimator_2(image, alpha, r=90):
|
||||
# Thanks to the source: https://github.com/Photoroom/fast-foreground-estimation
|
||||
alpha = alpha[:, :, None]
|
||||
F, blur_B = FB_blur_fusion_foreground_estimator(
|
||||
image, image, image, alpha, r)
|
||||
return FB_blur_fusion_foreground_estimator(image, F, blur_B, alpha, r=6)[0]
|
||||
|
||||
|
||||
def FB_blur_fusion_foreground_estimator(image, F, B, alpha, r=90):
|
||||
if isinstance(image, Image.Image):
|
||||
image = np.array(image) / 255.0
|
||||
blurred_alpha = cv2.blur(alpha, (r, r))[:, :, None]
|
||||
|
||||
blurred_FA = cv2.blur(F * alpha, (r, r))
|
||||
blurred_F = blurred_FA / (blurred_alpha + 1e-5)
|
||||
|
||||
blurred_B1A = cv2.blur(B * (1 - alpha), (r, r))
|
||||
blurred_B = blurred_B1A / ((1 - blurred_alpha) + 1e-5)
|
||||
F = blurred_F + alpha * \
|
||||
(image - alpha * blurred_F - (1 - alpha) * blurred_B)
|
||||
F = np.clip(F, 0, 1)
|
||||
return F, blurred_B
|
||||
|
||||
|
||||
def preproc(image, label, preproc_methods=['flip']):
|
||||
if 'flip' in preproc_methods:
|
||||
image, label = cv_random_flip(image, label)
|
||||
if 'crop' in preproc_methods:
|
||||
image, label = random_crop(image, label)
|
||||
if 'rotate' in preproc_methods:
|
||||
image, label = random_rotate(image, label)
|
||||
if 'enhance' in preproc_methods:
|
||||
image = color_enhance(image)
|
||||
if 'pepper' in preproc_methods:
|
||||
image = random_pepper(image)
|
||||
return image, label
|
||||
|
||||
|
||||
def cv_random_flip(img, label):
|
||||
if random.random() > 0.5:
|
||||
img = img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
label = label.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
return img, label
|
||||
|
||||
|
||||
def random_crop(image, label):
|
||||
border = 30
|
||||
image_width = image.size[0]
|
||||
image_height = image.size[1]
|
||||
border = int(min(image_width, image_height) * 0.1)
|
||||
crop_win_width = np.random.randint(image_width - border, image_width)
|
||||
crop_win_height = np.random.randint(image_height - border, image_height)
|
||||
random_region = (
|
||||
(image_width - crop_win_width) >> 1, (image_height - crop_win_height) >> 1, (image_width + crop_win_width) >> 1,
|
||||
(image_height + crop_win_height) >> 1)
|
||||
return image.crop(random_region), label.crop(random_region)
|
||||
|
||||
|
||||
def random_rotate(image, label, angle=15):
|
||||
mode = Image.BICUBIC
|
||||
if random.random() > 0.8:
|
||||
random_angle = np.random.randint(-angle, angle)
|
||||
image = image.rotate(random_angle, mode)
|
||||
label = label.rotate(random_angle, mode)
|
||||
return image, label
|
||||
|
||||
|
||||
def color_enhance(image):
|
||||
bright_intensity = random.randint(5, 15) / 10.0
|
||||
image = ImageEnhance.Brightness(image).enhance(bright_intensity)
|
||||
contrast_intensity = random.randint(5, 15) / 10.0
|
||||
image = ImageEnhance.Contrast(image).enhance(contrast_intensity)
|
||||
color_intensity = random.randint(0, 20) / 10.0
|
||||
image = ImageEnhance.Color(image).enhance(color_intensity)
|
||||
sharp_intensity = random.randint(0, 30) / 10.0
|
||||
image = ImageEnhance.Sharpness(image).enhance(sharp_intensity)
|
||||
return image
|
||||
|
||||
|
||||
def random_gaussian(image, mean=0.1, sigma=0.35):
|
||||
def gaussianNoisy(im, mean=mean, sigma=sigma):
|
||||
for _i in range(len(im)):
|
||||
im[_i] += random.gauss(mean, sigma)
|
||||
return im
|
||||
|
||||
img = np.asarray(image)
|
||||
width, height = img.shape
|
||||
img = gaussianNoisy(img[:].flatten(), mean, sigma)
|
||||
img = img.reshape([width, height])
|
||||
return Image.fromarray(np.uint8(img))
|
||||
|
||||
|
||||
def random_pepper(img, N=0.0015):
|
||||
img = np.array(img)
|
||||
noiseNum = int(N * img.shape[0] * img.shape[1])
|
||||
for i in range(noiseNum):
|
||||
randX = random.randint(0, img.shape[0] - 1)
|
||||
randY = random.randint(0, img.shape[1] - 1)
|
||||
img[randX, randY] = random.randint(0, 1) * 255
|
||||
return Image.fromarray(img)
|
||||
@@ -0,0 +1,44 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from collections import OrderedDict
|
||||
from torchvision.models import vgg16, vgg16_bn, VGG16_Weights, VGG16_BN_Weights, resnet50, ResNet50_Weights
|
||||
from ..backbones.pvt_v2 import pvt_v2_b0, pvt_v2_b1, pvt_v2_b2, pvt_v2_b5
|
||||
from ..backbones.swin_v1 import swin_v1_t, swin_v1_s, swin_v1_b, swin_v1_l
|
||||
from ...config import Config
|
||||
|
||||
|
||||
config = Config()
|
||||
|
||||
def build_backbone(bb_name, pretrained=True, params_settings=''):
|
||||
if bb_name == 'vgg16':
|
||||
bb_net = list(vgg16(pretrained=VGG16_Weights.DEFAULT if pretrained else None).children())[0]
|
||||
bb = nn.Sequential(OrderedDict({'conv1': bb_net[:4], 'conv2': bb_net[4:9], 'conv3': bb_net[9:16], 'conv4': bb_net[16:23]}))
|
||||
elif bb_name == 'vgg16bn':
|
||||
bb_net = list(vgg16_bn(pretrained=VGG16_BN_Weights.DEFAULT if pretrained else None).children())[0]
|
||||
bb = nn.Sequential(OrderedDict({'conv1': bb_net[:6], 'conv2': bb_net[6:13], 'conv3': bb_net[13:23], 'conv4': bb_net[23:33]}))
|
||||
elif bb_name == 'resnet50':
|
||||
bb_net = list(resnet50(pretrained=ResNet50_Weights.DEFAULT if pretrained else None).children())
|
||||
bb = nn.Sequential(OrderedDict({'conv1': nn.Sequential(*bb_net[0:3]), 'conv2': bb_net[4], 'conv3': bb_net[5], 'conv4': bb_net[6]}))
|
||||
else:
|
||||
bb = eval('{}({})'.format(bb_name, params_settings))
|
||||
if pretrained:
|
||||
bb = load_weights(bb, bb_name)
|
||||
return bb
|
||||
|
||||
def load_weights(model, model_name):
|
||||
save_model = torch.load(config.weights[model_name], map_location='cpu', weights_only=True)
|
||||
model_dict = model.state_dict()
|
||||
state_dict = {k: v if v.size() == model_dict[k].size() else model_dict[k] for k, v in save_model.items() if k in model_dict.keys()}
|
||||
# to ignore the weights with mismatched size when I modify the backbone itself.
|
||||
if not state_dict:
|
||||
save_model_keys = list(save_model.keys())
|
||||
sub_item = save_model_keys[0] if len(save_model_keys) == 1 else None
|
||||
state_dict = {k: v if v.size() == model_dict[k].size() else model_dict[k] for k, v in save_model[sub_item].items() if k in model_dict.keys()}
|
||||
if not state_dict or not sub_item:
|
||||
print('Weights are not successfully loaded. Check the state dict of weights file.')
|
||||
return None
|
||||
else:
|
||||
print('Found correct weights in the "{}" item of loaded state_dict.'.format(sub_item))
|
||||
model_dict.update(state_dict)
|
||||
model.load_state_dict(model_dict)
|
||||
return model
|
||||
@@ -0,0 +1,433 @@
|
||||
import math
|
||||
from functools import partial
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
try:
|
||||
# version > 0.6.13
|
||||
from timm.layers import DropPath, to_2tuple, trunc_normal_
|
||||
except Exception:
|
||||
from timm.models.layers import DropPath, to_2tuple, trunc_normal_
|
||||
|
||||
from ...config import Config
|
||||
|
||||
config = Config()
|
||||
|
||||
class Mlp(nn.Module):
|
||||
def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
|
||||
super().__init__()
|
||||
out_features = out_features or in_features
|
||||
hidden_features = hidden_features or in_features
|
||||
self.fc1 = nn.Linear(in_features, hidden_features)
|
||||
self.dwconv = DWConv(hidden_features)
|
||||
self.act = act_layer()
|
||||
self.fc2 = nn.Linear(hidden_features, out_features)
|
||||
self.drop = nn.Dropout(drop)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x, H, W):
|
||||
x = self.fc1(x)
|
||||
x = self.dwconv(x, H, W)
|
||||
x = self.act(x)
|
||||
x = self.drop(x)
|
||||
x = self.fc2(x)
|
||||
x = self.drop(x)
|
||||
return x
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0., sr_ratio=1):
|
||||
super().__init__()
|
||||
assert dim % num_heads == 0, f"dim {dim} should be divided by num_heads {num_heads}."
|
||||
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim // num_heads
|
||||
self.scale = qk_scale or head_dim ** -0.5
|
||||
|
||||
self.q = nn.Linear(dim, dim, bias=qkv_bias)
|
||||
self.kv = nn.Linear(dim, dim * 2, bias=qkv_bias)
|
||||
self.attn_drop_prob = attn_drop
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
|
||||
self.sr_ratio = sr_ratio
|
||||
if sr_ratio > 1:
|
||||
self.sr = nn.Conv2d(dim, dim, kernel_size=sr_ratio, stride=sr_ratio)
|
||||
self.norm = nn.LayerNorm(dim)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x, H, W):
|
||||
B, N, C = x.shape
|
||||
q = self.q(x).reshape(B, N, self.num_heads, C // self.num_heads).permute(0, 2, 1, 3)
|
||||
|
||||
if self.sr_ratio > 1:
|
||||
x_ = x.permute(0, 2, 1).reshape(B, C, H, W)
|
||||
x_ = self.sr(x_).reshape(B, C, -1).permute(0, 2, 1)
|
||||
x_ = self.norm(x_)
|
||||
kv = self.kv(x_).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
|
||||
else:
|
||||
kv = self.kv(x).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
|
||||
k, v = kv[0], kv[1]
|
||||
|
||||
if config.SDPA_enabled:
|
||||
x = torch.nn.functional.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
attn_mask=None, dropout_p=self.attn_drop_prob, is_causal=False
|
||||
).transpose(1, 2).reshape(B, N, C)
|
||||
else:
|
||||
attn = (q @ k.transpose(-2, -1)) * self.scale
|
||||
attn = attn.softmax(dim=-1)
|
||||
attn = self.attn_drop(attn)
|
||||
|
||||
x = (attn @ v).transpose(1, 2).reshape(B, N, C)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class Block(nn.Module):
|
||||
|
||||
def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,
|
||||
drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm, sr_ratio=1):
|
||||
super().__init__()
|
||||
self.norm1 = norm_layer(dim)
|
||||
self.attn = Attention(
|
||||
dim,
|
||||
num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale,
|
||||
attn_drop=attn_drop, proj_drop=drop, sr_ratio=sr_ratio)
|
||||
# NOTE: drop path for stochastic depth, we shall see if this is better than dropout here
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
|
||||
self.norm2 = norm_layer(dim)
|
||||
mlp_hidden_dim = int(dim * mlp_ratio)
|
||||
self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x, H, W):
|
||||
x = x + self.drop_path(self.attn(self.norm1(x), H, W))
|
||||
x = x + self.drop_path(self.mlp(self.norm2(x), H, W))
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class OverlapPatchEmbed(nn.Module):
|
||||
""" Image to Patch Embedding
|
||||
"""
|
||||
|
||||
def __init__(self, img_size=224, patch_size=7, stride=4, in_channels=3, embed_dim=768):
|
||||
super().__init__()
|
||||
img_size = to_2tuple(img_size)
|
||||
patch_size = to_2tuple(patch_size)
|
||||
|
||||
self.img_size = img_size
|
||||
self.patch_size = patch_size
|
||||
self.H, self.W = img_size[0] // patch_size[0], img_size[1] // patch_size[1]
|
||||
self.num_patches = self.H * self.W
|
||||
self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=stride,
|
||||
padding=(patch_size[0] // 2, patch_size[1] // 2))
|
||||
self.norm = nn.LayerNorm(embed_dim)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.proj(x)
|
||||
_, _, H, W = x.shape
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
x = self.norm(x)
|
||||
|
||||
return x, H, W
|
||||
|
||||
|
||||
class PyramidVisionTransformerImpr(nn.Module):
|
||||
def __init__(self, img_size=224, patch_size=16, in_channels=3, num_classes=1000, embed_dims=[64, 128, 256, 512],
|
||||
num_heads=[1, 2, 4, 8], mlp_ratios=[4, 4, 4, 4], qkv_bias=False, qk_scale=None, drop_rate=0.,
|
||||
attn_drop_rate=0., drop_path_rate=0., norm_layer=nn.LayerNorm,
|
||||
depths=[3, 4, 6, 3], sr_ratios=[8, 4, 2, 1]):
|
||||
super().__init__()
|
||||
self.num_classes = num_classes
|
||||
self.depths = depths
|
||||
|
||||
# patch_embed
|
||||
self.patch_embed1 = OverlapPatchEmbed(img_size=img_size, patch_size=7, stride=4, in_channels=in_channels,
|
||||
embed_dim=embed_dims[0])
|
||||
self.patch_embed2 = OverlapPatchEmbed(img_size=img_size // 4, patch_size=3, stride=2, in_channels=embed_dims[0],
|
||||
embed_dim=embed_dims[1])
|
||||
self.patch_embed3 = OverlapPatchEmbed(img_size=img_size // 8, patch_size=3, stride=2, in_channels=embed_dims[1],
|
||||
embed_dim=embed_dims[2])
|
||||
self.patch_embed4 = OverlapPatchEmbed(img_size=img_size // 16, patch_size=3, stride=2, in_channels=embed_dims[2],
|
||||
embed_dim=embed_dims[3])
|
||||
|
||||
# transformer encoder
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule
|
||||
cur = 0
|
||||
self.block1 = nn.ModuleList([Block(
|
||||
dim=embed_dims[0], num_heads=num_heads[0], mlp_ratio=mlp_ratios[0], qkv_bias=qkv_bias, qk_scale=qk_scale,
|
||||
drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
|
||||
sr_ratio=sr_ratios[0])
|
||||
for i in range(depths[0])])
|
||||
self.norm1 = norm_layer(embed_dims[0])
|
||||
|
||||
cur += depths[0]
|
||||
self.block2 = nn.ModuleList([Block(
|
||||
dim=embed_dims[1], num_heads=num_heads[1], mlp_ratio=mlp_ratios[1], qkv_bias=qkv_bias, qk_scale=qk_scale,
|
||||
drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
|
||||
sr_ratio=sr_ratios[1])
|
||||
for i in range(depths[1])])
|
||||
self.norm2 = norm_layer(embed_dims[1])
|
||||
|
||||
cur += depths[1]
|
||||
self.block3 = nn.ModuleList([Block(
|
||||
dim=embed_dims[2], num_heads=num_heads[2], mlp_ratio=mlp_ratios[2], qkv_bias=qkv_bias, qk_scale=qk_scale,
|
||||
drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
|
||||
sr_ratio=sr_ratios[2])
|
||||
for i in range(depths[2])])
|
||||
self.norm3 = norm_layer(embed_dims[2])
|
||||
|
||||
cur += depths[2]
|
||||
self.block4 = nn.ModuleList([Block(
|
||||
dim=embed_dims[3], num_heads=num_heads[3], mlp_ratio=mlp_ratios[3], qkv_bias=qkv_bias, qk_scale=qk_scale,
|
||||
drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
|
||||
sr_ratio=sr_ratios[3])
|
||||
for i in range(depths[3])])
|
||||
self.norm4 = norm_layer(embed_dims[3])
|
||||
|
||||
# classification head
|
||||
# self.head = nn.Linear(embed_dims[3], num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def init_weights(self, pretrained=None):
|
||||
if isinstance(pretrained, str):
|
||||
logger = 1
|
||||
#load_checkpoint(self, pretrained, map_location='cpu', strict=False, logger=logger)
|
||||
|
||||
def reset_drop_path(self, drop_path_rate):
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(self.depths))]
|
||||
cur = 0
|
||||
for i in range(self.depths[0]):
|
||||
self.block1[i].drop_path.drop_prob = dpr[cur + i]
|
||||
|
||||
cur += self.depths[0]
|
||||
for i in range(self.depths[1]):
|
||||
self.block2[i].drop_path.drop_prob = dpr[cur + i]
|
||||
|
||||
cur += self.depths[1]
|
||||
for i in range(self.depths[2]):
|
||||
self.block3[i].drop_path.drop_prob = dpr[cur + i]
|
||||
|
||||
cur += self.depths[2]
|
||||
for i in range(self.depths[3]):
|
||||
self.block4[i].drop_path.drop_prob = dpr[cur + i]
|
||||
|
||||
def freeze_patch_emb(self):
|
||||
self.patch_embed1.requires_grad = False
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'pos_embed1', 'pos_embed2', 'pos_embed3', 'pos_embed4', 'cls_token'} # has pos_embed may be better
|
||||
|
||||
def get_classifier(self):
|
||||
return self.head
|
||||
|
||||
def reset_classifier(self, num_classes, global_pool=''):
|
||||
self.num_classes = num_classes
|
||||
self.head = nn.Linear(self.embed_dim, num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
def forward_features(self, x):
|
||||
B = x.shape[0]
|
||||
outs = []
|
||||
|
||||
# stage 1
|
||||
x, H, W = self.patch_embed1(x)
|
||||
for i, blk in enumerate(self.block1):
|
||||
x = blk(x, H, W)
|
||||
x = self.norm1(x)
|
||||
x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
|
||||
outs.append(x)
|
||||
|
||||
# stage 2
|
||||
x, H, W = self.patch_embed2(x)
|
||||
for i, blk in enumerate(self.block2):
|
||||
x = blk(x, H, W)
|
||||
x = self.norm2(x)
|
||||
x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
|
||||
outs.append(x)
|
||||
|
||||
# stage 3
|
||||
x, H, W = self.patch_embed3(x)
|
||||
for i, blk in enumerate(self.block3):
|
||||
x = blk(x, H, W)
|
||||
x = self.norm3(x)
|
||||
x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
|
||||
outs.append(x)
|
||||
|
||||
# stage 4
|
||||
x, H, W = self.patch_embed4(x)
|
||||
for i, blk in enumerate(self.block4):
|
||||
x = blk(x, H, W)
|
||||
x = self.norm4(x)
|
||||
x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
|
||||
outs.append(x)
|
||||
|
||||
return outs
|
||||
|
||||
# return x.mean(dim=1)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.forward_features(x)
|
||||
# x = self.head(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class DWConv(nn.Module):
|
||||
def __init__(self, dim=768):
|
||||
super(DWConv, self).__init__()
|
||||
self.dwconv = nn.Conv2d(dim, dim, 3, 1, 1, bias=True, groups=dim)
|
||||
|
||||
def forward(self, x, H, W):
|
||||
B, N, C = x.shape
|
||||
x = x.transpose(1, 2).view(B, C, H, W).contiguous()
|
||||
x = self.dwconv(x)
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
def _conv_filter(state_dict, patch_size=16):
|
||||
""" convert patch embedding weight from manual patchify + linear proj to conv"""
|
||||
out_dict = {}
|
||||
for k, v in state_dict.items():
|
||||
if 'patch_embed.proj.weight' in k:
|
||||
v = v.reshape((v.shape[0], 3, patch_size, patch_size))
|
||||
out_dict[k] = v
|
||||
|
||||
return out_dict
|
||||
|
||||
|
||||
class pvt_v2_b0(PyramidVisionTransformerImpr):
|
||||
def __init__(self, **kwargs):
|
||||
super(pvt_v2_b0, self).__init__(
|
||||
patch_size=4, embed_dims=[32, 64, 160, 256], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
|
||||
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[2, 2, 2, 2], sr_ratios=[8, 4, 2, 1],
|
||||
drop_rate=0.0, drop_path_rate=0.1)
|
||||
|
||||
|
||||
class pvt_v2_b1(PyramidVisionTransformerImpr):
|
||||
def __init__(self, **kwargs):
|
||||
super(pvt_v2_b1, self).__init__(
|
||||
patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
|
||||
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[2, 2, 2, 2], sr_ratios=[8, 4, 2, 1],
|
||||
drop_rate=0.0, drop_path_rate=0.1)
|
||||
|
||||
|
||||
class pvt_v2_b2(PyramidVisionTransformerImpr):
|
||||
def __init__(self, in_channels=3, **kwargs):
|
||||
super(pvt_v2_b2, self).__init__(
|
||||
patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
|
||||
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 4, 6, 3], sr_ratios=[8, 4, 2, 1],
|
||||
drop_rate=0.0, drop_path_rate=0.1, in_channels=in_channels)
|
||||
|
||||
|
||||
class pvt_v2_b3(PyramidVisionTransformerImpr):
|
||||
def __init__(self, **kwargs):
|
||||
super(pvt_v2_b3, self).__init__(
|
||||
patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
|
||||
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 4, 18, 3], sr_ratios=[8, 4, 2, 1],
|
||||
drop_rate=0.0, drop_path_rate=0.1)
|
||||
|
||||
|
||||
class pvt_v2_b4(PyramidVisionTransformerImpr):
|
||||
def __init__(self, **kwargs):
|
||||
super(pvt_v2_b4, self).__init__(
|
||||
patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
|
||||
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 8, 27, 3], sr_ratios=[8, 4, 2, 1],
|
||||
drop_rate=0.0, drop_path_rate=0.1)
|
||||
|
||||
|
||||
class pvt_v2_b5(PyramidVisionTransformerImpr):
|
||||
def __init__(self, **kwargs):
|
||||
super(pvt_v2_b5, self).__init__(
|
||||
patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[4, 4, 4, 4],
|
||||
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 6, 40, 3], sr_ratios=[8, 4, 2, 1],
|
||||
drop_rate=0.0, drop_path_rate=0.1)
|
||||
@@ -0,0 +1,632 @@
|
||||
# --------------------------------------------------------
|
||||
# Swin Transformer
|
||||
# Copyright (c) 2021 Microsoft
|
||||
# Licensed under The MIT License [see LICENSE for details]
|
||||
# Written by Ze Liu, Yutong Lin, Yixuan Wei
|
||||
# --------------------------------------------------------
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.utils.checkpoint as checkpoint
|
||||
import numpy as np
|
||||
try:
|
||||
# version > 0.6.13
|
||||
from timm.layers import DropPath, to_2tuple, trunc_normal_
|
||||
except Exception:
|
||||
from timm.models.layers import DropPath, to_2tuple, trunc_normal_
|
||||
|
||||
from ...config import Config
|
||||
|
||||
|
||||
config = Config()
|
||||
|
||||
class Mlp(nn.Module):
|
||||
""" Multilayer perceptron."""
|
||||
|
||||
def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
|
||||
super().__init__()
|
||||
out_features = out_features or in_features
|
||||
hidden_features = hidden_features or in_features
|
||||
self.fc1 = nn.Linear(in_features, hidden_features)
|
||||
self.act = act_layer()
|
||||
self.fc2 = nn.Linear(hidden_features, out_features)
|
||||
self.drop = nn.Dropout(drop)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.fc1(x)
|
||||
x = self.act(x)
|
||||
x = self.drop(x)
|
||||
x = self.fc2(x)
|
||||
x = self.drop(x)
|
||||
return x
|
||||
|
||||
|
||||
def window_partition(x, window_size):
|
||||
"""
|
||||
Args:
|
||||
x: (B, H, W, C)
|
||||
window_size (int): window size
|
||||
|
||||
Returns:
|
||||
windows: (num_windows*B, window_size, window_size, C)
|
||||
"""
|
||||
B, H, W, C = x.shape
|
||||
x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)
|
||||
windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)
|
||||
return windows
|
||||
|
||||
|
||||
def window_reverse(windows, window_size, H, W):
|
||||
"""
|
||||
Args:
|
||||
windows: (num_windows*B, window_size, window_size, C)
|
||||
window_size (int): Window size
|
||||
H (int): Height of image
|
||||
W (int): Width of image
|
||||
|
||||
Returns:
|
||||
x: (B, H, W, C)
|
||||
"""
|
||||
C = int(windows.shape[-1])
|
||||
x = windows.view(-1, H // window_size, W // window_size, window_size, window_size, C)
|
||||
x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, H, W, C)
|
||||
return x
|
||||
|
||||
|
||||
class WindowAttention(nn.Module):
|
||||
""" Window based multi-head self attention (W-MSA) module with relative position bias.
|
||||
It supports both of shifted and non-shifted window.
|
||||
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
window_size (tuple[int]): The height and width of the window.
|
||||
num_heads (int): Number of attention heads.
|
||||
qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
|
||||
qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set
|
||||
attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0
|
||||
proj_drop (float, optional): Dropout ratio of output. Default: 0.0
|
||||
"""
|
||||
|
||||
def __init__(self, dim, window_size, num_heads, qkv_bias=True, qk_scale=None, attn_drop=0., proj_drop=0.):
|
||||
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.window_size = window_size # Wh, Ww
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim // num_heads
|
||||
self.scale = qk_scale or head_dim ** -0.5
|
||||
|
||||
# define a parameter table of relative position bias
|
||||
self.relative_position_bias_table = nn.Parameter(
|
||||
torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads)) # 2*Wh-1 * 2*Ww-1, nH
|
||||
|
||||
# get pair-wise relative position index for each token inside the window
|
||||
coords_h = torch.arange(self.window_size[0])
|
||||
coords_w = torch.arange(self.window_size[1])
|
||||
coords = torch.stack(torch.meshgrid([coords_h, coords_w], indexing='ij')) # 2, Wh, Ww
|
||||
coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww
|
||||
relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 2, Wh*Ww, Wh*Ww
|
||||
relative_coords = relative_coords.permute(1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2
|
||||
relative_coords[:, :, 0] += self.window_size[0] - 1 # shift to start from 0
|
||||
relative_coords[:, :, 1] += self.window_size[1] - 1
|
||||
relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1
|
||||
relative_position_index = relative_coords.sum(-1) # Wh*Ww, Wh*Ww
|
||||
self.register_buffer("relative_position_index", relative_position_index)
|
||||
|
||||
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
||||
self.attn_drop_prob = attn_drop
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
|
||||
trunc_normal_(self.relative_position_bias_table, std=.02)
|
||||
self.softmax = nn.Softmax(dim=-1)
|
||||
|
||||
def forward(self, x, mask=None):
|
||||
""" Forward function.
|
||||
|
||||
Args:
|
||||
x: input features with shape of (num_windows*B, N, C)
|
||||
mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None
|
||||
"""
|
||||
B_, N, C = x.shape
|
||||
qkv = self.qkv(x).reshape(B_, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
|
||||
q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple)
|
||||
|
||||
q = q * self.scale
|
||||
|
||||
if config.SDPA_enabled:
|
||||
x = torch.nn.functional.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
attn_mask=None, dropout_p=self.attn_drop_prob, is_causal=False
|
||||
).transpose(1, 2).reshape(B_, N, C)
|
||||
else:
|
||||
attn = (q @ k.transpose(-2, -1))
|
||||
|
||||
relative_position_bias = self.relative_position_bias_table[self.relative_position_index.view(-1)].view(
|
||||
self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1) # Wh*Ww,Wh*Ww,nH
|
||||
relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous() # nH, Wh*Ww, Wh*Ww
|
||||
attn = attn + relative_position_bias.unsqueeze(0)
|
||||
|
||||
if mask is not None:
|
||||
nW = mask.shape[0]
|
||||
attn = attn.view(B_ // nW, nW, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0)
|
||||
attn = attn.view(-1, self.num_heads, N, N)
|
||||
attn = self.softmax(attn)
|
||||
else:
|
||||
attn = self.softmax(attn)
|
||||
|
||||
attn = self.attn_drop(attn)
|
||||
|
||||
x = (attn @ v).transpose(1, 2).reshape(B_, N, C)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
return x
|
||||
|
||||
|
||||
class SwinTransformerBlock(nn.Module):
|
||||
""" Swin Transformer Block.
|
||||
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
num_heads (int): Number of attention heads.
|
||||
window_size (int): Window size.
|
||||
shift_size (int): Shift size for SW-MSA.
|
||||
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
|
||||
qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
|
||||
qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
|
||||
drop (float, optional): Dropout rate. Default: 0.0
|
||||
attn_drop (float, optional): Attention dropout rate. Default: 0.0
|
||||
drop_path (float, optional): Stochastic depth rate. Default: 0.0
|
||||
act_layer (nn.Module, optional): Activation layer. Default: nn.GELU
|
||||
norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
|
||||
"""
|
||||
|
||||
def __init__(self, dim, num_heads, window_size=7, shift_size=0,
|
||||
mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., drop_path=0.,
|
||||
act_layer=nn.GELU, norm_layer=nn.LayerNorm):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.window_size = window_size
|
||||
self.shift_size = shift_size
|
||||
self.mlp_ratio = mlp_ratio
|
||||
assert 0 <= self.shift_size < self.window_size, "shift_size must in 0-window_size"
|
||||
|
||||
self.norm1 = norm_layer(dim)
|
||||
self.attn = WindowAttention(
|
||||
dim, window_size=to_2tuple(self.window_size), num_heads=num_heads,
|
||||
qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop)
|
||||
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
|
||||
self.norm2 = norm_layer(dim)
|
||||
mlp_hidden_dim = int(dim * mlp_ratio)
|
||||
self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
|
||||
|
||||
self.H = None
|
||||
self.W = None
|
||||
|
||||
def forward(self, x, mask_matrix):
|
||||
""" Forward function.
|
||||
|
||||
Args:
|
||||
x: Input feature, tensor size (B, H*W, C).
|
||||
H, W: Spatial resolution of the input feature.
|
||||
mask_matrix: Attention mask for cyclic shift.
|
||||
"""
|
||||
B, L, C = x.shape
|
||||
H, W = self.H, self.W
|
||||
assert L == H * W, "input feature has wrong size"
|
||||
|
||||
shortcut = x
|
||||
x = self.norm1(x)
|
||||
x = x.view(B, H, W, C)
|
||||
|
||||
# pad feature maps to multiples of window size
|
||||
pad_l = pad_t = 0
|
||||
pad_r = (self.window_size - W % self.window_size) % self.window_size
|
||||
pad_b = (self.window_size - H % self.window_size) % self.window_size
|
||||
x = F.pad(x, (0, 0, pad_l, pad_r, pad_t, pad_b))
|
||||
_, Hp, Wp, _ = x.shape
|
||||
|
||||
# cyclic shift
|
||||
if self.shift_size > 0:
|
||||
shifted_x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2))
|
||||
attn_mask = mask_matrix
|
||||
else:
|
||||
shifted_x = x
|
||||
attn_mask = None
|
||||
|
||||
# partition windows
|
||||
x_windows = window_partition(shifted_x, self.window_size) # nW*B, window_size, window_size, C
|
||||
x_windows = x_windows.view(-1, self.window_size * self.window_size, C) # nW*B, window_size*window_size, C
|
||||
|
||||
# W-MSA/SW-MSA
|
||||
attn_windows = self.attn(x_windows, mask=attn_mask) # nW*B, window_size*window_size, C
|
||||
|
||||
# merge windows
|
||||
attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C)
|
||||
shifted_x = window_reverse(attn_windows, self.window_size, Hp, Wp) # B H' W' C
|
||||
|
||||
# reverse cyclic shift
|
||||
if self.shift_size > 0:
|
||||
x = torch.roll(shifted_x, shifts=(self.shift_size, self.shift_size), dims=(1, 2))
|
||||
else:
|
||||
x = shifted_x
|
||||
|
||||
if pad_r > 0 or pad_b > 0:
|
||||
x = x[:, :H, :W, :].contiguous()
|
||||
|
||||
x = x.view(B, H * W, C)
|
||||
|
||||
# FFN
|
||||
x = shortcut + self.drop_path(x)
|
||||
x = x + self.drop_path(self.mlp(self.norm2(x)))
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class PatchMerging(nn.Module):
|
||||
""" Patch Merging Layer
|
||||
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
|
||||
"""
|
||||
def __init__(self, dim, norm_layer=nn.LayerNorm):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)
|
||||
self.norm = norm_layer(4 * dim)
|
||||
|
||||
def forward(self, x, H, W):
|
||||
""" Forward function.
|
||||
|
||||
Args:
|
||||
x: Input feature, tensor size (B, H*W, C).
|
||||
H, W: Spatial resolution of the input feature.
|
||||
"""
|
||||
B, L, C = x.shape
|
||||
assert L == H * W, "input feature has wrong size"
|
||||
|
||||
x = x.view(B, H, W, C)
|
||||
|
||||
# padding
|
||||
pad_input = (H % 2 == 1) or (W % 2 == 1)
|
||||
if pad_input:
|
||||
x = F.pad(x, (0, 0, 0, W % 2, 0, H % 2))
|
||||
|
||||
x0 = x[:, 0::2, 0::2, :] # B H/2 W/2 C
|
||||
x1 = x[:, 1::2, 0::2, :] # B H/2 W/2 C
|
||||
x2 = x[:, 0::2, 1::2, :] # B H/2 W/2 C
|
||||
x3 = x[:, 1::2, 1::2, :] # B H/2 W/2 C
|
||||
x = torch.cat([x0, x1, x2, x3], -1) # B H/2 W/2 4*C
|
||||
x = x.view(B, -1, 4 * C) # B H/2*W/2 4*C
|
||||
|
||||
x = self.norm(x)
|
||||
x = self.reduction(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class BasicLayer(nn.Module):
|
||||
""" A basic Swin Transformer layer for one stage.
|
||||
|
||||
Args:
|
||||
dim (int): Number of feature channels
|
||||
depth (int): Depths of this stage.
|
||||
num_heads (int): Number of attention head.
|
||||
window_size (int): Local window size. Default: 7.
|
||||
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
|
||||
qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
|
||||
qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
|
||||
drop (float, optional): Dropout rate. Default: 0.0
|
||||
attn_drop (float, optional): Attention dropout rate. Default: 0.0
|
||||
drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0
|
||||
norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
|
||||
downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None
|
||||
use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
depth,
|
||||
num_heads,
|
||||
window_size=7,
|
||||
mlp_ratio=4.,
|
||||
qkv_bias=True,
|
||||
qk_scale=None,
|
||||
drop=0.,
|
||||
attn_drop=0.,
|
||||
drop_path=0.,
|
||||
norm_layer=nn.LayerNorm,
|
||||
downsample=None,
|
||||
use_checkpoint=False):
|
||||
super().__init__()
|
||||
self.window_size = window_size
|
||||
self.shift_size = window_size // 2
|
||||
self.depth = depth
|
||||
self.use_checkpoint = use_checkpoint
|
||||
|
||||
# build blocks
|
||||
self.blocks = nn.ModuleList([
|
||||
SwinTransformerBlock(
|
||||
dim=dim,
|
||||
num_heads=num_heads,
|
||||
window_size=window_size,
|
||||
shift_size=0 if (i % 2 == 0) else window_size // 2,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop,
|
||||
attn_drop=attn_drop,
|
||||
drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path,
|
||||
norm_layer=norm_layer)
|
||||
for i in range(depth)])
|
||||
|
||||
# patch merging layer
|
||||
if downsample is not None:
|
||||
self.downsample = downsample(dim=dim, norm_layer=norm_layer)
|
||||
else:
|
||||
self.downsample = None
|
||||
|
||||
def forward(self, x, H, W):
|
||||
""" Forward function.
|
||||
|
||||
Args:
|
||||
x: Input feature, tensor size (B, H*W, C).
|
||||
H, W: Spatial resolution of the input feature.
|
||||
"""
|
||||
|
||||
# calculate attention mask for SW-MSA
|
||||
# Turn int to torch.tensor for the compatiability with torch.compile in PyTorch 2.5.
|
||||
Hp = torch.ceil(torch.tensor(H) / self.window_size).to(torch.int64) * self.window_size
|
||||
Wp = torch.ceil(torch.tensor(W) / self.window_size).to(torch.int64) * self.window_size
|
||||
img_mask = torch.zeros((1, Hp, Wp, 1), device=x.device) # 1 Hp Wp 1
|
||||
h_slices = (slice(0, -self.window_size),
|
||||
slice(-self.window_size, -self.shift_size),
|
||||
slice(-self.shift_size, None))
|
||||
w_slices = (slice(0, -self.window_size),
|
||||
slice(-self.window_size, -self.shift_size),
|
||||
slice(-self.shift_size, None))
|
||||
cnt = 0
|
||||
for h in h_slices:
|
||||
for w in w_slices:
|
||||
img_mask[:, h, w, :] = cnt
|
||||
cnt += 1
|
||||
|
||||
mask_windows = window_partition(img_mask, self.window_size) # nW, window_size, window_size, 1
|
||||
mask_windows = mask_windows.view(-1, self.window_size * self.window_size)
|
||||
attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
|
||||
attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0)).to(x.dtype)
|
||||
|
||||
for blk in self.blocks:
|
||||
blk.H, blk.W = H, W
|
||||
if self.use_checkpoint:
|
||||
x = checkpoint.checkpoint(blk, x, attn_mask)
|
||||
else:
|
||||
x = blk(x, attn_mask)
|
||||
if self.downsample is not None:
|
||||
x_down = self.downsample(x, H, W)
|
||||
Wh, Ww = (H + 1) // 2, (W + 1) // 2
|
||||
return x, H, W, x_down, Wh, Ww
|
||||
else:
|
||||
return x, H, W, x, H, W
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
""" Image to Patch Embedding
|
||||
|
||||
Args:
|
||||
patch_size (int): Patch token size. Default: 4.
|
||||
in_channels (int): Number of input image channels. Default: 3.
|
||||
embed_dim (int): Number of linear projection output channels. Default: 96.
|
||||
norm_layer (nn.Module, optional): Normalization layer. Default: None
|
||||
"""
|
||||
|
||||
def __init__(self, patch_size=4, in_channels=3, embed_dim=96, norm_layer=None):
|
||||
super().__init__()
|
||||
patch_size = to_2tuple(patch_size)
|
||||
self.patch_size = patch_size
|
||||
|
||||
self.in_channels = in_channels
|
||||
self.embed_dim = embed_dim
|
||||
|
||||
self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size)
|
||||
if norm_layer is not None:
|
||||
self.norm = norm_layer(embed_dim)
|
||||
else:
|
||||
self.norm = None
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward function."""
|
||||
# padding
|
||||
_, _, H, W = x.size()
|
||||
if W % self.patch_size[1] != 0:
|
||||
x = F.pad(x, (0, self.patch_size[1] - W % self.patch_size[1]))
|
||||
if H % self.patch_size[0] != 0:
|
||||
x = F.pad(x, (0, 0, 0, self.patch_size[0] - H % self.patch_size[0]))
|
||||
|
||||
x = self.proj(x) # B C Wh Ww
|
||||
if self.norm is not None:
|
||||
Wh, Ww = x.size(2), x.size(3)
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
x = self.norm(x)
|
||||
x = x.transpose(1, 2).view(-1, self.embed_dim, Wh, Ww)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class SwinTransformer(nn.Module):
|
||||
""" Swin Transformer backbone.
|
||||
A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows` -
|
||||
https://arxiv.org/pdf/2103.14030
|
||||
|
||||
Args:
|
||||
pretrain_img_size (int): Input image size for training the pretrained model,
|
||||
used in absolute postion embedding. Default 224.
|
||||
patch_size (int | tuple(int)): Patch size. Default: 4.
|
||||
in_channels (int): Number of input image channels. Default: 3.
|
||||
embed_dim (int): Number of linear projection output channels. Default: 96.
|
||||
depths (tuple[int]): Depths of each Swin Transformer stage.
|
||||
num_heads (tuple[int]): Number of attention head of each stage.
|
||||
window_size (int): Window size. Default: 7.
|
||||
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
|
||||
qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True
|
||||
qk_scale (float): Override default qk scale of head_dim ** -0.5 if set.
|
||||
drop_rate (float): Dropout rate.
|
||||
attn_drop_rate (float): Attention dropout rate. Default: 0.
|
||||
drop_path_rate (float): Stochastic depth rate. Default: 0.2.
|
||||
norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.
|
||||
ape (bool): If True, add absolute position embedding to the patch embedding. Default: False.
|
||||
patch_norm (bool): If True, add normalization after patch embedding. Default: True.
|
||||
out_indices (Sequence[int]): Output from which stages.
|
||||
frozen_stages (int): Stages to be frozen (stop grad and set eval mode).
|
||||
-1 means not freezing any parameters.
|
||||
use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
pretrain_img_size=224,
|
||||
patch_size=4,
|
||||
in_channels=3,
|
||||
embed_dim=96,
|
||||
depths=[2, 2, 6, 2],
|
||||
num_heads=[3, 6, 12, 24],
|
||||
window_size=7,
|
||||
mlp_ratio=4.,
|
||||
qkv_bias=True,
|
||||
qk_scale=None,
|
||||
drop_rate=0.,
|
||||
attn_drop_rate=0.,
|
||||
drop_path_rate=0.2,
|
||||
norm_layer=nn.LayerNorm,
|
||||
ape=False,
|
||||
patch_norm=True,
|
||||
out_indices=(0, 1, 2, 3),
|
||||
frozen_stages=-1,
|
||||
use_checkpoint=False):
|
||||
super().__init__()
|
||||
|
||||
self.pretrain_img_size = pretrain_img_size
|
||||
self.num_layers = len(depths)
|
||||
self.embed_dim = embed_dim
|
||||
self.ape = ape
|
||||
self.patch_norm = patch_norm
|
||||
self.out_indices = out_indices
|
||||
self.frozen_stages = frozen_stages
|
||||
|
||||
# split image into non-overlapping patches
|
||||
self.patch_embed = PatchEmbed(
|
||||
patch_size=patch_size, in_channels=in_channels, embed_dim=embed_dim,
|
||||
norm_layer=norm_layer if self.patch_norm else None)
|
||||
|
||||
# absolute position embedding
|
||||
if self.ape:
|
||||
pretrain_img_size = to_2tuple(pretrain_img_size)
|
||||
patch_size = to_2tuple(patch_size)
|
||||
patches_resolution = [pretrain_img_size[0] // patch_size[0], pretrain_img_size[1] // patch_size[1]]
|
||||
|
||||
self.absolute_pos_embed = nn.Parameter(torch.zeros(1, embed_dim, patches_resolution[0], patches_resolution[1]))
|
||||
trunc_normal_(self.absolute_pos_embed, std=.02)
|
||||
|
||||
self.pos_drop = nn.Dropout(p=drop_rate)
|
||||
|
||||
# stochastic depth
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule
|
||||
|
||||
# build layers
|
||||
self.layers = nn.ModuleList()
|
||||
for i_layer in range(self.num_layers):
|
||||
layer = BasicLayer(
|
||||
dim=int(embed_dim * 2 ** i_layer),
|
||||
depth=depths[i_layer],
|
||||
num_heads=num_heads[i_layer],
|
||||
window_size=window_size,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop_rate,
|
||||
attn_drop=attn_drop_rate,
|
||||
drop_path=dpr[sum(depths[:i_layer]):sum(depths[:i_layer + 1])],
|
||||
norm_layer=norm_layer,
|
||||
downsample=PatchMerging if (i_layer < self.num_layers - 1) else None,
|
||||
use_checkpoint=use_checkpoint)
|
||||
self.layers.append(layer)
|
||||
|
||||
num_features = [int(embed_dim * 2 ** i) for i in range(self.num_layers)]
|
||||
self.num_features = num_features
|
||||
|
||||
# add a norm layer for each output
|
||||
for i_layer in out_indices:
|
||||
layer = norm_layer(num_features[i_layer])
|
||||
layer_name = f'norm{i_layer}'
|
||||
self.add_module(layer_name, layer)
|
||||
|
||||
self._freeze_stages()
|
||||
|
||||
def _freeze_stages(self):
|
||||
if self.frozen_stages >= 0:
|
||||
self.patch_embed.eval()
|
||||
for param in self.patch_embed.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
if self.frozen_stages >= 1 and self.ape:
|
||||
self.absolute_pos_embed.requires_grad = False
|
||||
|
||||
if self.frozen_stages >= 2:
|
||||
self.pos_drop.eval()
|
||||
for i in range(0, self.frozen_stages - 1):
|
||||
m = self.layers[i]
|
||||
m.eval()
|
||||
for param in m.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward function."""
|
||||
x = self.patch_embed(x)
|
||||
|
||||
Wh, Ww = x.size(2), x.size(3)
|
||||
if self.ape:
|
||||
# interpolate the position embedding to the corresponding size
|
||||
absolute_pos_embed = F.interpolate(self.absolute_pos_embed, size=(Wh, Ww), mode='bicubic')
|
||||
x = (x + absolute_pos_embed) # B Wh*Ww C
|
||||
|
||||
outs = []#x.contiguous()]
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
x = self.pos_drop(x)
|
||||
for i in range(self.num_layers):
|
||||
layer = self.layers[i]
|
||||
x_out, H, W, x, Wh, Ww = layer(x, Wh, Ww)
|
||||
|
||||
if i in self.out_indices:
|
||||
norm_layer = getattr(self, f'norm{i}')
|
||||
x_out = norm_layer(x_out)
|
||||
|
||||
out = x_out.view(-1, H, W, self.num_features[i]).permute(0, 3, 1, 2).contiguous()
|
||||
outs.append(out)
|
||||
|
||||
return tuple(outs)
|
||||
|
||||
def train(self, mode=True):
|
||||
"""Convert the model into training mode while keep layers freezed."""
|
||||
super(SwinTransformer, self).train(mode)
|
||||
self._freeze_stages()
|
||||
|
||||
def swin_v1_t():
|
||||
model = SwinTransformer(embed_dim=96, depths=[2, 2, 6, 2], num_heads=[3, 6, 12, 24], window_size=7)
|
||||
return model
|
||||
|
||||
def swin_v1_s():
|
||||
model = SwinTransformer(embed_dim=96, depths=[2, 2, 18, 2], num_heads=[3, 6, 12, 24], window_size=7)
|
||||
return model
|
||||
|
||||
def swin_v1_b():
|
||||
model = SwinTransformer(embed_dim=128, depths=[2, 2, 18, 2], num_heads=[4, 8, 16, 32], window_size=12)
|
||||
return model
|
||||
|
||||
def swin_v1_l():
|
||||
model = SwinTransformer(embed_dim=192, depths=[2, 2, 18, 2], num_heads=[6, 12, 24, 48], window_size=12)
|
||||
return model
|
||||
@@ -0,0 +1,338 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
from kornia.filters import laplacian
|
||||
from huggingface_hub import PyTorchModelHubMixin
|
||||
|
||||
from ..config import Config
|
||||
from ..dataset import class_labels_TR_sorted
|
||||
from .backbones.build_backbone import build_backbone
|
||||
from .modules.decoder_blocks import BasicDecBlk, ResBlk
|
||||
from .modules.lateral_blocks import BasicLatBlk
|
||||
from .modules.aspp import ASPP, ASPPDeformable
|
||||
from .refinement.refiner import Refiner, RefinerPVTInChannels4, RefUNet
|
||||
from .refinement.stem_layer import StemLayer
|
||||
|
||||
|
||||
def image2patches(image, grid_h=2, grid_w=2, patch_ref=None, transformation='b c (hg h) (wg w) -> (b hg wg) c h w'):
|
||||
if patch_ref is not None:
|
||||
grid_h, grid_w = image.shape[-2] // patch_ref.shape[-2], image.shape[-1] // patch_ref.shape[-1]
|
||||
patches = rearrange(image, transformation, hg=grid_h, wg=grid_w)
|
||||
return patches
|
||||
|
||||
def patches2image(patches, grid_h=2, grid_w=2, patch_ref=None, transformation='(b hg wg) c h w -> b c (hg h) (wg w)'):
|
||||
if patch_ref is not None:
|
||||
grid_h, grid_w = patch_ref.shape[-2] // patches[0].shape[-2], patch_ref.shape[-1] // patches[0].shape[-1]
|
||||
image = rearrange(patches, transformation, hg=grid_h, wg=grid_w)
|
||||
return image
|
||||
|
||||
class BiRefNet(nn.Module):
|
||||
def __init__(self, bb_pretrained=True, bb_index=6):
|
||||
super(BiRefNet, self).__init__()
|
||||
self.config = Config(bb_index)
|
||||
self.epoch = 1
|
||||
self.bb = build_backbone(self.config.bb, pretrained=bb_pretrained)
|
||||
|
||||
channels = self.config.lateral_channels_in_collection
|
||||
|
||||
if self.config.auxiliary_classification:
|
||||
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
||||
self.cls_head = nn.Sequential(
|
||||
nn.Linear(channels[0], len(class_labels_TR_sorted))
|
||||
)
|
||||
|
||||
if self.config.squeeze_block:
|
||||
self.squeeze_module = nn.Sequential(*[
|
||||
eval(self.config.squeeze_block.split('_x')[0])(channels[0]+sum(self.config.cxt), channels[0])
|
||||
for _ in range(eval(self.config.squeeze_block.split('_x')[1]))
|
||||
])
|
||||
|
||||
self.decoder = Decoder(channels)
|
||||
|
||||
if self.config.ender:
|
||||
self.dec_end = nn.Sequential(
|
||||
nn.Conv2d(1, 16, 3, 1, 1),
|
||||
nn.Conv2d(16, 1, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
)
|
||||
|
||||
# refine patch-level segmentation
|
||||
if self.config.refine:
|
||||
if self.config.refine == 'itself':
|
||||
self.stem_layer = StemLayer(in_channels=3+1, inter_channels=48, out_channels=3, norm_layer='BN' if self.config.batch_size > 1 else 'LN')
|
||||
else:
|
||||
self.refiner = eval('{}({})'.format(self.config.refine, 'in_channels=3+1'))
|
||||
|
||||
if self.config.freeze_bb:
|
||||
# Freeze the backbone...
|
||||
print(self.named_parameters())
|
||||
for key, value in self.named_parameters():
|
||||
if 'bb.' in key and 'refiner.' not in key:
|
||||
value.requires_grad = False
|
||||
|
||||
def forward_enc(self, x):
|
||||
if self.config.bb in ['vgg16', 'vgg16bn', 'resnet50']:
|
||||
x1 = self.bb.conv1(x); x2 = self.bb.conv2(x1); x3 = self.bb.conv3(x2); x4 = self.bb.conv4(x3)
|
||||
else:
|
||||
x1, x2, x3, x4 = self.bb(x)
|
||||
if self.config.mul_scl_ipt == 'cat':
|
||||
B, C, H, W = x.shape
|
||||
x1_, x2_, x3_, x4_ = self.bb(F.interpolate(x, size=(H//2, W//2), mode='bilinear', align_corners=True))
|
||||
x1 = torch.cat([x1, F.interpolate(x1_, size=x1.shape[2:], mode='bilinear', align_corners=True)], dim=1)
|
||||
x2 = torch.cat([x2, F.interpolate(x2_, size=x2.shape[2:], mode='bilinear', align_corners=True)], dim=1)
|
||||
x3 = torch.cat([x3, F.interpolate(x3_, size=x3.shape[2:], mode='bilinear', align_corners=True)], dim=1)
|
||||
x4 = torch.cat([x4, F.interpolate(x4_, size=x4.shape[2:], mode='bilinear', align_corners=True)], dim=1)
|
||||
elif self.config.mul_scl_ipt == 'add':
|
||||
B, C, H, W = x.shape
|
||||
x1_, x2_, x3_, x4_ = self.bb(F.interpolate(x, size=(H//2, W//2), mode='bilinear', align_corners=True))
|
||||
x1 = x1 + F.interpolate(x1_, size=x1.shape[2:], mode='bilinear', align_corners=True)
|
||||
x2 = x2 + F.interpolate(x2_, size=x2.shape[2:], mode='bilinear', align_corners=True)
|
||||
x3 = x3 + F.interpolate(x3_, size=x3.shape[2:], mode='bilinear', align_corners=True)
|
||||
x4 = x4 + F.interpolate(x4_, size=x4.shape[2:], mode='bilinear', align_corners=True)
|
||||
class_preds = self.cls_head(self.avgpool(x4).view(x4.shape[0], -1)) if self.training and self.config.auxiliary_classification else None
|
||||
if self.config.cxt:
|
||||
x4 = torch.cat(
|
||||
(
|
||||
*[
|
||||
F.interpolate(x1, size=x4.shape[2:], mode='bilinear', align_corners=True),
|
||||
F.interpolate(x2, size=x4.shape[2:], mode='bilinear', align_corners=True),
|
||||
F.interpolate(x3, size=x4.shape[2:], mode='bilinear', align_corners=True),
|
||||
][-len(self.config.cxt):],
|
||||
x4
|
||||
),
|
||||
dim=1
|
||||
)
|
||||
return (x1, x2, x3, x4), class_preds
|
||||
|
||||
def forward_ori(self, x):
|
||||
########## Encoder ##########
|
||||
(x1, x2, x3, x4), class_preds = self.forward_enc(x)
|
||||
if self.config.squeeze_block:
|
||||
x4 = self.squeeze_module(x4)
|
||||
########## Decoder ##########
|
||||
features = [x, x1, x2, x3, x4]
|
||||
if self.training and self.config.out_ref:
|
||||
features.append(laplacian(torch.mean(x, dim=1).unsqueeze(1), kernel_size=5))
|
||||
scaled_preds = self.decoder(features)
|
||||
return scaled_preds, class_preds
|
||||
|
||||
def forward(self, x):
|
||||
scaled_preds, class_preds = self.forward_ori(x)
|
||||
class_preds_lst = [class_preds]
|
||||
return [scaled_preds, class_preds_lst] if self.training else scaled_preds
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(self, channels):
|
||||
super(Decoder, self).__init__()
|
||||
self.config = Config()
|
||||
DecoderBlock = eval(self.config.dec_blk)
|
||||
LateralBlock = eval(self.config.lat_blk)
|
||||
|
||||
if self.config.dec_ipt:
|
||||
self.split = self.config.dec_ipt_split
|
||||
N_dec_ipt = 64
|
||||
DBlock = SimpleConvs
|
||||
ic = 64
|
||||
ipt_cha_opt = 1
|
||||
self.ipt_blk5 = DBlock(2**10*3 if self.split else 3, [N_dec_ipt, channels[0]//8][ipt_cha_opt], inter_channels=ic)
|
||||
self.ipt_blk4 = DBlock(2**8*3 if self.split else 3, [N_dec_ipt, channels[0]//8][ipt_cha_opt], inter_channels=ic)
|
||||
self.ipt_blk3 = DBlock(2**6*3 if self.split else 3, [N_dec_ipt, channels[1]//8][ipt_cha_opt], inter_channels=ic)
|
||||
self.ipt_blk2 = DBlock(2**4*3 if self.split else 3, [N_dec_ipt, channels[2]//8][ipt_cha_opt], inter_channels=ic)
|
||||
self.ipt_blk1 = DBlock(2**0*3 if self.split else 3, [N_dec_ipt, channels[3]//8][ipt_cha_opt], inter_channels=ic)
|
||||
else:
|
||||
self.split = None
|
||||
|
||||
self.decoder_block4 = DecoderBlock(channels[0]+([N_dec_ipt, channels[0]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[1])
|
||||
self.decoder_block3 = DecoderBlock(channels[1]+([N_dec_ipt, channels[0]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[2])
|
||||
self.decoder_block2 = DecoderBlock(channels[2]+([N_dec_ipt, channels[1]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[3])
|
||||
self.decoder_block1 = DecoderBlock(channels[3]+([N_dec_ipt, channels[2]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[3]//2)
|
||||
self.conv_out1 = nn.Sequential(nn.Conv2d(channels[3]//2+([N_dec_ipt, channels[3]//8][ipt_cha_opt] if self.config.dec_ipt else 0), 1, 1, 1, 0))
|
||||
|
||||
self.lateral_block4 = LateralBlock(channels[1], channels[1])
|
||||
self.lateral_block3 = LateralBlock(channels[2], channels[2])
|
||||
self.lateral_block2 = LateralBlock(channels[3], channels[3])
|
||||
|
||||
if self.config.ms_supervision:
|
||||
self.conv_ms_spvn_4 = nn.Conv2d(channels[1], 1, 1, 1, 0)
|
||||
self.conv_ms_spvn_3 = nn.Conv2d(channels[2], 1, 1, 1, 0)
|
||||
self.conv_ms_spvn_2 = nn.Conv2d(channels[3], 1, 1, 1, 0)
|
||||
|
||||
if self.config.out_ref:
|
||||
_N = 16
|
||||
self.gdt_convs_4 = nn.Sequential(nn.Conv2d(channels[1], _N, 3, 1, 1), nn.BatchNorm2d(_N) if self.config.batch_size > 1 else nn.Identity(), nn.ReLU(inplace=True))
|
||||
self.gdt_convs_3 = nn.Sequential(nn.Conv2d(channels[2], _N, 3, 1, 1), nn.BatchNorm2d(_N) if self.config.batch_size > 1 else nn.Identity(), nn.ReLU(inplace=True))
|
||||
self.gdt_convs_2 = nn.Sequential(nn.Conv2d(channels[3], _N, 3, 1, 1), nn.BatchNorm2d(_N) if self.config.batch_size > 1 else nn.Identity(), nn.ReLU(inplace=True))
|
||||
|
||||
self.gdt_convs_pred_4 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
|
||||
self.gdt_convs_pred_3 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
|
||||
self.gdt_convs_pred_2 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
|
||||
|
||||
self.gdt_convs_attn_4 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
|
||||
self.gdt_convs_attn_3 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
|
||||
self.gdt_convs_attn_2 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
|
||||
|
||||
def forward(self, features):
|
||||
if self.training and self.config.out_ref:
|
||||
outs_gdt_pred = []
|
||||
outs_gdt_label = []
|
||||
x, x1, x2, x3, x4, gdt_gt = features
|
||||
else:
|
||||
x, x1, x2, x3, x4 = features
|
||||
outs = []
|
||||
|
||||
if self.config.dec_ipt:
|
||||
patches_batch = image2patches(x, patch_ref=x4, transformation='b c (hg h) (wg w) -> b (c hg wg) h w') if self.split else x
|
||||
x4 = torch.cat((x4, self.ipt_blk5(F.interpolate(patches_batch, size=x4.shape[2:], mode='bilinear', align_corners=True))), 1)
|
||||
p4 = self.decoder_block4(x4)
|
||||
m4 = self.conv_ms_spvn_4(p4) if self.config.ms_supervision and self.training else None
|
||||
if self.config.out_ref:
|
||||
p4_gdt = self.gdt_convs_4(p4)
|
||||
if self.training:
|
||||
# >> GT:
|
||||
m4_dia = m4
|
||||
gdt_label_main_4 = gdt_gt * F.interpolate(m4_dia, size=gdt_gt.shape[2:], mode='bilinear', align_corners=True)
|
||||
outs_gdt_label.append(gdt_label_main_4)
|
||||
# >> Pred:
|
||||
gdt_pred_4 = self.gdt_convs_pred_4(p4_gdt)
|
||||
outs_gdt_pred.append(gdt_pred_4)
|
||||
gdt_attn_4 = self.gdt_convs_attn_4(p4_gdt).sigmoid()
|
||||
# >> Finally:
|
||||
p4 = p4 * gdt_attn_4
|
||||
_p4 = F.interpolate(p4, size=x3.shape[2:], mode='bilinear', align_corners=True)
|
||||
_p3 = _p4 + self.lateral_block4(x3)
|
||||
|
||||
if self.config.dec_ipt:
|
||||
patches_batch = image2patches(x, patch_ref=_p3, transformation='b c (hg h) (wg w) -> b (c hg wg) h w') if self.split else x
|
||||
_p3 = torch.cat((_p3, self.ipt_blk4(F.interpolate(patches_batch, size=x3.shape[2:], mode='bilinear', align_corners=True))), 1)
|
||||
p3 = self.decoder_block3(_p3)
|
||||
m3 = self.conv_ms_spvn_3(p3) if self.config.ms_supervision and self.training else None
|
||||
if self.config.out_ref:
|
||||
p3_gdt = self.gdt_convs_3(p3)
|
||||
if self.training:
|
||||
# >> GT:
|
||||
# m3 --dilation--> m3_dia
|
||||
# G_3^gt * m3_dia --> G_3^m, which is the label of gradient
|
||||
m3_dia = m3
|
||||
gdt_label_main_3 = gdt_gt * F.interpolate(m3_dia, size=gdt_gt.shape[2:], mode='bilinear', align_corners=True)
|
||||
outs_gdt_label.append(gdt_label_main_3)
|
||||
# >> Pred:
|
||||
# p3 --conv--BN--> F_3^G, where F_3^G predicts the \hat{G_3} with xx
|
||||
# F_3^G --sigmoid--> A_3^G
|
||||
gdt_pred_3 = self.gdt_convs_pred_3(p3_gdt)
|
||||
outs_gdt_pred.append(gdt_pred_3)
|
||||
gdt_attn_3 = self.gdt_convs_attn_3(p3_gdt).sigmoid()
|
||||
# >> Finally:
|
||||
# p3 = p3 * A_3^G
|
||||
p3 = p3 * gdt_attn_3
|
||||
_p3 = F.interpolate(p3, size=x2.shape[2:], mode='bilinear', align_corners=True)
|
||||
_p2 = _p3 + self.lateral_block3(x2)
|
||||
|
||||
if self.config.dec_ipt:
|
||||
patches_batch = image2patches(x, patch_ref=_p2, transformation='b c (hg h) (wg w) -> b (c hg wg) h w') if self.split else x
|
||||
_p2 = torch.cat((_p2, self.ipt_blk3(F.interpolate(patches_batch, size=x2.shape[2:], mode='bilinear', align_corners=True))), 1)
|
||||
p2 = self.decoder_block2(_p2)
|
||||
m2 = self.conv_ms_spvn_2(p2) if self.config.ms_supervision and self.training else None
|
||||
if self.config.out_ref:
|
||||
p2_gdt = self.gdt_convs_2(p2)
|
||||
if self.training:
|
||||
# >> GT:
|
||||
m2_dia = m2
|
||||
gdt_label_main_2 = gdt_gt * F.interpolate(m2_dia, size=gdt_gt.shape[2:], mode='bilinear', align_corners=True)
|
||||
outs_gdt_label.append(gdt_label_main_2)
|
||||
# >> Pred:
|
||||
gdt_pred_2 = self.gdt_convs_pred_2(p2_gdt)
|
||||
outs_gdt_pred.append(gdt_pred_2)
|
||||
gdt_attn_2 = self.gdt_convs_attn_2(p2_gdt).sigmoid()
|
||||
# >> Finally:
|
||||
p2 = p2 * gdt_attn_2
|
||||
_p2 = F.interpolate(p2, size=x1.shape[2:], mode='bilinear', align_corners=True)
|
||||
_p1 = _p2 + self.lateral_block2(x1)
|
||||
|
||||
if self.config.dec_ipt:
|
||||
patches_batch = image2patches(x, patch_ref=_p1, transformation='b c (hg h) (wg w) -> b (c hg wg) h w') if self.split else x
|
||||
_p1 = torch.cat((_p1, self.ipt_blk2(F.interpolate(patches_batch, size=x1.shape[2:], mode='bilinear', align_corners=True))), 1)
|
||||
_p1 = self.decoder_block1(_p1)
|
||||
_p1 = F.interpolate(_p1, size=x.shape[2:], mode='bilinear', align_corners=True)
|
||||
|
||||
if self.config.dec_ipt:
|
||||
patches_batch = image2patches(x, patch_ref=_p1, transformation='b c (hg h) (wg w) -> b (c hg wg) h w') if self.split else x
|
||||
_p1 = torch.cat((_p1, self.ipt_blk1(F.interpolate(patches_batch, size=x.shape[2:], mode='bilinear', align_corners=True))), 1)
|
||||
p1_out = self.conv_out1(_p1)
|
||||
|
||||
if self.config.ms_supervision and self.training:
|
||||
outs.append(m4)
|
||||
outs.append(m3)
|
||||
outs.append(m2)
|
||||
outs.append(p1_out)
|
||||
return outs if not (self.config.out_ref and self.training) else ([outs_gdt_pred, outs_gdt_label], outs)
|
||||
|
||||
|
||||
class SimpleConvs(nn.Module):
|
||||
def __init__(
|
||||
self, in_channels: int, out_channels: int, inter_channels=64
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.conv1 = nn.Conv2d(in_channels, inter_channels, 3, 1, 1)
|
||||
self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, 1)
|
||||
|
||||
def forward(self, x):
|
||||
return self.conv_out(self.conv1(x))
|
||||
|
||||
|
||||
###########
|
||||
|
||||
|
||||
class BiRefNetC2F(
|
||||
nn.Module,
|
||||
PyTorchModelHubMixin,
|
||||
library_name="birefnet_c2f",
|
||||
repo_url="https://github.com/ZhengPeng7/BiRefNet_C2F",
|
||||
tags=['Image Segmentation', 'Background Removal', 'Mask Generation', 'Dichotomous Image Segmentation', 'Camouflaged Object Detection', 'Salient Object Detection']
|
||||
):
|
||||
def __init__(self, bb_pretrained=True):
|
||||
super(BiRefNetC2F, self).__init__()
|
||||
self.config = Config()
|
||||
self.epoch = 1
|
||||
self.grid = 4
|
||||
self.model_coarse = BiRefNet(bb_pretrained=True)
|
||||
self.model_fine = BiRefNet(bb_pretrained=True)
|
||||
self.input_mixer = nn.Conv2d(4, 3, 1, 1, 0)
|
||||
self.output_mixer_merge_post = nn.Sequential(nn.Conv2d(1, 16, 3, 1, 1), nn.Conv2d(16, 1, 3, 1, 1))
|
||||
|
||||
def forward(self, x):
|
||||
x_ori = x.clone()
|
||||
########## Coarse ##########
|
||||
x = F.interpolate(x, size=[s//self.grid for s in self.config.size[::-1]], mode='bilinear', align_corners=True)
|
||||
|
||||
if self.training:
|
||||
scaled_preds, class_preds_lst = self.model_coarse(x)
|
||||
else:
|
||||
scaled_preds = self.model_coarse(x)
|
||||
########## Fine ##########
|
||||
x_HR_patches = image2patches(x_ori, patch_ref=x, transformation='b c (hg h) (wg w) -> (b hg wg) c h w')
|
||||
pred = F.interpolate(scaled_preds[-1] if not (self.config.out_ref and self.training) else scaled_preds[1][-1], size=x_ori.shape[2:], mode='bilinear', align_corners=True)
|
||||
pred_patches = image2patches(pred, patch_ref=x, transformation='b c (hg h) (wg w) -> (b hg wg) c h w')
|
||||
t = torch.cat([x_HR_patches, pred_patches], dim=1)
|
||||
x_HR = self.input_mixer(t)
|
||||
|
||||
pred_patches = image2patches(pred, patch_ref=x_HR, transformation='b c (hg h) (wg w) -> b (c hg wg) h w')
|
||||
if self.training:
|
||||
scaled_preds_HR, class_preds_lst_HR = self.model_fine(x_HR)
|
||||
else:
|
||||
scaled_preds_HR = self.model_fine(x_HR)
|
||||
if self.training:
|
||||
if self.config.out_ref:
|
||||
[outs_gdt_pred, outs_gdt_label], outs = scaled_preds
|
||||
[outs_gdt_pred_HR, outs_gdt_label_HR], outs_HR = scaled_preds_HR
|
||||
for idx_out, out_HR in enumerate(outs_HR):
|
||||
outs_HR[idx_out] = self.output_mixer_merge_post(patches2image(out_HR, grid_h=self.grid, grid_w=self.grid, transformation='(b hg wg) c h w -> b c (hg h) (wg w)'))
|
||||
return [([outs_gdt_pred + outs_gdt_pred_HR, outs_gdt_label + outs_gdt_label_HR], outs + outs_HR), class_preds_lst] # handle gt here
|
||||
else:
|
||||
return [
|
||||
scaled_preds + [self.output_mixer_merge_post(patches2image(scaled_pred_HR, grid_h=self.grid, grid_w=self.grid, transformation='(b hg wg) c h w -> b c (hg h) (wg w)')) for scaled_pred_HR in scaled_preds_HR],
|
||||
class_preds_lst
|
||||
]
|
||||
else:
|
||||
return scaled_preds + [self.output_mixer_merge_post(patches2image(scaled_pred_HR, grid_h=self.grid, grid_w=self.grid, transformation='(b hg wg) c h w -> b c (hg h) (wg w)')) for scaled_pred_HR in scaled_preds_HR]
|
||||
@@ -0,0 +1,119 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from ..modules.deform_conv import DeformableConv2d
|
||||
from ...config import Config
|
||||
|
||||
|
||||
config = Config()
|
||||
|
||||
|
||||
class _ASPPModule(nn.Module):
|
||||
def __init__(self, in_channels, planes, kernel_size, padding, dilation):
|
||||
super(_ASPPModule, self).__init__()
|
||||
self.atrous_conv = nn.Conv2d(in_channels, planes, kernel_size=kernel_size,
|
||||
stride=1, padding=padding, dilation=dilation, bias=False)
|
||||
self.bn = nn.BatchNorm2d(planes) if config.batch_size > 1 else nn.Identity()
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.atrous_conv(x)
|
||||
x = self.bn(x)
|
||||
|
||||
return self.relu(x)
|
||||
|
||||
|
||||
class ASPP(nn.Module):
|
||||
def __init__(self, in_channels=64, out_channels=None, output_stride=16):
|
||||
super(ASPP, self).__init__()
|
||||
self.down_scale = 1
|
||||
if out_channels is None:
|
||||
out_channels = in_channels
|
||||
self.in_channelster = 256 // self.down_scale
|
||||
if output_stride == 16:
|
||||
dilations = [1, 6, 12, 18]
|
||||
elif output_stride == 8:
|
||||
dilations = [1, 12, 24, 36]
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
self.aspp1 = _ASPPModule(in_channels, self.in_channelster, 1, padding=0, dilation=dilations[0])
|
||||
self.aspp2 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[1], dilation=dilations[1])
|
||||
self.aspp3 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[2], dilation=dilations[2])
|
||||
self.aspp4 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[3], dilation=dilations[3])
|
||||
|
||||
self.global_avg_pool = nn.Sequential(nn.AdaptiveAvgPool2d((1, 1)),
|
||||
nn.Conv2d(in_channels, self.in_channelster, 1, stride=1, bias=False),
|
||||
nn.BatchNorm2d(self.in_channelster) if config.batch_size > 1 else nn.Identity(),
|
||||
nn.ReLU(inplace=True))
|
||||
self.conv1 = nn.Conv2d(self.in_channelster * 5, out_channels, 1, bias=False)
|
||||
self.bn1 = nn.BatchNorm2d(out_channels) if config.batch_size > 1 else nn.Identity()
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
self.dropout = nn.Dropout(0.5)
|
||||
|
||||
def forward(self, x):
|
||||
x1 = self.aspp1(x)
|
||||
x2 = self.aspp2(x)
|
||||
x3 = self.aspp3(x)
|
||||
x4 = self.aspp4(x)
|
||||
x5 = self.global_avg_pool(x)
|
||||
x5 = F.interpolate(x5, size=x1.size()[2:], mode='bilinear', align_corners=True)
|
||||
x = torch.cat((x1, x2, x3, x4, x5), dim=1)
|
||||
|
||||
x = self.conv1(x)
|
||||
x = self.bn1(x)
|
||||
x = self.relu(x)
|
||||
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
##################### Deformable
|
||||
class _ASPPModuleDeformable(nn.Module):
|
||||
def __init__(self, in_channels, planes, kernel_size, padding):
|
||||
super(_ASPPModuleDeformable, self).__init__()
|
||||
self.atrous_conv = DeformableConv2d(in_channels, planes, kernel_size=kernel_size,
|
||||
stride=1, padding=padding, bias=False)
|
||||
self.bn = nn.BatchNorm2d(planes) if config.batch_size > 1 else nn.Identity()
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.atrous_conv(x)
|
||||
x = self.bn(x)
|
||||
|
||||
return self.relu(x)
|
||||
|
||||
|
||||
class ASPPDeformable(nn.Module):
|
||||
def __init__(self, in_channels, out_channels=None, parallel_block_sizes=[1, 3, 7]):
|
||||
super(ASPPDeformable, self).__init__()
|
||||
self.down_scale = 1
|
||||
if out_channels is None:
|
||||
out_channels = in_channels
|
||||
self.in_channelster = 256 // self.down_scale
|
||||
|
||||
self.aspp1 = _ASPPModuleDeformable(in_channels, self.in_channelster, 1, padding=0)
|
||||
self.aspp_deforms = nn.ModuleList([
|
||||
_ASPPModuleDeformable(in_channels, self.in_channelster, conv_size, padding=int(conv_size//2)) for conv_size in parallel_block_sizes
|
||||
])
|
||||
|
||||
self.global_avg_pool = nn.Sequential(nn.AdaptiveAvgPool2d((1, 1)),
|
||||
nn.Conv2d(in_channels, self.in_channelster, 1, stride=1, bias=False),
|
||||
nn.BatchNorm2d(self.in_channelster) if config.batch_size > 1 else nn.Identity(),
|
||||
nn.ReLU(inplace=True))
|
||||
self.conv1 = nn.Conv2d(self.in_channelster * (2 + len(self.aspp_deforms)), out_channels, 1, bias=False)
|
||||
self.bn1 = nn.BatchNorm2d(out_channels) if config.batch_size > 1 else nn.Identity()
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
self.dropout = nn.Dropout(0.5)
|
||||
|
||||
def forward(self, x):
|
||||
x1 = self.aspp1(x)
|
||||
x_aspp_deforms = [aspp_deform(x) for aspp_deform in self.aspp_deforms]
|
||||
x5 = self.global_avg_pool(x)
|
||||
x5 = F.interpolate(x5, size=x1.size()[2:], mode='bilinear', align_corners=True)
|
||||
x = torch.cat((x1, *x_aspp_deforms, x5), dim=1)
|
||||
|
||||
x = self.conv1(x)
|
||||
x = self.bn1(x)
|
||||
x = self.relu(x)
|
||||
|
||||
return self.dropout(x)
|
||||
@@ -0,0 +1,65 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from ..modules.aspp import ASPP, ASPPDeformable
|
||||
from ...config import Config
|
||||
|
||||
|
||||
config = Config()
|
||||
|
||||
|
||||
class BasicDecBlk(nn.Module):
|
||||
def __init__(self, in_channels=64, out_channels=64, inter_channels=64):
|
||||
super(BasicDecBlk, self).__init__()
|
||||
inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
|
||||
self.conv_in = nn.Conv2d(in_channels, inter_channels, 3, 1, padding=1)
|
||||
self.relu_in = nn.ReLU(inplace=True)
|
||||
if config.dec_att == 'ASPP':
|
||||
self.dec_att = ASPP(in_channels=inter_channels)
|
||||
elif config.dec_att == 'ASPPDeformable':
|
||||
self.dec_att = ASPPDeformable(in_channels=inter_channels)
|
||||
self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, padding=1)
|
||||
self.bn_in = nn.BatchNorm2d(inter_channels) if config.batch_size > 1 else nn.Identity()
|
||||
self.bn_out = nn.BatchNorm2d(out_channels) if config.batch_size > 1 else nn.Identity()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv_in(x)
|
||||
x = self.bn_in(x)
|
||||
x = self.relu_in(x)
|
||||
if hasattr(self, 'dec_att'):
|
||||
x = self.dec_att(x)
|
||||
x = self.conv_out(x)
|
||||
x = self.bn_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class ResBlk(nn.Module):
|
||||
def __init__(self, in_channels=64, out_channels=None, inter_channels=64):
|
||||
super(ResBlk, self).__init__()
|
||||
if out_channels is None:
|
||||
out_channels = in_channels
|
||||
inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
|
||||
|
||||
self.conv_in = nn.Conv2d(in_channels, inter_channels, 3, 1, padding=1)
|
||||
self.bn_in = nn.BatchNorm2d(inter_channels) if config.batch_size > 1 else nn.Identity()
|
||||
self.relu_in = nn.ReLU(inplace=True)
|
||||
|
||||
if config.dec_att == 'ASPP':
|
||||
self.dec_att = ASPP(in_channels=inter_channels)
|
||||
elif config.dec_att == 'ASPPDeformable':
|
||||
self.dec_att = ASPPDeformable(in_channels=inter_channels)
|
||||
|
||||
self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, padding=1)
|
||||
self.bn_out = nn.BatchNorm2d(out_channels) if config.batch_size > 1 else nn.Identity()
|
||||
|
||||
self.conv_resi = nn.Conv2d(in_channels, out_channels, 1, 1, 0)
|
||||
|
||||
def forward(self, x):
|
||||
_x = self.conv_resi(x)
|
||||
x = self.conv_in(x)
|
||||
x = self.bn_in(x)
|
||||
x = self.relu_in(x)
|
||||
if hasattr(self, 'dec_att'):
|
||||
x = self.dec_att(x)
|
||||
x = self.conv_out(x)
|
||||
x = self.bn_out(x)
|
||||
return x + _x
|
||||
@@ -0,0 +1,66 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torchvision.ops import deform_conv2d
|
||||
|
||||
|
||||
class DeformableConv2d(nn.Module):
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1,
|
||||
bias=False):
|
||||
|
||||
super(DeformableConv2d, self).__init__()
|
||||
|
||||
assert type(kernel_size) == tuple or type(kernel_size) == int
|
||||
|
||||
kernel_size = kernel_size if type(kernel_size) == tuple else (kernel_size, kernel_size)
|
||||
self.stride = stride if type(stride) == tuple else (stride, stride)
|
||||
self.padding = padding
|
||||
|
||||
self.offset_conv = nn.Conv2d(in_channels,
|
||||
2 * kernel_size[0] * kernel_size[1],
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=self.padding,
|
||||
bias=True)
|
||||
|
||||
nn.init.constant_(self.offset_conv.weight, 0.)
|
||||
nn.init.constant_(self.offset_conv.bias, 0.)
|
||||
|
||||
self.modulator_conv = nn.Conv2d(in_channels,
|
||||
1 * kernel_size[0] * kernel_size[1],
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=self.padding,
|
||||
bias=True)
|
||||
|
||||
nn.init.constant_(self.modulator_conv.weight, 0.)
|
||||
nn.init.constant_(self.modulator_conv.bias, 0.)
|
||||
|
||||
self.regular_conv = nn.Conv2d(in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=self.padding,
|
||||
bias=bias)
|
||||
|
||||
def forward(self, x):
|
||||
#h, w = x.shape[2:]
|
||||
#max_offset = max(h, w)/4.
|
||||
|
||||
offset = self.offset_conv(x)#.clamp(-max_offset, max_offset)
|
||||
modulator = 2. * torch.sigmoid(self.modulator_conv(x))
|
||||
|
||||
x = deform_conv2d(
|
||||
input=x,
|
||||
offset=offset,
|
||||
weight=self.regular_conv.weight,
|
||||
bias=self.regular_conv.bias,
|
||||
padding=self.padding,
|
||||
mask=modulator,
|
||||
stride=self.stride,
|
||||
)
|
||||
return x
|
||||
@@ -0,0 +1,21 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from functools import partial
|
||||
|
||||
from ...config import Config
|
||||
|
||||
|
||||
config = Config()
|
||||
|
||||
|
||||
class BasicLatBlk(nn.Module):
|
||||
def __init__(self, in_channels=64, out_channels=64, inter_channels=64):
|
||||
super(BasicLatBlk, self).__init__()
|
||||
inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
|
||||
self.conv = nn.Conv2d(in_channels, out_channels, 1, 1, 0)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
return x
|
||||
@@ -0,0 +1,222 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from typing import Any, Optional, Tuple, Type
|
||||
|
||||
|
||||
class PromptEncoder(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
embed_dim=256,
|
||||
image_embedding_size=1024,
|
||||
input_image_size=(1024, 1024),
|
||||
mask_in_chans=16,
|
||||
activation=nn.GELU
|
||||
) -> None:
|
||||
super().__init__()
|
||||
"""
|
||||
Codes are partially from SAM: https://github.com/facebookresearch/segment-anything/blob/6fdee8f2727f4506cfbbe553e23b895e27956588/segment_anything/modeling/prompt_encoder.py.
|
||||
|
||||
Arguments:
|
||||
embed_dim (int): The prompts' embedding dimension
|
||||
image_embedding_size (tuple(int, int)): The spatial size of the
|
||||
image embedding, as (H, W).
|
||||
input_image_size (int): The padded size of the image as input
|
||||
to the image encoder, as (H, W).
|
||||
mask_in_chans (int): The number of hidden channels used for
|
||||
encoding input masks.
|
||||
activation (nn.Module): The activation to use when encoding
|
||||
input masks.
|
||||
"""
|
||||
super().__init__()
|
||||
self.embed_dim = embed_dim
|
||||
self.input_image_size = input_image_size
|
||||
self.image_embedding_size = image_embedding_size
|
||||
self.pe_layer = PositionEmbeddingRandom(embed_dim // 2)
|
||||
|
||||
self.num_point_embeddings: int = 4 # pos/neg point + 2 box corners
|
||||
point_embeddings = [nn.Embedding(1, embed_dim) for i in range(self.num_point_embeddings)]
|
||||
self.point_embeddings = nn.ModuleList(point_embeddings)
|
||||
self.not_a_point_embed = nn.Embedding(1, embed_dim)
|
||||
|
||||
self.mask_input_size = (4 * image_embedding_size[0], 4 * image_embedding_size[1])
|
||||
self.mask_downscaling = nn.Sequential(
|
||||
nn.Conv2d(1, mask_in_chans // 4, kernel_size=2, stride=2),
|
||||
LayerNorm2d(mask_in_chans // 4),
|
||||
activation(),
|
||||
nn.Conv2d(mask_in_chans // 4, mask_in_chans, kernel_size=2, stride=2),
|
||||
LayerNorm2d(mask_in_chans),
|
||||
activation(),
|
||||
nn.Conv2d(mask_in_chans, embed_dim, kernel_size=1),
|
||||
)
|
||||
self.no_mask_embed = nn.Embedding(1, embed_dim)
|
||||
|
||||
def get_dense_pe(self) -> torch.Tensor:
|
||||
"""
|
||||
Returns the positional encoding used to encode point prompts,
|
||||
applied to a dense set of points the shape of the image encoding.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Positional encoding with shape
|
||||
1x(embed_dim)x(embedding_h)x(embedding_w)
|
||||
"""
|
||||
return self.pe_layer(self.image_embedding_size).unsqueeze(0)
|
||||
|
||||
def _embed_points(
|
||||
self,
|
||||
points: torch.Tensor,
|
||||
labels: torch.Tensor,
|
||||
pad: bool,
|
||||
) -> torch.Tensor:
|
||||
"""Embeds point prompts."""
|
||||
points = points + 0.5 # Shift to center of pixel
|
||||
if pad:
|
||||
padding_point = torch.zeros((points.shape[0], 1, 2), device=points.device)
|
||||
padding_label = -torch.ones((labels.shape[0], 1), device=labels.device)
|
||||
points = torch.cat([points, padding_point], dim=1)
|
||||
labels = torch.cat([labels, padding_label], dim=1)
|
||||
point_embedding = self.pe_layer.forward_with_coords(points, self.input_image_size)
|
||||
point_embedding[labels == -1] = 0.0
|
||||
point_embedding[labels == -1] += self.not_a_point_embed.weight
|
||||
point_embedding[labels == 0] += self.point_embeddings[0].weight
|
||||
point_embedding[labels == 1] += self.point_embeddings[1].weight
|
||||
return point_embedding
|
||||
|
||||
def _embed_boxes(self, boxes: torch.Tensor) -> torch.Tensor:
|
||||
"""Embeds box prompts."""
|
||||
boxes = boxes + 0.5 # Shift to center of pixel
|
||||
coords = boxes.reshape(-1, 2, 2)
|
||||
corner_embedding = self.pe_layer.forward_with_coords(coords, self.input_image_size)
|
||||
corner_embedding[:, 0, :] += self.point_embeddings[2].weight
|
||||
corner_embedding[:, 1, :] += self.point_embeddings[3].weight
|
||||
return corner_embedding
|
||||
|
||||
def _embed_masks(self, masks: torch.Tensor) -> torch.Tensor:
|
||||
"""Embeds mask inputs."""
|
||||
mask_embedding = self.mask_downscaling(masks)
|
||||
return mask_embedding
|
||||
|
||||
def _get_batch_size(
|
||||
self,
|
||||
points: Optional[Tuple[torch.Tensor, torch.Tensor]],
|
||||
boxes: Optional[torch.Tensor],
|
||||
masks: Optional[torch.Tensor],
|
||||
) -> int:
|
||||
"""
|
||||
Gets the batch size of the output given the batch size of the input prompts.
|
||||
"""
|
||||
if points is not None:
|
||||
return points[0].shape[0]
|
||||
elif boxes is not None:
|
||||
return boxes.shape[0]
|
||||
elif masks is not None:
|
||||
return masks.shape[0]
|
||||
else:
|
||||
return 1
|
||||
|
||||
def _get_device(self) -> torch.device:
|
||||
return self.point_embeddings[0].weight.device
|
||||
|
||||
def forward(
|
||||
self,
|
||||
points: Optional[Tuple[torch.Tensor, torch.Tensor]],
|
||||
boxes: Optional[torch.Tensor],
|
||||
masks: Optional[torch.Tensor],
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Embeds different types of prompts, returning both sparse and dense
|
||||
embeddings.
|
||||
|
||||
Arguments:
|
||||
points (tuple(torch.Tensor, torch.Tensor) or none): point coordinates
|
||||
and labels to embed.
|
||||
boxes (torch.Tensor or none): boxes to embed
|
||||
masks (torch.Tensor or none): masks to embed
|
||||
|
||||
Returns:
|
||||
torch.Tensor: sparse embeddings for the points and boxes, with shape
|
||||
BxNx(embed_dim), where N is determined by the number of input points
|
||||
and boxes.
|
||||
torch.Tensor: dense embeddings for the masks, in the shape
|
||||
Bx(embed_dim)x(embed_H)x(embed_W)
|
||||
"""
|
||||
bs = self._get_batch_size(points, boxes, masks)
|
||||
sparse_embeddings = torch.empty((bs, 0, self.embed_dim), device=self._get_device())
|
||||
if points is not None:
|
||||
coords, labels = points
|
||||
point_embeddings = self._embed_points(coords, labels, pad=(boxes is None))
|
||||
sparse_embeddings = torch.cat([sparse_embeddings, point_embeddings], dim=1)
|
||||
if boxes is not None:
|
||||
box_embeddings = self._embed_boxes(boxes)
|
||||
sparse_embeddings = torch.cat([sparse_embeddings, box_embeddings], dim=1)
|
||||
|
||||
if masks is not None:
|
||||
dense_embeddings = self._embed_masks(masks)
|
||||
else:
|
||||
dense_embeddings = self.no_mask_embed.weight.reshape(1, -1, 1, 1).expand(
|
||||
bs, -1, self.image_embedding_size[0], self.image_embedding_size[1]
|
||||
)
|
||||
|
||||
return sparse_embeddings, dense_embeddings
|
||||
|
||||
|
||||
class PositionEmbeddingRandom(nn.Module):
|
||||
"""
|
||||
Positional encoding using random spatial frequencies.
|
||||
"""
|
||||
|
||||
def __init__(self, num_pos_feats: int = 64, scale: Optional[float] = None) -> None:
|
||||
super().__init__()
|
||||
if scale is None or scale <= 0.0:
|
||||
scale = 1.0
|
||||
self.register_buffer(
|
||||
"positional_encoding_gaussian_matrix",
|
||||
scale * torch.randn((2, num_pos_feats)),
|
||||
)
|
||||
|
||||
def _pe_encoding(self, coords: torch.Tensor) -> torch.Tensor:
|
||||
"""Positionally encode points that are normalized to [0,1]."""
|
||||
# assuming coords are in [0, 1]^2 square and have d_1 x ... x d_n x 2 shape
|
||||
coords = 2 * coords - 1
|
||||
coords = coords @ self.positional_encoding_gaussian_matrix
|
||||
coords = 2 * np.pi * coords
|
||||
# outputs d_1 x ... x d_n x C shape
|
||||
return torch.cat([torch.sin(coords), torch.cos(coords)], dim=-1)
|
||||
|
||||
def forward(self, size: Tuple[int, int]) -> torch.Tensor:
|
||||
"""Generate positional encoding for a grid of the specified size."""
|
||||
h, w = size
|
||||
device: Any = self.positional_encoding_gaussian_matrix.device
|
||||
grid = torch.ones((h, w), device=device, dtype=torch.float32)
|
||||
y_embed = grid.cumsum(dim=0) - 0.5
|
||||
x_embed = grid.cumsum(dim=1) - 0.5
|
||||
y_embed = y_embed / h
|
||||
x_embed = x_embed / w
|
||||
|
||||
pe = self._pe_encoding(torch.stack([x_embed, y_embed], dim=-1))
|
||||
return pe.permute(2, 0, 1) # C x H x W
|
||||
|
||||
def forward_with_coords(
|
||||
self, coords_input: torch.Tensor, image_size: Tuple[int, int]
|
||||
) -> torch.Tensor:
|
||||
"""Positionally encode points that are not normalized to [0,1]."""
|
||||
coords = coords_input.clone()
|
||||
coords[:, :, 0] = coords[:, :, 0] / image_size[1]
|
||||
coords[:, :, 1] = coords[:, :, 1] / image_size[0]
|
||||
return self._pe_encoding(coords.to(torch.float)) # B x N x C
|
||||
|
||||
|
||||
class LayerNorm2d(nn.Module):
|
||||
def __init__(self, num_channels: int, eps: float = 1e-6) -> None:
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(num_channels))
|
||||
self.bias = nn.Parameter(torch.zeros(num_channels))
|
||||
self.eps = eps
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
u = x.mean(1, keepdim=True)
|
||||
s = (x - u).pow(2).mean(1, keepdim=True)
|
||||
x = (x - u) / torch.sqrt(s + self.eps)
|
||||
x = self.weight[:, None, None] * x + self.bias[:, None, None]
|
||||
return x
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
def build_act_layer(act_layer):
|
||||
if act_layer == 'ReLU':
|
||||
return nn.ReLU(inplace=True)
|
||||
elif act_layer == 'SiLU':
|
||||
return nn.SiLU(inplace=True)
|
||||
elif act_layer == 'GELU':
|
||||
return nn.GELU()
|
||||
|
||||
raise NotImplementedError(f'build_act_layer does not support {act_layer}')
|
||||
|
||||
|
||||
def build_norm_layer(dim,
|
||||
norm_layer,
|
||||
in_format='channels_last',
|
||||
out_format='channels_last',
|
||||
eps=1e-6):
|
||||
layers = []
|
||||
if norm_layer == 'BN':
|
||||
if in_format == 'channels_last':
|
||||
layers.append(to_channels_first())
|
||||
layers.append(nn.BatchNorm2d(dim))
|
||||
if out_format == 'channels_last':
|
||||
layers.append(to_channels_last())
|
||||
elif norm_layer == 'LN':
|
||||
if in_format == 'channels_first':
|
||||
layers.append(to_channels_last())
|
||||
layers.append(nn.LayerNorm(dim, eps=eps))
|
||||
if out_format == 'channels_first':
|
||||
layers.append(to_channels_first())
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'build_norm_layer does not support {norm_layer}')
|
||||
return nn.Sequential(*layers)
|
||||
|
||||
|
||||
class to_channels_first(nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, x):
|
||||
return x.permute(0, 3, 1, 2)
|
||||
|
||||
|
||||
class to_channels_last(nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, x):
|
||||
return x.permute(0, 2, 3, 1)
|
||||
@@ -0,0 +1,252 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from collections import OrderedDict
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torchvision.models import vgg16, vgg16_bn
|
||||
from torchvision.models import resnet50
|
||||
|
||||
from ...config import Config
|
||||
from ...dataset import class_labels_TR_sorted
|
||||
from ..backbones.build_backbone import build_backbone
|
||||
from ..modules.decoder_blocks import BasicDecBlk
|
||||
from ..modules.lateral_blocks import BasicLatBlk
|
||||
from ..refinement.stem_layer import StemLayer
|
||||
|
||||
|
||||
class RefinerPVTInChannels4(nn.Module):
|
||||
def __init__(self, in_channels=3+1):
|
||||
super(RefinerPVTInChannels4, self).__init__()
|
||||
self.config = Config()
|
||||
self.epoch = 1
|
||||
self.bb = build_backbone(self.config.bb, params_settings='in_channels=4')
|
||||
|
||||
lateral_channels_in_collection = {
|
||||
'vgg16': [512, 256, 128, 64], 'vgg16bn': [512, 256, 128, 64], 'resnet50': [1024, 512, 256, 64],
|
||||
'pvt_v2_b2': [512, 320, 128, 64], 'pvt_v2_b5': [512, 320, 128, 64],
|
||||
'swin_v1_b': [1024, 512, 256, 128], 'swin_v1_l': [1536, 768, 384, 192],
|
||||
}
|
||||
channels = lateral_channels_in_collection[self.config.bb]
|
||||
self.squeeze_module = BasicDecBlk(channels[0], channels[0])
|
||||
|
||||
self.decoder = Decoder(channels)
|
||||
|
||||
if 0:
|
||||
for key, value in self.named_parameters():
|
||||
if 'bb.' in key:
|
||||
value.requires_grad = False
|
||||
|
||||
def forward(self, x):
|
||||
if isinstance(x, list):
|
||||
x = torch.cat(x, dim=1)
|
||||
########## Encoder ##########
|
||||
if self.config.bb in ['vgg16', 'vgg16bn', 'resnet50']:
|
||||
x1 = self.bb.conv1(x)
|
||||
x2 = self.bb.conv2(x1)
|
||||
x3 = self.bb.conv3(x2)
|
||||
x4 = self.bb.conv4(x3)
|
||||
else:
|
||||
x1, x2, x3, x4 = self.bb(x)
|
||||
|
||||
x4 = self.squeeze_module(x4)
|
||||
|
||||
########## Decoder ##########
|
||||
|
||||
features = [x, x1, x2, x3, x4]
|
||||
scaled_preds = self.decoder(features)
|
||||
|
||||
return scaled_preds
|
||||
|
||||
|
||||
class Refiner(nn.Module):
|
||||
def __init__(self, in_channels=3+1):
|
||||
super(Refiner, self).__init__()
|
||||
self.config = Config()
|
||||
self.epoch = 1
|
||||
self.stem_layer = StemLayer(in_channels=in_channels, inter_channels=48, out_channels=3, norm_layer='BN' if self.config.batch_size > 1 else 'LN')
|
||||
self.bb = build_backbone(self.config.bb)
|
||||
|
||||
lateral_channels_in_collection = {
|
||||
'vgg16': [512, 256, 128, 64], 'vgg16bn': [512, 256, 128, 64], 'resnet50': [1024, 512, 256, 64],
|
||||
'pvt_v2_b2': [512, 320, 128, 64], 'pvt_v2_b5': [512, 320, 128, 64],
|
||||
'swin_v1_b': [1024, 512, 256, 128], 'swin_v1_l': [1536, 768, 384, 192],
|
||||
}
|
||||
channels = lateral_channels_in_collection[self.config.bb]
|
||||
self.squeeze_module = BasicDecBlk(channels[0], channels[0])
|
||||
|
||||
self.decoder = Decoder(channels)
|
||||
|
||||
if 0:
|
||||
for key, value in self.named_parameters():
|
||||
if 'bb.' in key:
|
||||
value.requires_grad = False
|
||||
|
||||
def forward(self, x):
|
||||
if isinstance(x, list):
|
||||
x = torch.cat(x, dim=1)
|
||||
x = self.stem_layer(x)
|
||||
########## Encoder ##########
|
||||
if self.config.bb in ['vgg16', 'vgg16bn', 'resnet50']:
|
||||
x1 = self.bb.conv1(x)
|
||||
x2 = self.bb.conv2(x1)
|
||||
x3 = self.bb.conv3(x2)
|
||||
x4 = self.bb.conv4(x3)
|
||||
else:
|
||||
x1, x2, x3, x4 = self.bb(x)
|
||||
|
||||
x4 = self.squeeze_module(x4)
|
||||
|
||||
########## Decoder ##########
|
||||
|
||||
features = [x, x1, x2, x3, x4]
|
||||
scaled_preds = self.decoder(features)
|
||||
|
||||
return scaled_preds
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(self, channels):
|
||||
super(Decoder, self).__init__()
|
||||
self.config = Config()
|
||||
DecoderBlock = eval('BasicDecBlk')
|
||||
LateralBlock = eval('BasicLatBlk')
|
||||
|
||||
self.decoder_block4 = DecoderBlock(channels[0], channels[1])
|
||||
self.decoder_block3 = DecoderBlock(channels[1], channels[2])
|
||||
self.decoder_block2 = DecoderBlock(channels[2], channels[3])
|
||||
self.decoder_block1 = DecoderBlock(channels[3], channels[3]//2)
|
||||
|
||||
self.lateral_block4 = LateralBlock(channels[1], channels[1])
|
||||
self.lateral_block3 = LateralBlock(channels[2], channels[2])
|
||||
self.lateral_block2 = LateralBlock(channels[3], channels[3])
|
||||
|
||||
if self.config.ms_supervision:
|
||||
self.conv_ms_spvn_4 = nn.Conv2d(channels[1], 1, 1, 1, 0)
|
||||
self.conv_ms_spvn_3 = nn.Conv2d(channels[2], 1, 1, 1, 0)
|
||||
self.conv_ms_spvn_2 = nn.Conv2d(channels[3], 1, 1, 1, 0)
|
||||
self.conv_out1 = nn.Sequential(nn.Conv2d(channels[3]//2, 1, 1, 1, 0))
|
||||
|
||||
def forward(self, features):
|
||||
x, x1, x2, x3, x4 = features
|
||||
outs = []
|
||||
p4 = self.decoder_block4(x4)
|
||||
_p4 = F.interpolate(p4, size=x3.shape[2:], mode='bilinear', align_corners=True)
|
||||
_p3 = _p4 + self.lateral_block4(x3)
|
||||
|
||||
p3 = self.decoder_block3(_p3)
|
||||
_p3 = F.interpolate(p3, size=x2.shape[2:], mode='bilinear', align_corners=True)
|
||||
_p2 = _p3 + self.lateral_block3(x2)
|
||||
|
||||
p2 = self.decoder_block2(_p2)
|
||||
_p2 = F.interpolate(p2, size=x1.shape[2:], mode='bilinear', align_corners=True)
|
||||
_p1 = _p2 + self.lateral_block2(x1)
|
||||
|
||||
_p1 = self.decoder_block1(_p1)
|
||||
_p1 = F.interpolate(_p1, size=x.shape[2:], mode='bilinear', align_corners=True)
|
||||
p1_out = self.conv_out1(_p1)
|
||||
|
||||
if self.config.ms_supervision:
|
||||
outs.append(self.conv_ms_spvn_4(p4))
|
||||
outs.append(self.conv_ms_spvn_3(p3))
|
||||
outs.append(self.conv_ms_spvn_2(p2))
|
||||
outs.append(p1_out)
|
||||
return outs
|
||||
|
||||
|
||||
class RefUNet(nn.Module):
|
||||
# Refinement
|
||||
def __init__(self, in_channels=3+1):
|
||||
super(RefUNet, self).__init__()
|
||||
self.encoder_1 = nn.Sequential(
|
||||
nn.Conv2d(in_channels, 64, 3, 1, 1),
|
||||
nn.Conv2d(64, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.encoder_2 = nn.Sequential(
|
||||
nn.MaxPool2d(2, 2, ceil_mode=True),
|
||||
nn.Conv2d(64, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.encoder_3 = nn.Sequential(
|
||||
nn.MaxPool2d(2, 2, ceil_mode=True),
|
||||
nn.Conv2d(64, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.encoder_4 = nn.Sequential(
|
||||
nn.MaxPool2d(2, 2, ceil_mode=True),
|
||||
nn.Conv2d(64, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.pool4 = nn.MaxPool2d(2, 2, ceil_mode=True)
|
||||
#####
|
||||
self.decoder_5 = nn.Sequential(
|
||||
nn.Conv2d(64, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
#####
|
||||
self.decoder_4 = nn.Sequential(
|
||||
nn.Conv2d(128, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.decoder_3 = nn.Sequential(
|
||||
nn.Conv2d(128, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.decoder_2 = nn.Sequential(
|
||||
nn.Conv2d(128, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.decoder_1 = nn.Sequential(
|
||||
nn.Conv2d(128, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.conv_d0 = nn.Conv2d(64, 1, 3, 1, 1)
|
||||
|
||||
self.upscore2 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
|
||||
|
||||
def forward(self, x):
|
||||
outs = []
|
||||
if isinstance(x, list):
|
||||
x = torch.cat(x, dim=1)
|
||||
hx = x
|
||||
|
||||
hx1 = self.encoder_1(hx)
|
||||
hx2 = self.encoder_2(hx1)
|
||||
hx3 = self.encoder_3(hx2)
|
||||
hx4 = self.encoder_4(hx3)
|
||||
|
||||
hx = self.decoder_5(self.pool4(hx4))
|
||||
hx = torch.cat((self.upscore2(hx), hx4), 1)
|
||||
|
||||
d4 = self.decoder_4(hx)
|
||||
hx = torch.cat((self.upscore2(d4), hx3), 1)
|
||||
|
||||
d3 = self.decoder_3(hx)
|
||||
hx = torch.cat((self.upscore2(d3), hx2), 1)
|
||||
|
||||
d2 = self.decoder_2(hx)
|
||||
hx = torch.cat((self.upscore2(d2), hx1), 1)
|
||||
|
||||
d1 = self.decoder_1(hx)
|
||||
|
||||
x = self.conv_d0(d1)
|
||||
outs.append(x)
|
||||
return outs
|
||||
@@ -0,0 +1,45 @@
|
||||
import torch.nn as nn
|
||||
from ..modules.utils import build_act_layer, build_norm_layer
|
||||
|
||||
|
||||
class StemLayer(nn.Module):
|
||||
r""" Stem layer of InternImage
|
||||
Args:
|
||||
in_channels (int): number of input channels
|
||||
out_channels (int): number of output channels
|
||||
act_layer (str): activation layer
|
||||
norm_layer (str): normalization layer
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
in_channels=3+1,
|
||||
inter_channels=48,
|
||||
out_channels=96,
|
||||
act_layer='GELU',
|
||||
norm_layer='BN'):
|
||||
super().__init__()
|
||||
self.conv1 = nn.Conv2d(in_channels,
|
||||
inter_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
self.norm1 = build_norm_layer(
|
||||
inter_channels, norm_layer, 'channels_first', 'channels_first'
|
||||
)
|
||||
self.act = build_act_layer(act_layer)
|
||||
self.conv2 = nn.Conv2d(inter_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
self.norm2 = build_norm_layer(
|
||||
out_channels, norm_layer, 'channels_first', 'channels_first'
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv1(x)
|
||||
x = self.norm1(x)
|
||||
x = self.act(x)
|
||||
x = self.conv2(x)
|
||||
x = self.norm2(x)
|
||||
return x
|
||||
+100
@@ -0,0 +1,100 @@
|
||||
import logging
|
||||
import os
|
||||
import torch
|
||||
from torchvision import transforms
|
||||
import numpy as np
|
||||
import random
|
||||
import cv2
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def path_to_image(path, size=(1024, 1024), color_type=['rgb', 'gray'][0]):
|
||||
if color_type.lower() == 'rgb':
|
||||
image = cv2.imread(path)
|
||||
elif color_type.lower() == 'gray':
|
||||
image = cv2.imread(path, cv2.IMREAD_GRAYSCALE)
|
||||
else:
|
||||
print('Select the color_type to return, either to RGB or gray image.')
|
||||
return
|
||||
if size:
|
||||
image = cv2.resize(image, size, interpolation=cv2.INTER_LINEAR)
|
||||
if color_type.lower() == 'rgb':
|
||||
image = Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGR2RGB)).convert('RGB')
|
||||
else:
|
||||
image = Image.fromarray(image).convert('L')
|
||||
return image
|
||||
|
||||
|
||||
|
||||
def check_state_dict(state_dict, unwanted_prefixes=['module.', '_orig_mod.']):
|
||||
for k, v in list(state_dict.items()):
|
||||
prefix_length = 0
|
||||
for unwanted_prefix in unwanted_prefixes:
|
||||
if k[prefix_length:].startswith(unwanted_prefix):
|
||||
prefix_length += len(unwanted_prefix)
|
||||
state_dict[k[prefix_length:]] = state_dict.pop(k)
|
||||
return state_dict
|
||||
|
||||
|
||||
def generate_smoothed_gt(gts):
|
||||
epsilon = 0.001
|
||||
new_gts = (1-epsilon)*gts+epsilon/2
|
||||
return new_gts
|
||||
|
||||
|
||||
class Logger():
|
||||
def __init__(self, path="log.txt"):
|
||||
self.logger = logging.getLogger('BiRefNet')
|
||||
self.file_handler = logging.FileHandler(path, "w")
|
||||
self.stdout_handler = logging.StreamHandler()
|
||||
self.stdout_handler.setFormatter(logging.Formatter('%(asctime)s %(levelname)s %(message)s'))
|
||||
self.file_handler.setFormatter(logging.Formatter('%(asctime)s %(levelname)s %(message)s'))
|
||||
self.logger.addHandler(self.file_handler)
|
||||
self.logger.addHandler(self.stdout_handler)
|
||||
self.logger.setLevel(logging.INFO)
|
||||
self.logger.propagate = False
|
||||
|
||||
def info(self, txt):
|
||||
self.logger.info(txt)
|
||||
|
||||
def close(self):
|
||||
self.file_handler.close()
|
||||
self.stdout_handler.close()
|
||||
|
||||
|
||||
class AverageMeter(object):
|
||||
"""Computes and stores the average and current value"""
|
||||
def __init__(self):
|
||||
self.reset()
|
||||
|
||||
def reset(self):
|
||||
self.val = 0.0
|
||||
self.avg = 0.0
|
||||
self.sum = 0.0
|
||||
self.count = 0.0
|
||||
|
||||
def update(self, val, n=1):
|
||||
self.val = val
|
||||
self.sum += val * n
|
||||
self.count += n
|
||||
self.avg = self.sum / self.count
|
||||
|
||||
|
||||
def save_checkpoint(state, path, filename="latest.pth"):
|
||||
torch.save(state, os.path.join(path, filename))
|
||||
|
||||
|
||||
def save_tensor_img(tenor_im, path):
|
||||
im = tenor_im.cpu().clone()
|
||||
im = im.squeeze(0)
|
||||
tensor2pil = transforms.ToPILImage()
|
||||
im = tensor2pil(im)
|
||||
im.save(path)
|
||||
|
||||
|
||||
def set_seed(seed):
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
np.random.seed(seed)
|
||||
random.seed(seed)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
+135
@@ -0,0 +1,135 @@
|
||||
import os
|
||||
import math
|
||||
import torch
|
||||
import folder_paths
|
||||
|
||||
|
||||
class Config:
|
||||
def __init__(self) -> None:
|
||||
self.ms_supervision = True
|
||||
self.out_ref = self.ms_supervision and True
|
||||
self.dec_ipt = True
|
||||
self.dec_ipt_split = True
|
||||
self.locate_head = False
|
||||
self.cxt_num = [0, 3][1] # multi-scale skip connections from encoder
|
||||
self.mul_scl_ipt = ['', 'add', 'cat'][2]
|
||||
self.refine = ['', 'itself', 'RefUNet', 'Refiner', 'RefinerPVTInChannels4'][0]
|
||||
self.progressive_ref = self.refine and True
|
||||
self.ender = self.progressive_ref and False
|
||||
self.scale = self.progressive_ref and 2
|
||||
self.dec_att = ['', 'ASPP', 'ASPPDeformable'][2]
|
||||
self.squeeze_block = ['', 'BasicDecBlk_x1', 'ResBlk_x4', 'ASPP_x3', 'ASPPDeformable_x3'][1]
|
||||
self.dec_blk = ['BasicDecBlk', 'ResBlk', 'HierarAttDecBlk'][0]
|
||||
self.auxiliary_classification = False
|
||||
self.refine_iteration = 1
|
||||
self.freeze_bb = False
|
||||
self.precisionHigh = True
|
||||
self.compile = True
|
||||
self.load_all = True
|
||||
self.verbose_eval = True
|
||||
|
||||
self.size = 1024
|
||||
self.batch_size = 2
|
||||
self.IoU_finetune_last_epochs = [0, -20][1] # choose 0 to skip
|
||||
if self.dec_blk == 'HierarAttDecBlk':
|
||||
self.batch_size = 2 ** [0, 1, 2, 3, 4][2]
|
||||
self.model = [
|
||||
'BiRefNet',
|
||||
][0]
|
||||
|
||||
# Components
|
||||
self.lat_blk = ['BasicLatBlk'][0]
|
||||
self.dec_channels_inter = ['fixed', 'adap'][0]
|
||||
|
||||
# Backbone
|
||||
self.bb = [
|
||||
'vgg16', 'vgg16bn', 'resnet50', # 0, 1, 2
|
||||
'pvt_v2_b2', 'pvt_v2_b5', # 3-bs10, 4-bs5
|
||||
'swin_v1_b', 'swin_v1_l', # 5-bs9, 6-bs6
|
||||
'swin_v1_t', 'swin_v1_s', # 7, 8
|
||||
'pvt_v2_b0', 'pvt_v2_b1', # 9, 10
|
||||
][6]
|
||||
self.lateral_channels_in_collection = {
|
||||
'vgg16': [512, 256, 128, 64], 'vgg16bn': [512, 256, 128, 64], 'resnet50': [1024, 512, 256, 64],
|
||||
'pvt_v2_b2': [512, 320, 128, 64], 'pvt_v2_b5': [512, 320, 128, 64],
|
||||
'swin_v1_b': [1024, 512, 256, 128], 'swin_v1_l': [1536, 768, 384, 192],
|
||||
'swin_v1_t': [768, 384, 192, 96], 'swin_v1_s': [768, 384, 192, 96],
|
||||
'pvt_v2_b0': [256, 160, 64, 32], 'pvt_v2_b1': [512, 320, 128, 64],
|
||||
}[self.bb]
|
||||
if self.mul_scl_ipt == 'cat':
|
||||
self.lateral_channels_in_collection = [channel * 2 for channel in self.lateral_channels_in_collection]
|
||||
self.cxt = self.lateral_channels_in_collection[1:][::-1][-self.cxt_num:] if self.cxt_num else []
|
||||
# self.sys_home_dir = '/root/autodl-tmp'
|
||||
# self.weights_root_dir = os.path.join(self.sys_home_dir, 'weights')
|
||||
# self.weights = {
|
||||
# 'pvt_v2_b2': os.path.join(self.weights_root_dir, 'pvt_v2_b2.pth'),
|
||||
# 'pvt_v2_b5': os.path.join(self.weights_root_dir, ['pvt_v2_b5.pth', 'pvt_v2_b5_22k.pth'][0]),
|
||||
# 'swin_v1_b': os.path.join(self.weights_root_dir, ['swin_base_patch4_window12_384_22kto1k.pth', 'swin_base_patch4_window12_384_22k.pth'][0]),
|
||||
# 'swin_v1_l': os.path.join(self.weights_root_dir, ['swin_large_patch4_window12_384_22kto1k.pth', 'swin_large_patch4_window12_384_22k.pth'][0]),
|
||||
# 'swin_v1_t': os.path.join(self.weights_root_dir, ['swin_tiny_patch4_window7_224_22kto1k_finetune.pth'][0]),
|
||||
# 'swin_v1_s': os.path.join(self.weights_root_dir, ['swin_small_patch4_window7_224_22kto1k_finetune.pth'][0]),
|
||||
# 'pvt_v2_b0': os.path.join(self.weights_root_dir, ['pvt_v2_b0.pth'][0]),
|
||||
# 'pvt_v2_b1': os.path.join(self.weights_root_dir, ['pvt_v2_b1.pth'][0]),
|
||||
# }
|
||||
weight_paths_name = "birefnet"
|
||||
self.weights = {
|
||||
'pvt_v2_b2': folder_paths.get_full_path(weight_paths_name, 'pvt_v2_b2.pth'),
|
||||
'pvt_v2_b5': folder_paths.get_full_path(weight_paths_name, ['pvt_v2_b5.pth', 'pvt_v2_b5_22k.pth'][0]),
|
||||
'swin_v1_b': folder_paths.get_full_path(weight_paths_name, ['swin_base_patch4_window12_384_22kto1k.pth', 'swin_base_patch4_window12_384_22k.pth'][0]),
|
||||
'swin_v1_l': folder_paths.get_full_path(weight_paths_name, ['swin_large_patch4_window12_384_22kto1k.pth', 'swin_large_patch4_window12_384_22k.pth'][0]),
|
||||
'swin_v1_t': folder_paths.get_full_path(weight_paths_name, ['swin_tiny_patch4_window7_224_22kto1k_finetune.pth'][0]),
|
||||
'swin_v1_s': folder_paths.get_full_path(weight_paths_name, ['swin_small_patch4_window7_224_22kto1k_finetune.pth'][0]),
|
||||
'pvt_v2_b0': folder_paths.get_full_path(weight_paths_name, ['pvt_v2_b0.pth'][0]),
|
||||
'pvt_v2_b1': folder_paths.get_full_path(weight_paths_name, ['pvt_v2_b1.pth'][0]),
|
||||
}
|
||||
|
||||
# Training
|
||||
self.num_workers = 5 # will be decrease to min(it, batch_size) at the initialization of the data_loader
|
||||
self.optimizer = ['Adam', 'AdamW'][0]
|
||||
self.lr = 1e-5 * math.sqrt(self.batch_size / 5) # adapt the lr linearly
|
||||
self.lr_decay_epochs = [1e4] # Set to negative N to decay the lr in the last N-th epoch.
|
||||
self.lr_decay_rate = 0.5
|
||||
self.only_S_MAE = False
|
||||
self.SDPA_enabled = False # Bug. Slower and errors occur in multi-GPUs
|
||||
|
||||
# Data
|
||||
# self.data_root_dir = os.path.join(self.sys_home_dir, 'datasets/dis')
|
||||
self.task = ['DIS5K', 'COD', 'HRSOD'][0]
|
||||
self.training_set = {
|
||||
'DIS5K': 'DIS-TR',
|
||||
'COD': 'TR-COD10K+TR-CAMO',
|
||||
'HRSOD': ['TR-DUTS', 'TR-HRSOD+TR-UHRSD', 'TR-DUTS+TR-HRSOD+TR-UHRSD'][1]
|
||||
}[self.task]
|
||||
self.preproc_methods = ['flip', 'enhance', 'rotate', 'pepper', 'crop'][:4]
|
||||
|
||||
# Loss
|
||||
self.lambdas_pix_last = {
|
||||
# not 0 means opening this loss
|
||||
# original rate -- 1 : 30 : 1.5 : 0.2, bce x 30
|
||||
'bce': 30 * 1, # high performance
|
||||
'iou': 0.5 * 1, # 0 / 255
|
||||
'iou_patch': 0.5 * 0, # 0 / 255, win_size = (64, 64)
|
||||
'mse': 150 * 0, # can smooth the saliency map
|
||||
'triplet': 3 * 0,
|
||||
'reg': 100 * 0,
|
||||
'ssim': 10 * 1, # help contours,
|
||||
'cnt': 5 * 0, # help contours
|
||||
}
|
||||
self.lambdas_cls = {
|
||||
'ce': 5.0
|
||||
}
|
||||
# Adv
|
||||
self.lambda_adv_g = 10. * 0 # turn to 0 to avoid adv training
|
||||
self.lambda_adv_d = 3. * (self.lambda_adv_g > 0)
|
||||
|
||||
# others
|
||||
self.device = [0, 'cpu'][0 if torch.cuda.is_available() else 1] # .to(0) == .to('cuda:0')
|
||||
|
||||
self.batch_size_valid = 1
|
||||
self.rand_seed = 7
|
||||
# run_sh_file = [f for f in os.listdir('.') if 'train.sh' == f] + [os.path.join('..', f) for f in os.listdir('..') if 'train.sh' == f]
|
||||
# with open(run_sh_file[0], 'r') as f:
|
||||
# lines = f.readlines()
|
||||
# self.save_last = int([l.strip() for l in lines if 'val_last=' in l][0].split('=')[-1])
|
||||
# self.save_step = int([l.strip() for l in lines if 'step=' in l][0].split('=')[-1])
|
||||
# self.val_step = [0, self.save_step][0]
|
||||
@@ -0,0 +1,94 @@
|
||||
import os
|
||||
# import cv2
|
||||
from tqdm import tqdm
|
||||
from PIL import Image
|
||||
from torch.utils import data
|
||||
from torchvision import transforms
|
||||
|
||||
from birefnet_old.preproc import preproc
|
||||
from birefnet_old.config import Config
|
||||
from birefnet_old.utils import path_to_image
|
||||
|
||||
|
||||
Image.MAX_IMAGE_PIXELS = None # remove DecompressionBombWarning
|
||||
config = Config()
|
||||
_class_labels_TR_sorted = 'Airplane, Ant, Antenna, Archery, Axe, BabyCarriage, Bag, BalanceBeam, Balcony, Balloon, Basket, BasketballHoop, Beatle, Bed, Bee, Bench, Bicycle, BicycleFrame, BicycleStand, Boat, Bonsai, BoomLift, Bridge, BunkBed, Butterfly, Button, Cable, CableLift, Cage, Camcorder, Cannon, Canoe, Car, CarParkDropArm, Carriage, Cart, Caterpillar, CeilingLamp, Centipede, Chair, Clip, Clock, Clothes, CoatHanger, Comb, ConcretePumpTruck, Crack, Crane, Cup, DentalChair, Desk, DeskChair, Diagram, DishRack, DoorHandle, Dragonfish, Dragonfly, Drum, Earphone, Easel, ElectricIron, Excavator, Eyeglasses, Fan, Fence, Fencing, FerrisWheel, FireExtinguisher, Fishing, Flag, FloorLamp, Forklift, GasStation, Gate, Gear, Goal, Golf, GymEquipment, Hammock, Handcart, Handcraft, Handrail, HangGlider, Harp, Harvester, Headset, Helicopter, Helmet, Hook, HorizontalBar, Hydrovalve, IroningTable, Jewelry, Key, KidsPlayground, Kitchenware, Kite, Knife, Ladder, LaundryRack, Lightning, Lobster, Locust, Machine, MachineGun, MagazineRack, Mantis, Medal, MemorialArchway, Microphone, Missile, MobileHolder, Monitor, Mosquito, Motorcycle, MovingTrolley, Mower, MusicPlayer, MusicStand, ObservationTower, Octopus, OilWell, OlympicLogo, OperatingTable, OutdoorFitnessEquipment, Parachute, Pavilion, Piano, Pipe, PlowHarrow, PoleVault, Punchbag, Rack, Racket, Rifle, Ring, Robot, RockClimbing, Rope, Sailboat, Satellite, Scaffold, Scale, Scissor, Scooter, Sculpture, Seadragon, Seahorse, Seal, SewingMachine, Ship, Shoe, ShoppingCart, ShoppingTrolley, Shower, Shrimp, Signboard, Skateboarding, Skeleton, Skiing, Spade, SpeedBoat, Spider, Spoon, Stair, Stand, Stationary, SteeringWheel, Stethoscope, Stool, Stove, StreetLamp, SweetStand, Swing, Sword, TV, Table, TableChair, TableLamp, TableTennis, Tank, Tapeline, Teapot, Telescope, Tent, TobaccoPipe, Toy, Tractor, TrafficLight, TrafficSign, Trampoline, TransmissionTower, Tree, Tricycle, TrimmerCover, Tripod, Trombone, Truck, Trumpet, Tuba, UAV, Umbrella, UnevenBars, UtilityPole, VacuumCleaner, Violin, Wakesurfing, Watch, WaterTower, WateringPot, Well, WellLid, Wheel, Wheelchair, WindTurbine, Windmill, WineGlass, WireWhisk, Yacht'
|
||||
class_labels_TR_sorted = _class_labels_TR_sorted.split(', ')
|
||||
|
||||
|
||||
class MyData(data.Dataset):
|
||||
def __init__(self, datasets, image_size, is_train=True):
|
||||
self.size_train = image_size
|
||||
self.size_test = image_size
|
||||
self.keep_size = not config.size
|
||||
self.data_size = (config.size, config.size)
|
||||
self.is_train = is_train
|
||||
self.load_all = config.load_all
|
||||
self.device = config.device
|
||||
if self.is_train and config.auxiliary_classification:
|
||||
self.cls_name2id = {_name: _id for _id, _name in enumerate(class_labels_TR_sorted)}
|
||||
self.transform_image = transforms.Compose([
|
||||
transforms.Resize(self.data_size),
|
||||
transforms.ToTensor(),
|
||||
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
|
||||
][self.load_all or self.keep_size:])
|
||||
self.transform_label = transforms.Compose([
|
||||
transforms.Resize(self.data_size),
|
||||
transforms.ToTensor(),
|
||||
][self.load_all or self.keep_size:])
|
||||
dataset_root = os.path.join(config.data_root_dir, config.task)
|
||||
# datasets can be a list of different datasets for training on combined sets.
|
||||
self.image_paths = []
|
||||
for dataset in datasets.split('+'):
|
||||
image_root = os.path.join(dataset_root, dataset, 'im')
|
||||
self.image_paths += [os.path.join(image_root, p) for p in os.listdir(image_root)]
|
||||
self.label_paths = []
|
||||
for p in self.image_paths:
|
||||
for ext in ['.png', '.jpg', '.PNG', '.JPG', '.JPEG']:
|
||||
## 'im' and 'gt' may need modifying
|
||||
p_gt = p.replace('/im/', '/gt/').replace('.'+p.split('.')[-1], ext)
|
||||
if os.path.exists(p_gt):
|
||||
self.label_paths.append(p_gt)
|
||||
break
|
||||
if self.load_all:
|
||||
self.images_loaded, self.labels_loaded = [], []
|
||||
self.class_labels_loaded = []
|
||||
# for image_path, label_path in zip(self.image_paths, self.label_paths):
|
||||
for image_path, label_path in tqdm(zip(self.image_paths, self.label_paths), total=len(self.image_paths)):
|
||||
_image = path_to_image(image_path, size=(config.size, config.size), color_type='rgb')
|
||||
_label = path_to_image(label_path, size=(config.size, config.size), color_type='gray')
|
||||
self.images_loaded.append(_image)
|
||||
self.labels_loaded.append(_label)
|
||||
self.class_labels_loaded.append(
|
||||
self.cls_name2id[label_path.split('/')[-1].split('#')[3]] if self.is_train and config.auxiliary_classification else -1
|
||||
)
|
||||
|
||||
|
||||
def __getitem__(self, index):
|
||||
|
||||
if self.load_all:
|
||||
image = self.images_loaded[index]
|
||||
label = self.labels_loaded[index]
|
||||
class_label = self.class_labels_loaded[index] if self.is_train and config.auxiliary_classification else -1
|
||||
else:
|
||||
image = path_to_image(self.image_paths[index], size=(config.size, config.size), color_type='rgb')
|
||||
label = path_to_image(self.label_paths[index], size=(config.size, config.size), color_type='gray')
|
||||
class_label = self.cls_name2id[self.label_paths[index].split('/')[-1].split('#')[3]] if self.is_train and config.auxiliary_classification else -1
|
||||
|
||||
# loading image and label
|
||||
if self.is_train:
|
||||
image, label = preproc(image, label, preproc_methods=config.preproc_methods)
|
||||
# else:
|
||||
# if _label.shape[0] > 2048 or _label.shape[1] > 2048:
|
||||
# _image = cv2.resize(_image, (2048, 2048), interpolation=cv2.INTER_LINEAR)
|
||||
# _label = cv2.resize(_label, (2048, 2048), interpolation=cv2.INTER_LINEAR)
|
||||
|
||||
image, label = self.transform_image(image), self.transform_label(label)
|
||||
|
||||
if self.is_train:
|
||||
return image, label, class_label
|
||||
else:
|
||||
return image, label, self.label_paths[index]
|
||||
|
||||
def __len__(self):
|
||||
return len(self.image_paths)
|
||||
@@ -0,0 +1,44 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from collections import OrderedDict
|
||||
from torchvision.models import vgg16, vgg16_bn, VGG16_Weights, VGG16_BN_Weights, resnet50, ResNet50_Weights
|
||||
from ..backbones.pvt_v2 import pvt_v2_b0, pvt_v2_b1, pvt_v2_b2, pvt_v2_b5
|
||||
from ..backbones.swin_v1 import swin_v1_t, swin_v1_s, swin_v1_b, swin_v1_l
|
||||
from ...config import Config
|
||||
|
||||
|
||||
config = Config()
|
||||
|
||||
def build_backbone(bb_name, pretrained=True, params_settings=''):
|
||||
if bb_name == 'vgg16':
|
||||
bb_net = list(vgg16(pretrained=VGG16_Weights.DEFAULT if pretrained else None).children())[0]
|
||||
bb = nn.Sequential(OrderedDict({'conv1': bb_net[:4], 'conv2': bb_net[4:9], 'conv3': bb_net[9:16], 'conv4': bb_net[16:23]}))
|
||||
elif bb_name == 'vgg16bn':
|
||||
bb_net = list(vgg16_bn(pretrained=VGG16_BN_Weights.DEFAULT if pretrained else None).children())[0]
|
||||
bb = nn.Sequential(OrderedDict({'conv1': bb_net[:6], 'conv2': bb_net[6:13], 'conv3': bb_net[13:23], 'conv4': bb_net[23:33]}))
|
||||
elif bb_name == 'resnet50':
|
||||
bb_net = list(resnet50(pretrained=ResNet50_Weights.DEFAULT if pretrained else None).children())
|
||||
bb = nn.Sequential(OrderedDict({'conv1': nn.Sequential(*bb_net[0:3]), 'conv2': bb_net[4], 'conv3': bb_net[5], 'conv4': bb_net[6]}))
|
||||
else:
|
||||
bb = eval('{}({})'.format(bb_name, params_settings))
|
||||
if pretrained:
|
||||
bb = load_weights(bb, bb_name)
|
||||
return bb
|
||||
|
||||
def load_weights(model, model_name):
|
||||
save_model = torch.load(config.weights[model_name])
|
||||
model_dict = model.state_dict()
|
||||
state_dict = {k: v if v.size() == model_dict[k].size() else model_dict[k] for k, v in save_model.items() if k in model_dict.keys()}
|
||||
# to ignore the weights with mismatched size when I modify the backbone itself.
|
||||
if not state_dict:
|
||||
save_model_keys = list(save_model.keys())
|
||||
sub_item = save_model_keys[0] if len(save_model_keys) == 1 else None
|
||||
state_dict = {k: v if v.size() == model_dict[k].size() else model_dict[k] for k, v in save_model[sub_item].items() if k in model_dict.keys()}
|
||||
if not state_dict or not sub_item:
|
||||
print('Weights are not successully loaded. Check the state dict of weights file.')
|
||||
return None
|
||||
else:
|
||||
print('Found correct weights in the "{}" item of loaded state_dict.'.format(sub_item))
|
||||
model_dict.update(state_dict)
|
||||
model.load_state_dict(model_dict)
|
||||
return model
|
||||
@@ -0,0 +1,438 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from functools import partial
|
||||
|
||||
try:
|
||||
# version > 0.6.13
|
||||
from timm.layers import DropPath, to_2tuple, trunc_normal_
|
||||
except Exception:
|
||||
from timm.models.layers import DropPath, to_2tuple, trunc_normal_
|
||||
|
||||
import math
|
||||
|
||||
from ...config import Config
|
||||
|
||||
config = Config()
|
||||
|
||||
class Mlp(nn.Module):
|
||||
def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
|
||||
super().__init__()
|
||||
out_features = out_features or in_features
|
||||
hidden_features = hidden_features or in_features
|
||||
self.fc1 = nn.Linear(in_features, hidden_features)
|
||||
self.dwconv = DWConv(hidden_features)
|
||||
self.act = act_layer()
|
||||
self.fc2 = nn.Linear(hidden_features, out_features)
|
||||
self.drop = nn.Dropout(drop)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x, H, W):
|
||||
x = self.fc1(x)
|
||||
x = self.dwconv(x, H, W)
|
||||
x = self.act(x)
|
||||
x = self.drop(x)
|
||||
x = self.fc2(x)
|
||||
x = self.drop(x)
|
||||
return x
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0., sr_ratio=1):
|
||||
super().__init__()
|
||||
assert dim % num_heads == 0, f"dim {dim} should be divided by num_heads {num_heads}."
|
||||
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim // num_heads
|
||||
self.scale = qk_scale or head_dim ** -0.5
|
||||
|
||||
self.q = nn.Linear(dim, dim, bias=qkv_bias)
|
||||
self.kv = nn.Linear(dim, dim * 2, bias=qkv_bias)
|
||||
self.attn_drop_prob = attn_drop
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
|
||||
self.sr_ratio = sr_ratio
|
||||
if sr_ratio > 1:
|
||||
self.sr = nn.Conv2d(dim, dim, kernel_size=sr_ratio, stride=sr_ratio)
|
||||
self.norm = nn.LayerNorm(dim)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x, H, W):
|
||||
B, N, C = x.shape
|
||||
q = self.q(x).reshape(B, N, self.num_heads, C // self.num_heads).permute(0, 2, 1, 3)
|
||||
|
||||
if self.sr_ratio > 1:
|
||||
x_ = x.permute(0, 2, 1).reshape(B, C, H, W)
|
||||
x_ = self.sr(x_).reshape(B, C, -1).permute(0, 2, 1)
|
||||
x_ = self.norm(x_)
|
||||
kv = self.kv(x_).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
|
||||
else:
|
||||
kv = self.kv(x).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
|
||||
k, v = kv[0], kv[1]
|
||||
|
||||
if config.SDPA_enabled:
|
||||
x = torch.nn.functional.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
attn_mask=None, dropout_p=self.attn_drop_prob, is_causal=False
|
||||
).transpose(1, 2).reshape(B, N, C)
|
||||
else:
|
||||
attn = (q @ k.transpose(-2, -1)) * self.scale
|
||||
attn = attn.softmax(dim=-1)
|
||||
attn = self.attn_drop(attn)
|
||||
|
||||
x = (attn @ v).transpose(1, 2).reshape(B, N, C)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class Block(nn.Module):
|
||||
|
||||
def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,
|
||||
drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm, sr_ratio=1):
|
||||
super().__init__()
|
||||
self.norm1 = norm_layer(dim)
|
||||
self.attn = Attention(
|
||||
dim,
|
||||
num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale,
|
||||
attn_drop=attn_drop, proj_drop=drop, sr_ratio=sr_ratio)
|
||||
# NOTE: drop path for stochastic depth, we shall see if this is better than dropout here
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
|
||||
self.norm2 = norm_layer(dim)
|
||||
mlp_hidden_dim = int(dim * mlp_ratio)
|
||||
self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x, H, W):
|
||||
x = x + self.drop_path(self.attn(self.norm1(x), H, W))
|
||||
x = x + self.drop_path(self.mlp(self.norm2(x), H, W))
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class OverlapPatchEmbed(nn.Module):
|
||||
""" Image to Patch Embedding
|
||||
"""
|
||||
|
||||
def __init__(self, img_size=224, patch_size=7, stride=4, in_channels=3, embed_dim=768):
|
||||
super().__init__()
|
||||
img_size = to_2tuple(img_size)
|
||||
patch_size = to_2tuple(patch_size)
|
||||
|
||||
self.img_size = img_size
|
||||
self.patch_size = patch_size
|
||||
self.H, self.W = img_size[0] // patch_size[0], img_size[1] // patch_size[1]
|
||||
self.num_patches = self.H * self.W
|
||||
self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=stride,
|
||||
padding=(patch_size[0] // 2, patch_size[1] // 2))
|
||||
self.norm = nn.LayerNorm(embed_dim)
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.proj(x)
|
||||
_, _, H, W = x.shape
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
x = self.norm(x)
|
||||
|
||||
return x, H, W
|
||||
|
||||
|
||||
class PyramidVisionTransformerImpr(nn.Module):
|
||||
def __init__(self, img_size=224, patch_size=16, in_channels=3, num_classes=1000, embed_dims=[64, 128, 256, 512],
|
||||
num_heads=[1, 2, 4, 8], mlp_ratios=[4, 4, 4, 4], qkv_bias=False, qk_scale=None, drop_rate=0.,
|
||||
attn_drop_rate=0., drop_path_rate=0., norm_layer=nn.LayerNorm,
|
||||
depths=[3, 4, 6, 3], sr_ratios=[8, 4, 2, 1]):
|
||||
super().__init__()
|
||||
self.num_classes = num_classes
|
||||
self.depths = depths
|
||||
|
||||
# patch_embed
|
||||
self.patch_embed1 = OverlapPatchEmbed(img_size=img_size, patch_size=7, stride=4, in_channels=in_channels,
|
||||
embed_dim=embed_dims[0])
|
||||
self.patch_embed2 = OverlapPatchEmbed(img_size=img_size // 4, patch_size=3, stride=2, in_channels=embed_dims[0],
|
||||
embed_dim=embed_dims[1])
|
||||
self.patch_embed3 = OverlapPatchEmbed(img_size=img_size // 8, patch_size=3, stride=2, in_channels=embed_dims[1],
|
||||
embed_dim=embed_dims[2])
|
||||
self.patch_embed4 = OverlapPatchEmbed(img_size=img_size // 16, patch_size=3, stride=2, in_channels=embed_dims[2],
|
||||
embed_dim=embed_dims[3])
|
||||
|
||||
# transformer encoder
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule
|
||||
cur = 0
|
||||
self.block1 = nn.ModuleList([Block(
|
||||
dim=embed_dims[0], num_heads=num_heads[0], mlp_ratio=mlp_ratios[0], qkv_bias=qkv_bias, qk_scale=qk_scale,
|
||||
drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
|
||||
sr_ratio=sr_ratios[0])
|
||||
for i in range(depths[0])])
|
||||
self.norm1 = norm_layer(embed_dims[0])
|
||||
|
||||
cur += depths[0]
|
||||
self.block2 = nn.ModuleList([Block(
|
||||
dim=embed_dims[1], num_heads=num_heads[1], mlp_ratio=mlp_ratios[1], qkv_bias=qkv_bias, qk_scale=qk_scale,
|
||||
drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
|
||||
sr_ratio=sr_ratios[1])
|
||||
for i in range(depths[1])])
|
||||
self.norm2 = norm_layer(embed_dims[1])
|
||||
|
||||
cur += depths[1]
|
||||
self.block3 = nn.ModuleList([Block(
|
||||
dim=embed_dims[2], num_heads=num_heads[2], mlp_ratio=mlp_ratios[2], qkv_bias=qkv_bias, qk_scale=qk_scale,
|
||||
drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
|
||||
sr_ratio=sr_ratios[2])
|
||||
for i in range(depths[2])])
|
||||
self.norm3 = norm_layer(embed_dims[2])
|
||||
|
||||
cur += depths[2]
|
||||
self.block4 = nn.ModuleList([Block(
|
||||
dim=embed_dims[3], num_heads=num_heads[3], mlp_ratio=mlp_ratios[3], qkv_bias=qkv_bias, qk_scale=qk_scale,
|
||||
drop=drop_rate, attn_drop=attn_drop_rate, drop_path=dpr[cur + i], norm_layer=norm_layer,
|
||||
sr_ratio=sr_ratios[3])
|
||||
for i in range(depths[3])])
|
||||
self.norm4 = norm_layer(embed_dims[3])
|
||||
|
||||
# classification head
|
||||
# self.head = nn.Linear(embed_dims[3], num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
self.apply(self._init_weights)
|
||||
|
||||
def _init_weights(self, m):
|
||||
if isinstance(m, nn.Linear):
|
||||
trunc_normal_(m.weight, std=.02)
|
||||
if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
nn.init.constant_(m.bias, 0)
|
||||
elif isinstance(m, nn.LayerNorm):
|
||||
nn.init.constant_(m.bias, 0)
|
||||
nn.init.constant_(m.weight, 1.0)
|
||||
elif isinstance(m, nn.Conv2d):
|
||||
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
|
||||
fan_out //= m.groups
|
||||
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
|
||||
if m.bias is not None:
|
||||
m.bias.data.zero_()
|
||||
|
||||
def init_weights(self, pretrained=None):
|
||||
if isinstance(pretrained, str):
|
||||
logger = 1
|
||||
#load_checkpoint(self, pretrained, map_location='cpu', strict=False, logger=logger)
|
||||
|
||||
def reset_drop_path(self, drop_path_rate):
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(self.depths))]
|
||||
cur = 0
|
||||
for i in range(self.depths[0]):
|
||||
self.block1[i].drop_path.drop_prob = dpr[cur + i]
|
||||
|
||||
cur += self.depths[0]
|
||||
for i in range(self.depths[1]):
|
||||
self.block2[i].drop_path.drop_prob = dpr[cur + i]
|
||||
|
||||
cur += self.depths[1]
|
||||
for i in range(self.depths[2]):
|
||||
self.block3[i].drop_path.drop_prob = dpr[cur + i]
|
||||
|
||||
cur += self.depths[2]
|
||||
for i in range(self.depths[3]):
|
||||
self.block4[i].drop_path.drop_prob = dpr[cur + i]
|
||||
|
||||
def freeze_patch_emb(self):
|
||||
self.patch_embed1.requires_grad = False
|
||||
|
||||
@torch.jit.ignore
|
||||
def no_weight_decay(self):
|
||||
return {'pos_embed1', 'pos_embed2', 'pos_embed3', 'pos_embed4', 'cls_token'} # has pos_embed may be better
|
||||
|
||||
def get_classifier(self):
|
||||
return self.head
|
||||
|
||||
def reset_classifier(self, num_classes, global_pool=''):
|
||||
self.num_classes = num_classes
|
||||
self.head = nn.Linear(self.embed_dim, num_classes) if num_classes > 0 else nn.Identity()
|
||||
|
||||
def forward_features(self, x):
|
||||
B = x.shape[0]
|
||||
outs = []
|
||||
|
||||
# stage 1
|
||||
x, H, W = self.patch_embed1(x)
|
||||
for i, blk in enumerate(self.block1):
|
||||
x = blk(x, H, W)
|
||||
x = self.norm1(x)
|
||||
x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
|
||||
outs.append(x)
|
||||
|
||||
# stage 2
|
||||
x, H, W = self.patch_embed2(x)
|
||||
for i, blk in enumerate(self.block2):
|
||||
x = blk(x, H, W)
|
||||
x = self.norm2(x)
|
||||
x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
|
||||
outs.append(x)
|
||||
|
||||
# stage 3
|
||||
x, H, W = self.patch_embed3(x)
|
||||
for i, blk in enumerate(self.block3):
|
||||
x = blk(x, H, W)
|
||||
x = self.norm3(x)
|
||||
x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
|
||||
outs.append(x)
|
||||
|
||||
# stage 4
|
||||
x, H, W = self.patch_embed4(x)
|
||||
for i, blk in enumerate(self.block4):
|
||||
x = blk(x, H, W)
|
||||
x = self.norm4(x)
|
||||
x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
|
||||
outs.append(x)
|
||||
|
||||
return outs
|
||||
|
||||
# return x.mean(dim=1)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.forward_features(x)
|
||||
# x = self.head(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class DWConv(nn.Module):
|
||||
def __init__(self, dim=768):
|
||||
super(DWConv, self).__init__()
|
||||
self.dwconv = nn.Conv2d(dim, dim, 3, 1, 1, bias=True, groups=dim)
|
||||
|
||||
def forward(self, x, H, W):
|
||||
B, N, C = x.shape
|
||||
x = x.transpose(1, 2).view(B, C, H, W).contiguous()
|
||||
x = self.dwconv(x)
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
def _conv_filter(state_dict, patch_size=16):
|
||||
""" convert patch embedding weight from manual patchify + linear proj to conv"""
|
||||
out_dict = {}
|
||||
for k, v in state_dict.items():
|
||||
if 'patch_embed.proj.weight' in k:
|
||||
v = v.reshape((v.shape[0], 3, patch_size, patch_size))
|
||||
out_dict[k] = v
|
||||
|
||||
return out_dict
|
||||
|
||||
|
||||
## @register_model
|
||||
class pvt_v2_b0(PyramidVisionTransformerImpr):
|
||||
def __init__(self, **kwargs):
|
||||
super(pvt_v2_b0, self).__init__(
|
||||
patch_size=4, embed_dims=[32, 64, 160, 256], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
|
||||
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[2, 2, 2, 2], sr_ratios=[8, 4, 2, 1],
|
||||
drop_rate=0.0, drop_path_rate=0.1)
|
||||
|
||||
|
||||
|
||||
## @register_model
|
||||
class pvt_v2_b1(PyramidVisionTransformerImpr):
|
||||
def __init__(self, **kwargs):
|
||||
super(pvt_v2_b1, self).__init__(
|
||||
patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
|
||||
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[2, 2, 2, 2], sr_ratios=[8, 4, 2, 1],
|
||||
drop_rate=0.0, drop_path_rate=0.1)
|
||||
|
||||
## @register_model
|
||||
class pvt_v2_b2(PyramidVisionTransformerImpr):
|
||||
def __init__(self, in_channels=3, **kwargs):
|
||||
super(pvt_v2_b2, self).__init__(
|
||||
patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
|
||||
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 4, 6, 3], sr_ratios=[8, 4, 2, 1],
|
||||
drop_rate=0.0, drop_path_rate=0.1, in_channels=in_channels)
|
||||
|
||||
## @register_model
|
||||
class pvt_v2_b3(PyramidVisionTransformerImpr):
|
||||
def __init__(self, **kwargs):
|
||||
super(pvt_v2_b3, self).__init__(
|
||||
patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
|
||||
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 4, 18, 3], sr_ratios=[8, 4, 2, 1],
|
||||
drop_rate=0.0, drop_path_rate=0.1)
|
||||
|
||||
## @register_model
|
||||
class pvt_v2_b4(PyramidVisionTransformerImpr):
|
||||
def __init__(self, **kwargs):
|
||||
super(pvt_v2_b4, self).__init__(
|
||||
patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[8, 8, 4, 4],
|
||||
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 8, 27, 3], sr_ratios=[8, 4, 2, 1],
|
||||
drop_rate=0.0, drop_path_rate=0.1)
|
||||
|
||||
|
||||
## @register_model
|
||||
class pvt_v2_b5(PyramidVisionTransformerImpr):
|
||||
def __init__(self, **kwargs):
|
||||
super(pvt_v2_b5, self).__init__(
|
||||
patch_size=4, embed_dims=[64, 128, 320, 512], num_heads=[1, 2, 5, 8], mlp_ratios=[4, 4, 4, 4],
|
||||
qkv_bias=True, norm_layer=partial(nn.LayerNorm, eps=1e-6), depths=[3, 6, 40, 3], sr_ratios=[8, 4, 2, 1],
|
||||
drop_rate=0.0, drop_path_rate=0.1)
|
||||
@@ -0,0 +1,657 @@
|
||||
# --------------------------------------------------------
|
||||
# Swin Transformer
|
||||
# Copyright (c) 2021 Microsoft
|
||||
# Licensed under The MIT License [see LICENSE for details]
|
||||
# Written by Ze Liu, Yutong Lin, Yixuan Wei
|
||||
# --------------------------------------------------------
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.utils.checkpoint as checkpoint
|
||||
import numpy as np
|
||||
try:
|
||||
# version > 0.6.13
|
||||
from timm.layers import DropPath, to_2tuple, trunc_normal_
|
||||
except Exception:
|
||||
from timm.models.layers import DropPath, to_2tuple, trunc_normal_
|
||||
|
||||
from ...config import Config
|
||||
|
||||
|
||||
config = Config()
|
||||
|
||||
class Mlp(nn.Module):
|
||||
""" Multilayer perceptron."""
|
||||
|
||||
def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
|
||||
super().__init__()
|
||||
out_features = out_features or in_features
|
||||
hidden_features = hidden_features or in_features
|
||||
self.fc1 = nn.Linear(in_features, hidden_features)
|
||||
self.act = act_layer()
|
||||
self.fc2 = nn.Linear(hidden_features, out_features)
|
||||
self.drop = nn.Dropout(drop)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.fc1(x)
|
||||
x = self.act(x)
|
||||
x = self.drop(x)
|
||||
x = self.fc2(x)
|
||||
x = self.drop(x)
|
||||
return x
|
||||
|
||||
|
||||
def window_partition(x, window_size):
|
||||
"""
|
||||
Args:
|
||||
x: (B, H, W, C)
|
||||
window_size (int): window size
|
||||
|
||||
Returns:
|
||||
windows: (num_windows*B, window_size, window_size, C)
|
||||
"""
|
||||
B, H, W, C = x.shape
|
||||
x = x.view(B, H // window_size, window_size, W // window_size, window_size, C)
|
||||
windows = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(-1, window_size, window_size, C)
|
||||
return windows
|
||||
|
||||
|
||||
def window_reverse(windows, window_size, H, W):
|
||||
"""
|
||||
Args:
|
||||
windows: (num_windows*B, window_size, window_size, C)
|
||||
window_size (int): Window size
|
||||
H (int): Height of image
|
||||
W (int): Width of image
|
||||
|
||||
Returns:
|
||||
x: (B, H, W, C)
|
||||
"""
|
||||
B = int(windows.shape[0] / (H * W / window_size / window_size))
|
||||
x = windows.view(B, H // window_size, W // window_size, window_size, window_size, -1)
|
||||
x = x.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, H, W, -1)
|
||||
return x
|
||||
|
||||
|
||||
class WindowAttention(nn.Module):
|
||||
""" Window based multi-head self attention (W-MSA) module with relative position bias.
|
||||
It supports both of shifted and non-shifted window.
|
||||
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
window_size (tuple[int]): The height and width of the window.
|
||||
num_heads (int): Number of attention heads.
|
||||
qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
|
||||
qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set
|
||||
attn_drop (float, optional): Dropout ratio of attention weight. Default: 0.0
|
||||
proj_drop (float, optional): Dropout ratio of output. Default: 0.0
|
||||
"""
|
||||
|
||||
def __init__(self, dim, window_size, num_heads, qkv_bias=True, qk_scale=None, attn_drop=0., proj_drop=0.):
|
||||
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.window_size = window_size # Wh, Ww
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim // num_heads
|
||||
self.scale = qk_scale or head_dim ** -0.5
|
||||
|
||||
# define a parameter table of relative position bias
|
||||
self.relative_position_bias_table = nn.Parameter(
|
||||
torch.zeros((2 * window_size[0] - 1) * (2 * window_size[1] - 1), num_heads)) # 2*Wh-1 * 2*Ww-1, nH
|
||||
|
||||
# get pair-wise relative position index for each token inside the window
|
||||
coords_h = torch.arange(self.window_size[0])
|
||||
coords_w = torch.arange(self.window_size[1])
|
||||
coords = torch.stack(torch.meshgrid([coords_h, coords_w], indexing='ij')) # 2, Wh, Ww
|
||||
coords_flatten = torch.flatten(coords, 1) # 2, Wh*Ww
|
||||
relative_coords = coords_flatten[:, :, None] - coords_flatten[:, None, :] # 2, Wh*Ww, Wh*Ww
|
||||
relative_coords = relative_coords.permute(1, 2, 0).contiguous() # Wh*Ww, Wh*Ww, 2
|
||||
relative_coords[:, :, 0] += self.window_size[0] - 1 # shift to start from 0
|
||||
relative_coords[:, :, 1] += self.window_size[1] - 1
|
||||
relative_coords[:, :, 0] *= 2 * self.window_size[1] - 1
|
||||
relative_position_index = relative_coords.sum(-1) # Wh*Ww, Wh*Ww
|
||||
self.register_buffer("relative_position_index", relative_position_index)
|
||||
|
||||
self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)
|
||||
self.attn_drop_prob = attn_drop
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
|
||||
trunc_normal_(self.relative_position_bias_table, std=.02)
|
||||
self.softmax = nn.Softmax(dim=-1)
|
||||
|
||||
def forward(self, x, mask=None):
|
||||
""" Forward function.
|
||||
|
||||
Args:
|
||||
x: input features with shape of (num_windows*B, N, C)
|
||||
mask: (0/-inf) mask with shape of (num_windows, Wh*Ww, Wh*Ww) or None
|
||||
"""
|
||||
B_, N, C = x.shape
|
||||
qkv = self.qkv(x).reshape(B_, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
|
||||
q, k, v = qkv[0], qkv[1], qkv[2] # make torchscript happy (cannot use tensor as tuple)
|
||||
|
||||
q = q * self.scale
|
||||
|
||||
if config.SDPA_enabled:
|
||||
x = torch.nn.functional.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
attn_mask=None, dropout_p=self.attn_drop_prob, is_causal=False
|
||||
).transpose(1, 2).reshape(B_, N, C)
|
||||
else:
|
||||
attn = (q @ k.transpose(-2, -1))
|
||||
|
||||
relative_position_bias = self.relative_position_bias_table[self.relative_position_index.view(-1)].view(
|
||||
self.window_size[0] * self.window_size[1], self.window_size[0] * self.window_size[1], -1) # Wh*Ww,Wh*Ww,nH
|
||||
relative_position_bias = relative_position_bias.permute(2, 0, 1).contiguous() # nH, Wh*Ww, Wh*Ww
|
||||
attn = attn + relative_position_bias.unsqueeze(0)
|
||||
|
||||
if mask is not None:
|
||||
nW = mask.shape[0]
|
||||
attn = attn.view(B_ // nW, nW, self.num_heads, N, N) + mask.unsqueeze(1).unsqueeze(0)
|
||||
attn = attn.view(-1, self.num_heads, N, N)
|
||||
attn = self.softmax(attn)
|
||||
else:
|
||||
attn = self.softmax(attn)
|
||||
|
||||
attn = self.attn_drop(attn)
|
||||
|
||||
x = (attn @ v).transpose(1, 2).reshape(B_, N, C)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
return x
|
||||
|
||||
|
||||
class SwinTransformerBlock(nn.Module):
|
||||
""" Swin Transformer Block.
|
||||
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
num_heads (int): Number of attention heads.
|
||||
window_size (int): Window size.
|
||||
shift_size (int): Shift size for SW-MSA.
|
||||
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.
|
||||
qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
|
||||
qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
|
||||
drop (float, optional): Dropout rate. Default: 0.0
|
||||
attn_drop (float, optional): Attention dropout rate. Default: 0.0
|
||||
drop_path (float, optional): Stochastic depth rate. Default: 0.0
|
||||
act_layer (nn.Module, optional): Activation layer. Default: nn.GELU
|
||||
norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
|
||||
"""
|
||||
|
||||
def __init__(self, dim, num_heads, window_size=7, shift_size=0,
|
||||
mlp_ratio=4., qkv_bias=True, qk_scale=None, drop=0., attn_drop=0., drop_path=0.,
|
||||
act_layer=nn.GELU, norm_layer=nn.LayerNorm):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.window_size = window_size
|
||||
self.shift_size = shift_size
|
||||
self.mlp_ratio = mlp_ratio
|
||||
assert 0 <= self.shift_size < self.window_size, "shift_size must in 0-window_size"
|
||||
|
||||
self.norm1 = norm_layer(dim)
|
||||
self.attn = WindowAttention(
|
||||
dim, window_size=to_2tuple(self.window_size), num_heads=num_heads,
|
||||
qkv_bias=qkv_bias, qk_scale=qk_scale, attn_drop=attn_drop, proj_drop=drop)
|
||||
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
|
||||
self.norm2 = norm_layer(dim)
|
||||
mlp_hidden_dim = int(dim * mlp_ratio)
|
||||
self.mlp = Mlp(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
|
||||
|
||||
self.H = None
|
||||
self.W = None
|
||||
|
||||
def forward(self, x, mask_matrix):
|
||||
""" Forward function.
|
||||
|
||||
Args:
|
||||
x: Input feature, tensor size (B, H*W, C).
|
||||
H, W: Spatial resolution of the input feature.
|
||||
mask_matrix: Attention mask for cyclic shift.
|
||||
"""
|
||||
B, L, C = x.shape
|
||||
H, W = self.H, self.W
|
||||
assert L == H * W, "input feature has wrong size"
|
||||
|
||||
shortcut = x
|
||||
x = self.norm1(x)
|
||||
x = x.view(B, H, W, C)
|
||||
|
||||
# pad feature maps to multiples of window size
|
||||
pad_l = pad_t = 0
|
||||
pad_r = (self.window_size - W % self.window_size) % self.window_size
|
||||
pad_b = (self.window_size - H % self.window_size) % self.window_size
|
||||
x = F.pad(x, (0, 0, pad_l, pad_r, pad_t, pad_b))
|
||||
_, Hp, Wp, _ = x.shape
|
||||
|
||||
# cyclic shift
|
||||
if self.shift_size > 0:
|
||||
shifted_x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2))
|
||||
attn_mask = mask_matrix
|
||||
else:
|
||||
shifted_x = x
|
||||
attn_mask = None
|
||||
|
||||
# partition windows
|
||||
x_windows = window_partition(shifted_x, self.window_size) # nW*B, window_size, window_size, C
|
||||
x_windows = x_windows.view(-1, self.window_size * self.window_size, C) # nW*B, window_size*window_size, C
|
||||
|
||||
# W-MSA/SW-MSA
|
||||
attn_windows = self.attn(x_windows, mask=attn_mask) # nW*B, window_size*window_size, C
|
||||
|
||||
# merge windows
|
||||
attn_windows = attn_windows.view(-1, self.window_size, self.window_size, C)
|
||||
shifted_x = window_reverse(attn_windows, self.window_size, Hp, Wp) # B H' W' C
|
||||
|
||||
# reverse cyclic shift
|
||||
if self.shift_size > 0:
|
||||
x = torch.roll(shifted_x, shifts=(self.shift_size, self.shift_size), dims=(1, 2))
|
||||
else:
|
||||
x = shifted_x
|
||||
|
||||
if pad_r > 0 or pad_b > 0:
|
||||
x = x[:, :H, :W, :].contiguous()
|
||||
|
||||
x = x.view(B, H * W, C)
|
||||
|
||||
# FFN
|
||||
x = shortcut + self.drop_path(x)
|
||||
x = x + self.drop_path(self.mlp(self.norm2(x)))
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class PatchMerging(nn.Module):
|
||||
""" Patch Merging Layer
|
||||
|
||||
Args:
|
||||
dim (int): Number of input channels.
|
||||
norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
|
||||
"""
|
||||
def __init__(self, dim, norm_layer=nn.LayerNorm):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.reduction = nn.Linear(4 * dim, 2 * dim, bias=False)
|
||||
self.norm = norm_layer(4 * dim)
|
||||
|
||||
def forward(self, x, H, W):
|
||||
""" Forward function.
|
||||
|
||||
Args:
|
||||
x: Input feature, tensor size (B, H*W, C).
|
||||
H, W: Spatial resolution of the input feature.
|
||||
"""
|
||||
B, L, C = x.shape
|
||||
assert L == H * W, "input feature has wrong size"
|
||||
|
||||
x = x.view(B, H, W, C)
|
||||
|
||||
# padding
|
||||
pad_input = (H % 2 == 1) or (W % 2 == 1)
|
||||
if pad_input:
|
||||
x = F.pad(x, (0, 0, 0, W % 2, 0, H % 2))
|
||||
|
||||
x0 = x[:, 0::2, 0::2, :] # B H/2 W/2 C
|
||||
x1 = x[:, 1::2, 0::2, :] # B H/2 W/2 C
|
||||
x2 = x[:, 0::2, 1::2, :] # B H/2 W/2 C
|
||||
x3 = x[:, 1::2, 1::2, :] # B H/2 W/2 C
|
||||
x = torch.cat([x0, x1, x2, x3], -1) # B H/2 W/2 4*C
|
||||
x = x.view(B, -1, 4 * C) # B H/2*W/2 4*C
|
||||
|
||||
x = self.norm(x)
|
||||
x = self.reduction(x)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class BasicLayer(nn.Module):
|
||||
""" A basic Swin Transformer layer for one stage.
|
||||
|
||||
Args:
|
||||
dim (int): Number of feature channels
|
||||
depth (int): Depths of this stage.
|
||||
num_heads (int): Number of attention head.
|
||||
window_size (int): Local window size. Default: 7.
|
||||
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
|
||||
qkv_bias (bool, optional): If True, add a learnable bias to query, key, value. Default: True
|
||||
qk_scale (float | None, optional): Override default qk scale of head_dim ** -0.5 if set.
|
||||
drop (float, optional): Dropout rate. Default: 0.0
|
||||
attn_drop (float, optional): Attention dropout rate. Default: 0.0
|
||||
drop_path (float | tuple[float], optional): Stochastic depth rate. Default: 0.0
|
||||
norm_layer (nn.Module, optional): Normalization layer. Default: nn.LayerNorm
|
||||
downsample (nn.Module | None, optional): Downsample layer at the end of the layer. Default: None
|
||||
use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
dim,
|
||||
depth,
|
||||
num_heads,
|
||||
window_size=7,
|
||||
mlp_ratio=4.,
|
||||
qkv_bias=True,
|
||||
qk_scale=None,
|
||||
drop=0.,
|
||||
attn_drop=0.,
|
||||
drop_path=0.,
|
||||
norm_layer=nn.LayerNorm,
|
||||
downsample=None,
|
||||
use_checkpoint=False):
|
||||
super().__init__()
|
||||
self.window_size = window_size
|
||||
self.shift_size = window_size // 2
|
||||
self.depth = depth
|
||||
self.use_checkpoint = use_checkpoint
|
||||
|
||||
# build blocks
|
||||
self.blocks = nn.ModuleList([
|
||||
SwinTransformerBlock(
|
||||
dim=dim,
|
||||
num_heads=num_heads,
|
||||
window_size=window_size,
|
||||
shift_size=0 if (i % 2 == 0) else window_size // 2,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop,
|
||||
attn_drop=attn_drop,
|
||||
drop_path=drop_path[i] if isinstance(drop_path, list) else drop_path,
|
||||
norm_layer=norm_layer)
|
||||
for i in range(depth)])
|
||||
|
||||
# patch merging layer
|
||||
if downsample is not None:
|
||||
self.downsample = downsample(dim=dim, norm_layer=norm_layer)
|
||||
else:
|
||||
self.downsample = None
|
||||
|
||||
def forward(self, x, H, W):
|
||||
""" Forward function.
|
||||
|
||||
Args:
|
||||
x: Input feature, tensor size (B, H*W, C).
|
||||
H, W: Spatial resolution of the input feature.
|
||||
"""
|
||||
|
||||
# calculate attention mask for SW-MSA
|
||||
# Turn int to torch.tensor for the compatiability with torch.compile in PyTorch 2.5.
|
||||
Hp = torch.ceil(torch.tensor(H) / self.window_size).to(torch.int64) * self.window_size
|
||||
Wp = torch.ceil(torch.tensor(W) / self.window_size).to(torch.int64) * self.window_size
|
||||
img_mask = torch.zeros((1, Hp, Wp, 1), device=x.device) # 1 Hp Wp 1
|
||||
h_slices = (slice(0, -self.window_size),
|
||||
slice(-self.window_size, -self.shift_size),
|
||||
slice(-self.shift_size, None))
|
||||
w_slices = (slice(0, -self.window_size),
|
||||
slice(-self.window_size, -self.shift_size),
|
||||
slice(-self.shift_size, None))
|
||||
cnt = 0
|
||||
for h in h_slices:
|
||||
for w in w_slices:
|
||||
img_mask[:, h, w, :] = cnt
|
||||
cnt += 1
|
||||
|
||||
mask_windows = window_partition(img_mask, self.window_size) # nW, window_size, window_size, 1
|
||||
mask_windows = mask_windows.view(-1, self.window_size * self.window_size)
|
||||
attn_mask = mask_windows.unsqueeze(1) - mask_windows.unsqueeze(2)
|
||||
attn_mask = attn_mask.masked_fill(attn_mask != 0, float(-100.0)).masked_fill(attn_mask == 0, float(0.0))
|
||||
|
||||
for blk in self.blocks:
|
||||
blk.H, blk.W = H, W
|
||||
if self.use_checkpoint:
|
||||
x = checkpoint.checkpoint(blk, x, attn_mask)
|
||||
else:
|
||||
x = blk(x, attn_mask)
|
||||
if self.downsample is not None:
|
||||
x_down = self.downsample(x, H, W)
|
||||
Wh, Ww = (H + 1) // 2, (W + 1) // 2
|
||||
return x, H, W, x_down, Wh, Ww
|
||||
else:
|
||||
return x, H, W, x, H, W
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
""" Image to Patch Embedding
|
||||
|
||||
Args:
|
||||
patch_size (int): Patch token size. Default: 4.
|
||||
in_channels (int): Number of input image channels. Default: 3.
|
||||
embed_dim (int): Number of linear projection output channels. Default: 96.
|
||||
norm_layer (nn.Module, optional): Normalization layer. Default: None
|
||||
"""
|
||||
|
||||
def __init__(self, patch_size=4, in_channels=3, embed_dim=96, norm_layer=None):
|
||||
super().__init__()
|
||||
patch_size = to_2tuple(patch_size)
|
||||
self.patch_size = patch_size
|
||||
|
||||
self.in_channels = in_channels
|
||||
self.embed_dim = embed_dim
|
||||
|
||||
self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size)
|
||||
if norm_layer is not None:
|
||||
self.norm = norm_layer(embed_dim)
|
||||
else:
|
||||
self.norm = None
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward function."""
|
||||
# padding
|
||||
_, _, H, W = x.size()
|
||||
if W % self.patch_size[1] != 0:
|
||||
x = F.pad(x, (0, self.patch_size[1] - W % self.patch_size[1]))
|
||||
if H % self.patch_size[0] != 0:
|
||||
x = F.pad(x, (0, 0, 0, self.patch_size[0] - H % self.patch_size[0]))
|
||||
|
||||
x = self.proj(x) # B C Wh Ww
|
||||
if self.norm is not None:
|
||||
Wh, Ww = x.size(2), x.size(3)
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
x = self.norm(x)
|
||||
x = x.transpose(1, 2).view(-1, self.embed_dim, Wh, Ww)
|
||||
|
||||
return x
|
||||
|
||||
|
||||
class SwinTransformer(nn.Module):
|
||||
""" Swin Transformer backbone.
|
||||
A PyTorch impl of : `Swin Transformer: Hierarchical Vision Transformer using Shifted Windows` -
|
||||
https://arxiv.org/pdf/2103.14030
|
||||
|
||||
Args:
|
||||
pretrain_img_size (int): Input image size for training the pretrained model,
|
||||
used in absolute postion embedding. Default 224.
|
||||
patch_size (int | tuple(int)): Patch size. Default: 4.
|
||||
in_channels (int): Number of input image channels. Default: 3.
|
||||
embed_dim (int): Number of linear projection output channels. Default: 96.
|
||||
depths (tuple[int]): Depths of each Swin Transformer stage.
|
||||
num_heads (tuple[int]): Number of attention head of each stage.
|
||||
window_size (int): Window size. Default: 7.
|
||||
mlp_ratio (float): Ratio of mlp hidden dim to embedding dim. Default: 4.
|
||||
qkv_bias (bool): If True, add a learnable bias to query, key, value. Default: True
|
||||
qk_scale (float): Override default qk scale of head_dim ** -0.5 if set.
|
||||
drop_rate (float): Dropout rate.
|
||||
attn_drop_rate (float): Attention dropout rate. Default: 0.
|
||||
drop_path_rate (float): Stochastic depth rate. Default: 0.2.
|
||||
norm_layer (nn.Module): Normalization layer. Default: nn.LayerNorm.
|
||||
ape (bool): If True, add absolute position embedding to the patch embedding. Default: False.
|
||||
patch_norm (bool): If True, add normalization after patch embedding. Default: True.
|
||||
out_indices (Sequence[int]): Output from which stages.
|
||||
frozen_stages (int): Stages to be frozen (stop grad and set eval mode).
|
||||
-1 means not freezing any parameters.
|
||||
use_checkpoint (bool): Whether to use checkpointing to save memory. Default: False.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
pretrain_img_size=224,
|
||||
patch_size=4,
|
||||
in_channels=3,
|
||||
embed_dim=96,
|
||||
depths=[2, 2, 6, 2],
|
||||
num_heads=[3, 6, 12, 24],
|
||||
window_size=7,
|
||||
mlp_ratio=4.,
|
||||
qkv_bias=True,
|
||||
qk_scale=None,
|
||||
drop_rate=0.,
|
||||
attn_drop_rate=0.,
|
||||
drop_path_rate=0.2,
|
||||
norm_layer=nn.LayerNorm,
|
||||
ape=False,
|
||||
patch_norm=True,
|
||||
out_indices=(0, 1, 2, 3),
|
||||
frozen_stages=-1,
|
||||
use_checkpoint=False):
|
||||
super().__init__()
|
||||
|
||||
self.pretrain_img_size = pretrain_img_size
|
||||
self.num_layers = len(depths)
|
||||
self.embed_dim = embed_dim
|
||||
self.ape = ape
|
||||
self.patch_norm = patch_norm
|
||||
self.out_indices = out_indices
|
||||
self.frozen_stages = frozen_stages
|
||||
|
||||
# split image into non-overlapping patches
|
||||
self.patch_embed = PatchEmbed(
|
||||
patch_size=patch_size, in_channels=in_channels, embed_dim=embed_dim,
|
||||
norm_layer=norm_layer if self.patch_norm else None)
|
||||
|
||||
# absolute position embedding
|
||||
if self.ape:
|
||||
pretrain_img_size = to_2tuple(pretrain_img_size)
|
||||
patch_size = to_2tuple(patch_size)
|
||||
patches_resolution = [pretrain_img_size[0] // patch_size[0], pretrain_img_size[1] // patch_size[1]]
|
||||
|
||||
self.absolute_pos_embed = nn.Parameter(torch.zeros(1, embed_dim, patches_resolution[0], patches_resolution[1]))
|
||||
trunc_normal_(self.absolute_pos_embed, std=.02)
|
||||
|
||||
self.pos_drop = nn.Dropout(p=drop_rate)
|
||||
|
||||
# stochastic depth
|
||||
dpr = [x.item() for x in torch.linspace(0, drop_path_rate, sum(depths))] # stochastic depth decay rule
|
||||
|
||||
# build layers
|
||||
self.layers = nn.ModuleList()
|
||||
for i_layer in range(self.num_layers):
|
||||
layer = BasicLayer(
|
||||
dim=int(embed_dim * 2 ** i_layer),
|
||||
depth=depths[i_layer],
|
||||
num_heads=num_heads[i_layer],
|
||||
window_size=window_size,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qkv_bias=qkv_bias,
|
||||
qk_scale=qk_scale,
|
||||
drop=drop_rate,
|
||||
attn_drop=attn_drop_rate,
|
||||
drop_path=dpr[sum(depths[:i_layer]):sum(depths[:i_layer + 1])],
|
||||
norm_layer=norm_layer,
|
||||
downsample=PatchMerging if (i_layer < self.num_layers - 1) else None,
|
||||
use_checkpoint=use_checkpoint)
|
||||
self.layers.append(layer)
|
||||
|
||||
num_features = [int(embed_dim * 2 ** i) for i in range(self.num_layers)]
|
||||
self.num_features = num_features
|
||||
|
||||
# add a norm layer for each output
|
||||
for i_layer in out_indices:
|
||||
layer = norm_layer(num_features[i_layer])
|
||||
layer_name = f'norm{i_layer}'
|
||||
self.add_module(layer_name, layer)
|
||||
|
||||
self._freeze_stages()
|
||||
|
||||
def _freeze_stages(self):
|
||||
if self.frozen_stages >= 0:
|
||||
self.patch_embed.eval()
|
||||
for param in self.patch_embed.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
if self.frozen_stages >= 1 and self.ape:
|
||||
self.absolute_pos_embed.requires_grad = False
|
||||
|
||||
if self.frozen_stages >= 2:
|
||||
self.pos_drop.eval()
|
||||
for i in range(0, self.frozen_stages - 1):
|
||||
m = self.layers[i]
|
||||
m.eval()
|
||||
for param in m.parameters():
|
||||
param.requires_grad = False
|
||||
|
||||
# def init_weights(self, pretrained=None):
|
||||
# """Initialize the weights in backbone.
|
||||
#
|
||||
# Args:
|
||||
# pretrained (str, optional): Path to pre-trained weights.
|
||||
# Defaults to None.
|
||||
# """
|
||||
#
|
||||
# def _init_weights(m):
|
||||
# if isinstance(m, nn.Linear):
|
||||
# trunc_normal_(m.weight, std=.02)
|
||||
# if isinstance(m, nn.Linear) and m.bias is not None:
|
||||
# nn.init.constant_(m.bias, 0)
|
||||
# elif isinstance(m, nn.LayerNorm):
|
||||
# nn.init.constant_(m.bias, 0)
|
||||
# nn.init.constant_(m.weight, 1.0)
|
||||
#
|
||||
# if isinstance(pretrained, str):
|
||||
# self.apply(_init_weights)
|
||||
# logger = get_root_logger()
|
||||
# load_checkpoint(self, pretrained, strict=False, logger=logger)
|
||||
# elif pretrained is None:
|
||||
# self.apply(_init_weights)
|
||||
# else:
|
||||
# raise TypeError('pretrained must be a str or None')
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward function."""
|
||||
x = self.patch_embed(x)
|
||||
|
||||
Wh, Ww = x.size(2), x.size(3)
|
||||
if self.ape:
|
||||
# interpolate the position embedding to the corresponding size
|
||||
absolute_pos_embed = F.interpolate(self.absolute_pos_embed, size=(Wh, Ww), mode='bicubic')
|
||||
x = (x + absolute_pos_embed) # B Wh*Ww C
|
||||
|
||||
outs = []#x.contiguous()]
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
x = self.pos_drop(x)
|
||||
for i in range(self.num_layers):
|
||||
layer = self.layers[i]
|
||||
x_out, H, W, x, Wh, Ww = layer(x, Wh, Ww)
|
||||
|
||||
if i in self.out_indices:
|
||||
norm_layer = getattr(self, f'norm{i}')
|
||||
x_out = norm_layer(x_out)
|
||||
|
||||
out = x_out.view(-1, H, W, self.num_features[i]).permute(0, 3, 1, 2).contiguous()
|
||||
outs.append(out)
|
||||
|
||||
return tuple(outs)
|
||||
|
||||
def train(self, mode=True):
|
||||
"""Convert the model into training mode while keep layers freezed."""
|
||||
super(SwinTransformer, self).train(mode)
|
||||
self._freeze_stages()
|
||||
|
||||
def swin_v1_t():
|
||||
model = SwinTransformer(embed_dim=96, depths=[2, 2, 6, 2], num_heads=[3, 6, 12, 24], window_size=7)
|
||||
return model
|
||||
|
||||
def swin_v1_s():
|
||||
model = SwinTransformer(embed_dim=96, depths=[2, 2, 18, 2], num_heads=[3, 6, 12, 24], window_size=7)
|
||||
return model
|
||||
|
||||
def swin_v1_b():
|
||||
model = SwinTransformer(embed_dim=128, depths=[2, 2, 18, 2], num_heads=[4, 8, 16, 32], window_size=12)
|
||||
return model
|
||||
|
||||
def swin_v1_l():
|
||||
model = SwinTransformer(embed_dim=192, depths=[2, 2, 18, 2], num_heads=[6, 12, 24, 48], window_size=12)
|
||||
return model
|
||||
@@ -0,0 +1,291 @@
|
||||
# import torch
|
||||
# import torch.nn as nn
|
||||
from collections import OrderedDict
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torchvision.models import vgg16, vgg16_bn
|
||||
from torchvision.models import resnet50
|
||||
from kornia.filters import laplacian
|
||||
|
||||
from ..config import Config
|
||||
from ..dataset import class_labels_TR_sorted
|
||||
from .backbones.build_backbone import build_backbone
|
||||
from .modules.decoder_blocks import BasicDecBlk, ResBlk, HierarAttDecBlk
|
||||
from .modules.lateral_blocks import BasicLatBlk
|
||||
from .modules.aspp import ASPP, ASPPDeformable
|
||||
from .modules.ing import *
|
||||
from .refinement.refiner import Refiner, RefinerPVTInChannels4, RefUNet
|
||||
from .refinement.stem_layer import StemLayer
|
||||
|
||||
|
||||
class BiRefNet(nn.Module):
|
||||
def __init__(self, bb_pretrained=True):
|
||||
super(BiRefNet, self).__init__()
|
||||
self.config = Config()
|
||||
self.epoch = 1
|
||||
self.bb = build_backbone(self.config.bb, pretrained=bb_pretrained)
|
||||
|
||||
channels = self.config.lateral_channels_in_collection
|
||||
|
||||
if self.config.auxiliary_classification:
|
||||
self.avgpool = nn.AdaptiveAvgPool2d((1, 1))
|
||||
self.cls_head = nn.Sequential(
|
||||
nn.Linear(channels[0], len(class_labels_TR_sorted))
|
||||
)
|
||||
|
||||
if self.config.squeeze_block:
|
||||
self.squeeze_module = nn.Sequential(*[
|
||||
eval(self.config.squeeze_block.split('_x')[0])(channels[0]+sum(self.config.cxt), channels[0])
|
||||
for _ in range(eval(self.config.squeeze_block.split('_x')[1]))
|
||||
])
|
||||
|
||||
self.decoder = Decoder(channels)
|
||||
|
||||
if self.config.locate_head:
|
||||
self.locate_header = nn.ModuleList([
|
||||
BasicDecBlk(channels[0], channels[-1]),
|
||||
nn.Sequential(
|
||||
nn.Conv2d(channels[-1], 1, 1, 1, 0),
|
||||
)
|
||||
])
|
||||
|
||||
if self.config.ender:
|
||||
self.dec_end = nn.Sequential(
|
||||
nn.Conv2d(1, 16, 3, 1, 1),
|
||||
nn.Conv2d(16, 1, 3, 1, 1),
|
||||
nn.ReLU(inplace=True),
|
||||
)
|
||||
|
||||
# refine patch-level segmentation
|
||||
if self.config.refine:
|
||||
if self.config.refine == 'itself':
|
||||
self.stem_layer = StemLayer(in_channels=3+1, inter_channels=48, out_channels=3)
|
||||
else:
|
||||
self.refiner = eval('{}({})'.format(self.config.refine, 'in_channels=3+1'))
|
||||
|
||||
if self.config.freeze_bb:
|
||||
# Freeze the backbone...
|
||||
print(self.named_parameters())
|
||||
for key, value in self.named_parameters():
|
||||
if 'bb.' in key and 'refiner.' not in key:
|
||||
value.requires_grad = False
|
||||
|
||||
def forward_enc(self, x):
|
||||
if self.config.bb in ['vgg16', 'vgg16bn', 'resnet50']:
|
||||
x1 = self.bb.conv1(x); x2 = self.bb.conv2(x1); x3 = self.bb.conv3(x2); x4 = self.bb.conv4(x3)
|
||||
else:
|
||||
x1, x2, x3, x4 = self.bb(x)
|
||||
if self.config.mul_scl_ipt == 'cat':
|
||||
B, C, H, W = x.shape
|
||||
x1_, x2_, x3_, x4_ = self.bb(F.interpolate(x, size=(H//2, W//2), mode='bilinear', align_corners=True))
|
||||
x1 = torch.cat([x1, F.interpolate(x1_, size=x1.shape[2:], mode='bilinear', align_corners=True)], dim=1)
|
||||
x2 = torch.cat([x2, F.interpolate(x2_, size=x2.shape[2:], mode='bilinear', align_corners=True)], dim=1)
|
||||
x3 = torch.cat([x3, F.interpolate(x3_, size=x3.shape[2:], mode='bilinear', align_corners=True)], dim=1)
|
||||
x4 = torch.cat([x4, F.interpolate(x4_, size=x4.shape[2:], mode='bilinear', align_corners=True)], dim=1)
|
||||
elif self.config.mul_scl_ipt == 'add':
|
||||
B, C, H, W = x.shape
|
||||
x1_, x2_, x3_, x4_ = self.bb(F.interpolate(x, size=(H//2, W//2), mode='bilinear', align_corners=True))
|
||||
x1 = x1 + F.interpolate(x1_, size=x1.shape[2:], mode='bilinear', align_corners=True)
|
||||
x2 = x2 + F.interpolate(x2_, size=x2.shape[2:], mode='bilinear', align_corners=True)
|
||||
x3 = x3 + F.interpolate(x3_, size=x3.shape[2:], mode='bilinear', align_corners=True)
|
||||
x4 = x4 + F.interpolate(x4_, size=x4.shape[2:], mode='bilinear', align_corners=True)
|
||||
class_preds = self.cls_head(self.avgpool(x4).view(x4.shape[0], -1)) if self.training and self.config.auxiliary_classification else None
|
||||
if self.config.cxt:
|
||||
x4 = torch.cat(
|
||||
(
|
||||
*[
|
||||
F.interpolate(x1, size=x4.shape[2:], mode='bilinear', align_corners=True),
|
||||
F.interpolate(x2, size=x4.shape[2:], mode='bilinear', align_corners=True),
|
||||
F.interpolate(x3, size=x4.shape[2:], mode='bilinear', align_corners=True),
|
||||
][-len(self.config.cxt):],
|
||||
x4
|
||||
),
|
||||
dim=1
|
||||
)
|
||||
return (x1, x2, x3, x4), class_preds
|
||||
|
||||
def forward_ori(self, x):
|
||||
########## Encoder ##########
|
||||
(x1, x2, x3, x4), class_preds = self.forward_enc(x)
|
||||
if self.config.squeeze_block:
|
||||
x4 = self.squeeze_module(x4)
|
||||
########## Decoder ##########
|
||||
features = [x, x1, x2, x3, x4]
|
||||
if self.config.out_ref:
|
||||
features.append(laplacian(torch.mean(x, dim=1).unsqueeze(1), kernel_size=5))
|
||||
scaled_preds = self.decoder(features)
|
||||
return scaled_preds, class_preds
|
||||
|
||||
def forward_ref(self, x, pred):
|
||||
# refine patch-level segmentation
|
||||
if pred.shape[2:] != x.shape[2:]:
|
||||
pred = F.interpolate(pred, size=x.shape[2:], mode='bilinear', align_corners=True)
|
||||
# pred = pred.sigmoid()
|
||||
if self.config.refine == 'itself':
|
||||
x = self.stem_layer(torch.cat([x, pred], dim=1))
|
||||
scaled_preds, class_preds = self.forward_ori(x)
|
||||
else:
|
||||
scaled_preds = self.refiner([x, pred])
|
||||
class_preds = None
|
||||
return scaled_preds, class_preds
|
||||
|
||||
def forward_ref_end(self, x):
|
||||
# remove the grids of concatenated preds
|
||||
return self.dec_end(x) if self.config.ender else x
|
||||
|
||||
|
||||
def forward(self, x):
|
||||
scaled_preds, class_preds = self.forward_ori(x)
|
||||
class_preds_lst = [class_preds]
|
||||
return [scaled_preds, class_preds_lst] if self.training else scaled_preds
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(self, channels):
|
||||
super(Decoder, self).__init__()
|
||||
self.config = Config()
|
||||
DecoderBlock = eval(self.config.dec_blk)
|
||||
LateralBlock = eval(self.config.lat_blk)
|
||||
|
||||
if self.config.dec_ipt:
|
||||
self.split = self.config.dec_ipt_split
|
||||
N_dec_ipt = 64
|
||||
DBlock = SimpleConvs
|
||||
ic = 64
|
||||
ipt_cha_opt = 1
|
||||
self.ipt_blk4 = DBlock(2**8*3 if self.split else 3, [N_dec_ipt, channels[0]//8][ipt_cha_opt], inter_channels=ic)
|
||||
self.ipt_blk3 = DBlock(2**6*3 if self.split else 3, [N_dec_ipt, channels[1]//8][ipt_cha_opt], inter_channels=ic)
|
||||
self.ipt_blk2 = DBlock(2**4*3 if self.split else 3, [N_dec_ipt, channels[2]//8][ipt_cha_opt], inter_channels=ic)
|
||||
self.ipt_blk1 = DBlock(2**0*3 if self.split else 3, [N_dec_ipt, channels[3]//8][ipt_cha_opt], inter_channels=ic)
|
||||
else:
|
||||
self.split = None
|
||||
|
||||
self.decoder_block4 = DecoderBlock(channels[0], channels[1])
|
||||
self.decoder_block3 = DecoderBlock(channels[1]+([N_dec_ipt, channels[0]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[2])
|
||||
self.decoder_block2 = DecoderBlock(channels[2]+([N_dec_ipt, channels[1]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[3])
|
||||
self.decoder_block1 = DecoderBlock(channels[3]+([N_dec_ipt, channels[2]//8][ipt_cha_opt] if self.config.dec_ipt else 0), channels[3]//2)
|
||||
self.conv_out1 = nn.Sequential(nn.Conv2d(channels[3]//2+([N_dec_ipt, channels[3]//8][ipt_cha_opt] if self.config.dec_ipt else 0), 1, 1, 1, 0))
|
||||
|
||||
self.lateral_block4 = LateralBlock(channels[1], channels[1])
|
||||
self.lateral_block3 = LateralBlock(channels[2], channels[2])
|
||||
self.lateral_block2 = LateralBlock(channels[3], channels[3])
|
||||
|
||||
if self.config.ms_supervision:
|
||||
self.conv_ms_spvn_4 = nn.Conv2d(channels[1], 1, 1, 1, 0)
|
||||
self.conv_ms_spvn_3 = nn.Conv2d(channels[2], 1, 1, 1, 0)
|
||||
self.conv_ms_spvn_2 = nn.Conv2d(channels[3], 1, 1, 1, 0)
|
||||
|
||||
if self.config.out_ref:
|
||||
_N = 16
|
||||
# self.gdt_convs_4 = nn.Sequential(nn.Conv2d(channels[1], _N, 3, 1, 1), nn.BatchNorm2d(_N), nn.ReLU(inplace=True))
|
||||
self.gdt_convs_3 = nn.Sequential(nn.Conv2d(channels[2], _N, 3, 1, 1), nn.BatchNorm2d(_N), nn.ReLU(inplace=True))
|
||||
self.gdt_convs_2 = nn.Sequential(nn.Conv2d(channels[3], _N, 3, 1, 1), nn.BatchNorm2d(_N), nn.ReLU(inplace=True))
|
||||
|
||||
# self.gdt_convs_pred_4 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
|
||||
self.gdt_convs_pred_3 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
|
||||
self.gdt_convs_pred_2 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
|
||||
|
||||
# self.gdt_convs_attn_4 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
|
||||
self.gdt_convs_attn_3 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
|
||||
self.gdt_convs_attn_2 = nn.Sequential(nn.Conv2d(_N, 1, 1, 1, 0))
|
||||
|
||||
|
||||
def get_patches_batch(self, x, p):
|
||||
_size_h, _size_w = p.shape[2:]
|
||||
patches_batch = []
|
||||
for idx in range(x.shape[0]):
|
||||
columns_x = torch.split(x[idx], split_size_or_sections=_size_w, dim=-1)
|
||||
patches_x = []
|
||||
for column_x in columns_x:
|
||||
patches_x += [p.unsqueeze(0) for p in torch.split(column_x, split_size_or_sections=_size_h, dim=-2)]
|
||||
patch_sample = torch.cat(patches_x, dim=1)
|
||||
patches_batch.append(patch_sample)
|
||||
return torch.cat(patches_batch, dim=0)
|
||||
|
||||
def forward(self, features):
|
||||
if self.config.out_ref:
|
||||
outs_gdt_pred = []
|
||||
outs_gdt_label = []
|
||||
x, x1, x2, x3, x4, gdt_gt = features
|
||||
else:
|
||||
x, x1, x2, x3, x4 = features
|
||||
outs = []
|
||||
p4 = self.decoder_block4(x4)
|
||||
m4 = self.conv_ms_spvn_4(p4) if self.config.ms_supervision else None
|
||||
_p4 = F.interpolate(p4, size=x3.shape[2:], mode='bilinear', align_corners=True)
|
||||
_p3 = _p4 + self.lateral_block4(x3)
|
||||
if self.config.dec_ipt:
|
||||
patches_batch = self.get_patches_batch(x, _p3) if self.split else x
|
||||
_p3 = torch.cat((_p3, self.ipt_blk4(F.interpolate(patches_batch, size=x3.shape[2:], mode='bilinear', align_corners=True))), 1)
|
||||
|
||||
p3 = self.decoder_block3(_p3)
|
||||
m3 = self.conv_ms_spvn_3(p3) if self.config.ms_supervision else None
|
||||
if self.config.out_ref:
|
||||
# >> GT:
|
||||
# m3 --dilation--> m3_dia
|
||||
# G_3^gt * m3_dia --> G_3^m, which is the label of gradient
|
||||
m3_dia = m3
|
||||
gdt_label_main_3 = gdt_gt * F.interpolate(m3_dia, size=gdt_gt.shape[2:], mode='bilinear', align_corners=True)
|
||||
outs_gdt_label.append(gdt_label_main_3)
|
||||
# >> Pred:
|
||||
# p3 --conv--BN--> F_3^G, where F_3^G predicts the \hat{G_3} with xx
|
||||
# F_3^G --sigmoid--> A_3^G
|
||||
p3_gdt = self.gdt_convs_3(p3)
|
||||
gdt_pred_3 = self.gdt_convs_pred_3(p3_gdt)
|
||||
outs_gdt_pred.append(gdt_pred_3)
|
||||
gdt_attn_3 = self.gdt_convs_attn_3(p3_gdt).sigmoid()
|
||||
# >> Finally:
|
||||
# p3 = p3 * A_3^G
|
||||
p3 = p3 * gdt_attn_3
|
||||
_p3 = F.interpolate(p3, size=x2.shape[2:], mode='bilinear', align_corners=True)
|
||||
_p2 = _p3 + self.lateral_block3(x2)
|
||||
if self.config.dec_ipt:
|
||||
patches_batch = self.get_patches_batch(x, _p2) if self.split else x
|
||||
_p2 = torch.cat((_p2, self.ipt_blk3(F.interpolate(patches_batch, size=x2.shape[2:], mode='bilinear', align_corners=True))), 1)
|
||||
|
||||
p2 = self.decoder_block2(_p2)
|
||||
m2 = self.conv_ms_spvn_2(p2) if self.config.ms_supervision else None
|
||||
if self.config.out_ref:
|
||||
# >> GT:
|
||||
m2_dia = m2
|
||||
gdt_label_main_2 = gdt_gt * F.interpolate(m2_dia, size=gdt_gt.shape[2:], mode='bilinear', align_corners=True)
|
||||
outs_gdt_label.append(gdt_label_main_2)
|
||||
# >> Pred:
|
||||
p2_gdt = self.gdt_convs_2(p2)
|
||||
gdt_pred_2 = self.gdt_convs_pred_2(p2_gdt)
|
||||
outs_gdt_pred.append(gdt_pred_2)
|
||||
gdt_attn_2 = self.gdt_convs_attn_2(p2_gdt).sigmoid()
|
||||
# >> Finally:
|
||||
p2 = p2 * gdt_attn_2
|
||||
_p2 = F.interpolate(p2, size=x1.shape[2:], mode='bilinear', align_corners=True)
|
||||
_p1 = _p2 + self.lateral_block2(x1)
|
||||
if self.config.dec_ipt:
|
||||
patches_batch = self.get_patches_batch(x, _p1) if self.split else x
|
||||
_p1 = torch.cat((_p1, self.ipt_blk2(F.interpolate(patches_batch, size=x1.shape[2:], mode='bilinear', align_corners=True))), 1)
|
||||
|
||||
_p1 = self.decoder_block1(_p1)
|
||||
_p1 = F.interpolate(_p1, size=x.shape[2:], mode='bilinear', align_corners=True)
|
||||
if self.config.dec_ipt:
|
||||
patches_batch = self.get_patches_batch(x, _p1) if self.split else x
|
||||
_p1 = torch.cat((_p1, self.ipt_blk1(F.interpolate(patches_batch, size=x.shape[2:], mode='bilinear', align_corners=True))), 1)
|
||||
p1_out = self.conv_out1(_p1)
|
||||
|
||||
if self.config.ms_supervision:
|
||||
outs.append(m4)
|
||||
outs.append(m3)
|
||||
outs.append(m2)
|
||||
outs.append(p1_out)
|
||||
return outs if not (self.config.out_ref and self.training) else ([outs_gdt_pred, outs_gdt_label], outs)
|
||||
|
||||
|
||||
class SimpleConvs(nn.Module):
|
||||
def __init__(
|
||||
self, in_channels: int, out_channels: int, inter_channels=64
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.conv1 = nn.Conv2d(in_channels, inter_channels, 3, 1, 1)
|
||||
self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, 1)
|
||||
|
||||
def forward(self, x):
|
||||
return self.conv_out(self.conv1(x))
|
||||
@@ -0,0 +1,163 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from ..modules.deform_conv import DeformableConv2d
|
||||
from ...config import Config
|
||||
|
||||
|
||||
config = Config()
|
||||
|
||||
|
||||
class ASPPComplex(nn.Module):
|
||||
def __init__(self, in_channels=64, out_channels=None, output_stride=16):
|
||||
super(ASPPComplex, self).__init__()
|
||||
self.down_scale = 1
|
||||
if out_channels is None:
|
||||
out_channels = in_channels
|
||||
self.in_channelster = 256 // self.down_scale
|
||||
if output_stride == 16:
|
||||
dilations = [1, 6, 12, 18]
|
||||
elif output_stride == 8:
|
||||
dilations = [1, 12, 24, 36]
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
self.aspp1 = _ASPPModule(in_channels, self.in_channelster, 1, padding=0, dilation=dilations[0])
|
||||
self.aspp2 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[1], dilation=dilations[1])
|
||||
self.aspp3 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[2], dilation=dilations[2])
|
||||
self.aspp4 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[3], dilation=dilations[3])
|
||||
|
||||
self.global_avg_pool = nn.Sequential(nn.AdaptiveAvgPool2d((1, 1)),
|
||||
nn.Conv2d(in_channels, self.in_channelster, 1, stride=1, bias=False),
|
||||
nn.BatchNorm2d(self.in_channelster) if config.batch_size > 1 else nn.Identity(),
|
||||
nn.ReLU(inplace=True))
|
||||
self.conv1 = nn.Conv2d(self.in_channelster * 5, out_channels, 1, bias=False)
|
||||
self.bn1 = nn.BatchNorm2d(out_channels)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
self.dropout = nn.Dropout(0.5)
|
||||
|
||||
def forward(self, x):
|
||||
x1 = self.aspp1(x)
|
||||
x2 = self.aspp2(x)
|
||||
x3 = self.aspp3(x)
|
||||
x4 = self.aspp4(x)
|
||||
x5 = self.global_avg_pool(x)
|
||||
x5 = F.interpolate(x5, size=x1.size()[2:], mode='bilinear', align_corners=True)
|
||||
x = torch.cat((x1, x2, x3, x4, x5), dim=1)
|
||||
|
||||
x = self.conv1(x)
|
||||
x = self.bn1(x)
|
||||
x = self.relu(x)
|
||||
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
class _ASPPModule(nn.Module):
|
||||
def __init__(self, in_channels, planes, kernel_size, padding, dilation):
|
||||
super(_ASPPModule, self).__init__()
|
||||
self.atrous_conv = nn.Conv2d(in_channels, planes, kernel_size=kernel_size,
|
||||
stride=1, padding=padding, dilation=dilation, bias=False)
|
||||
self.bn = nn.BatchNorm2d(planes)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.atrous_conv(x)
|
||||
x = self.bn(x)
|
||||
|
||||
return self.relu(x)
|
||||
|
||||
|
||||
class ASPP(nn.Module):
|
||||
def __init__(self, in_channels=64, out_channels=None, output_stride=16):
|
||||
super(ASPP, self).__init__()
|
||||
self.down_scale = 1
|
||||
if out_channels is None:
|
||||
out_channels = in_channels
|
||||
self.in_channelster = 256 // self.down_scale
|
||||
if output_stride == 16:
|
||||
dilations = [1, 6, 12, 18]
|
||||
elif output_stride == 8:
|
||||
dilations = [1, 12, 24, 36]
|
||||
else:
|
||||
raise NotImplementedError
|
||||
|
||||
self.aspp1 = _ASPPModule(in_channels, self.in_channelster, 1, padding=0, dilation=dilations[0])
|
||||
self.aspp2 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[1], dilation=dilations[1])
|
||||
self.aspp3 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[2], dilation=dilations[2])
|
||||
self.aspp4 = _ASPPModule(in_channels, self.in_channelster, 3, padding=dilations[3], dilation=dilations[3])
|
||||
|
||||
self.global_avg_pool = nn.Sequential(nn.AdaptiveAvgPool2d((1, 1)),
|
||||
nn.Conv2d(in_channels, self.in_channelster, 1, stride=1, bias=False),
|
||||
nn.BatchNorm2d(self.in_channelster) if config.batch_size > 1 else nn.Identity(),
|
||||
nn.ReLU(inplace=True))
|
||||
self.conv1 = nn.Conv2d(self.in_channelster * 5, out_channels, 1, bias=False)
|
||||
self.bn1 = nn.BatchNorm2d(out_channels)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
self.dropout = nn.Dropout(0.5)
|
||||
|
||||
def forward(self, x):
|
||||
x1 = self.aspp1(x)
|
||||
x2 = self.aspp2(x)
|
||||
x3 = self.aspp3(x)
|
||||
x4 = self.aspp4(x)
|
||||
x5 = self.global_avg_pool(x)
|
||||
x5 = F.interpolate(x5, size=x1.size()[2:], mode='bilinear', align_corners=True)
|
||||
x = torch.cat((x1, x2, x3, x4, x5), dim=1)
|
||||
|
||||
x = self.conv1(x)
|
||||
x = self.bn1(x)
|
||||
x = self.relu(x)
|
||||
|
||||
return self.dropout(x)
|
||||
|
||||
|
||||
##################### Deformable
|
||||
class _ASPPModuleDeformable(nn.Module):
|
||||
def __init__(self, in_channels, planes, kernel_size, padding):
|
||||
super(_ASPPModuleDeformable, self).__init__()
|
||||
self.atrous_conv = DeformableConv2d(in_channels, planes, kernel_size=kernel_size,
|
||||
stride=1, padding=padding, bias=False)
|
||||
self.bn = nn.BatchNorm2d(planes)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.atrous_conv(x)
|
||||
x = self.bn(x)
|
||||
|
||||
return self.relu(x)
|
||||
|
||||
|
||||
class ASPPDeformable(nn.Module):
|
||||
def __init__(self, in_channels, out_channels=None, num_parallel_block=1):
|
||||
super(ASPPDeformable, self).__init__()
|
||||
self.down_scale = 1
|
||||
if out_channels is None:
|
||||
out_channels = in_channels
|
||||
self.in_channelster = 256 // self.down_scale
|
||||
|
||||
self.aspp1 = _ASPPModuleDeformable(in_channels, self.in_channelster, 1, padding=0)
|
||||
self.aspp_deforms = nn.ModuleList([
|
||||
_ASPPModuleDeformable(in_channels, self.in_channelster, 3, padding=1) for _ in range(num_parallel_block)
|
||||
])
|
||||
|
||||
self.global_avg_pool = nn.Sequential(nn.AdaptiveAvgPool2d((1, 1)),
|
||||
nn.Conv2d(in_channels, self.in_channelster, 1, stride=1, bias=False),
|
||||
nn.BatchNorm2d(self.in_channelster) if config.batch_size > 1 else nn.Identity(),
|
||||
nn.ReLU(inplace=True))
|
||||
self.conv1 = nn.Conv2d(self.in_channelster * (2 + len(self.aspp_deforms)), out_channels, 1, bias=False)
|
||||
self.bn1 = nn.BatchNorm2d(out_channels)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
self.dropout = nn.Dropout(0.5)
|
||||
|
||||
def forward(self, x):
|
||||
x1 = self.aspp1(x)
|
||||
x_aspp_deforms = [aspp_deform(x) for aspp_deform in self.aspp_deforms]
|
||||
x5 = self.global_avg_pool(x)
|
||||
x5 = F.interpolate(x5, size=x1.size()[2:], mode='bilinear', align_corners=True)
|
||||
x = torch.cat((x1, *x_aspp_deforms, x5), dim=1)
|
||||
|
||||
x = self.conv1(x)
|
||||
x = self.bn1(x)
|
||||
x = self.relu(x)
|
||||
|
||||
return self.dropout(x)
|
||||
@@ -0,0 +1,93 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.nn import init
|
||||
|
||||
|
||||
class SEWeightModule(nn.Module):
|
||||
def __init__(self, channels, reduction=16):
|
||||
super(SEWeightModule, self).__init__()
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d(1)
|
||||
self.fc1 = nn.Conv2d(channels, channels//reduction, kernel_size=1, padding=0)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
self.fc2 = nn.Conv2d(channels//reduction, channels, kernel_size=1, padding=0)
|
||||
self.sigmoid = nn.Sigmoid()
|
||||
|
||||
def forward(self, x):
|
||||
out = self.avg_pool(x)
|
||||
out = self.fc1(out)
|
||||
out = self.relu(out)
|
||||
out = self.fc2(out)
|
||||
weight = self.sigmoid(out)
|
||||
return weight
|
||||
|
||||
|
||||
class PSA(nn.Module):
|
||||
|
||||
def __init__(self, in_channels, S=4, reduction=4):
|
||||
super().__init__()
|
||||
self.S = S
|
||||
|
||||
_convs = []
|
||||
for i in range(S):
|
||||
_convs.append(nn.Conv2d(in_channels//S, in_channels//S, kernel_size=2*(i+1)+1, padding=i+1))
|
||||
self.convs = nn.ModuleList(_convs)
|
||||
|
||||
self.se_block = SEWeightModule(in_channels//S, reduction=S*reduction)
|
||||
|
||||
self.softmax = nn.Softmax(dim=1)
|
||||
|
||||
def forward(self, x):
|
||||
b, c, h, w = x.size()
|
||||
|
||||
# Step1: SPC module
|
||||
SPC_out = x.view(b, self.S, c//self.S, h, w) #bs,s,ci,h,w
|
||||
for idx, conv in enumerate(self.convs):
|
||||
SPC_out[:,idx,:,:,:] = conv(SPC_out[:,idx,:,:,:].clone())
|
||||
|
||||
# Step2: SE weight
|
||||
se_out=[]
|
||||
for idx in range(self.S):
|
||||
se_out.append(self.se_block(SPC_out[:, idx, :, :, :]))
|
||||
SE_out = torch.stack(se_out, dim=1)
|
||||
SE_out = SE_out.expand_as(SPC_out)
|
||||
|
||||
# Step3: Softmax
|
||||
softmax_out = self.softmax(SE_out)
|
||||
|
||||
# Step4: SPA
|
||||
PSA_out = SPC_out * softmax_out
|
||||
PSA_out = PSA_out.view(b, -1, h, w)
|
||||
|
||||
return PSA_out
|
||||
|
||||
|
||||
class SGE(nn.Module):
|
||||
|
||||
def __init__(self, groups):
|
||||
super().__init__()
|
||||
self.groups=groups
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d(1)
|
||||
self.weight=nn.Parameter(torch.zeros(1,groups,1,1))
|
||||
self.bias=nn.Parameter(torch.zeros(1,groups,1,1))
|
||||
self.sig=nn.Sigmoid()
|
||||
|
||||
def forward(self, x):
|
||||
b, c, h,w=x.shape
|
||||
x=x.view(b*self.groups,-1,h,w) #bs*g,dim//g,h,w
|
||||
xn=x*self.avg_pool(x) #bs*g,dim//g,h,w
|
||||
xn=xn.sum(dim=1,keepdim=True) #bs*g,1,h,w
|
||||
t=xn.view(b*self.groups,-1) #bs*g,h*w
|
||||
|
||||
t=t-t.mean(dim=1,keepdim=True) #bs*g,h*w
|
||||
std=t.std(dim=1,keepdim=True)+1e-5
|
||||
t=t/std #bs*g,h*w
|
||||
t=t.view(b,self.groups,h,w) #bs,g,h*w
|
||||
|
||||
t=t*self.weight+self.bias #bs,g,h*w
|
||||
t=t.view(b*self.groups,1,h,w) #bs*g,1,h*w
|
||||
x=x*self.sig(t)
|
||||
x=x.view(b,c,h,w)
|
||||
|
||||
return x
|
||||
|
||||
@@ -0,0 +1,101 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from ..modules.aspp import ASPP, ASPPDeformable
|
||||
from ..modules.attentions import PSA, SGE
|
||||
from ...config import Config
|
||||
|
||||
|
||||
config = Config()
|
||||
|
||||
|
||||
class BasicDecBlk(nn.Module):
|
||||
def __init__(self, in_channels=64, out_channels=64, inter_channels=64):
|
||||
super(BasicDecBlk, self).__init__()
|
||||
inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
|
||||
self.conv_in = nn.Conv2d(in_channels, inter_channels, 3, 1, padding=1)
|
||||
self.relu_in = nn.ReLU(inplace=True)
|
||||
if config.dec_att == 'ASPP':
|
||||
self.dec_att = ASPP(in_channels=inter_channels)
|
||||
elif config.dec_att == 'ASPPDeformable':
|
||||
self.dec_att = ASPPDeformable(in_channels=inter_channels)
|
||||
self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, padding=1)
|
||||
self.bn_in = nn.BatchNorm2d(inter_channels)
|
||||
self.bn_out = nn.BatchNorm2d(out_channels)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv_in(x)
|
||||
x = self.bn_in(x)
|
||||
x = self.relu_in(x)
|
||||
if hasattr(self, 'dec_att'):
|
||||
x = self.dec_att(x)
|
||||
x = self.conv_out(x)
|
||||
x = self.bn_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class ResBlk(nn.Module):
|
||||
def __init__(self, in_channels=64, out_channels=None, inter_channels=64):
|
||||
super(ResBlk, self).__init__()
|
||||
if out_channels is None:
|
||||
out_channels = in_channels
|
||||
inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
|
||||
|
||||
self.conv_in = nn.Conv2d(in_channels, inter_channels, 3, 1, padding=1)
|
||||
self.bn_in = nn.BatchNorm2d(inter_channels)
|
||||
self.relu_in = nn.ReLU(inplace=True)
|
||||
|
||||
if config.dec_att == 'ASPP':
|
||||
self.dec_att = ASPP(in_channels=inter_channels)
|
||||
elif config.dec_att == 'ASPPDeformable':
|
||||
self.dec_att = ASPPDeformable(in_channels=inter_channels)
|
||||
|
||||
self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, padding=1)
|
||||
self.bn_out = nn.BatchNorm2d(out_channels)
|
||||
|
||||
self.conv_resi = nn.Conv2d(in_channels, out_channels, 1, 1, 0)
|
||||
|
||||
def forward(self, x):
|
||||
_x = self.conv_resi(x)
|
||||
x = self.conv_in(x)
|
||||
x = self.bn_in(x)
|
||||
x = self.relu_in(x)
|
||||
if hasattr(self, 'dec_att'):
|
||||
x = self.dec_att(x)
|
||||
x = self.conv_out(x)
|
||||
x = self.bn_out(x)
|
||||
return x + _x
|
||||
|
||||
|
||||
class HierarAttDecBlk(nn.Module):
|
||||
def __init__(self, in_channels=64, out_channels=None, inter_channels=64):
|
||||
super(HierarAttDecBlk, self).__init__()
|
||||
if out_channels is None:
|
||||
out_channels = in_channels
|
||||
inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
|
||||
self.split_y = 8 # must be divided by channels of all intermediate features
|
||||
self.split_x = 8
|
||||
|
||||
self.conv_in = nn.Conv2d(in_channels, inter_channels, 3, 1, 1)
|
||||
|
||||
self.psa = PSA(inter_channels*self.split_y*self.split_x, S=config.batch_size)
|
||||
self.sge = SGE(groups=config.batch_size)
|
||||
|
||||
if config.dec_att == 'ASPP':
|
||||
self.dec_att = ASPP(in_channels=inter_channels)
|
||||
elif config.dec_att == 'ASPPDeformable':
|
||||
self.dec_att = ASPPDeformable(in_channels=inter_channels)
|
||||
self.conv_out = nn.Conv2d(inter_channels, out_channels, 3, 1, 1)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv_in(x)
|
||||
N, C, H, W = x.shape
|
||||
x_patchs = x.reshape(N, -1, H//self.split_y, W//self.split_x)
|
||||
|
||||
# Hierarchical attention: group attention X patch spatial attention
|
||||
x_patchs = self.psa(x_patchs) # Group Channel Attention -- each group is a single image
|
||||
x_patchs = self.sge(x_patchs) # Patch Spatial Attention
|
||||
x = x.reshape(N, C, H, W)
|
||||
if hasattr(self, 'dec_att'):
|
||||
x = self.dec_att(x)
|
||||
x = self.conv_out(x)
|
||||
return x
|
||||
@@ -0,0 +1,66 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torchvision.ops import deform_conv2d
|
||||
|
||||
|
||||
class DeformableConv2d(nn.Module):
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1,
|
||||
bias=False):
|
||||
|
||||
super(DeformableConv2d, self).__init__()
|
||||
|
||||
assert type(kernel_size) == tuple or type(kernel_size) == int
|
||||
|
||||
kernel_size = kernel_size if type(kernel_size) == tuple else (kernel_size, kernel_size)
|
||||
self.stride = stride if type(stride) == tuple else (stride, stride)
|
||||
self.padding = padding
|
||||
|
||||
self.offset_conv = nn.Conv2d(in_channels,
|
||||
2 * kernel_size[0] * kernel_size[1],
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=self.padding,
|
||||
bias=True)
|
||||
|
||||
nn.init.constant_(self.offset_conv.weight, 0.)
|
||||
nn.init.constant_(self.offset_conv.bias, 0.)
|
||||
|
||||
self.modulator_conv = nn.Conv2d(in_channels,
|
||||
1 * kernel_size[0] * kernel_size[1],
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=self.padding,
|
||||
bias=True)
|
||||
|
||||
nn.init.constant_(self.modulator_conv.weight, 0.)
|
||||
nn.init.constant_(self.modulator_conv.bias, 0.)
|
||||
|
||||
self.regular_conv = nn.Conv2d(in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=kernel_size,
|
||||
stride=stride,
|
||||
padding=self.padding,
|
||||
bias=bias)
|
||||
|
||||
def forward(self, x):
|
||||
#h, w = x.shape[2:]
|
||||
#max_offset = max(h, w)/4.
|
||||
|
||||
offset = self.offset_conv(x)#.clamp(-max_offset, max_offset)
|
||||
modulator = 2. * torch.sigmoid(self.modulator_conv(x))
|
||||
|
||||
x = deform_conv2d(
|
||||
input=x,
|
||||
offset=offset,
|
||||
weight=self.regular_conv.weight,
|
||||
bias=self.regular_conv.bias,
|
||||
padding=self.padding,
|
||||
mask=modulator,
|
||||
stride=self.stride,
|
||||
)
|
||||
return x
|
||||
@@ -0,0 +1,29 @@
|
||||
import torch.nn as nn
|
||||
from ..modules.mlp import MLPLayer
|
||||
|
||||
|
||||
class BlockA(nn.Module):
|
||||
def __init__(self, in_channels=64, out_channels=64, inter_channels=64, mlp_ratio=4.):
|
||||
super(BlockA, self).__init__()
|
||||
inter_channels = in_channels
|
||||
self.conv = nn.Conv2d(in_channels, inter_channels, 3, 1, 1)
|
||||
self.norm1 = nn.LayerNorm(inter_channels)
|
||||
self.ffn = MLPLayer(in_features=inter_channels,
|
||||
hidden_features=int(inter_channels * mlp_ratio),
|
||||
act_layer=nn.GELU,
|
||||
drop=0.)
|
||||
self.norm2 = nn.LayerNorm(inter_channels)
|
||||
|
||||
def forward(self, x):
|
||||
B, C, H, W = x.shape
|
||||
_x = self.conv(x)
|
||||
_x = _x.flatten(2).transpose(1, 2)
|
||||
_x = self.norm1(_x)
|
||||
x = x.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
|
||||
|
||||
x = x + _x
|
||||
_x1 = self.ffn(x)
|
||||
_x1 = self.norm2(_x1)
|
||||
_x1 = _x1.reshape(B, H, W, -1).permute(0, 3, 1, 2).contiguous()
|
||||
x = x + _x1
|
||||
return x
|
||||
@@ -0,0 +1,21 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from functools import partial
|
||||
|
||||
from ...config import Config
|
||||
|
||||
|
||||
config = Config()
|
||||
|
||||
|
||||
class BasicLatBlk(nn.Module):
|
||||
def __init__(self, in_channels=64, out_channels=64, inter_channels=64):
|
||||
super(BasicLatBlk, self).__init__()
|
||||
inter_channels = in_channels // 4 if config.dec_channels_inter == 'adap' else 64
|
||||
self.conv = nn.Conv2d(in_channels, out_channels, 1, 1, 0)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv(x)
|
||||
return x
|
||||
@@ -0,0 +1,121 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from functools import partial
|
||||
|
||||
try:
|
||||
# version > 0.6.13
|
||||
from timm.layers import DropPath, to_2tuple, trunc_normal_
|
||||
except Exception:
|
||||
from timm.models.layers import DropPath, to_2tuple, trunc_normal_
|
||||
|
||||
import math
|
||||
|
||||
|
||||
class MLPLayer(nn.Module):
|
||||
def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.):
|
||||
super().__init__()
|
||||
out_features = out_features or in_features
|
||||
hidden_features = hidden_features or in_features
|
||||
self.fc1 = nn.Linear(in_features, hidden_features)
|
||||
self.act = act_layer()
|
||||
self.fc2 = nn.Linear(hidden_features, out_features)
|
||||
self.drop = nn.Dropout(drop)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.fc1(x)
|
||||
x = self.act(x)
|
||||
x = self.drop(x)
|
||||
x = self.fc2(x)
|
||||
x = self.drop(x)
|
||||
return x
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0., proj_drop=0., sr_ratio=1):
|
||||
super().__init__()
|
||||
assert dim % num_heads == 0, f"dim {dim} should be divided by num_heads {num_heads}."
|
||||
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
head_dim = dim // num_heads
|
||||
self.scale = qk_scale or head_dim ** -0.5
|
||||
|
||||
self.q = nn.Linear(dim, dim, bias=qkv_bias)
|
||||
self.kv = nn.Linear(dim, dim * 2, bias=qkv_bias)
|
||||
self.attn_drop = nn.Dropout(attn_drop)
|
||||
self.proj = nn.Linear(dim, dim)
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
|
||||
self.sr_ratio = sr_ratio
|
||||
if sr_ratio > 1:
|
||||
self.sr = nn.Conv2d(dim, dim, kernel_size=sr_ratio, stride=sr_ratio)
|
||||
self.norm = nn.LayerNorm(dim)
|
||||
|
||||
def forward(self, x, H, W):
|
||||
B, N, C = x.shape
|
||||
q = self.q(x).reshape(B, N, self.num_heads, C // self.num_heads).permute(0, 2, 1, 3)
|
||||
|
||||
if self.sr_ratio > 1:
|
||||
x_ = x.permute(0, 2, 1).reshape(B, C, H, W)
|
||||
x_ = self.sr(x_).reshape(B, C, -1).permute(0, 2, 1)
|
||||
x_ = self.norm(x_)
|
||||
kv = self.kv(x_).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
|
||||
else:
|
||||
kv = self.kv(x).reshape(B, -1, 2, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4)
|
||||
k, v = kv[0], kv[1]
|
||||
|
||||
attn = (q @ k.transpose(-2, -1)) * self.scale
|
||||
attn = attn.softmax(dim=-1)
|
||||
attn = self.attn_drop(attn)
|
||||
|
||||
x = (attn @ v).transpose(1, 2).reshape(B, N, C)
|
||||
x = self.proj(x)
|
||||
x = self.proj_drop(x)
|
||||
return x
|
||||
|
||||
|
||||
class Block(nn.Module):
|
||||
def __init__(self, dim, num_heads, mlp_ratio=4., qkv_bias=False, qk_scale=None, drop=0., attn_drop=0.,
|
||||
drop_path=0., act_layer=nn.GELU, norm_layer=nn.LayerNorm, sr_ratio=1):
|
||||
super().__init__()
|
||||
self.norm1 = norm_layer(dim)
|
||||
self.attn = Attention(
|
||||
dim,
|
||||
num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale,
|
||||
attn_drop=attn_drop, proj_drop=drop, sr_ratio=sr_ratio)
|
||||
# NOTE: drop path for stochastic depth, we shall see if this is better than dropout here
|
||||
self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
|
||||
self.norm2 = norm_layer(dim)
|
||||
mlp_hidden_dim = int(dim * mlp_ratio)
|
||||
self.mlp = MLPLayer(in_features=dim, hidden_features=mlp_hidden_dim, act_layer=act_layer, drop=drop)
|
||||
|
||||
def forward(self, x, H, W):
|
||||
x = x + self.drop_path(self.attn(self.norm1(x), H, W))
|
||||
x = x + self.drop_path(self.mlp(self.norm2(x), H, W))
|
||||
return x
|
||||
|
||||
|
||||
class OverlapPatchEmbed(nn.Module):
|
||||
""" Image to Patch Embedding
|
||||
"""
|
||||
|
||||
def __init__(self, img_size=224, patch_size=7, stride=4, in_channels=3, embed_dim=768):
|
||||
super().__init__()
|
||||
img_size = to_2tuple(img_size)
|
||||
patch_size = to_2tuple(patch_size)
|
||||
|
||||
self.img_size = img_size
|
||||
self.patch_size = patch_size
|
||||
self.H, self.W = img_size[0] // patch_size[0], img_size[1] // patch_size[1]
|
||||
self.num_patches = self.H * self.W
|
||||
self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=stride,
|
||||
padding=(patch_size[0] // 2, patch_size[1] // 2))
|
||||
self.norm = nn.LayerNorm(embed_dim)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.proj(x)
|
||||
_, _, H, W = x.shape
|
||||
x = x.flatten(2).transpose(1, 2)
|
||||
x = self.norm(x)
|
||||
return x, H, W
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
def build_act_layer(act_layer):
|
||||
if act_layer == 'ReLU':
|
||||
return nn.ReLU(inplace=True)
|
||||
elif act_layer == 'SiLU':
|
||||
return nn.SiLU(inplace=True)
|
||||
elif act_layer == 'GELU':
|
||||
return nn.GELU()
|
||||
|
||||
raise NotImplementedError(f'build_act_layer does not support {act_layer}')
|
||||
|
||||
|
||||
def build_norm_layer(dim,
|
||||
norm_layer,
|
||||
in_format='channels_last',
|
||||
out_format='channels_last',
|
||||
eps=1e-6):
|
||||
layers = []
|
||||
if norm_layer == 'BN':
|
||||
if in_format == 'channels_last':
|
||||
layers.append(to_channels_first())
|
||||
layers.append(nn.BatchNorm2d(dim))
|
||||
if out_format == 'channels_last':
|
||||
layers.append(to_channels_last())
|
||||
elif norm_layer == 'LN':
|
||||
if in_format == 'channels_first':
|
||||
layers.append(to_channels_last())
|
||||
layers.append(nn.LayerNorm(dim, eps=eps))
|
||||
if out_format == 'channels_first':
|
||||
layers.append(to_channels_first())
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f'build_norm_layer does not support {norm_layer}')
|
||||
return nn.Sequential(*layers)
|
||||
|
||||
|
||||
class to_channels_first(nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, x):
|
||||
return x.permute(0, 3, 1, 2)
|
||||
|
||||
|
||||
class to_channels_last(nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, x):
|
||||
return x.permute(0, 2, 3, 1)
|
||||
@@ -0,0 +1,253 @@
|
||||
# import torch
|
||||
# import torch.nn as nn
|
||||
# from collections import OrderedDict
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
# from torchvision.models import vgg16, vgg16_bn
|
||||
# from torchvision.models import resnet50
|
||||
|
||||
from birefnet_old.config import Config
|
||||
from birefnet_old.dataset import class_labels_TR_sorted
|
||||
from birefnet_old.models.backbones.build_backbone import build_backbone
|
||||
from birefnet_old.models.modules.decoder_blocks import BasicDecBlk
|
||||
from birefnet_old.models.modules.lateral_blocks import BasicLatBlk
|
||||
from birefnet_old.models.modules.ing import *
|
||||
from birefnet_old.models.refinement.stem_layer import StemLayer
|
||||
|
||||
|
||||
class RefinerPVTInChannels4(nn.Module):
|
||||
def __init__(self, in_channels=3+1):
|
||||
super(RefinerPVTInChannels4, self).__init__()
|
||||
self.config = Config()
|
||||
self.epoch = 1
|
||||
self.bb = build_backbone(self.config.bb, params_settings='in_channels=4')
|
||||
|
||||
lateral_channels_in_collection = {
|
||||
'vgg16': [512, 256, 128, 64], 'vgg16bn': [512, 256, 128, 64], 'resnet50': [1024, 512, 256, 64],
|
||||
'pvt_v2_b2': [512, 320, 128, 64], 'pvt_v2_b5': [512, 320, 128, 64],
|
||||
'swin_v1_b': [1024, 512, 256, 128], 'swin_v1_l': [1536, 768, 384, 192],
|
||||
}
|
||||
channels = lateral_channels_in_collection[self.config.bb]
|
||||
self.squeeze_module = BasicDecBlk(channels[0], channels[0])
|
||||
|
||||
self.decoder = Decoder(channels)
|
||||
|
||||
if 0:
|
||||
for key, value in self.named_parameters():
|
||||
if 'bb.' in key:
|
||||
value.requires_grad = False
|
||||
|
||||
def forward(self, x):
|
||||
if isinstance(x, list):
|
||||
x = torch.cat(x, dim=1)
|
||||
########## Encoder ##########
|
||||
if self.config.bb in ['vgg16', 'vgg16bn', 'resnet50']:
|
||||
x1 = self.bb.conv1(x)
|
||||
x2 = self.bb.conv2(x1)
|
||||
x3 = self.bb.conv3(x2)
|
||||
x4 = self.bb.conv4(x3)
|
||||
else:
|
||||
x1, x2, x3, x4 = self.bb(x)
|
||||
|
||||
x4 = self.squeeze_module(x4)
|
||||
|
||||
########## Decoder ##########
|
||||
|
||||
features = [x, x1, x2, x3, x4]
|
||||
scaled_preds = self.decoder(features)
|
||||
|
||||
return scaled_preds
|
||||
|
||||
|
||||
class Refiner(nn.Module):
|
||||
def __init__(self, in_channels=3+1):
|
||||
super(Refiner, self).__init__()
|
||||
self.config = Config()
|
||||
self.epoch = 1
|
||||
self.stem_layer = StemLayer(in_channels=in_channels, inter_channels=48, out_channels=3)
|
||||
self.bb = build_backbone(self.config.bb)
|
||||
|
||||
lateral_channels_in_collection = {
|
||||
'vgg16': [512, 256, 128, 64], 'vgg16bn': [512, 256, 128, 64], 'resnet50': [1024, 512, 256, 64],
|
||||
'pvt_v2_b2': [512, 320, 128, 64], 'pvt_v2_b5': [512, 320, 128, 64],
|
||||
'swin_v1_b': [1024, 512, 256, 128], 'swin_v1_l': [1536, 768, 384, 192],
|
||||
}
|
||||
channels = lateral_channels_in_collection[self.config.bb]
|
||||
self.squeeze_module = BasicDecBlk(channels[0], channels[0])
|
||||
|
||||
self.decoder = Decoder(channels)
|
||||
|
||||
if 0:
|
||||
for key, value in self.named_parameters():
|
||||
if 'bb.' in key:
|
||||
value.requires_grad = False
|
||||
|
||||
def forward(self, x):
|
||||
if isinstance(x, list):
|
||||
x = torch.cat(x, dim=1)
|
||||
x = self.stem_layer(x)
|
||||
########## Encoder ##########
|
||||
if self.config.bb in ['vgg16', 'vgg16bn', 'resnet50']:
|
||||
x1 = self.bb.conv1(x)
|
||||
x2 = self.bb.conv2(x1)
|
||||
x3 = self.bb.conv3(x2)
|
||||
x4 = self.bb.conv4(x3)
|
||||
else:
|
||||
x1, x2, x3, x4 = self.bb(x)
|
||||
|
||||
x4 = self.squeeze_module(x4)
|
||||
|
||||
########## Decoder ##########
|
||||
|
||||
features = [x, x1, x2, x3, x4]
|
||||
scaled_preds = self.decoder(features)
|
||||
|
||||
return scaled_preds
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(self, channels):
|
||||
super(Decoder, self).__init__()
|
||||
self.config = Config()
|
||||
DecoderBlock = eval('BasicDecBlk')
|
||||
LateralBlock = eval('BasicLatBlk')
|
||||
|
||||
self.decoder_block4 = DecoderBlock(channels[0], channels[1])
|
||||
self.decoder_block3 = DecoderBlock(channels[1], channels[2])
|
||||
self.decoder_block2 = DecoderBlock(channels[2], channels[3])
|
||||
self.decoder_block1 = DecoderBlock(channels[3], channels[3]//2)
|
||||
|
||||
self.lateral_block4 = LateralBlock(channels[1], channels[1])
|
||||
self.lateral_block3 = LateralBlock(channels[2], channels[2])
|
||||
self.lateral_block2 = LateralBlock(channels[3], channels[3])
|
||||
|
||||
if self.config.ms_supervision:
|
||||
self.conv_ms_spvn_4 = nn.Conv2d(channels[1], 1, 1, 1, 0)
|
||||
self.conv_ms_spvn_3 = nn.Conv2d(channels[2], 1, 1, 1, 0)
|
||||
self.conv_ms_spvn_2 = nn.Conv2d(channels[3], 1, 1, 1, 0)
|
||||
self.conv_out1 = nn.Sequential(nn.Conv2d(channels[3]//2, 1, 1, 1, 0))
|
||||
|
||||
def forward(self, features):
|
||||
x, x1, x2, x3, x4 = features
|
||||
outs = []
|
||||
p4 = self.decoder_block4(x4)
|
||||
_p4 = F.interpolate(p4, size=x3.shape[2:], mode='bilinear', align_corners=True)
|
||||
_p3 = _p4 + self.lateral_block4(x3)
|
||||
|
||||
p3 = self.decoder_block3(_p3)
|
||||
_p3 = F.interpolate(p3, size=x2.shape[2:], mode='bilinear', align_corners=True)
|
||||
_p2 = _p3 + self.lateral_block3(x2)
|
||||
|
||||
p2 = self.decoder_block2(_p2)
|
||||
_p2 = F.interpolate(p2, size=x1.shape[2:], mode='bilinear', align_corners=True)
|
||||
_p1 = _p2 + self.lateral_block2(x1)
|
||||
|
||||
_p1 = self.decoder_block1(_p1)
|
||||
_p1 = F.interpolate(_p1, size=x.shape[2:], mode='bilinear', align_corners=True)
|
||||
p1_out = self.conv_out1(_p1)
|
||||
|
||||
if self.config.ms_supervision:
|
||||
outs.append(self.conv_ms_spvn_4(p4))
|
||||
outs.append(self.conv_ms_spvn_3(p3))
|
||||
outs.append(self.conv_ms_spvn_2(p2))
|
||||
outs.append(p1_out)
|
||||
return outs
|
||||
|
||||
|
||||
class RefUNet(nn.Module):
|
||||
# Refinement
|
||||
def __init__(self, in_channels=3+1):
|
||||
super(RefUNet, self).__init__()
|
||||
self.encoder_1 = nn.Sequential(
|
||||
nn.Conv2d(in_channels, 64, 3, 1, 1),
|
||||
nn.Conv2d(64, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.encoder_2 = nn.Sequential(
|
||||
nn.MaxPool2d(2, 2, ceil_mode=True),
|
||||
nn.Conv2d(64, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.encoder_3 = nn.Sequential(
|
||||
nn.MaxPool2d(2, 2, ceil_mode=True),
|
||||
nn.Conv2d(64, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.encoder_4 = nn.Sequential(
|
||||
nn.MaxPool2d(2, 2, ceil_mode=True),
|
||||
nn.Conv2d(64, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.pool4 = nn.MaxPool2d(2, 2, ceil_mode=True)
|
||||
#####
|
||||
self.decoder_5 = nn.Sequential(
|
||||
nn.Conv2d(64, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
#####
|
||||
self.decoder_4 = nn.Sequential(
|
||||
nn.Conv2d(128, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.decoder_3 = nn.Sequential(
|
||||
nn.Conv2d(128, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.decoder_2 = nn.Sequential(
|
||||
nn.Conv2d(128, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.decoder_1 = nn.Sequential(
|
||||
nn.Conv2d(128, 64, 3, 1, 1),
|
||||
nn.BatchNorm2d(64),
|
||||
nn.ReLU(inplace=True)
|
||||
)
|
||||
|
||||
self.conv_d0 = nn.Conv2d(64, 1, 3, 1, 1)
|
||||
|
||||
self.upscore2 = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
|
||||
|
||||
def forward(self, x):
|
||||
outs = []
|
||||
if isinstance(x, list):
|
||||
x = torch.cat(x, dim=1)
|
||||
hx = x
|
||||
|
||||
hx1 = self.encoder_1(hx)
|
||||
hx2 = self.encoder_2(hx1)
|
||||
hx3 = self.encoder_3(hx2)
|
||||
hx4 = self.encoder_4(hx3)
|
||||
|
||||
hx = self.decoder_5(self.pool4(hx4))
|
||||
hx = torch.cat((self.upscore2(hx), hx4), 1)
|
||||
|
||||
d4 = self.decoder_4(hx)
|
||||
hx = torch.cat((self.upscore2(d4), hx3), 1)
|
||||
|
||||
d3 = self.decoder_3(hx)
|
||||
hx = torch.cat((self.upscore2(d3), hx2), 1)
|
||||
|
||||
d2 = self.decoder_2(hx)
|
||||
hx = torch.cat((self.upscore2(d2), hx1), 1)
|
||||
|
||||
d1 = self.decoder_1(hx)
|
||||
|
||||
x = self.conv_d0(d1)
|
||||
outs.append(x)
|
||||
return outs
|
||||
@@ -0,0 +1,45 @@
|
||||
import torch.nn as nn
|
||||
from birefnet_old.models.modules.utils import build_act_layer, build_norm_layer
|
||||
|
||||
|
||||
class StemLayer(nn.Module):
|
||||
r""" Stem layer of InternImage
|
||||
Args:
|
||||
in_channels (int): number of input channels
|
||||
out_channels (int): number of output channels
|
||||
act_layer (str): activation layer
|
||||
norm_layer (str): normalization layer
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
in_channels=3+1,
|
||||
inter_channels=48,
|
||||
out_channels=96,
|
||||
act_layer='GELU',
|
||||
norm_layer='BN'):
|
||||
super().__init__()
|
||||
self.conv1 = nn.Conv2d(in_channels,
|
||||
inter_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
self.norm1 = build_norm_layer(
|
||||
inter_channels, norm_layer, 'channels_first', 'channels_first'
|
||||
)
|
||||
self.act = build_act_layer(act_layer)
|
||||
self.conv2 = nn.Conv2d(inter_channels,
|
||||
out_channels,
|
||||
kernel_size=3,
|
||||
stride=1,
|
||||
padding=1)
|
||||
self.norm2 = build_norm_layer(
|
||||
out_channels, norm_layer, 'channels_first', 'channels_first'
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv1(x)
|
||||
x = self.norm1(x)
|
||||
x = self.act(x)
|
||||
x = self.conv2(x)
|
||||
x = self.norm2(x)
|
||||
return x
|
||||
@@ -0,0 +1,85 @@
|
||||
from PIL import Image, ImageEnhance
|
||||
import random
|
||||
import numpy as np
|
||||
import random
|
||||
|
||||
|
||||
def preproc(image, label, preproc_methods=['flip']):
|
||||
if 'flip' in preproc_methods:
|
||||
image, label = cv_random_flip(image, label)
|
||||
if 'crop' in preproc_methods:
|
||||
image, label = random_crop(image, label)
|
||||
if 'rotate' in preproc_methods:
|
||||
image, label = random_rotate(image, label)
|
||||
if 'enhance' in preproc_methods:
|
||||
image = color_enhance(image)
|
||||
if 'pepper' in preproc_methods:
|
||||
label = random_pepper(label)
|
||||
return image, label
|
||||
|
||||
|
||||
def cv_random_flip(img, label):
|
||||
if random.random() > 0.5:
|
||||
img = img.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
label = label.transpose(Image.FLIP_LEFT_RIGHT)
|
||||
return img, label
|
||||
|
||||
|
||||
def random_crop(image, label):
|
||||
border = 30
|
||||
image_width = image.size[0]
|
||||
image_height = image.size[1]
|
||||
border = int(min(image_width, image_height) * 0.1)
|
||||
crop_win_width = np.random.randint(image_width - border, image_width)
|
||||
crop_win_height = np.random.randint(image_height - border, image_height)
|
||||
random_region = (
|
||||
(image_width - crop_win_width) >> 1, (image_height - crop_win_height) >> 1, (image_width + crop_win_width) >> 1,
|
||||
(image_height + crop_win_height) >> 1)
|
||||
return image.crop(random_region), label.crop(random_region)
|
||||
|
||||
|
||||
def random_rotate(image, label, angle=15):
|
||||
mode = Image.BICUBIC
|
||||
if random.random() > 0.8:
|
||||
random_angle = np.random.randint(-angle, angle)
|
||||
image = image.rotate(random_angle, mode)
|
||||
label = label.rotate(random_angle, mode)
|
||||
return image, label
|
||||
|
||||
|
||||
def color_enhance(image):
|
||||
bright_intensity = random.randint(5, 15) / 10.0
|
||||
image = ImageEnhance.Brightness(image).enhance(bright_intensity)
|
||||
contrast_intensity = random.randint(5, 15) / 10.0
|
||||
image = ImageEnhance.Contrast(image).enhance(contrast_intensity)
|
||||
color_intensity = random.randint(0, 20) / 10.0
|
||||
image = ImageEnhance.Color(image).enhance(color_intensity)
|
||||
sharp_intensity = random.randint(0, 30) / 10.0
|
||||
image = ImageEnhance.Sharpness(image).enhance(sharp_intensity)
|
||||
return image
|
||||
|
||||
|
||||
def random_gaussian(image, mean=0.1, sigma=0.35):
|
||||
def gaussianNoisy(im, mean=mean, sigma=sigma):
|
||||
for _i in range(len(im)):
|
||||
im[_i] += random.gauss(mean, sigma)
|
||||
return im
|
||||
|
||||
img = np.asarray(image)
|
||||
width, height = img.shape
|
||||
img = gaussianNoisy(img[:].flatten(), mean, sigma)
|
||||
img = img.reshape([width, height])
|
||||
return Image.fromarray(np.uint8(img))
|
||||
|
||||
|
||||
def random_pepper(img, N=0.0015):
|
||||
img = np.array(img)
|
||||
noiseNum = int(N * img.shape[0] * img.shape[1])
|
||||
for i in range(noiseNum):
|
||||
randX = random.randint(0, img.shape[0] - 1)
|
||||
randY = random.randint(0, img.shape[1] - 1)
|
||||
if random.randint(0, 1) == 0:
|
||||
img[randX, randY] = 0
|
||||
else:
|
||||
img[randX, randY] = 255
|
||||
return Image.fromarray(img)
|
||||
@@ -0,0 +1,97 @@
|
||||
import logging
|
||||
import os
|
||||
import torch
|
||||
from torchvision import transforms
|
||||
import numpy as np
|
||||
import random
|
||||
import cv2
|
||||
from PIL import Image
|
||||
|
||||
|
||||
def path_to_image(path, size=(1024, 1024), color_type=['rgb', 'gray'][0]):
|
||||
if color_type.lower() == 'rgb':
|
||||
image = cv2.imread(path)
|
||||
elif color_type.lower() == 'gray':
|
||||
image = cv2.imread(path, cv2.IMREAD_GRAYSCALE)
|
||||
else:
|
||||
print('Select the color_type to return, either to RGB or gray image.')
|
||||
return
|
||||
if size:
|
||||
image = cv2.resize(image, size, interpolation=cv2.INTER_LINEAR)
|
||||
if color_type.lower() == 'rgb':
|
||||
image = Image.fromarray(cv2.cvtColor(image, cv2.COLOR_BGR2RGB)).convert('RGB')
|
||||
else:
|
||||
image = Image.fromarray(image).convert('L')
|
||||
return image
|
||||
|
||||
|
||||
|
||||
def check_state_dict(state_dict, unwanted_prefix='_orig_mod.'):
|
||||
for k, v in list(state_dict.items()):
|
||||
if k.startswith(unwanted_prefix):
|
||||
state_dict[k[len(unwanted_prefix):]] = state_dict.pop(k)
|
||||
return state_dict
|
||||
|
||||
|
||||
def generate_smoothed_gt(gts):
|
||||
epsilon = 0.001
|
||||
new_gts = (1-epsilon)*gts+epsilon/2
|
||||
return new_gts
|
||||
|
||||
|
||||
class Logger():
|
||||
def __init__(self, path="log.txt"):
|
||||
self.logger = logging.getLogger('BiRefNet')
|
||||
self.file_handler = logging.FileHandler(path, "w")
|
||||
self.stdout_handler = logging.StreamHandler()
|
||||
self.stdout_handler.setFormatter(logging.Formatter('%(asctime)s %(levelname)s %(message)s'))
|
||||
self.file_handler.setFormatter(logging.Formatter('%(asctime)s %(levelname)s %(message)s'))
|
||||
self.logger.addHandler(self.file_handler)
|
||||
self.logger.addHandler(self.stdout_handler)
|
||||
self.logger.setLevel(logging.INFO)
|
||||
self.logger.propagate = False
|
||||
|
||||
def info(self, txt):
|
||||
self.logger.info(txt)
|
||||
|
||||
def close(self):
|
||||
self.file_handler.close()
|
||||
self.stdout_handler.close()
|
||||
|
||||
|
||||
class AverageMeter(object):
|
||||
"""Computes and stores the average and current value"""
|
||||
def __init__(self):
|
||||
self.reset()
|
||||
|
||||
def reset(self):
|
||||
self.val = 0.0
|
||||
self.avg = 0.0
|
||||
self.sum = 0.0
|
||||
self.count = 0.0
|
||||
|
||||
def update(self, val, n=1):
|
||||
self.val = val
|
||||
self.sum += val * n
|
||||
self.count += n
|
||||
self.avg = self.sum / self.count
|
||||
|
||||
|
||||
def save_checkpoint(state, path, filename="latest.pth"):
|
||||
torch.save(state, os.path.join(path, filename))
|
||||
|
||||
|
||||
def save_tensor_img(tenor_im, path):
|
||||
im = tenor_im.cpu().clone()
|
||||
im = im.squeeze(0)
|
||||
tensor2pil = transforms.ToPILImage()
|
||||
im = tensor2pil(im)
|
||||
im.save(path)
|
||||
|
||||
|
||||
def set_seed(seed):
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
np.random.seed(seed)
|
||||
random.seed(seed)
|
||||
torch.backends.cudnn.deterministic = True
|
||||
+15
@@ -0,0 +1,15 @@
|
||||
[project]
|
||||
name = "comfyui_birefnet_ll"
|
||||
description = "Sync with version of BiRefNet. NODES:AutoDownloadBiRefNetModel, LoadRembgByBiRefNetModel, RembgByBiRefNet, RembgByBiRefNetAdvanced, GetMaskByBiRefNet, BlurFusionForegroundEstimation."
|
||||
version = "1.1.4"
|
||||
license = {file = "LICENSE"}
|
||||
dependencies = ["numpy", "opencv-python", "timm"]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/lldacing/ComfyUI_BiRefNet_ll"
|
||||
# Used by Comfy Registry https://comfyregistry.org
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "lldacing"
|
||||
DisplayName = "ComfyUI_BiRefNet_ll"
|
||||
Icon = ""
|
||||
+3
@@ -0,0 +1,3 @@
|
||||
numpy
|
||||
opencv-python
|
||||
timm
|
||||
+1127
File diff suppressed because it is too large
Load Diff
+267
@@ -0,0 +1,267 @@
|
||||
<!DOCTYPE html>
|
||||
<html lang="zh-CN">
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||
<title>BiRefNet WebUI · 智能抠图</title>
|
||||
<link rel="icon" href="data:image/svg+xml,%3Csvg xmlns='http://www.w3.org/2000/svg' viewBox='0 0 32 32'%3E%3Crect width='32' height='32' rx='8' fill='%234f8cff'/%3E%3Cpath d='M9 21c4-1 5-10 14-10-4 1-5 10-14 10z' fill='white'/%3E%3C/svg%3E">
|
||||
<link rel="stylesheet" href="/static/style.css">
|
||||
</head>
|
||||
<body>
|
||||
<header class="topbar">
|
||||
<div class="brand">
|
||||
<span class="logo" aria-hidden="true">
|
||||
<svg viewBox="0 0 32 32" width="26" height="26"><rect width="32" height="32" rx="8" fill="var(--accent)"/><path d="M9 21c4-1 5-10 14-10-4 1-5 10-14 10z" fill="#fff"/></svg>
|
||||
</span>
|
||||
<h1>BiRefNet <em>WebUI</em></h1>
|
||||
</div>
|
||||
|
||||
<dl class="meta" id="meta">
|
||||
<div><dt>设备</dt><dd id="meta-device">检测中…</dd></div>
|
||||
<div><dt>模型</dt><dd id="meta-models">—</dd></div>
|
||||
<div><dt>模型代码</dt><dd id="meta-node" class="path" title="">—</dd></div>
|
||||
</dl>
|
||||
|
||||
<div class="topbar-actions">
|
||||
<button class="btn ghost sm" id="btn-reload-models" type="button">重新扫描模型</button>
|
||||
<button class="btn ghost sm" id="btn-unload" type="button">卸载显存</button>
|
||||
</div>
|
||||
</header>
|
||||
|
||||
<main class="layout">
|
||||
<!-- ============================= 参数面板 ============================= -->
|
||||
<aside class="sidebar">
|
||||
<section class="block">
|
||||
<header class="block-head"><h2>1 · 图片</h2><span class="hint" id="file-count">未选择</span></header>
|
||||
<div class="dropzone compact" id="dropzone" tabindex="0" role="button" aria-label="选择或拖入图片">
|
||||
<p>拖入图片 / 点击选择 / <kbd>Ctrl</kbd>+<kbd>V</kbd> 粘贴</p>
|
||||
<p class="sub">支持 PNG · JPG · WebP · BMP · TIFF · 一次一张,新图会替换上一张</p>
|
||||
<input type="file" id="file-input" accept="image/*" hidden>
|
||||
</div>
|
||||
<ul class="thumbs" id="thumbs"></ul>
|
||||
<p class="pending-note" id="pending-note"></p>
|
||||
</section>
|
||||
|
||||
<section class="block">
|
||||
<header class="block-head"><h2>2 · 模型</h2><span class="hint" id="model-hint"></span></header>
|
||||
<label class="field">
|
||||
<span>权重文件</span>
|
||||
<select id="opt-model"></select>
|
||||
</label>
|
||||
<div class="field-row">
|
||||
<label class="field">
|
||||
<span>设备</span>
|
||||
<select id="opt-device">
|
||||
<option value="auto">自动 (GPU)</option>
|
||||
<option value="cpu">CPU</option>
|
||||
</select>
|
||||
</label>
|
||||
<label class="field">
|
||||
<span>精度</span>
|
||||
<select id="opt-dtype">
|
||||
<option value="auto">自动 (fp16 加速)</option>
|
||||
<option value="float32">float32 · 最稳</option>
|
||||
<option value="float16">float16 · 最省显存</option>
|
||||
<option value="bfloat16">bfloat16</option>
|
||||
</select>
|
||||
</label>
|
||||
</div>
|
||||
<label class="field">
|
||||
<span>架构</span>
|
||||
<select id="opt-arch">
|
||||
<option value="auto">自动识别</option>
|
||||
<option value="v1">v1(新版,safetensors)</option>
|
||||
<option value="old">old(旧版 .pth)</option>
|
||||
</select>
|
||||
</label>
|
||||
</section>
|
||||
|
||||
<section class="block">
|
||||
<header class="block-head"><h2>3 · 参数</h2><span class="hint">按需调整</span></header>
|
||||
|
||||
<label class="field">
|
||||
<span>预处理尺寸</span>
|
||||
<select id="opt-resolution-mode">
|
||||
<option value="square">固定 1024×1024(推荐)</option>
|
||||
<option value="longest">长边自适应(保宽高比)</option>
|
||||
<option value="custom">自定义宽高</option>
|
||||
</select>
|
||||
</label>
|
||||
|
||||
<div class="field-row" id="row-longest">
|
||||
<label class="field">
|
||||
<span>长边像素</span>
|
||||
<input type="number" id="opt-longest-side" value="1024" min="32" max="4096" step="32">
|
||||
</label>
|
||||
</div>
|
||||
|
||||
<div class="field-row hidden" id="row-wh">
|
||||
<label class="field">
|
||||
<span>宽</span>
|
||||
<input type="number" id="opt-width" value="1024" min="32" max="4096" step="32">
|
||||
</label>
|
||||
<label class="field">
|
||||
<span>高</span>
|
||||
<input type="number" id="opt-height" value="1024" min="32" max="4096" step="32">
|
||||
</label>
|
||||
</div>
|
||||
|
||||
<div class="field-row">
|
||||
<label class="field">
|
||||
<span>插值方式</span>
|
||||
<select id="opt-upscale">
|
||||
<option value="bilinear">bilinear</option>
|
||||
<option value="bicubic">bicubic</option>
|
||||
<option value="nearest">nearest</option>
|
||||
<option value="nearest-exact">nearest-exact</option>
|
||||
</select>
|
||||
</label>
|
||||
<label class="field">
|
||||
<span>遮罩阈值</span>
|
||||
<input type="number" id="opt-threshold" value="0" min="0" max="1" step="0.004">
|
||||
</label>
|
||||
</div>
|
||||
|
||||
<label class="field">
|
||||
<span>背景</span>
|
||||
<select id="opt-background">
|
||||
<option value="transparent">透明(PNG Alpha)</option>
|
||||
<option value="color">纯色背景</option>
|
||||
</select>
|
||||
</label>
|
||||
|
||||
<div class="field-row hidden" id="row-bgcolor">
|
||||
<label class="field">
|
||||
<span>背景色</span>
|
||||
<input type="color" id="opt-bgcolor" value="#ffffff">
|
||||
</label>
|
||||
<div class="swatches" id="swatches">
|
||||
<button type="button" data-color="#ffffff" style="--c:#ffffff" title="白"></button>
|
||||
<button type="button" data-color="#000000" style="--c:#000000" title="黑"></button>
|
||||
<button type="button" data-color="#00ff00" style="--c:#00ff00" title="绿幕"></button>
|
||||
<button type="button" data-color="#4f8cff" style="--c:#4f8cff" title="蓝"></button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<label class="check">
|
||||
<input type="checkbox" id="opt-refine" checked>
|
||||
<span>前景精修(去除边缘残留背景色)</span>
|
||||
</label>
|
||||
<div class="field-row" id="row-blur">
|
||||
<label class="field">
|
||||
<span>大核 r1</span>
|
||||
<input type="number" id="opt-blur1" value="90" min="1" max="255">
|
||||
</label>
|
||||
<label class="field">
|
||||
<span>小核 r2</span>
|
||||
<input type="number" id="opt-blur2" value="6" min="1" max="255">
|
||||
</label>
|
||||
</div>
|
||||
|
||||
<label class="check">
|
||||
<input type="checkbox" id="opt-mask">
|
||||
<span>同时输出遮罩(灰度 PNG)</span>
|
||||
</label>
|
||||
|
||||
<label class="field">
|
||||
<span>最终尺寸 · 最长边(0 = 保持原图尺寸)</span>
|
||||
<input type="number" id="opt-final-side" value="0" min="0" max="32768" step="16">
|
||||
</label>
|
||||
<p class="field-tip">非 0 时按原图比例等比缩放,使结果最长边等于该值。例:2000×3000 填 1440 → 960×1440</p>
|
||||
</section>
|
||||
|
||||
<section class="block actions">
|
||||
<button class="btn primary" id="btn-run" type="button" disabled>开始抠图</button>
|
||||
<div class="progress" id="progress" hidden>
|
||||
<div class="bar"><i id="progress-bar"></i></div>
|
||||
<div class="progress-meta">
|
||||
<span id="progress-stage">等待中</span>
|
||||
<button class="btn link" id="btn-cancel" type="button">取消</button>
|
||||
</div>
|
||||
</div>
|
||||
<p class="error-text" id="error-text" hidden></p>
|
||||
</section>
|
||||
|
||||
<section class="block">
|
||||
<header class="block-head"><h2>4 · 批量处理</h2><span class="hint">按目录批量抠图</span></header>
|
||||
|
||||
<label class="field">
|
||||
<span>图片目录(模板路径)</span>
|
||||
<input type="text" id="batch-input-dir" placeholder="例:F:\photos\待抠图" spellcheck="false" autocomplete="off">
|
||||
</label>
|
||||
|
||||
<label class="field">
|
||||
<span>输出目录</span>
|
||||
<input type="text" id="batch-output-dir" placeholder="默认:项目 outputs 目录" spellcheck="false" autocomplete="off">
|
||||
</label>
|
||||
|
||||
<label class="check">
|
||||
<input type="checkbox" id="batch-recursive">
|
||||
<span>包含子目录</span>
|
||||
</label>
|
||||
|
||||
<button class="btn" id="btn-batch" type="button">开始批量抠图</button>
|
||||
|
||||
<div class="progress" id="batch-progress" hidden>
|
||||
<div class="bar"><i id="batch-progress-bar"></i></div>
|
||||
<div class="progress-meta">
|
||||
<span id="batch-progress-stage">准备中…</span>
|
||||
<button class="btn link" id="btn-batch-cancel" type="button">取消</button>
|
||||
</div>
|
||||
</div>
|
||||
<p class="error-text" id="batch-error" hidden></p>
|
||||
<p class="batch-note">
|
||||
两个目录都必须<strong>已存在</strong>(本工具不会自动创建,路径不对会直接报错)。结果写入输出目录,命名
|
||||
<code>RMBG_原名_时间戳.后缀</code>,原名超 20 字符截断;抠图参数沿用左侧当前设置。
|
||||
</p>
|
||||
</section>
|
||||
</aside>
|
||||
|
||||
<!-- ============================= 结果区 ============================= -->
|
||||
<section class="stage">
|
||||
<header class="stage-head">
|
||||
<h2>结果 <span class="count" id="result-count">0</span></h2>
|
||||
<div class="stage-tools">
|
||||
<div class="segmented" id="view-switch">
|
||||
<button type="button" data-view="cutout" class="active">抠图</button>
|
||||
<button type="button" data-view="mask">遮罩</button>
|
||||
<button type="button" data-view="original">原图</button>
|
||||
</div>
|
||||
<button class="btn ghost sm" id="btn-zip" type="button" disabled>打包下载 ZIP</button>
|
||||
<button class="btn ghost sm" id="btn-clear" type="button" title="清空整个结果列表" disabled>清空</button>
|
||||
</div>
|
||||
</header>
|
||||
|
||||
<p class="stage-tip">结果会一直保留在列表里(传新图、改参数都不会清空)· 点文件名或「重新处理」会套用它上次使用的参数 · 点卡片选中 · 点 × 移除单条</p>
|
||||
|
||||
<section class="batch-report" id="batch-report" hidden>
|
||||
<header class="batch-report-head">
|
||||
<h3>批量任务 <span class="badge" id="batch-badge">排队中</span></h3>
|
||||
<span class="batch-report-path" id="batch-report-path" title=""></span>
|
||||
<button class="btn ghost sm" id="btn-batch-hide" type="button">收起</button>
|
||||
</header>
|
||||
<p class="batch-report-sum" id="batch-report-sum"></p>
|
||||
<ul class="batch-rows" id="batch-rows"></ul>
|
||||
</section>
|
||||
|
||||
<div class="empty-state" id="empty">
|
||||
<div class="empty-art" aria-hidden="true">
|
||||
<svg viewBox="0 0 240 140" width="240" height="140">
|
||||
<rect x="1" y="1" width="238" height="138" rx="14" fill="none" stroke="var(--border)" stroke-dasharray="6 6"/>
|
||||
<circle cx="70" cy="62" r="22" fill="none" stroke="var(--accent)" stroke-width="2"/>
|
||||
<path d="M104 92c14-4 18-40 52-40-14 4-18 40-52 40z" fill="var(--accent)" opacity=".85"/>
|
||||
<path d="M46 108h148" stroke="var(--border)" stroke-width="2" stroke-linecap="round"/>
|
||||
</svg>
|
||||
</div>
|
||||
<h3>还没有结果</h3>
|
||||
<p>左侧拖入图片 → 选择模型 → 点击「开始抠图」。处理完成后可拖拽分割线对比原图与抠图效果,记录会累积在右侧供随时重新处理。</p>
|
||||
</div>
|
||||
|
||||
<div class="grid" id="results"></div>
|
||||
</section>
|
||||
</main>
|
||||
|
||||
<script src="/static/zip.js"></script>
|
||||
<script src="/static/app.js"></script>
|
||||
</body>
|
||||
</html>
|
||||
+401
@@ -0,0 +1,401 @@
|
||||
/* =========================================================================
|
||||
BiRefNet WebUI · 深色主题
|
||||
设计约束:字号层级 4 档 · 字体 2 种(系统 UI + 等宽)· 色彩角色 4 个
|
||||
(accent / ok / warn / danger)· 单焦点(开始抠图)
|
||||
========================================================================= */
|
||||
:root {
|
||||
--bg: #0d0f13;
|
||||
--surface: #141821;
|
||||
--surface-2: #1a1f2a;
|
||||
--surface-3: #202634;
|
||||
--border: #252b38;
|
||||
--border-soft: #1e2430;
|
||||
--text: #e9edf5;
|
||||
--text-dim: #95a0b3;
|
||||
--text-faint: #6b7688;
|
||||
|
||||
--accent: #4f8cff;
|
||||
--accent-hover: #6ba0ff;
|
||||
--accent-soft: rgba(79, 140, 255, 0.14);
|
||||
--ok: #35c98a;
|
||||
--warn: #f0b429;
|
||||
--danger: #f2555a;
|
||||
|
||||
--font: -apple-system, BlinkMacSystemFont, "Segoe UI", "Microsoft YaHei", "PingFang SC", system-ui, sans-serif;
|
||||
--mono: ui-monospace, SFMono-Regular, Consolas, "Courier New", monospace;
|
||||
|
||||
--radius: 12px;
|
||||
--radius-sm: 8px;
|
||||
--shadow: 0 8px 28px rgba(0, 0, 0, 0.35);
|
||||
}
|
||||
|
||||
* { box-sizing: border-box; }
|
||||
|
||||
html, body {
|
||||
margin: 0;
|
||||
height: 100%;
|
||||
background: var(--bg);
|
||||
color: var(--text);
|
||||
font-family: var(--font);
|
||||
font-size: 14px;
|
||||
line-height: 1.5;
|
||||
-webkit-font-smoothing: antialiased;
|
||||
}
|
||||
|
||||
body { display: flex; flex-direction: column; overflow: hidden; }
|
||||
|
||||
h1, h2, h3 { margin: 0; font-weight: 600; }
|
||||
button, input, select { font: inherit; color: inherit; }
|
||||
::-webkit-scrollbar { width: 10px; height: 10px; }
|
||||
::-webkit-scrollbar-thumb { background: #2b3242; border-radius: 8px; border: 2px solid var(--bg); }
|
||||
::-webkit-scrollbar-track { background: transparent; }
|
||||
|
||||
/* ------------------------------- 顶栏 ------------------------------- */
|
||||
.topbar {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 24px;
|
||||
padding: 14px 22px;
|
||||
background: var(--surface);
|
||||
border-bottom: 1px solid var(--border);
|
||||
flex: 0 0 auto;
|
||||
}
|
||||
.brand { display: flex; align-items: center; gap: 12px; }
|
||||
.brand .logo { display: grid; place-items: center; }
|
||||
.brand h1 { font-size: 28px; letter-spacing: -0.5px; line-height: 1.1; }
|
||||
.brand h1 em { font-style: normal; font-weight: 300; color: var(--text-dim); }
|
||||
|
||||
.meta { display: flex; gap: 26px; margin: 0; margin-left: auto; }
|
||||
.meta > div { display: flex; flex-direction: column; gap: 2px; min-width: 0; }
|
||||
.meta dt { font-size: 11px; letter-spacing: 0.08em; text-transform: uppercase; color: var(--text-faint); }
|
||||
.meta dd { margin: 0; font-size: 13px; color: var(--text); font-family: var(--mono); }
|
||||
.meta dd.path { max-width: 320px; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; color: var(--text-dim); }
|
||||
.topbar-actions { display: flex; gap: 8px; }
|
||||
|
||||
/* ------------------------------- 主体 ------------------------------- */
|
||||
.layout { flex: 1 1 auto; display: grid; grid-template-columns: 348px 1fr; min-height: 0; }
|
||||
|
||||
.sidebar {
|
||||
background: var(--surface);
|
||||
border-right: 1px solid var(--border);
|
||||
overflow-y: auto;
|
||||
padding: 16px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 14px;
|
||||
}
|
||||
|
||||
.block {
|
||||
background: var(--surface-2);
|
||||
border: 1px solid var(--border-soft);
|
||||
border-radius: var(--radius);
|
||||
padding: 14px;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 10px;
|
||||
}
|
||||
.block-head { display: flex; align-items: baseline; justify-content: space-between; gap: 8px; }
|
||||
.block-head h2 { font-size: 13px; letter-spacing: 0.02em; color: var(--text); }
|
||||
.hint { font-size: 11px; color: var(--text-faint); font-family: var(--mono); }
|
||||
|
||||
/* 表单 */
|
||||
.field { display: flex; flex-direction: column; gap: 5px; min-width: 0; flex: 1; }
|
||||
.field > span { font-size: 12px; color: var(--text-dim); }
|
||||
.field-row { display: flex; gap: 10px; }
|
||||
select, input[type="number"], input[type="text"] {
|
||||
width: 100%;
|
||||
padding: 7px 9px;
|
||||
background: var(--bg);
|
||||
border: 1px solid var(--border);
|
||||
border-radius: var(--radius-sm);
|
||||
font-family: var(--mono);
|
||||
font-size: 12.5px;
|
||||
transition: border-color .15s, box-shadow .15s;
|
||||
}
|
||||
select:focus, input:focus { outline: none; border-color: var(--accent); box-shadow: 0 0 0 3px var(--accent-soft); }
|
||||
input[type="color"] {
|
||||
width: 100%; height: 34px; padding: 3px; background: var(--bg);
|
||||
border: 1px solid var(--border); border-radius: var(--radius-sm); cursor: pointer;
|
||||
}
|
||||
.check { display: flex; align-items: center; gap: 8px; font-size: 12.5px; color: var(--text-dim); cursor: pointer; }
|
||||
.check input { accent-color: var(--accent); width: 15px; height: 15px; }
|
||||
.hidden { display: none !important; }
|
||||
|
||||
.swatches { display: flex; align-items: center; gap: 6px; }
|
||||
.swatches button {
|
||||
width: 30px; height: 34px; border-radius: var(--radius-sm); cursor: pointer;
|
||||
background: var(--c); border: 1px solid var(--border);
|
||||
}
|
||||
.swatches button:hover { border-color: var(--accent); }
|
||||
|
||||
/* 上传区 */
|
||||
.dropzone {
|
||||
border: 1.5px dashed var(--border);
|
||||
border-radius: var(--radius);
|
||||
background: linear-gradient(180deg, rgba(79, 140, 255, 0.04), transparent);
|
||||
text-align: center;
|
||||
cursor: pointer;
|
||||
transition: border-color .15s, background .15s;
|
||||
}
|
||||
.dropzone.compact { padding: 18px 12px; }
|
||||
.dropzone:hover, .dropzone:focus-visible, .dropzone.over {
|
||||
border-color: var(--accent); background: var(--accent-soft); outline: none;
|
||||
}
|
||||
.dropzone p { margin: 0; font-size: 13px; }
|
||||
.dropzone .sub { margin-top: 4px; font-size: 11.5px; color: var(--text-faint); }
|
||||
kbd {
|
||||
font-family: var(--mono); font-size: 11px; padding: 1px 5px;
|
||||
border: 1px solid var(--border); border-radius: 4px; background: var(--bg);
|
||||
}
|
||||
|
||||
/* 待处理槽(单张) */
|
||||
.thumbs { list-style: none; margin: 0; padding: 0; display: flex; flex-direction: column; gap: 6px; }
|
||||
.thumbs:empty { display: none; }
|
||||
.pending-item {
|
||||
display: grid; grid-template-columns: 76px 1fr auto; align-items: center; gap: 10px;
|
||||
padding: 8px; background: var(--bg);
|
||||
border: 1px solid var(--accent); border-radius: var(--radius-sm);
|
||||
box-shadow: 0 0 0 3px var(--accent-soft);
|
||||
}
|
||||
.thumb-wrap {
|
||||
width: 76px; height: 76px; border-radius: 8px; overflow: hidden;
|
||||
background-color: #0a0c10;
|
||||
background-image:
|
||||
linear-gradient(45deg, #171b23 25%, transparent 25%),
|
||||
linear-gradient(-45deg, #171b23 25%, transparent 25%),
|
||||
linear-gradient(45deg, transparent 75%, #171b23 75%),
|
||||
linear-gradient(-45deg, transparent 75%, #171b23 75%);
|
||||
background-size: 14px 14px;
|
||||
background-position: 0 0, 0 7px, 7px -7px, -7px 0;
|
||||
}
|
||||
.thumb-wrap img { width: 100%; height: 100%; object-fit: contain; display: block; }
|
||||
.pending-meta { min-width: 0; display: flex; flex-direction: column; gap: 4px; }
|
||||
.pending-meta .tname { font-size: 12.5px; color: var(--text); word-break: break-all; }
|
||||
.pending-meta .tsize { font-size: 11px; color: var(--text-faint); font-family: var(--mono); }
|
||||
.thumbs .tremove {
|
||||
border: 0; background: transparent; color: var(--text-faint); cursor: pointer;
|
||||
font-size: 17px; line-height: 1; padding: 3px 8px; border-radius: 6px;
|
||||
}
|
||||
.thumbs .tremove:hover { color: var(--danger); background: rgba(242, 85, 90, .12); }
|
||||
.pending-note { margin: 0; font-size: 11.5px; line-height: 1.5; color: var(--accent); }
|
||||
.pending-note:empty { display: none; }
|
||||
|
||||
/* 按钮 */
|
||||
.btn {
|
||||
border: 1px solid var(--border);
|
||||
background: var(--surface-3);
|
||||
color: var(--text);
|
||||
border-radius: var(--radius-sm);
|
||||
padding: 9px 14px;
|
||||
cursor: pointer;
|
||||
transition: background .15s, border-color .15s, transform .05s;
|
||||
white-space: nowrap;
|
||||
}
|
||||
.btn:hover:not(:disabled) { border-color: var(--accent); }
|
||||
.btn:active:not(:disabled) { transform: translateY(1px); }
|
||||
.btn:disabled { opacity: .42; cursor: not-allowed; }
|
||||
.btn.sm { padding: 6px 11px; font-size: 12.5px; }
|
||||
.btn.ghost { background: transparent; }
|
||||
.btn.primary {
|
||||
background: var(--accent); border-color: var(--accent); color: #fff;
|
||||
font-weight: 600; font-size: 15px; padding: 12px 16px;
|
||||
}
|
||||
.btn.primary:hover:not(:disabled) { background: var(--accent-hover); border-color: var(--accent-hover); }
|
||||
.btn.link { background: none; border: 0; color: var(--accent); padding: 0; font-size: 12.5px; }
|
||||
|
||||
.actions { gap: 12px; }
|
||||
.progress { display: flex; flex-direction: column; gap: 6px; }
|
||||
.bar { height: 6px; border-radius: 99px; background: var(--bg); overflow: hidden; }
|
||||
.bar > i { display: block; height: 100%; width: 0; background: var(--accent); transition: width .25s ease; }
|
||||
.progress-meta { display: flex; justify-content: space-between; align-items: center; font-size: 12px; color: var(--text-dim); }
|
||||
.error-text { margin: 0; font-size: 12.5px; color: var(--danger); word-break: break-all; }
|
||||
|
||||
/* 批量处理说明 */
|
||||
.batch-note { margin: 0; font-size: 11px; line-height: 1.65; color: var(--text-faint); }
|
||||
.field-tip { margin: -2px 0 0; font-size: 11px; line-height: 1.6; color: var(--text-faint); }
|
||||
.batch-note strong { color: var(--text-dim); font-weight: 500; }
|
||||
.batch-note code {
|
||||
font-family: var(--mono); font-size: 10.5px; color: var(--text-dim);
|
||||
background: var(--bg); border: 1px solid var(--border-soft); border-radius: 4px; padding: 0 4px;
|
||||
}
|
||||
|
||||
/* ------------------------------- 结果区 ------------------------------- */
|
||||
.stage { display: flex; flex-direction: column; min-width: 0; min-height: 0; }
|
||||
.stage-head {
|
||||
display: flex; align-items: center; justify-content: space-between;
|
||||
padding: 14px 22px; border-bottom: 1px solid var(--border); gap: 16px;
|
||||
}
|
||||
.stage-head h2 { font-size: 15px; }
|
||||
.stage-head .count { font-family: var(--mono); color: var(--text-dim); font-weight: 400; }
|
||||
.stage-tools { display: flex; align-items: center; gap: 10px; }
|
||||
.stage-tip {
|
||||
flex: 0 0 auto; margin: 0; padding: 7px 22px; font-size: 11.5px; color: var(--text-faint);
|
||||
border-bottom: 1px solid var(--border-soft); background: rgba(79, 140, 255, .04);
|
||||
}
|
||||
|
||||
.segmented { display: flex; background: var(--surface-2); border: 1px solid var(--border); border-radius: var(--radius-sm); padding: 2px; }.segmented button {
|
||||
border: 0; background: transparent; color: var(--text-dim);
|
||||
padding: 5px 12px; border-radius: 6px; cursor: pointer; font-size: 12.5px;
|
||||
}
|
||||
.segmented button.active { background: var(--accent-soft); color: var(--accent); }
|
||||
|
||||
/* 批量任务报告(结果区顶部,仅在跑过 / 存在批量任务时出现) */
|
||||
.batch-report {
|
||||
flex: 0 0 auto; margin: 14px 22px 0; padding: 12px 14px;
|
||||
background: var(--surface); border: 1px solid var(--border-soft); border-radius: var(--radius);
|
||||
display: flex; flex-direction: column; gap: 8px; max-height: 40vh;
|
||||
}
|
||||
.batch-report[hidden] { display: none; }
|
||||
.batch-report-head { display: flex; align-items: center; gap: 10px; }
|
||||
.batch-report-head h3 { font-size: 13px; display: flex; align-items: center; gap: 8px; }
|
||||
.batch-report-path {
|
||||
flex: 1; min-width: 0; font-family: var(--mono); font-size: 11px; color: var(--text-faint);
|
||||
overflow: hidden; text-overflow: ellipsis; white-space: nowrap;
|
||||
}
|
||||
.batch-report-sum { margin: 0; font-family: var(--mono); font-size: 11.5px; color: var(--text-dim); }
|
||||
.batch-rows {
|
||||
list-style: none; margin: 0; padding: 0; overflow-y: auto;
|
||||
display: flex; flex-direction: column; gap: 4px;
|
||||
}
|
||||
.batch-rows:empty { display: none; }
|
||||
.batch-row {
|
||||
display: grid; grid-template-columns: 34px minmax(0, 1fr) minmax(0, 1.4fr) auto auto;
|
||||
align-items: center; gap: 10px; padding: 5px 9px;
|
||||
background: var(--surface-2); border: 1px solid var(--border-soft); border-radius: var(--radius-sm);
|
||||
font-family: var(--mono); font-size: 11.5px; color: var(--text-dim);
|
||||
}
|
||||
.batch-row .bi { text-align: right; color: var(--text-faint); }
|
||||
.batch-row .bs,
|
||||
.batch-row .bo { min-width: 0; overflow: hidden; text-overflow: ellipsis; white-space: nowrap; }
|
||||
.batch-row .bo { color: var(--text); }
|
||||
.batch-row .bo::before { content: "→ "; color: var(--text-faint); }
|
||||
.batch-row .bo:empty::before { content: ""; }
|
||||
.batch-row .bt { color: var(--text-faint); font-size: 11px; }
|
||||
.batch-row.running { border-color: rgba(79, 140, 255, .4); background: var(--accent-soft); }
|
||||
.batch-row.failed { border-color: rgba(242, 85, 90, .35); }
|
||||
.batch-row.canceled { opacity: .55; }
|
||||
|
||||
.grid {
|
||||
flex: 1 1 auto; overflow-y: auto; padding: 18px 22px 28px;
|
||||
display: grid; gap: 18px; grid-template-columns: repeat(auto-fill, minmax(360px, 1fr));
|
||||
align-content: start;
|
||||
}
|
||||
|
||||
.card {
|
||||
background: var(--surface); border: 1px solid var(--border-soft);
|
||||
border-radius: var(--radius); overflow: hidden; display: flex; flex-direction: column;
|
||||
transition: border-color .15s, box-shadow .15s;
|
||||
}
|
||||
.card-head { display: flex; align-items: center; gap: 8px; padding: 8px 10px; border-bottom: 1px solid var(--border-soft); cursor: pointer; }
|
||||
.card-head .name {
|
||||
flex: 1; min-width: 0; text-align: left; cursor: pointer;
|
||||
border: 0; background: transparent; color: var(--text);
|
||||
padding: 3px 6px; margin-left: -6px; border-radius: 6px;
|
||||
font-size: 12.5px; overflow: hidden; text-overflow: ellipsis; white-space: nowrap;
|
||||
transition: color .15s, background .15s;
|
||||
}
|
||||
.card-head .name:hover {
|
||||
color: var(--accent); background: var(--accent-soft);
|
||||
text-decoration: underline; text-underline-offset: 3px;
|
||||
}
|
||||
.card-head .name:focus-visible { outline: none; box-shadow: 0 0 0 2px var(--accent-soft); }
|
||||
.card-del {
|
||||
flex: 0 0 auto; border: 0; background: transparent; color: var(--text-faint);
|
||||
cursor: pointer; font-size: 17px; line-height: 1; padding: 2px 8px; border-radius: 6px;
|
||||
transition: color .15s, background .15s;
|
||||
}
|
||||
.card-del:hover { color: var(--danger); background: rgba(242, 85, 90, .12); }
|
||||
|
||||
/* 处理这条记录时用的参数:始终可见(便于对照历史),过长省略,完整 JSON 在 title 里 */
|
||||
.card-params {
|
||||
margin: 0; padding: 6px 12px;
|
||||
font-family: var(--mono); font-size: 11px; color: var(--text-faint);
|
||||
background: var(--surface-2); border-bottom: 1px solid var(--border-soft);
|
||||
overflow: hidden; text-overflow: ellipsis; white-space: nowrap;
|
||||
}
|
||||
.card-params::before { content: "参数 "; color: var(--text-dim); }
|
||||
.card-params:empty { display: none; }
|
||||
|
||||
/* 选中态:主题色一圈描边 + 头部同色底 */
|
||||
.card.selected { border-color: var(--accent); box-shadow: 0 0 0 3px var(--accent-soft); }
|
||||
.card.selected .card-head { background: var(--accent-soft); border-bottom-color: rgba(79, 140, 255, .35); }
|
||||
.badge {
|
||||
font-family: var(--mono); font-size: 11px; padding: 2px 7px; border-radius: 99px;
|
||||
background: var(--surface-3); color: var(--text-dim); border: 1px solid var(--border);
|
||||
}
|
||||
.badge.ok { color: var(--ok); border-color: rgba(53, 201, 138, .35); background: rgba(53, 201, 138, .1); }
|
||||
.badge.warn { color: var(--warn); border-color: rgba(240, 180, 41, .35); background: rgba(240, 180, 41, .1); }
|
||||
.badge.err { color: var(--danger); border-color: rgba(242, 85, 90, .35); background: rgba(242, 85, 90, .1); }
|
||||
|
||||
/* 对比视图 */
|
||||
.compare {
|
||||
position: relative; width: 100%; aspect-ratio: 4 / 3; overflow: hidden;
|
||||
background-color: #0a0c10;
|
||||
background-image:
|
||||
linear-gradient(45deg, #171b23 25%, transparent 25%),
|
||||
linear-gradient(-45deg, #171b23 25%, transparent 25%),
|
||||
linear-gradient(45deg, transparent 75%, #171b23 75%),
|
||||
linear-gradient(-45deg, transparent 75%, #171b23 75%);
|
||||
background-size: 18px 18px;
|
||||
background-position: 0 0, 0 9px, 9px -9px, -9px 0;
|
||||
}
|
||||
.compare img {
|
||||
position: absolute; inset: 0; width: 100%; height: 100%;
|
||||
object-fit: contain; user-select: none; -webkit-user-drag: none; pointer-events: none;
|
||||
}
|
||||
.compare .after { clip-path: inset(0 calc(100% - var(--p, 50%)) 0 0); }
|
||||
.compare .divider {
|
||||
position: absolute; top: 0; bottom: 0; left: var(--p, 50%); width: 2px;
|
||||
background: var(--accent); box-shadow: 0 0 12px rgba(79, 140, 255, .55); pointer-events: none;
|
||||
}
|
||||
.compare .divider::after {
|
||||
content: ""; position: absolute; top: 50%; left: 50%; width: 26px; height: 26px;
|
||||
transform: translate(-50%, -50%); border-radius: 50%;
|
||||
background: var(--accent); border: 3px solid var(--surface);
|
||||
}
|
||||
.compare input[type="range"] {
|
||||
position: absolute; inset: 0; width: 100%; height: 100%; margin: 0;
|
||||
opacity: 0; cursor: ew-resize; -webkit-appearance: none; appearance: none;
|
||||
}
|
||||
.compare .tag {
|
||||
position: absolute; bottom: 8px; font-family: var(--mono); font-size: 11px;
|
||||
padding: 2px 7px; border-radius: 6px; background: rgba(8, 10, 14, .72); color: var(--text);
|
||||
border: 1px solid var(--border); pointer-events: none;
|
||||
}
|
||||
.compare .tag.left { left: 8px; }
|
||||
.compare .tag.right { right: 8px; color: var(--accent); }
|
||||
|
||||
.card-body { padding: 10px 12px; display: flex; flex-direction: column; gap: 8px; }
|
||||
.card-stats { display: flex; flex-wrap: wrap; gap: 12px; font-family: var(--mono); font-size: 11.5px; color: var(--text-faint); }
|
||||
.card-stats b { color: var(--text-dim); font-weight: 500; }
|
||||
.card-actions { display: flex; gap: 8px; flex-wrap: wrap; }
|
||||
.card-note { margin: 0; font-size: 11.5px; color: var(--warn); }
|
||||
.card-error { font-size: 12px; color: var(--danger); word-break: break-all; }
|
||||
|
||||
/* 卡片状态 */
|
||||
.card.running .card-body, .card.pending .card-body { display: none; }
|
||||
.card .status {
|
||||
padding: 26px 12px; display: flex; flex-direction: column; align-items: center; gap: 10px;
|
||||
color: var(--text-dim); font-size: 12.5px;
|
||||
}
|
||||
.spinner {
|
||||
width: 22px; height: 22px; border-radius: 50%;
|
||||
border: 2px solid var(--border); border-top-color: var(--accent);
|
||||
animation: spin .8s linear infinite;
|
||||
}
|
||||
@keyframes spin { to { transform: rotate(360deg); } }
|
||||
.card.done .status, .card.failed .status, .card.canceled .status { display: none; }
|
||||
|
||||
.empty-state {
|
||||
flex: 1 1 auto; display: flex; flex-direction: column; align-items: center; justify-content: center;
|
||||
gap: 12px; text-align: center; padding: 40px; color: var(--text-dim);
|
||||
}
|
||||
.empty-state h3 { font-size: 28px; font-weight: 600; color: var(--text); letter-spacing: -0.4px; }
|
||||
.empty-state p { margin: 0; max-width: 460px; font-size: 13px; line-height: 1.7; }
|
||||
.empty-state.hidden { display: none; }
|
||||
|
||||
@media (max-width: 1100px) {
|
||||
.layout { grid-template-columns: 1fr; }
|
||||
.sidebar { border-right: 0; border-bottom: 1px solid var(--border); }
|
||||
.meta { display: none; }
|
||||
}
|
||||
+587
@@ -0,0 +1,587 @@
|
||||
/* 前端回归测试:用最小假 DOM 在 Node 里跑 app.js,模拟完整交互链路。
|
||||
用法:node web/test_app.js
|
||||
覆盖:
|
||||
① 单槽待处理:替换 / 去重 / 一次多张
|
||||
② 结果历史:刷新后从服务端恢复、跨任务累积
|
||||
③ 回归重点:添加新图并重新处理时,右侧已有结果绝不能被清空
|
||||
④ 参数记录:每条记录带参数快照;「重新处理」把参数重新赋值到表单
|
||||
⑤ 同一张图多次处理 → 参数记录更新为最新
|
||||
⑥ 删除单条不被轮询复活 / 选中描边 / 打包 ZIP / 清空 */
|
||||
'use strict';
|
||||
const fs = require('fs');
|
||||
const path = require('path');
|
||||
|
||||
const appSrc = fs.readFileSync(process.env.APP_SRC || path.join(__dirname, 'app.js'), 'utf8');
|
||||
const zipSrc = fs.readFileSync(path.join(__dirname, 'zip.js'), 'utf8');
|
||||
|
||||
/* ---------------------------- 假 DOM ---------------------------- */
|
||||
function makeNode(tag = 'div') {
|
||||
const listeners = {};
|
||||
const classes = new Set();
|
||||
let text = '';
|
||||
const node = {
|
||||
tag,
|
||||
parent: null,
|
||||
children: [],
|
||||
dataset: {},
|
||||
style: { setProperty() {}, width: '' },
|
||||
value: '',
|
||||
checked: false,
|
||||
hidden: false,
|
||||
disabled: false,
|
||||
title: '',
|
||||
src: '',
|
||||
href: '',
|
||||
download: '',
|
||||
files: null,
|
||||
scrollTop: 0,
|
||||
// textContent 与真实 DOM 一致:赋值会清空所有子节点
|
||||
get textContent() { return text; },
|
||||
set textContent(v) {
|
||||
text = String(v);
|
||||
for (const c of node.children) c.parent = null;
|
||||
node.children.length = 0;
|
||||
},
|
||||
// className 与 classList 共享同一份类名集合
|
||||
get className() { return [...classes].join(' '); },
|
||||
set className(v) {
|
||||
classes.clear();
|
||||
for (const c of String(v).split(/\s+/)) if (c) classes.add(c);
|
||||
},
|
||||
classList: {
|
||||
add(c) { classes.add(c); },
|
||||
remove(c) { classes.delete(c); },
|
||||
toggle(c, on) {
|
||||
const want = on === undefined ? !classes.has(c) : !!on;
|
||||
want ? classes.add(c) : classes.delete(c);
|
||||
return want;
|
||||
},
|
||||
contains(c) { return classes.has(c); },
|
||||
},
|
||||
append(...kids) {
|
||||
for (const k of kids) {
|
||||
// 真实 DOM 里 append(undefined) 会塞进一个文本节点 —— 静默的视觉 bug,
|
||||
// 所以这里直接抛错,逼出这类问题(app.js 的 el() 已过滤空子节点)
|
||||
if (k === undefined || k === null) throw new Error('append() 收到了 undefined/null 子节点');
|
||||
k.parent = node;
|
||||
node.children.push(k);
|
||||
}
|
||||
},
|
||||
remove() {
|
||||
if (!node.parent) return;
|
||||
const i = node.parent.children.indexOf(node);
|
||||
if (i >= 0) node.parent.children.splice(i, 1);
|
||||
node.parent = null;
|
||||
},
|
||||
addEventListener(type, fn) { (listeners[type] = listeners[type] || []).push(fn); },
|
||||
dispatch(type, ev = {}) {
|
||||
const event = { target: node, preventDefault() {}, stopPropagation() {}, ...ev };
|
||||
for (const fn of listeners[type] || []) fn(event);
|
||||
},
|
||||
closest(sel) {
|
||||
const parts = sel.split(',').map((s) => s.trim());
|
||||
let cur = node;
|
||||
while (cur) {
|
||||
for (const p of parts) {
|
||||
if (p.includes('[')) { // 形如 button[data-view]
|
||||
const m = /^([\w-]*)\[([\w-]+)\]$/.exec(p);
|
||||
if (m && (!m[1] || cur.tag === m[1]) && cur.dataset[m[2]] !== undefined) return cur;
|
||||
} else if (p === cur.tag) {
|
||||
return cur;
|
||||
} else if (p.startsWith('.') && cur.className.split(/\s+/).includes(p.slice(1))) {
|
||||
return cur;
|
||||
}
|
||||
}
|
||||
cur = cur.parent;
|
||||
}
|
||||
return null;
|
||||
},
|
||||
querySelectorAll() { return []; },
|
||||
setAttribute() {},
|
||||
click() { node.dispatch('click'); },
|
||||
scrollIntoView() {},
|
||||
};
|
||||
return node;
|
||||
}
|
||||
|
||||
const byId = new Map();
|
||||
const viewSwitch = makeNode('div');
|
||||
global.document = {
|
||||
body: makeNode('body'),
|
||||
getElementById(id) {
|
||||
if (!byId.has(id)) byId.set(id, id === 'view-switch' ? viewSwitch : makeNode());
|
||||
return byId.get(id);
|
||||
},
|
||||
createElement(tag) { return makeNode(tag); },
|
||||
createTextNode(text) { const n = makeNode('#text'); n.textContent = String(text); return n; },
|
||||
addEventListener() {},
|
||||
querySelectorAll() { return []; },
|
||||
};
|
||||
|
||||
global.window = { location: { assign() {} } };
|
||||
new Function('window', zipSrc)(global.window); // 让 window.BiRefNetZip 可用
|
||||
|
||||
const lsData = new Map();
|
||||
global.localStorage = {
|
||||
getItem: (k) => (lsData.has(k) ? lsData.get(k) : null),
|
||||
setItem: (k, v) => lsData.set(k, String(v)),
|
||||
removeItem: (k) => lsData.delete(k),
|
||||
};
|
||||
Object.defineProperty(global, 'navigator', { value: { clipboard: {} }, configurable: true });
|
||||
Object.defineProperty(global, 'File', {
|
||||
configurable: true,
|
||||
value: class File {
|
||||
constructor(parts, name, opts = {}) {
|
||||
const first = Array.isArray(parts) ? parts[0] : parts;
|
||||
this.name = name;
|
||||
this.type = opts.type || '';
|
||||
this.lastModified = opts.lastModified || 0;
|
||||
this.size = (first && first.size) || 0;
|
||||
}
|
||||
},
|
||||
});
|
||||
let lastBlob = null;
|
||||
Object.defineProperty(global, 'Blob', {
|
||||
configurable: true,
|
||||
value: class Blob {
|
||||
constructor(parts, opts = {}) {
|
||||
this.parts = parts;
|
||||
this.type = opts.type || '';
|
||||
lastBlob = this;
|
||||
}
|
||||
},
|
||||
});
|
||||
global.URL = { createObjectURL: () => `blob:${Math.random()}`, revokeObjectURL() {} };
|
||||
global.ClipboardItem = class {};
|
||||
global.FormData = class {
|
||||
constructor() { this.entries = []; }
|
||||
append(k, v, name) { this.entries.push([k, v, name]); }
|
||||
};
|
||||
|
||||
/* ------------------------- 假服务端 ------------------------- */
|
||||
const STATE = {
|
||||
ok: true,
|
||||
node_dir: 'F:\\BiRefNet_WebUI\\vendor\\comfyui_birefnet_ll',
|
||||
output_dir: 'F:\\BiRefNet_WebUI\\outputs',
|
||||
environment: {
|
||||
python: '3.12.10', torch: '2.13.0+cu126', cuda: true,
|
||||
devices: [{ index: 0, name: 'NVIDIA GeForce RTX 3070 Laptop GPU', free_mem: 7, total_mem: 8 }],
|
||||
},
|
||||
models: [{
|
||||
key: 'Portrait', name: 'Portrait', file: 'Portrait.safetensors', arch: 'v1',
|
||||
size_mb: 843.9, backbone: 'swin_v1_l', path: 'models/Portrait.safetensors',
|
||||
}],
|
||||
defaults: { model: 'Portrait' },
|
||||
};
|
||||
|
||||
/** 与 index.html 里的默认值保持一致,供测试手动铺一遍表单初始值。 */
|
||||
const FORM_DEFAULTS = {
|
||||
'opt-model': 'Portrait', 'opt-device': 'auto', 'opt-dtype': 'auto', 'opt-arch': 'auto',
|
||||
'opt-resolution-mode': 'square', 'opt-width': '1024', 'opt-height': '1024',
|
||||
'opt-longest-side': '1024', 'opt-upscale': 'bilinear', 'opt-threshold': '0',
|
||||
'opt-blur1': '90', 'opt-blur2': '6', 'opt-background': 'transparent',
|
||||
'opt-bgcolor': '#ffffff', 'opt-final-side': '0',
|
||||
};
|
||||
const FORM_CHECKS = { 'opt-refine': true, 'opt-mask': false, 'batch-recursive': false };
|
||||
|
||||
const P_SQUARE = {
|
||||
model: 'Portrait', device: 'auto', dtype: 'auto', arch: 'auto',
|
||||
resolution_mode: 'square', width: 1024, height: 1024, longest_side: 1024,
|
||||
upscale_method: 'bilinear', mask_threshold: 0, refine_foreground: true,
|
||||
blur_size: 90, blur_size_two: 6, background: 'transparent', bg_color: '#ffffff',
|
||||
output_mask: false, final_longest_side: 0,
|
||||
};
|
||||
const P_GREEN = {
|
||||
...P_SQUARE, model: 'birefnet', resolution_mode: 'longest', longest_side: 512,
|
||||
background: 'color', bg_color: '#00ff00', refine_foreground: false,
|
||||
};
|
||||
|
||||
const stub = { polls: 0, posts: 0, detailFetches: 0, fileHits: 0, originalHits: 0, lastOptions: null, batchPolls: 0, batchPosts: 0, lastBatch: null };
|
||||
const tasks = new Map();
|
||||
let clock = 1;
|
||||
|
||||
function mkImage(index, name, over = {}) {
|
||||
const done = (over.state || 'done') === 'done';
|
||||
return {
|
||||
index, name, label: name.replace(/\.[^.]+$/, ''),
|
||||
state: 'done', stage: '完成', progress: 1,
|
||||
width: 1920, height: 1080, input_size: [1024, 1024], coverage: 0.42, elapsed: 3.1,
|
||||
error: null, warning: null,
|
||||
urls: done
|
||||
? { original: `/api/tasks/${over.task}/file/${index}/original`, cutout: `/api/tasks/${over.task}/file/${index}/cutout`, mask: `/api/tasks/${over.task}/file/${index}/mask` }
|
||||
: {},
|
||||
...over,
|
||||
};
|
||||
}
|
||||
|
||||
function mkTask(id, options, images, state = 'done') {
|
||||
const task = { id, created: clock++, state, stage: state === 'done' ? '完成' : '排队中', progress: state === 'done' ? 1 : 0, options, images };
|
||||
tasks.set(id, task);
|
||||
return task;
|
||||
}
|
||||
|
||||
// 历史:T1(无遮罩输出 + 列表里不带 options,逼前端补拉详情)+ T2(列表里直接带)
|
||||
const T1 = mkTask('T100', P_SQUARE, [{ index: 0, name: 'old.jpg', label: 'old', state: 'done', stage: '完成', progress: 1, width: 800, height: 600, input_size: [1024, 1024], coverage: 0.3, elapsed: 2.2, error: null, warning: null, urls: {} }]);
|
||||
T1.images[0].urls = { original: '/api/tasks/T100/file/0/original', cutout: '/api/tasks/T100/file/0/cutout' };
|
||||
const T2 = mkTask('T200', P_GREEN, [mkImage(0, 'mid.jpg', { task: 'T200' })]);
|
||||
T2.options = P_GREEN;
|
||||
|
||||
const taskList = () => [...tasks.values()]
|
||||
.sort((a, b) => b.created - a.created)
|
||||
.map((t) => ({
|
||||
id: t.id, created: t.created, state: t.state, stage: t.stage, progress: t.progress,
|
||||
mode: t.mode || 'upload', input_dir: t.input_dir || null, output_dir: t.output_dir || null,
|
||||
// 模仿服务端列表:T100 不返回 options(触发补拉详情)
|
||||
...(t.id === 'T100' ? {} : { options: t.options }),
|
||||
images: t.images,
|
||||
}));
|
||||
|
||||
global.fetch = async (url, options = {}) => {
|
||||
const json = (body) => ({ ok: true, status: 200, text: async () => JSON.stringify(body) });
|
||||
const bin = (size) => ({
|
||||
ok: true, status: 200,
|
||||
blob: async () => ({ size, type: 'image/png' }),
|
||||
arrayBuffer: async () => new Uint8Array(size).buffer,
|
||||
});
|
||||
const method = options.method || 'GET';
|
||||
|
||||
if (url.includes('/api/state')) return json(STATE);
|
||||
|
||||
if (url.endsWith('/api/batch') && method === 'POST') {
|
||||
stub.batchPosts += 1;
|
||||
const body = JSON.parse(options.body);
|
||||
stub.lastBatch = body;
|
||||
const id = `B${400 + stub.batchPosts}`;
|
||||
const names = ['photo.jpg', 'a_pretty_long_original_filename_here.png', 'third.png'];
|
||||
const images = names.map((name, i) => ({
|
||||
index: i, name, label: name, state: 'pending', stage: '排队中', progress: 0,
|
||||
output_name: null, urls: {},
|
||||
}));
|
||||
const task = mkTask(id, body.options || {}, images, 'running');
|
||||
task.mode = 'batch';
|
||||
task.input_dir = body.input_dir;
|
||||
task.output_dir = body.output_dir || STATE.output_dir;
|
||||
return json({
|
||||
ok: true, task_id: id, mode: 'batch', count: images.length,
|
||||
input_dir: task.input_dir, output_dir: task.output_dir,
|
||||
});
|
||||
}
|
||||
|
||||
if (url.includes('/api/tasks') && method === 'POST') {
|
||||
stub.posts += 1;
|
||||
const opts = JSON.parse(options.body.entries.find((e) => e[0] === 'options')[1]);
|
||||
stub.lastOptions = opts;
|
||||
const name = options.body.entries.find((e) => e[0] === 'files')[2] || 'upload.png';
|
||||
const id = `T${300 + stub.posts}`;
|
||||
const images = [{ index: 0, name, label: name.replace(/\.[^.]+$/, ''), state: 'pending', stage: '排队中', progress: 0, urls: {} }];
|
||||
mkTask(id, opts, images, 'running');
|
||||
return json({ ok: true, task_id: id, count: 1 });
|
||||
}
|
||||
|
||||
if (/\/file\/\d+\/(cutout|mask|original)$/.test(url)) {
|
||||
stub.fileHits += 1;
|
||||
if (url.endsWith('/original')) stub.originalHits += 1;
|
||||
return bin(2048);
|
||||
}
|
||||
|
||||
const detail = /\/api\/tasks\/([^/?]+)(?:\?([^#]*))?$/.exec(url);
|
||||
if (detail && method === 'GET') {
|
||||
stub.detailFetches += 1;
|
||||
const task = tasks.get(detail[1]);
|
||||
if (!task) return { ok: false, status: 404, text: async () => JSON.stringify({ ok: false, error: '任务不存在' }) };
|
||||
const isBatch = task.mode === 'batch';
|
||||
if (isBatch) stub.batchPolls += 1;
|
||||
// 轮询时推进状态:running → done(第二次拿到该任务时)
|
||||
if (task.state === 'running') {
|
||||
task.polls = (task.polls || 0) + 1;
|
||||
if (task.polls >= 2) {
|
||||
task.state = 'done';
|
||||
task.stage = '完成';
|
||||
task.progress = 1;
|
||||
task.images = task.images.map((img, i) => (isBatch
|
||||
// 批量:只补上 output_name / 完成态(服务端就是按 RMBG_ 命名写盘的)
|
||||
? { ...img, state: 'done', stage: '完成', progress: 1, elapsed: 4.2, output_name: `RMBG_${img.label.slice(0, 20)}_1791360000.png` }
|
||||
: mkImage(i, img.name, { task: task.id })));
|
||||
} else {
|
||||
task.images = task.images.map((img) => ({ ...img, state: 'running', stage: '推理', progress: 0.5, urls: {} }));
|
||||
}
|
||||
stub.polls += 1;
|
||||
}
|
||||
// ?tail=N → 只回传最后 N 张(批量轮询用)
|
||||
const tail = Number((new URLSearchParams(detail[2] || '').get('tail')) || 0);
|
||||
const images = tail > 0 ? task.images.slice(Math.max(0, task.images.length - tail)) : task.images;
|
||||
return json({ ok: true, task: { id: task.id, created: task.created, mode: task.mode || 'upload', input_dir: task.input_dir || null, output_dir: task.output_dir || null, state: task.state, stage: task.stage, progress: task.progress, error: null, counts: { total: task.images.length, done: task.state === 'done' ? task.images.length : 0, failed: 0, pending: 0 }, images_total: task.images.length, images_from: task.images.length - images.length, options: task.options, images } });
|
||||
}
|
||||
|
||||
if (url.endsWith('/api/tasks')) return json({ ok: true, tasks: taskList() });
|
||||
return json({ ok: false, error: '未知接口 ' + url });
|
||||
};
|
||||
|
||||
/* ---------------------------- 执行 ---------------------------- */
|
||||
for (const [id, v] of Object.entries(FORM_DEFAULTS)) byId.set(id, Object.assign(makeNode('input'), { id, value: v }));
|
||||
for (const [id, v] of Object.entries(FORM_CHECKS)) byId.set(id, Object.assign(makeNode('input'), { id, checked: v }));
|
||||
byId.set('thumbs', makeNode('ul'));
|
||||
byId.set('results', makeNode('div'));
|
||||
byId.set('empty', Object.assign(makeNode('div'), { className: 'empty-state' }));
|
||||
byId.set('btn-run', makeNode('button'));
|
||||
byId.set('btn-zip', Object.assign(makeNode('button'), { textContent: '打包下载 ZIP' }));
|
||||
// 批量处理面板(index.html 里存在,但未被 app.js 主动触碰的节点要在这里铺好)
|
||||
byId.set('batch-rows', makeNode('ul'));
|
||||
byId.set('batch-report', Object.assign(makeNode('section'), { hidden: true }));
|
||||
byId.set('batch-progress-bar', makeNode('i'));
|
||||
|
||||
// 预置一份「老版本(v1)」表单存档:当年 output_mask 默认是 true,参数名还是 max_output_side。
|
||||
// 启动时的迁移逻辑必须把遮罩默认值修正掉、把旧尺寸参数搬到新输入框,且不动其它参数。
|
||||
lsData.set('birefnet.webui.options', JSON.stringify({
|
||||
background: 'transparent', output_mask: true, max_output_side: 2048,
|
||||
}));
|
||||
|
||||
new Function(appSrc)();
|
||||
|
||||
const sleep = (ms) => new Promise((r) => setTimeout(r, ms));
|
||||
async function waitFor(cond, timeout = 5000) {
|
||||
const t0 = Date.now();
|
||||
while (Date.now() - t0 < timeout) { if (cond()) return true; await sleep(15); }
|
||||
return false;
|
||||
}
|
||||
|
||||
const results = () => byId.get('results').children;
|
||||
const cards = () => byId.get('results').children;
|
||||
const cardState = (c) => c.className;
|
||||
const head = (c) => c.children[0];
|
||||
const nameBtnOf = (c) => head(c).children[0];
|
||||
const badgeOf = (c) => head(c).children[1].textContent;
|
||||
const delBtnOf = (c) => head(c).children[2];
|
||||
const paramsLineOf = (c) => find(c, 'card-params');
|
||||
const find = (node, cls) => {
|
||||
for (const child of node.children || []) {
|
||||
if (child.className.split(/\s+/).includes(cls)) return child;
|
||||
const hit = find(child, cls);
|
||||
if (hit) return hit;
|
||||
}
|
||||
return null;
|
||||
};
|
||||
const compareImgCount = (c) => ((find(c, 'compare')?.children || []).filter((n) => n.tag === 'img').length);
|
||||
const statsCount = (c) => ((find(c, 'card-stats')?.children || []).length);
|
||||
const actionsOf = (c) => find(c, 'card-actions');
|
||||
const reuseBtnOf = (c) => actionsOf(c).children[0];
|
||||
|
||||
function pendingName() {
|
||||
const ul = byId.get('thumbs');
|
||||
if (ul.children.length !== 1) return `li×${ul.children.length}`;
|
||||
const meta = ul.children[0].children.find((x) => x.className.includes('pending-meta'));
|
||||
return meta ? meta.children[0].textContent : '(无 meta)';
|
||||
}
|
||||
function upload(files) {
|
||||
const input = byId.get('file-input');
|
||||
input.files = files;
|
||||
input.dispatch('change');
|
||||
}
|
||||
const F = (name, size, mtime) => ({ name, size, lastModified: mtime, type: 'image/jpeg' });
|
||||
const formValue = (id) => (id === 'opt-refine' || id === 'opt-mask' ? byId.get(id).checked : byId.get(id).value);
|
||||
const setForm = (id, v) => { if (id === 'opt-refine' || id === 'opt-mask') byId.get(id).checked = v; else byId.get(id).value = v; };
|
||||
const savedParams = () => JSON.parse(lsData.get('birefnet.webui.imageParams') || '{}');
|
||||
|
||||
(async () => {
|
||||
await sleep(80); // init / loadState / loadHistory
|
||||
|
||||
const checks = [];
|
||||
const check = (name, ok, detail) => checks.push([name, ok, detail]);
|
||||
|
||||
/* ============ 场景 0:老版本 localStorage 存档迁移 ============ */
|
||||
check('⓪ 老存档的 output_mask 不覆盖新默认(默认不输出遮罩)', formValue('opt-mask') === false, `checked=${formValue('opt-mask')}`);
|
||||
check('⓪ 老参数 max_output_side 搬到新的「最终尺寸 · 最长边」', formValue('opt-final-side') === '2048', formValue('opt-final-side'));
|
||||
check('⓪ 其余老参数照常恢复', formValue('opt-background') === 'transparent', formValue('opt-background'));
|
||||
check('⓪ 存档被标记为 v2(迁移只做一次)',
|
||||
(JSON.parse(lsData.get('birefnet.webui.options') || '{}')._schema) === 2,
|
||||
lsData.get('birefnet.webui.options'));
|
||||
setForm('opt-final-side', '0'); // 复原,避免影响后续场景
|
||||
|
||||
/* ============ 场景 1:刷新后从服务端恢复历史结果 ============ */
|
||||
check('① 启动即恢复 2 条历史记录', results().length === 2, `len=${results().length}`);
|
||||
check('① 列表缺 options 时自动补拉详情', stub.detailFetches >= 1, `fetches=${stub.detailFetches}`);
|
||||
check('① 历史记录显示参数行', (paramsLineOf(cards()[0])?.textContent || '').includes('Portrait'), paramsLineOf(cards()[0])?.textContent);
|
||||
check('① 历史记录直接可用(done)', cardState(cards()[0]).includes('done') && compareImgCount(cards()[0]) === 2, cardState(cards()[0]));
|
||||
check('① 无遮罩记录不产生 null 子节点', actionsOf(cards()[0]).children.every((c) => c.tag && c.tag !== '#text'), actionsOf(cards()[0]).children.map((c) => c.tag));
|
||||
check('① 计数 = 2', byId.get('result-count').textContent === '2', byId.get('result-count').textContent);
|
||||
|
||||
/* ============ 场景 2:单槽待处理 ============ */
|
||||
upload([F('first.jpg', 1000, 111)]);
|
||||
await sleep(20);
|
||||
check('② 上传后待处理 1 张', pendingName() === 'first.jpg', pendingName());
|
||||
upload([F('second.jpg', 2000, 222)]);
|
||||
await sleep(20);
|
||||
check('② 新图替换旧图(仍只有 1 张)', pendingName() === 'second.jpg', pendingName());
|
||||
upload([F('second.jpg', 2000, 222)]);
|
||||
await sleep(20);
|
||||
check('② 完全相同的图被忽略', pendingName() === 'second.jpg' && byId.get('thumbs').children.length === 1, pendingName());
|
||||
upload([F('x.jpg', 10, 1), F('y.jpg', 20, 2), F('z.jpg', 30, 3)]);
|
||||
await sleep(20);
|
||||
check('② 一次多张只留最后一张', pendingName() === 'z.jpg', pendingName());
|
||||
|
||||
/* ============ 场景 3(回归重点):新任务不清空已有结果 ============ */
|
||||
const before = results().slice();
|
||||
upload([F('new.jpg', 5000, 9)]);
|
||||
await sleep(20);
|
||||
setForm('opt-background', 'transparent');
|
||||
setForm('opt-bgcolor', '#ffffff');
|
||||
setForm('opt-model', 'Portrait');
|
||||
setForm('opt-resolution-mode', 'square');
|
||||
byId.get('btn-run').dispatch('click');
|
||||
await waitFor(() => stub.polls >= 1);
|
||||
check('③ 新任务提交后旧结果仍在(核心回归)',
|
||||
results().length === 3 && results()[0] === before[0] && results()[1] === before[1],
|
||||
`len=${results().length}`);
|
||||
check('③ 提交时参数带上了当前表单值', stub.lastOptions && stub.lastOptions.model === 'Portrait' && stub.lastOptions.background === 'transparent', JSON.stringify(stub.lastOptions));
|
||||
check('③ 新卡片立刻显示参数行', (paramsLineOf(cards()[2])?.textContent || '').includes('透明'), paramsLineOf(cards()[2])?.textContent);
|
||||
|
||||
await waitFor(() => stub.polls >= 2);
|
||||
await sleep(40);
|
||||
check('③ 新卡片进入 done 且旧结果完好', cardState(cards()[2]).includes('done') && results().length === 3, `${cardState(cards()[2])} len=${results().length}`);
|
||||
check('③ 旧卡片未被重建(DOM 复用)', results()[0] === before[0] && results()[1] === before[1], '被替换了');
|
||||
check('③ 参数已写入 localStorage', savedParams()['new.jpg']?.background === 'transparent', JSON.stringify(savedParams()['new.jpg'] || null));
|
||||
|
||||
/* ============ 场景 4:重新处理 = 参数重新赋值 ============ */
|
||||
setForm('opt-background', 'color');
|
||||
setForm('opt-bgcolor', '#00ff00');
|
||||
setForm('opt-resolution-mode', 'longest');
|
||||
setForm('opt-longest-side', '512');
|
||||
setForm('opt-refine', false);
|
||||
await sleep(20);
|
||||
reuseBtnOf(cards()[2]).dispatch('click');
|
||||
await sleep(80);
|
||||
check('④ 重新处理把图片放回待处理', pendingName() === 'new.jpg', pendingName());
|
||||
check('④ 参数被重新赋值(背景透明)', formValue('opt-background') === 'transparent', formValue('opt-background'));
|
||||
check('④ 参数被重新赋值(颜色/尺寸/精修)',
|
||||
formValue('opt-bgcolor') === '#ffffff' && formValue('opt-resolution-mode') === 'square' && formValue('opt-refine') === true,
|
||||
`${formValue('opt-bgcolor')} / ${formValue('opt-resolution-mode')} / ${formValue('opt-refine')}`);
|
||||
|
||||
/* ============ 场景 5:同一张图再次处理 → 参数更新为最新 ============ */
|
||||
setForm('opt-background', 'color');
|
||||
setForm('opt-bgcolor', '#00ff00');
|
||||
byId.get('btn-run').dispatch('click');
|
||||
await waitFor(() => stub.polls >= 3);
|
||||
check('⑤ 同图二次处理生成新记录(历史保留)', results().length === 4, `len=${results().length}`);
|
||||
check('⑤ 新记录参数为最新(纯色 #00ff00)', (paramsLineOf(cards()[3])?.textContent || '').includes('纯色 #00ff00'), paramsLineOf(cards()[3])?.textContent);
|
||||
check('⑤ localStorage 参数被最新值覆盖', savedParams()['new.jpg']?.background === 'color' && savedParams()['new.jpg']?.bg_color === '#00ff00', JSON.stringify(savedParams()['new.jpg'] || null));
|
||||
// 等这一轮真正跑完再进下一个场景(否则按钮仍处于「处理中」态,点击会被忽略)
|
||||
await waitFor(() => stub.polls >= 4 && cardState(cards()[3]).includes('done'));
|
||||
await sleep(30);
|
||||
check('⑤ 本轮任务已结束(可再次提交)', cardState(cards()[3]).includes('done'), cardState(cards()[3]));
|
||||
|
||||
/* ============ 场景 6:删除单条 + 轮询不复活 ============ */
|
||||
delBtnOf(cards()[1]).dispatch('click');
|
||||
await sleep(20);
|
||||
check('⑥ 删除后剩 3 条', results().length === 3, `len=${results().length}`);
|
||||
check('⑥ 计数同步', byId.get('result-count').textContent === '3', byId.get('result-count').textContent);
|
||||
const pollsBefore = stub.polls;
|
||||
byId.get('btn-run').dispatch('click'); // 再起一个任务制造轮询
|
||||
await waitFor(() => stub.polls > pollsBefore + 1);
|
||||
await sleep(40);
|
||||
check('⑥ 被删记录未被轮询复活', results().length === 4 && !results().some((c) => nameBtnOf(c).textContent.includes('mid.jpg')), `len=${results().length}`);
|
||||
|
||||
/* ============ 场景 7:选中描边 ============ */
|
||||
cards()[0].dispatch('click');
|
||||
await sleep(10);
|
||||
check('⑦ 点击卡片出现选中态', cardState(cards()[0]).includes('selected'), cardState(cards()[0]));
|
||||
cards()[0].dispatch('click');
|
||||
await sleep(10);
|
||||
check('⑦ 再点一次取消选中', !cardState(cards()[0]).includes('selected'), cardState(cards()[0]));
|
||||
nameBtnOf(cards()[0]).dispatch('click');
|
||||
await sleep(60);
|
||||
check('⑦ 点文件名不触发选中', !cardState(cards()[0]).includes('selected'), cardState(cards()[0]));
|
||||
|
||||
/* ============ 场景 8:跨任务打包 ZIP ============ */
|
||||
const visible = results().length;
|
||||
byId.get('btn-zip').dispatch('click');
|
||||
await sleep(120);
|
||||
const zip = lastBlob ? Buffer.from(concat(lastBlob.parts)) : null;
|
||||
const eocd = zip ? zip.lastIndexOf(Buffer.from([0x50, 0x4b, 0x05, 0x06])) : -1;
|
||||
const entryCount = eocd >= 0 ? zip.readUInt16LE(eocd + 10) : -1;
|
||||
check('⑧ ZIP 条目数 = 可见结果数', entryCount === visible, `entries=${entryCount} visible=${visible}`);
|
||||
check('⑧ 打包后按钮恢复可用', byId.get('btn-zip').disabled === false, String(byId.get('btn-zip').disabled));
|
||||
if (zip) fs.writeFileSync(path.join(__dirname, '..', 'outputs', 'test_zip_out.zip'), zip);
|
||||
|
||||
/* ============ 场景 9:清空 ============ */
|
||||
byId.get('btn-clear').dispatch('click');
|
||||
await sleep(20);
|
||||
check('⑨ 清空后列表为空', results().length === 0, `len=${results().length}`);
|
||||
check('⑨ 空状态重新显示', !byId.get('empty').classList.contains('hidden'), 'hidden 仍存在');
|
||||
check('⑨ 清空不影响参数记忆', savedParams()['new.jpg']?.background === 'color', JSON.stringify(savedParams()['new.jpg'] || null));
|
||||
|
||||
/* ============ 场景 10:目录批量处理 ============ */
|
||||
check('⑩ 输出目录默认预填为项目 outputs', byId.get('batch-output-dir').value === STATE.output_dir, byId.get('batch-output-dir').value);
|
||||
setForm('batch-input-dir', 'F:\\photos\\待抠图');
|
||||
setForm('batch-output-dir', ''); // 留空 → 由服务端兜底默认目录
|
||||
setForm('opt-background', 'color');
|
||||
byId.get('batch-recursive').checked = true;
|
||||
byId.get('btn-batch').dispatch('click');
|
||||
await waitFor(() => stub.batchPosts >= 1);
|
||||
check('⑩ 提交了输入/输出目录与递归开关',
|
||||
stub.lastBatch?.input_dir === 'F:\\photos\\待抠图' && stub.lastBatch?.output_dir === '' && stub.lastBatch?.recursive === true,
|
||||
JSON.stringify(stub.lastBatch));
|
||||
check('⑩ 抠图参数沿用左侧表单',
|
||||
stub.lastBatch?.options?.background === 'color' && stub.lastBatch?.options?.model === 'Portrait',
|
||||
JSON.stringify(stub.lastBatch?.options));
|
||||
|
||||
await waitFor(() => byId.get('batch-rows').children.length === 3);
|
||||
check('⑩ 批量清单逐文件铺开', byId.get('batch-rows').children.length === 3, byId.get('batch-rows').children.length);
|
||||
check('⑩ 清单标出输入 → 输出目录',
|
||||
(byId.get('batch-report-path').textContent || '').includes('待抠图') && (byId.get('batch-report-path').textContent || '').includes('outputs'),
|
||||
byId.get('batch-report-path').textContent);
|
||||
check('⑩ 批量任务不铺单图结果卡片', results().length === 0, `cards=${results().length}`);
|
||||
check('⑩ 清单展开时不显示空状态', byId.get('empty').classList.contains('hidden'), '空状态可见');
|
||||
|
||||
// 收尾会再拉一次全量(清单里每行都要是终态),所以要等按钮恢复才断言
|
||||
await waitFor(() => byId.get('btn-batch').textContent === '开始批量抠图');
|
||||
await sleep(60);
|
||||
const rows = byId.get('batch-rows').children;
|
||||
check('⑩ 收尾每行显示 RMBG_ 输出名', rows.every((r) => (r.cells.to.textContent || '').startsWith('RMBG_')), rows.map((r) => r.cells.to.textContent));
|
||||
check('⑩ 收尾每行状态为完成', rows.every((r) => r.cells.badge.textContent === '完成'), rows.map((r) => r.cells.badge.textContent));
|
||||
check('⑩ 汇总行给出张数', (byId.get('batch-report-sum').textContent || '').includes('共 3 张'), byId.get('batch-report-sum').textContent);
|
||||
check('⑩ 按钮恢复可点', byId.get('btn-batch').disabled === false && byId.get('btn-batch').textContent === '开始批量抠图', byId.get('btn-batch').textContent);
|
||||
check('⑩ 没有页面错误', byId.get('batch-error').hidden === true, byId.get('batch-error').textContent);
|
||||
|
||||
byId.get('btn-batch-hide').dispatch('click');
|
||||
await sleep(20);
|
||||
check('⑩ 收起后清单隐藏、空状态回归',
|
||||
byId.get('batch-report').hidden === true && !byId.get('empty').classList.contains('hidden'), '收起逻辑异常');
|
||||
|
||||
/* ============ 场景 11:最终尺寸·最长边 + 遮罩默认关闭 ============ */
|
||||
check('⑪ 「同时输出遮罩」默认不勾选(默认不产出遮罩文件)', formValue('opt-mask') === false, `checked=${formValue('opt-mask')}`);
|
||||
check('⑪ 「最终尺寸 · 最长边」默认 0', formValue('opt-final-side') === '0', formValue('opt-final-side'));
|
||||
|
||||
setForm('opt-final-side', '1440');
|
||||
setForm('opt-mask', true);
|
||||
const posts11 = stub.posts;
|
||||
const polls11 = stub.polls;
|
||||
byId.get('btn-run').dispatch('click');
|
||||
await waitFor(() => stub.posts > posts11);
|
||||
check('⑪ 提交参数带 final_longest_side / output_mask',
|
||||
stub.lastOptions?.final_longest_side === 1440 && stub.lastOptions?.output_mask === true,
|
||||
JSON.stringify(stub.lastOptions));
|
||||
await waitFor(() => stub.polls > polls11 + 1);
|
||||
await sleep(50);
|
||||
check('⑪ 结果卡片参数摘要显示最长边',
|
||||
(paramsLineOf(cards()[0])?.textContent || '').includes('最长边 1440'),
|
||||
paramsLineOf(cards()[0])?.textContent);
|
||||
|
||||
/* ============ 汇总 ============ */
|
||||
const errs = byId.get('error-text').textContent;
|
||||
console.log('='.repeat(64));
|
||||
let fail = 0;
|
||||
for (const [name, ok, detail] of checks) {
|
||||
console.log(` ${ok ? '✔' : '✘'} ${name}${ok ? '' : ` -> 实际: ${JSON.stringify(detail)}`}`);
|
||||
if (!ok) fail++;
|
||||
}
|
||||
if (errs) console.log(` ⚠ 页面上报错误: ${errs}`);
|
||||
console.log('='.repeat(64));
|
||||
console.log(fail || errs ? `失败 ${fail} 项${errs ? '(含页面错误)' : ''}` : `前端回归测试全部通过 ✔(${checks.length} 项)`);
|
||||
console.log(`ZIP 样本:outputs/test_zip_out.zip(${zip ? zip.length : 0} bytes,${entryCount} 个条目)`);
|
||||
process.exit(fail || errs ? 1 : 0);
|
||||
})();
|
||||
|
||||
function concat(parts) {
|
||||
const list = parts.map((p) => Buffer.from(p));
|
||||
return Buffer.concat(list);
|
||||
}
|
||||
+108
@@ -0,0 +1,108 @@
|
||||
/* =========================================================================
|
||||
BiRefNet WebUI · 极简 ZIP 打包器(无第三方依赖)
|
||||
用途:结果列表里的记录可能来自多个不同任务,服务端的「按任务打包」不再适用,
|
||||
所以在前端把所有可见结果打成一个 ZIP。
|
||||
实现:ZIP 规范里最朴素的 STORE(不压缩)方式 —— 结果本来就是 PNG,
|
||||
已经过 deflate,再压一次收益几乎为零,直接存更快更简单。
|
||||
========================================================================= */
|
||||
(function (root) {
|
||||
'use strict';
|
||||
|
||||
const CRC_TABLE = (() => {
|
||||
const table = new Uint32Array(256);
|
||||
for (let i = 0; i < 256; i++) {
|
||||
let c = i;
|
||||
for (let k = 0; k < 8; k++) c = (c & 1) ? (0xedb88320 ^ (c >>> 1)) : (c >>> 1);
|
||||
table[i] = c >>> 0;
|
||||
}
|
||||
return table;
|
||||
})();
|
||||
|
||||
function crc32(bytes) {
|
||||
let c = 0xffffffff;
|
||||
for (let i = 0; i < bytes.length; i++) c = CRC_TABLE[(c ^ bytes[i]) & 0xff] ^ (c >>> 8);
|
||||
return (c ^ 0xffffffff) >>> 0;
|
||||
}
|
||||
|
||||
/** JS Date -> MS-DOS 时间/日期(ZIP 头里的老格式)。 */
|
||||
function dosStamp(date) {
|
||||
const time = ((date.getHours() & 0x1f) << 11) | ((date.getMinutes() & 0x3f) << 5) | ((date.getSeconds() >> 1) & 0x1f);
|
||||
const day = (((date.getFullYear() - 1980) & 0x7f) << 9) | (((date.getMonth() + 1) & 0x0f) << 5) | (date.getDate() & 0x1f);
|
||||
return { time, day };
|
||||
}
|
||||
|
||||
/**
|
||||
* 生成一个 ZIP 文件。
|
||||
* @param {Array<{name: string, data: Uint8Array}>} entries
|
||||
* @returns {Uint8Array}
|
||||
*/
|
||||
function buildZip(entries) {
|
||||
const list = entries || [];
|
||||
if (!list.length) throw new Error('没有可打包的文件');
|
||||
const encoder = new TextEncoder();
|
||||
const { time, day } = dosStamp(new Date());
|
||||
|
||||
const body = []; // 本地文件头 + 数据
|
||||
const central = []; // 中央目录
|
||||
let offset = 0;
|
||||
|
||||
for (const entry of list) {
|
||||
const name = encoder.encode(String(entry.name));
|
||||
const data = entry.data instanceof Uint8Array ? entry.data : new Uint8Array(entry.data);
|
||||
const crc = crc32(data);
|
||||
|
||||
const local = new Uint8Array(30 + name.length);
|
||||
const lv = new DataView(local.buffer);
|
||||
lv.setUint32(0, 0x04034b50, true); // 本地文件头签名
|
||||
lv.setUint16(4, 20, true); // 解压所需版本
|
||||
lv.setUint16(6, 0x0800, true); // 标志位:文件名为 UTF-8
|
||||
lv.setUint16(8, 0, true); // 压缩方式 0 = STORE
|
||||
lv.setUint16(10, time, true);
|
||||
lv.setUint16(12, day, true);
|
||||
lv.setUint32(14, crc, true);
|
||||
lv.setUint32(18, data.length, true); // 压缩后大小
|
||||
lv.setUint32(22, data.length, true); // 原始大小
|
||||
lv.setUint16(26, name.length, true);
|
||||
lv.setUint16(28, 0, true); // 扩展字段长度
|
||||
local.set(name, 30);
|
||||
body.push(local, data);
|
||||
|
||||
const cd = new Uint8Array(46 + name.length);
|
||||
const cv = new DataView(cd.buffer);
|
||||
cv.setUint32(0, 0x02014b50, true); // 中央目录签名
|
||||
cv.setUint16(4, 20, true); // 生成程序版本
|
||||
cv.setUint16(6, 20, true); // 解压所需版本
|
||||
cv.setUint16(8, 0x0800, true);
|
||||
cv.setUint16(10, 0, true); // STORE
|
||||
cv.setUint16(12, time, true);
|
||||
cv.setUint16(14, day, true);
|
||||
cv.setUint32(16, crc, true);
|
||||
cv.setUint32(20, data.length, true);
|
||||
cv.setUint32(24, data.length, true);
|
||||
cv.setUint16(28, name.length, true);
|
||||
cv.setUint32(42, offset, true); // 本地头偏移
|
||||
cd.set(name, 46);
|
||||
central.push(cd);
|
||||
|
||||
offset += local.length + data.length;
|
||||
}
|
||||
|
||||
const centralSize = central.reduce((n, c) => n + c.length, 0);
|
||||
const eocd = new Uint8Array(22);
|
||||
const ev = new DataView(eocd.buffer);
|
||||
ev.setUint32(0, 0x06054b50, true); // 中央目录结束标记
|
||||
ev.setUint16(8, list.length, true); // 本盘条目数
|
||||
ev.setUint16(10, list.length, true); // 总条目数
|
||||
ev.setUint32(12, centralSize, true);
|
||||
ev.setUint32(16, offset, true);
|
||||
|
||||
const out = new Uint8Array(offset + centralSize + 22);
|
||||
let p = 0;
|
||||
for (const part of body) { out.set(part, p); p += part.length; }
|
||||
for (const c of central) { out.set(c, p); p += c.length; }
|
||||
out.set(eocd, p);
|
||||
return out;
|
||||
}
|
||||
|
||||
root.BiRefNetZip = { buildZip, crc32 };
|
||||
})(typeof window !== 'undefined' ? window : globalThis);
|
||||
Reference in New Issue
Block a user