commit 5833a303bcf82f0e0aede6d5839abcd37e77c08b Author: Hardwell99 <3022256519@qq.com> Date: Thu Oct 8 09:46:47 2026 +0800 feat: initial project setup diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..d5ab54b --- /dev/null +++ b/.gitignore @@ -0,0 +1,6 @@ +/python +*.pyc +/.workbuddy +/outputs +/models +run.bat diff --git a/README.md b/README.md new file mode 100644 index 0000000..57175ef --- /dev/null +++ b/README.md @@ -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_<原文件名主干>_.<后缀>`;主干超过 **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/` | 任务状态与进度(前端 400ms 轮询);`?tail=N` 只回传最后 N 张,批量任务靠它压载荷 | +| POST | `/api/tasks//cancel` | 取消(阶段粒度) | +| GET | `/api/tasks//file//` | 取结果,kind = cutout / mask / original | +| GET | `/api/tasks//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)。 diff --git a/Webui.py b/Webui.py new file mode 100644 index 0000000..36bf708 --- /dev/null +++ b/Webui.py @@ -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 指定模型代码目录(默认用项目内 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()) diff --git a/birefnet_web/__init__.py b/birefnet_web/__init__.py new file mode 100644 index 0000000..a239df6 --- /dev/null +++ b/birefnet_web/__init__.py @@ -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" diff --git a/birefnet_web/batch.py b/birefnet_web/batch.py new file mode 100644 index 0000000..362d400 --- /dev/null +++ b/birefnet_web/batch.py @@ -0,0 +1,222 @@ +# -*- coding: utf-8 -*- +"""目录批量抠图:路径校验 / 目录扫描 / 输出文件命名。 + +三条硬约束(对应需求): + +1. **校验阶段只读不写**:输入目录、输出目录都必须**已经存在**,本模块 + 绝不 ``mkdir``。路径写错时直接报错,而不是顺手把目录树建出来 —— + 这样一次误输入(或网页上的恶意路径)不会在磁盘上留下任何东西。 +2. **命名固定**:``RMBG_<原主干>_.<后缀>``,原主干超过 + :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)>_.<后缀>`` 组装输出文件名。 + + 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}") diff --git a/birefnet_web/compat.py b/birefnet_web/compat.py new file mode 100644 index 0000000..669f5d4 --- /dev/null +++ b/birefnet_web/compat.py @@ -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 diff --git a/birefnet_web/engine.py b/birefnet_web/engine.py new file mode 100644 index 0000000..d508bd9 --- /dev/null +++ b/birefnet_web/engine.py @@ -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", + } diff --git a/birefnet_web/fsutil.py b/birefnet_web/fsutil.py new file mode 100644 index 0000000..c0c4a32 --- /dev/null +++ b/birefnet_web/fsutil.py @@ -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))) diff --git a/birefnet_web/imageops.py b/birefnet_web/imageops.py new file mode 100644 index 0000000..a8eb7a3 --- /dev/null +++ b/birefnet_web/imageops.py @@ -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 diff --git a/birefnet_web/server.py b/birefnet_web/server.py new file mode 100644 index 0000000..6893c15 --- /dev/null +++ b/birefnet_web/server.py @@ -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/);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//<序号>_<主干>_cutout.png``。 + * 批量模式:写进用户指定的输出目录,文件名按需求固定为 + ``RMBG_<原主干(≤20)>_.<原后缀>``;同名时递增 ``_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/[/...] + 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() diff --git a/config.json b/config.json new file mode 100644 index 0000000..7fc9fc6 --- /dev/null +++ b/config.json @@ -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 +} diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..dea12e2 --- /dev/null +++ b/requirements.txt @@ -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 diff --git a/selftest.py b/selftest.py new file mode 100644 index 0000000..6010baa --- /dev/null +++ b/selftest.py @@ -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()) diff --git a/selftest_batch.py b/selftest_batch.py new file mode 100644 index 0000000..46a9024 --- /dev/null +++ b/selftest_batch.py @@ -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()) diff --git a/vendor/comfyui_birefnet_ll/LICENSE b/vendor/comfyui_birefnet_ll/LICENSE new file mode 100644 index 0000000..4599e07 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/LICENSE @@ -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. \ No newline at end of file diff --git a/vendor/comfyui_birefnet_ll/birefnet/__init__.py b/vendor/comfyui_birefnet_ll/birefnet/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/vendor/comfyui_birefnet_ll/birefnet/config.py b/vendor/comfyui_birefnet_ll/birefnet/config.py new file mode 100644 index 0000000..1630b46 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/config.py @@ -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]) + + diff --git a/vendor/comfyui_birefnet_ll/birefnet/dataset.py b/vendor/comfyui_birefnet_ll/birefnet/dataset.py new file mode 100644 index 0000000..d226485 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/dataset.py @@ -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) diff --git a/vendor/comfyui_birefnet_ll/birefnet/image_proc.py b/vendor/comfyui_birefnet_ll/birefnet/image_proc.py new file mode 100644 index 0000000..d976d7e --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/image_proc.py @@ -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) \ No newline at end of file diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/__init__.py b/vendor/comfyui_birefnet_ll/birefnet/models/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/backbones/__init__.py b/vendor/comfyui_birefnet_ll/birefnet/models/backbones/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/backbones/build_backbone.py b/vendor/comfyui_birefnet_ll/birefnet/models/backbones/build_backbone.py new file mode 100644 index 0000000..08b25dd --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/models/backbones/build_backbone.py @@ -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 diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/backbones/pvt_v2.py b/vendor/comfyui_birefnet_ll/birefnet/models/backbones/pvt_v2.py new file mode 100644 index 0000000..3089c53 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/models/backbones/pvt_v2.py @@ -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) diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/backbones/swin_v1.py b/vendor/comfyui_birefnet_ll/birefnet/models/backbones/swin_v1.py new file mode 100644 index 0000000..47631fe --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/models/backbones/swin_v1.py @@ -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 diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/birefnet.py b/vendor/comfyui_birefnet_ll/birefnet/models/birefnet.py new file mode 100644 index 0000000..67feb0f --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/models/birefnet.py @@ -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] diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/modules/__init__.py b/vendor/comfyui_birefnet_ll/birefnet/models/modules/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/modules/aspp.py b/vendor/comfyui_birefnet_ll/birefnet/models/modules/aspp.py new file mode 100644 index 0000000..d910d98 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/models/modules/aspp.py @@ -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) diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/modules/decoder_blocks.py b/vendor/comfyui_birefnet_ll/birefnet/models/modules/decoder_blocks.py new file mode 100644 index 0000000..487182f --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/models/modules/decoder_blocks.py @@ -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 \ No newline at end of file diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/modules/deform_conv.py b/vendor/comfyui_birefnet_ll/birefnet/models/modules/deform_conv.py new file mode 100644 index 0000000..43f5e57 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/models/modules/deform_conv.py @@ -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 diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/modules/lateral_blocks.py b/vendor/comfyui_birefnet_ll/birefnet/models/modules/lateral_blocks.py new file mode 100644 index 0000000..de907ac --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/models/modules/lateral_blocks.py @@ -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 diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/modules/prompt_encoder.py b/vendor/comfyui_birefnet_ll/birefnet/models/modules/prompt_encoder.py new file mode 100644 index 0000000..23ce18c --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/models/modules/prompt_encoder.py @@ -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 + diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/modules/utils.py b/vendor/comfyui_birefnet_ll/birefnet/models/modules/utils.py new file mode 100644 index 0000000..59bd912 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/models/modules/utils.py @@ -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) diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/refinement/__init__.py b/vendor/comfyui_birefnet_ll/birefnet/models/refinement/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/refinement/refiner.py b/vendor/comfyui_birefnet_ll/birefnet/models/refinement/refiner.py new file mode 100644 index 0000000..c3e2211 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/models/refinement/refiner.py @@ -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 diff --git a/vendor/comfyui_birefnet_ll/birefnet/models/refinement/stem_layer.py b/vendor/comfyui_birefnet_ll/birefnet/models/refinement/stem_layer.py new file mode 100644 index 0000000..2dc7f0f --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/models/refinement/stem_layer.py @@ -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 diff --git a/vendor/comfyui_birefnet_ll/birefnet/utils.py b/vendor/comfyui_birefnet_ll/birefnet/utils.py new file mode 100644 index 0000000..f7ea873 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet/utils.py @@ -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 \ No newline at end of file diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/__init__.py b/vendor/comfyui_birefnet_ll/birefnet_old/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/config.py b/vendor/comfyui_birefnet_ll/birefnet_old/config.py new file mode 100644 index 0000000..f63caa0 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/config.py @@ -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] diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/dataset.py b/vendor/comfyui_birefnet_ll/birefnet_old/dataset.py new file mode 100644 index 0000000..5c30e90 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/dataset.py @@ -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) diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/backbones/__init__.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/backbones/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/backbones/build_backbone.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/backbones/build_backbone.py new file mode 100644 index 0000000..e236636 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/backbones/build_backbone.py @@ -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 diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/backbones/pvt_v2.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/backbones/pvt_v2.py new file mode 100644 index 0000000..947a122 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/backbones/pvt_v2.py @@ -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) diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/backbones/swin_v1.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/backbones/swin_v1.py new file mode 100644 index 0000000..62a7aea --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/backbones/swin_v1.py @@ -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 diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/birefnet.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/birefnet.py new file mode 100644 index 0000000..fba83f0 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/birefnet.py @@ -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)) diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/__init__.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/aspp.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/aspp.py new file mode 100644 index 0000000..ce842f7 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/aspp.py @@ -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) diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/attentions.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/attentions.py new file mode 100644 index 0000000..e1032af --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/attentions.py @@ -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 + diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/decoder_blocks.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/decoder_blocks.py new file mode 100644 index 0000000..3d32736 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/decoder_blocks.py @@ -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 diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/deform_conv.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/deform_conv.py new file mode 100644 index 0000000..43f5e57 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/deform_conv.py @@ -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 diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/ing.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/ing.py new file mode 100644 index 0000000..b0026e9 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/ing.py @@ -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 \ No newline at end of file diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/lateral_blocks.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/lateral_blocks.py new file mode 100644 index 0000000..de907ac --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/lateral_blocks.py @@ -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 diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/mlp.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/mlp.py new file mode 100644 index 0000000..506bfe4 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/mlp.py @@ -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 + diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/utils.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/utils.py new file mode 100644 index 0000000..59bd912 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/modules/utils.py @@ -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) diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/refinement/__init__.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/refinement/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/refinement/refiner.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/refinement/refiner.py new file mode 100644 index 0000000..15bbb9a --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/refinement/refiner.py @@ -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 diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/models/refinement/stem_layer.py b/vendor/comfyui_birefnet_ll/birefnet_old/models/refinement/stem_layer.py new file mode 100644 index 0000000..128e61f --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/models/refinement/stem_layer.py @@ -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 diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/preproc.py b/vendor/comfyui_birefnet_ll/birefnet_old/preproc.py new file mode 100644 index 0000000..a059c5d --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/preproc.py @@ -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) diff --git a/vendor/comfyui_birefnet_ll/birefnet_old/utils.py b/vendor/comfyui_birefnet_ll/birefnet_old/utils.py new file mode 100644 index 0000000..d44c7d2 --- /dev/null +++ b/vendor/comfyui_birefnet_ll/birefnet_old/utils.py @@ -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 \ No newline at end of file diff --git a/vendor/comfyui_birefnet_ll/pyproject.toml b/vendor/comfyui_birefnet_ll/pyproject.toml new file mode 100644 index 0000000..60cb73d --- /dev/null +++ b/vendor/comfyui_birefnet_ll/pyproject.toml @@ -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 = "" diff --git a/vendor/comfyui_birefnet_ll/requirements.txt b/vendor/comfyui_birefnet_ll/requirements.txt new file mode 100644 index 0000000..293d45b --- /dev/null +++ b/vendor/comfyui_birefnet_ll/requirements.txt @@ -0,0 +1,3 @@ +numpy +opencv-python +timm \ No newline at end of file diff --git a/web/app.js b/web/app.js new file mode 100644 index 0000000..8085dd9 --- /dev/null +++ b/web/app.js @@ -0,0 +1,1127 @@ +/* ========================================================================= + BiRefNet WebUI 前端逻辑 + - 上传:拖拽 / 点击 / 剪贴板粘贴,单槽(每次只处理一张,新图替换旧图) + - 任务:POST /api/tasks 建任务 → 轮询 GET /api/tasks/ 拿进度 + - 结果:跨任务累积成历史列表,任何操作(换参数 / 传新图 / 再次处理) + 都不会清空已有记录;仅「×」单条移除与「清空」会减少记录 + - 参数:每条记录携带处理时的参数快照;按文件名记住「最近一次」的参数, + 重新处理时自动重新赋值给表单 + ========================================================================= */ +(() => { + 'use strict'; + + const $ = (id) => document.getElementById(id); + const ACCEPTED_TYPES = ['image/png', 'image/jpeg', 'image/webp', 'image/bmp', 'image/tiff', 'image/gif']; + const IMAGE_EXT = /\.(png|jpe?g|webp|bmp|tiff?|gif)$/i; + const LS_KEY = 'birefnet.webui.options'; // 表单当前值 + const LS_PARAMS = 'birefnet.webui.imageParams'; // 文件名 -> 最近一次参数 + // 本地存档结构版本:v2 起「同时输出遮罩」默认关闭、输出尺寸参数改名 final_longest_side。 + // 老存档里被显式存下来的 output_mask 不能覆盖新默认值,否则用户永远关不掉。 + const LS_SCHEMA = 2; + const MAX_PARAM_ENTRIES = 200; + const FINAL_STATES = ['done', 'partial', 'failed', 'canceled']; + // 批量任务动辄几百张,轮询只拉最后 N 行:处理是按序串行的,窗口外的一定已是终态 + const BATCH_TAIL = 40; + + const state = { + pending: null, // 单槽待处理:{ file, url, key(文件名) } + models: [], + defaults: {}, + env: null, + activeTaskId: null, // 正在轮询的单图任务 + timer: null, + view: 'cutout', + records: new Map(), // uid -> record(跨任务累积,uid = taskId#index) + dismissed: new Set(), // uid,用户显式移除过,轮询不再推回来 + selected: null, // 当前选中的 uid + running: false, + paramsByImage: {}, // 文件名 -> 最近一次使用的参数 + taskOptions: new Map(), // taskId -> 该任务使用的参数 + // 目录批量任务:与单图任务共用同一个服务端串行队列,但前端独立进度与清单 + batch: { + taskId: null, + running: false, + timer: null, + rows: new Map(), // 图片序号 ->
  • 行节点 + reportVisible: false, // 批量清单是否展开(空状态判断要用它,不能嗅 DOM 属性) + }, + }; + + let noteTimer = null; + + /* ----------------------------- 工具函数 ----------------------------- */ + function el(tag, attrs = {}, ...children) { + const node = document.createElement(tag); + for (const [k, v] of Object.entries(attrs)) { + if (k === 'class') node.className = v; + else if (k === 'text') node.textContent = v; + else if (k.startsWith('on')) node.addEventListener(k.slice(2), v); + else if (v !== null && v !== undefined) node.setAttribute(k, v); + } + for (const child of children) if (child) node.append(child); + return node; + } + + function fmtBytes(n) { + if (!n && n !== 0) return '—'; + if (n < 1024) return `${n} B`; + if (n < 1024 * 1024) return `${(n / 1024).toFixed(0)} KB`; + return `${(n / 1024 / 1024).toFixed(1)} MB`; + } + + function setError(msg) { + const box = $('error-text'); + if (!msg) { box.hidden = true; box.textContent = ''; return; } + box.hidden = false; + box.textContent = msg; + } + + function setNote(msg) { + const box = $('pending-note'); + if (box) box.textContent = msg || ''; + if (noteTimer) clearTimeout(noteTimer); + if (msg) noteTimer = setTimeout(() => { if (box) box.textContent = ''; }, 6000); + } + + async function api(path, options) { + const res = await fetch(path, options); + const text = await res.text(); + let data; + try { data = text ? JSON.parse(text) : {}; } catch { data = { error: text }; } + if (!res.ok || data.ok === false) throw new Error(data.error || `请求失败 ${res.status}`); + return data; + } + + function scrollTo(node) { + if (node && typeof node.scrollIntoView === 'function') { + try { node.scrollIntoView({ block: 'nearest' }); } catch { /* 忽略 */ } + } + } + + /* ----------------------------- 初始化 ----------------------------- */ + async function init() { + bindEvents(); + restoreOptions(); + try { + await loadState(); + } catch (err) { + setError(`读取服务状态失败:${err.message}`); + } + try { + await loadHistory(); // 恢复历史结果(服务端仍在内存里的任务) + } catch (err) { + setError(`恢复历史结果失败:${err.message}`); + } + } + + async function loadState() { + const data = await api('/api/state'); + state.models = data.models || []; + state.defaults = data.defaults || {}; + state.env = data.environment || {}; + + const device = state.env.cuda + ? (state.env.devices?.[0]?.name || 'CUDA') + ` · ${state.env.devices?.[0]?.free_mem || '?'}G 空闲` + : `CPU · torch ${state.env.torch}`; + $('meta-device').textContent = device; + $('meta-models').textContent = `${state.models.length} 个权重`; + const nodePath = data.node_dir || ''; + const nodeEl = $('meta-node'); + nodeEl.textContent = nodePath; + nodeEl.title = nodePath; + + renderModelOptions(); + toggleDeviceOptions(); + + // 批量处理的默认输出目录 = 项目 outputs(服务端返回的绝对路径) + const outInput = $('batch-output-dir'); + if (data.output_dir) { + outInput.placeholder = `默认:${data.output_dir}`; + if (!outInput.value) outInput.value = data.output_dir; + } + } + + function renderModelOptions() { + const select = $('opt-model'); + const previous = select.value; + select.textContent = ''; + if (!state.models.length) { + select.append(el('option', { value: '', text: '未找到模型权重' })); + $('model-hint').textContent = '请将 *.safetensors 放入模型目录'; + return; + } + for (const m of state.models) { + const label = `${m.name} · ${m.size_mb}MB · ${m.backbone}`; + select.append(el('option', { value: m.key, text: label, title: m.path })); + } + const want = previous || state.defaults.model; + const exists = state.models.some((m) => m.key === want); + select.value = exists ? want : state.models[0].key; + onModelChange(); + } + + function onModelChange() { + const model = state.models.find((m) => m.key === $('opt-model').value); + $('model-hint').textContent = model ? `${model.file} · ${model.arch === 'old' ? '旧版' : '新版'}` : ''; + } + + function toggleDeviceOptions() { + const select = $('opt-device'); + for (const opt of select.querySelectorAll('option')) { + if (opt.value === 'auto') opt.disabled = !state.env?.cuda; + } + if (!state.env?.cuda && select.value === 'auto') select.value = 'cpu'; + } + + /* ----------------------------- 事件绑定 ----------------------------- */ + function bindEvents() { + const dropzone = $('dropzone'); + const input = $('file-input'); + + dropzone.addEventListener('click', () => input.click()); + dropzone.addEventListener('keydown', (e) => { + if (e.key === 'Enter' || e.key === ' ') { e.preventDefault(); input.click(); } + }); + input.addEventListener('change', () => { addFiles(input.files); input.value = ''; }); + + ['dragenter', 'dragover'].forEach((evt) => + dropzone.addEventListener(evt, (e) => { e.preventDefault(); dropzone.classList.add('over'); })); + ['dragleave', 'drop'].forEach((evt) => + dropzone.addEventListener(evt, (e) => { e.preventDefault(); dropzone.classList.remove('over'); })); + dropzone.addEventListener('drop', (e) => addFiles(e.dataTransfer?.files)); + + document.addEventListener('dragover', (e) => e.preventDefault()); + document.addEventListener('drop', (e) => e.preventDefault()); + document.addEventListener('paste', (e) => addFiles(e.clipboardData?.files)); + + $('btn-run').addEventListener('click', run); + $('btn-cancel').addEventListener('click', cancelTask); + $('btn-batch').addEventListener('click', batchRun); + $('btn-batch-cancel').addEventListener('click', () => cancelTaskById(state.batch.taskId)); + $('btn-batch-hide').addEventListener('click', () => { + $('batch-report').hidden = true; + state.batch.reportVisible = false; + updateStageTools(); + }); + for (const id of ['batch-input-dir', 'batch-output-dir']) { + $(id).addEventListener('keydown', (e) => { if (e.key === 'Enter') batchRun(); }); + } + $('btn-clear').addEventListener('click', () => clearResults(false)); + $('btn-zip').addEventListener('click', downloadZip); + $('btn-unload').addEventListener('click', async () => { + try { await api('/api/engine/unload', { method: 'POST' }); flash($('btn-unload'), '已卸载'); } + catch (err) { setError(err.message); } + }); + $('btn-reload-models').addEventListener('click', async () => { + try { + const data = await api('/api/models/reload', { method: 'POST' }); + state.models = data.models || []; + renderModelOptions(); + flash($('btn-reload-models'), '已刷新'); + } catch (err) { setError(err.message); } + }); + + $('opt-model').addEventListener('change', onModelChange); + $('opt-resolution-mode').addEventListener('change', syncSizeRows); + $('opt-background').addEventListener('change', syncBgRow); + $('opt-refine').addEventListener('change', syncBlurRow); + $('opt-bgcolor').addEventListener('input', () => saveOptions()); + $('swatches').addEventListener('click', (e) => { + const btn = e.target.closest('button[data-color]'); + if (!btn) return; + $('opt-bgcolor').value = btn.dataset.color; + saveOptions(); + }); + $('view-switch').addEventListener('click', (e) => { + const btn = e.target.closest('button[data-view]'); + if (!btn) return; + setView(btn.dataset.view); + }); + + for (const node of document.querySelectorAll('select, input')) { + if (node.id === 'file-input') continue; + node.addEventListener('change', saveOptions); + } + + document.addEventListener('keydown', (e) => { + if ((e.ctrlKey || e.metaKey) && e.key === 'Enter' && !$('btn-run').disabled) run(); + if (e.key === 'Escape' && state.selected !== null) toggleSelect(state.selected); + }); + + // 初始化只同步行的显隐,**不能**顺手写存档:那会把用户上次的参数覆盖成默认值, + // 后面的 restoreOptions() 就再也读不回来了。 + syncSizeRows(false); syncBgRow(false); syncBlurRow(false); + } + + function flash(btn, text) { + if (!btn) return; + if (btn.dataset.timer) clearTimeout(Number(btn.dataset.timer)); + const origin = btn.dataset.label || btn.textContent; + btn.dataset.label = origin; + btn.textContent = text; + btn.dataset.timer = String(setTimeout(() => { btn.textContent = origin; }, 1200)); + } + + /** 同步「预处理尺寸」相关行的显隐;persist=false 时只改界面不写存档。 */ + function syncSizeRows(persist = true) { + const mode = $('opt-resolution-mode').value; + $('row-longest').classList.toggle('hidden', mode !== 'longest'); + $('row-wh').classList.toggle('hidden', mode !== 'custom'); + if (persist) saveOptions(); + } + + function syncBgRow(persist = true) { + $('row-bgcolor').classList.toggle('hidden', $('opt-background').value !== 'color'); + if (persist) saveOptions(); + } + + function syncBlurRow(persist = true) { + $('row-blur').classList.toggle('hidden', !$('opt-refine').checked); + if (persist) saveOptions(); + } + + /* --------------------------- 待处理(单槽) --------------------------- */ + /** 图片身份 = 文件名(参数按此记忆;不同图片只要文件名不同就互不影响)。 */ + function imageKey(file) { + return (file && file.name) || ''; + } + + /** 待处理槽去重用:文件名 + 字节数。 */ + function fileSig(file) { + return `${file.name}|${file.size}`; + } + + function addFiles(list) { + if (!list || !list.length) return; + const items = Array.from(list); + if (items.length > 1) { + setNote(`一次只处理一张图片,已保留最后一张(忽略前 ${items.length - 1} 张)`); + } + offerImage(items[items.length - 1], { keepNote: items.length > 1 }); + } + + /** + * 放入待处理槽(单槽:直接替换旧图)。 + * @param {File} file + * @param {{keepNote?: boolean, params?: object|null}} opts + * params === undefined → 按文件名自动套用「最近一次」的参数 + * params === null → 不动表单(调用方自己决定) + * params === {…} → 套用给定参数 + * @returns {boolean} 是否放入成功(重复图片返回 false) + */ + function offerImage(file, opts = {}) { + const { keepNote = false, params = undefined } = opts; + if (!file) return false; + if (!ACCEPTED_TYPES.includes(file.type) && !IMAGE_EXT.test(file.name)) { + setError(`已忽略非图片文件:${file.name}`); + return false; + } + const sig = fileSig(file); + if (state.pending && state.pending.sig === sig) { + setNote(`「${file.name}」已在待处理中,无需重复添加`); + return false; + } + const replaced = !!state.pending; + if (state.pending) URL.revokeObjectURL(state.pending.url); + + const key = imageKey(file); + state.pending = { file, url: URL.createObjectURL(file), key, sig }; + setError(''); + + // 参数:默认自动套用这张图上一次用过的参数 + let applied = null; + if (params !== null) { + applied = params !== undefined ? params : lookupParams(key); + if (applied) applyOptions(applied); + } + + const notes = []; + if (replaced) notes.push(`已用「${file.name}」替换上一张待处理图片`); + if (applied) notes.push(`已套用「${file.name}」上次保存的参数`); + if (!keepNote && notes.length) setNote(notes.join(';')); + else if (!keepNote && !notes.length) setNote(''); + + renderPending(); + return true; + } + + function removePending() { + if (!state.pending) return; + URL.revokeObjectURL(state.pending.url); + state.pending = null; + setNote(''); + renderPending(); + } + + function renderPending() { + const ul = $('thumbs'); + ul.textContent = ''; + const p = state.pending; + if (p) { + ul.append(el('li', { class: 'pending-item' }, + el('div', { class: 'thumb-wrap' }, el('img', { src: p.url, alt: '' })), + el('div', { class: 'pending-meta' }, + el('span', { class: 'tname', text: p.file.name, title: p.file.name }), + el('span', { class: 'tsize', text: fmtBytes(p.file.size) }), + ), + el('button', { class: 'tremove', type: 'button', title: '移除待处理图片', text: '×', onclick: removePending }), + )); + } + $('file-count').textContent = p ? '待处理 1 张' : '未选择'; + $('btn-run').disabled = !p || state.running || !state.models.length; + } + + /* ----------------------------- 参数 ----------------------------- */ + function collectOptions() { + return { + model: $('opt-model').value, + device: $('opt-device').value, + dtype: $('opt-dtype').value, + arch: $('opt-arch').value, + resolution_mode: $('opt-resolution-mode').value, + width: Number($('opt-width').value) || 1024, + height: Number($('opt-height').value) || 1024, + longest_side: Number($('opt-longest-side').value) || 1024, + upscale_method: $('opt-upscale').value, + mask_threshold: Number($('opt-threshold').value) || 0, + refine_foreground: $('opt-refine').checked, + blur_size: Number($('opt-blur1').value) || 90, + blur_size_two: Number($('opt-blur2').value) || 6, + background: $('opt-background').value, + bg_color: $('opt-bgcolor').value, + output_mask: $('opt-mask').checked, + final_longest_side: Number($('opt-final-side').value) || 0, + }; + } + + const OPTION_FIELDS = { + 'opt-device': 'device', 'opt-dtype': 'dtype', 'opt-arch': 'arch', + 'opt-resolution-mode': 'resolution_mode', 'opt-width': 'width', 'opt-height': 'height', + 'opt-longest-side': 'longest_side', 'opt-upscale': 'upscale_method', + 'opt-threshold': 'mask_threshold', 'opt-blur1': 'blur_size', 'opt-blur2': 'blur_size_two', + 'opt-background': 'background', 'opt-bgcolor': 'bg_color', 'opt-final-side': 'final_longest_side', + }; + + /** 把一份参数重新赋值到表单上(重新处理的关键动作)。 */ + function applyOptions(src) { + if (!src || typeof src !== 'object') return false; + let touched = false; + for (const [id, key] of Object.entries(OPTION_FIELDS)) { + let v = src[key]; + // 兼容 v1 存档里的旧参数名(max_output_side = 输出长边上限) + if (v === undefined && key === 'final_longest_side') v = src.max_output_side; + if (v === undefined || v === null || v === '') continue; + $(id).value = String(v); + touched = true; + } + // 模型权重可能已经被删掉/改名,只在仍然存在时才切 + if (src.model && state.models.some((m) => m.key === src.model)) { + $('opt-model').value = src.model; + touched = true; + } + if (typeof src.refine_foreground === 'boolean') $('opt-refine').checked = src.refine_foreground; + if (typeof src.output_mask === 'boolean') $('opt-mask').checked = src.output_mask; + onModelChange(); + syncSizeRows(false); syncBgRow(false); syncBlurRow(false); + saveOptions(); + return touched; + } + + function saveOptions() { + try { + localStorage.setItem(LS_KEY, JSON.stringify({ ...collectOptions(), _schema: LS_SCHEMA })); + } catch { /* 忽略隐私模式 */ } + } + + function restoreOptions() { + const saved = readJson(LS_KEY); + if (!saved || typeof saved !== 'object') return; + // 老存档(v1,没有 _schema 字段):丢掉 output_mask,让它在 v2 里回落到「默认不输出遮罩」 + if ((Number(saved._schema) || 0) < LS_SCHEMA) delete saved.output_mask; + applyOptions(saved); + } + + function readJson(key) { + try { return JSON.parse(localStorage.getItem(key) || 'null'); } catch { return null; } + } + + /** 记住某张图「最近一次」使用的参数(同一张图多次处理 → 覆盖为最新)。 */ + function rememberParams(fileKey, options) { + if (!fileKey || !options) return; + state.paramsByImage[fileKey] = options; + const keys = Object.keys(state.paramsByImage); + if (keys.length > MAX_PARAM_ENTRIES) { + for (const stale of keys.slice(0, keys.length - MAX_PARAM_ENTRIES)) delete state.paramsByImage[stale]; + } + try { localStorage.setItem(LS_PARAMS, JSON.stringify(state.paramsByImage)); } catch { /* 忽略 */ } + } + + function lookupParams(fileKey) { + const hit = fileKey ? state.paramsByImage[fileKey] : null; + return hit && typeof hit === 'object' ? hit : null; + } + + function loadParams() { + const saved = readJson(LS_PARAMS); + if (saved && typeof saved === 'object') state.paramsByImage = saved; + } + + /** 参数摘要(卡片上那一行);完整 JSON 放在 title 里。 */ + function paramsSummary(o) { + if (!o) return ''; + const size = o.resolution_mode === 'custom' ? `${o.width}×${o.height}` + : o.resolution_mode === 'longest' ? `长边 ${o.longest_side}` + : '方形 1024'; + const bits = [o.model || '默认模型', size, o.background === 'color' ? `纯色 ${o.bg_color}` : '透明']; + if (o.refine_foreground) bits.push('精修'); + if (o.mask_threshold) bits.push(`阈值 ${o.mask_threshold}`); + if (o.dtype && o.dtype !== 'auto') bits.push(o.dtype); + if (o.device === 'cpu') bits.push('CPU'); + if (o.final_longest_side) bits.push(`最长边 ${o.final_longest_side}`); + return bits.join(' · '); + } + + /* ----------------------------- 任务 ----------------------------- */ + async function run() { + if (state.running) return; + if (!state.pending) { setError('请先选择图片'); return; } + if (!$('opt-model').value) { setError('没有可用模型,请检查模型目录'); return; } + + setError(''); + state.running = true; + $('btn-run').disabled = true; + $('btn-run').textContent = '处理中…'; + $('progress').hidden = false; + setProgress(0, '上传中…'); + // 注意:这里绝不清理结果区 —— 新任务的结果会追加到历史列表末尾 + + const options = collectOptions(); + const fileKey = state.pending.key; + const form = new FormData(); + form.append('files', state.pending.file, state.pending.file.name); + form.append('options', JSON.stringify(options)); + + try { + const data = await api('/api/tasks', { method: 'POST', body: form }); + rememberParams(fileKey, options); // 记录该图「最新」参数 + state.taskOptions.set(data.task_id, options); + state.activeTaskId = data.task_id; + preventSleep(false); + poll(); + } catch (err) { + finishRun(err.message); + } + } + + function poll() { + clearTimeout(state.timer); + state.timer = setTimeout(async () => { + const taskId = state.activeTaskId; + if (!taskId) return; + try { + const data = await api(`/api/tasks/${taskId}`); + renderTask(data.task); + if (FINAL_STATES.includes(data.task.state)) { + const counts = data.task.counts || {}; + const msg = data.task.state === 'canceled' ? '已取消' + : `完成 ${counts.done || 0} 张${counts.failed ? `,失败 ${counts.failed} 张` : ''}`; + finishRun(data.task.state === 'failed' ? (data.task.error || '处理失败') : null, msg); + return; + } + poll(); + } catch (err) { + finishRun(err.message); + } + }, 400); + } + + function finishRun(errorMsg, okMsg) { + state.running = false; + $('btn-run').disabled = !state.pending || !state.models.length; + $('btn-run').textContent = '开始抠图'; + preventSleep(true); + if (errorMsg) { setError(errorMsg); setProgress(0, '出错'); } + else if (okMsg) setProgress(1, okMsg); + } + + let wakeLock = null; + async function preventSleep(allow) { + try { + if (!allow) { + if (!wakeLock && 'wakeLock' in navigator) wakeLock = await navigator.wakeLock.request('screen'); + } else if (wakeLock) { + await wakeLock.release(); + wakeLock = null; + } + } catch { /* 浏览器不支持则忽略 */ } + } + + async function cancelTask() { + await cancelTaskById(state.activeTaskId); + } + + /** 取消一个任务(单图与批量共用同一套任务队列)。 */ + async function cancelTaskById(taskId) { + if (!taskId) return; + try { await api(`/api/tasks/${taskId}/cancel`, { method: 'POST' }); } catch { /* 忽略 */ } + } + + /* ----------------------------- 批量处理 ----------------------------- */ + /* 与单图模式共用「模型 / 参数」面板;差异只在输入来源(目录)与输出落点 + (用户指定目录 + RMBG_ 命名)。批量任务不铺成结果卡片,只在结果区顶部 + 显示一张逐文件清单 —— 几百张图铺卡片既拖慢页面也没有对比价值。 */ + + function setBatchError(msg) { + const box = $('batch-error'); + if (!msg) { box.hidden = true; box.textContent = ''; return; } + box.hidden = false; + box.textContent = msg; + } + + function setBatchProgress(ratio, stage) { + $('batch-progress-bar').style.width = `${Math.max(0, Math.min(1, ratio)) * 100}%`; + $('batch-progress-stage').textContent = stage || ''; + } + + async function batchRun() { + if (state.batch.running) return; + const inputDir = $('batch-input-dir').value.trim(); + const outputDir = $('batch-output-dir').value.trim(); // 留空 → 服务端用项目 outputs + if (!inputDir) { setBatchError('请先填写图片目录(模板路径)'); return; } + if (!state.models.length) { setBatchError('没有可用模型,请检查模型目录'); return; } + + setBatchError(''); + state.batch.running = true; + $('btn-batch').disabled = true; + $('btn-batch').textContent = '提交中…'; + $('batch-progress').hidden = false; + setBatchProgress(0, '扫描目录…'); + + const options = collectOptions(); + try { + const data = await api('/api/batch', { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ + input_dir: inputDir, + output_dir: outputDir, + recursive: $('batch-recursive').checked, + options, + }), + }); + resetBatchReport(data.input_dir || inputDir, data.output_dir || outputDir); + state.batch.taskId = data.task_id; + preventSleep(false); + setBatchProgress(0.01, `已提交 ${data.count} 张,排队中…`); + pollBatch(true); + } catch (err) { + batchFinish(err.message, null); + } + } + + /** 清空并显示批量清单(新的批量任务开始 / 刷新页面恢复时调用)。 */ + function resetBatchReport(inputDir, outputDir) { + state.batch.rows.clear(); + $('batch-rows').textContent = ''; + $('batch-report').hidden = false; + state.batch.reportVisible = true; + const badge = $('batch-badge'); + badge.textContent = '排队中'; + badge.className = 'badge'; + const path = $('batch-report-path'); + path.textContent = `${inputDir} → ${outputDir}`; + path.title = `输入目录:${inputDir}\n输出目录:${outputDir}`; + $('batch-report-sum').textContent = ''; + updateStageTools(); + } + + function pollBatch(full) { + clearTimeout(state.batch.timer); + state.batch.timer = setTimeout(async () => { + const taskId = state.batch.taskId; + if (!taskId) return; + try { + const data = await api(`/api/tasks/${taskId}?tail=${full ? 0 : BATCH_TAIL}`); + renderBatchTask(data.task); + if (FINAL_STATES.includes(data.task.state)) { + // 收尾时再拉一次全量,确保清单里每一行都是最终状态 + if (!full) { pollBatch(true); return; } + batchFinish(null, data.task); + return; + } + pollBatch(false); + } catch (err) { + batchFinish(err.message, null); + } + }, 400); + } + + const BATCH_BADGE = { + queued: ['排队中', ''], running: ['处理中', ''], done: ['完成', 'ok'], + partial: ['部分完成', 'warn'], failed: ['失败', 'err'], canceled: ['已取消', 'warn'], + }; + + function renderBatchTask(task) { + const [text, cls] = BATCH_BADGE[task.state] || [task.state, '']; + const badge = $('batch-badge'); + badge.textContent = text; + badge.className = `badge ${cls}`.trim(); + setBatchProgress(task.progress || 0, task.stage || ''); + + const counts = task.counts || {}; + const bits = [ + `共 ${task.images_total || 0} 张`, + `完成 ${counts.done || 0}`, + ]; + if (counts.failed) bits.push(`失败 ${counts.failed}`); + if (counts.pending) bits.push(`待处理 ${counts.pending}`); + if (task.elapsed) bits.push(`${Math.round(task.elapsed)}s`); + $('batch-report-sum').textContent = bits.join(' · '); + + for (const img of task.images || []) upsertBatchRow(img); + } + + function upsertBatchRow(img) { + let row = state.batch.rows.get(img.index); + if (!row) { + const cells = { + idx: el('span', { class: 'bi', text: String(img.index + 1) }), + from: el('span', { class: 'bs' }), + to: el('span', { class: 'bo' }), + time: el('span', { class: 'bt' }), + badge: el('span', { class: 'badge' }), + }; + row = el('li', { class: 'batch-row' }, cells.idx, cells.from, cells.to, cells.time, cells.badge); + row.cells = cells; + state.batch.rows.set(img.index, row); + $('batch-rows').append(row); + } + const cells = row.cells; + const [text, cls] = BATCH_BADGE[img.state] || [img.state, '']; + cells.from.textContent = img.name || ''; + cells.from.title = img.name || ''; + cells.to.textContent = img.output_name || ''; + cells.to.title = img.output_name || ''; + cells.time.textContent = img.state === 'done' ? `${img.elapsed || 0}s` : ''; + cells.badge.textContent = text; + cells.badge.className = `badge ${cls}`.trim(); + row.className = `batch-row ${img.state}`; + row.title = img.error ? `${img.name}:${img.error}` : (img.name || ''); + } + + function batchFinish(errorMsg, task) { + clearTimeout(state.batch.timer); + state.batch.running = false; + state.batch.taskId = null; + $('btn-batch').disabled = false; + $('btn-batch').textContent = '开始批量抠图'; + preventSleep(true); + if (errorMsg) { + setBatchError(errorMsg); + setBatchProgress(0, '出错'); + return; + } + const counts = task.counts || {}; + const okMsg = task.state === 'canceled' + ? '已取消' + : `已输出 ${counts.done || 0} 张${counts.failed ? `,失败 ${counts.failed} 张` : ''}`; + setBatchProgress(1, okMsg); + setBatchError(''); + flash($('btn-batch'), `已输出 ${counts.done || 0} 张`); + } + + function setProgress(ratio, stage) { + $('progress-bar').style.width = `${Math.max(0, Math.min(1, ratio)) * 100}%`; + $('progress-stage').textContent = stage || ''; + } + + /* --------------------------- 结果历史 --------------------------- */ + + function recordUid(taskId, index) { + return `${taskId}#${index}`; + } + + /** 服务端内存里还留着的历史任务 → 渲染成结果记录(刷新页面不丢)。 */ + async function loadHistory() { + let list = []; + try { + const data = await api('/api/tasks'); + list = data.tasks || []; + } catch { + return; // 拿不到就当没有历史,不打扰用户 + } + if (!list.length) return; + + // 批量任务不铺结果卡片(几百张图铺卡片既慢又没有对比价值), + // 只在结果区顶部恢复最近一次的批量清单 + restoreBatchReport(list.filter((task) => task.mode === 'batch')); + + const singles = list.filter((task) => task.mode !== 'batch'); + if (!singles.length) return; + + // 列表按创建时间倒序返回,这里反转为「旧 → 新」,追加式渲染 + const ordered = singles.slice().reverse(); + const details = await Promise.all(ordered.map(async (task) => { + if (task.options) return task; // 列表里已带参数 + try { return (await api(`/api/tasks/${task.id}`)).task || task; } catch { return task; } + })); + + for (const task of details) { + if (task.options) state.taskOptions.set(task.id, task.options); + for (const img of task.images || []) upsertCard(task.id, img); + } + updateStageTools(); + scrollResultsToBottom(); + + // 刷新页面时若还有任务在跑,接着轮询 + const newest = ordered[ordered.length - 1]; + if (newest && (newest.state === 'queued' || newest.state === 'running')) { + state.activeTaskId = newest.id; + state.running = true; + $('btn-run').disabled = true; + $('btn-run').textContent = '处理中…'; + $('progress').hidden = false; + poll(); + } + } + + /** 恢复最近一次批量任务的清单;若它还在跑,顺手接着轮询。 */ + function restoreBatchReport(batchTasks) { + if (!batchTasks.length) return; + const newest = batchTasks[0]; // 列表已按创建时间倒序 + resetBatchReport(newest.input_dir || '(未记录)', newest.output_dir || ''); + renderBatchTask(newest); + $('batch-progress').hidden = false; + if (newest.state === 'queued' || newest.state === 'running') { + state.batch.taskId = newest.id; + state.batch.running = true; + $('btn-batch').disabled = true; + $('btn-batch').textContent = '处理中…'; + preventSleep(false); + pollBatch(true); // 先拉全量把清单补齐,之后只拉尾巴 + } + } + + function renderTask(task) { + if (task.mode === 'batch') return; // 批量任务由批量清单自己渲染 + const options = state.taskOptions.get(task.id) || task.options || null; + if (options) state.taskOptions.set(task.id, options); + for (const img of task.images || []) upsertCard(task.id, img); + updateStageTools(); + setProgress(task.progress || 0, task.stage || ''); + } + + function upsertCard(taskId, img) { + const uid = recordUid(taskId, img.index); + if (state.dismissed.has(uid)) return null; // 用户已移除,轮询不再推回来 + + let rec = state.records.get(uid); + if (!rec) { + rec = buildRecord(taskId, img); + state.records.set(uid, rec); + $('results').append(rec.root); + if (img.state === 'pending' || img.state === 'running') scrollTo(rec.root); + } + rec.image = img; + if (!rec.options) rec.options = state.taskOptions.get(taskId) || null; + paintRecord(rec); + return rec; + } + + function paintRecord(rec) { + const img = rec.image; + const { refs } = rec; + + const badgeMap = { + pending: ['等待', ''], running: ['处理中', ''], done: ['完成', 'ok'], + failed: ['失败', 'err'], canceled: ['已取消', 'warn'], + }; + const [text, cls] = badgeMap[img.state] || [img.state, '']; + refs.badge.textContent = text; + refs.badge.className = `badge ${cls}`.trim(); + refs.statusText.textContent = `${img.stage || '处理中'} · ${Math.round((img.progress || 0) * 100)}%`; + refs.root.className = `card ${img.state}${state.selected === rec.uid ? ' selected' : ''}`; + + if (rec.options) { + refs.params.textContent = paramsSummary(rec.options); + refs.params.title = `任务 ${rec.taskId}\n参数:${JSON.stringify(rec.options, null, 2)}`; + } + + if (img.state === 'done' && !rec.built) { + rec.built = true; + buildCompare(rec); + refs.stats.append( + stat('尺寸', `${img.width}×${img.height}`), + stat('输入', `${img.input_size?.[0]}×${img.input_size?.[1]}`), + stat('前景', `${(img.coverage * 100).toFixed(1)}%`), + stat('耗时', `${img.elapsed}s`), + ); + const links = [ + downloadLink(img.urls.cutout, `${img.label}_cutout.png`, '下载抠图'), + img.urls.mask ? downloadLink(img.urls.mask, `${img.label}_mask.png`, '下载遮罩') : null, + el('button', { class: 'btn sm ghost', type: 'button', text: '复制到剪贴板', onclick: (e) => copyImage(rec, e.target) }), + ].filter(Boolean); + refs.actions.append(...links); + if (img.warning) refs.note.textContent = `⚠ ${img.warning}`; + } + if (img.state === 'failed' && img.error && !refs.error.textContent) { + refs.error.textContent = img.error; + } + } + + function buildRecord(taskId, img) { + const rec = { + uid: recordUid(taskId, img.index), + taskId, + index: img.index, + image: img, + options: state.taskOptions.get(taskId) || null, + built: false, + refs: {}, + }; + + const badge = el('span', { class: 'badge', text: '等待' }); + const statusText = el('span', { text: '排队中' }); + const params = el('p', { class: 'card-params' }); + const stats = el('div', { class: 'card-stats' }); + const actions = el('div', { class: 'card-actions' }); + const note = el('p', { class: 'card-note' }); + const error = el('p', { class: 'card-error' }); + const compare = el('div', { class: 'compare' }); + + // 文件名 = 回填待处理 + 套用该图参数的入口 + const nameBtn = el('button', { + class: 'name', type: 'button', + title: '点击把这张图的原图放回待处理,并套用记录下来的参数', + text: `${img.index + 1}. ${img.name}`, + onclick: (e) => { e.stopPropagation(); reuseRecord(rec, e.currentTarget); }, + }); + const delBtn = el('button', { + class: 'card-del', type: 'button', title: '从结果列表移除这一条', + text: '×', + onclick: (e) => { e.stopPropagation(); dismissCard(rec.uid); }, + }); + const reuseBtn = el('button', { + class: 'btn sm ghost', type: 'button', text: '重新处理', + title: '把这张图放回待处理,并把它上次用的参数重新赋值到左侧面板', + onclick: (e) => { e.stopPropagation(); reuseRecord(rec, e.currentTarget); }, + }); + actions.append(reuseBtn); + + const root = el('article', { class: 'card pending', tabindex: '0', 'aria-label': `结果 ${img.index + 1}:${img.name}` }, + el('div', { class: 'card-head' }, nameBtn, badge, delBtn), + params, + el('div', { class: 'status' }, el('span', { class: 'spinner' }), statusText), + compare, + el('div', { class: 'card-body' }, stats, actions, note, error), + ); + + // 点击卡片本体选中(按钮/链接/滑块等交互元素不触发) + root.addEventListener('click', (e) => { + if (e.target.closest('button, a, input, .compare')) return; + toggleSelect(rec.uid); + }); + root.addEventListener('keydown', (e) => { + if (e.key === 'Enter' || e.key === ' ') { e.preventDefault(); toggleSelect(rec.uid); } + }); + + rec.root = root; + rec.refs = { root, badge, statusText, params, stats, actions, note, error, compare, nameBtn, delBtn }; + return rec; + } + + /** 选中/取消选中一条结果记录(主题色描边)。 */ + function toggleSelect(uid) { + state.selected = state.selected === uid ? null : uid; + for (const [key, rec] of state.records) { + rec.root.classList.toggle('selected', key === state.selected); + } + } + + /** 从结果列表删除一条记录(文件留在磁盘上)。 */ + function dismissCard(uid) { + const rec = state.records.get(uid); + if (!rec) return; + rec.root.remove(); + state.records.delete(uid); + state.dismissed.add(uid); + if (state.selected === uid) state.selected = null; + updateStageTools(); + } + + /** 把某条结果的原图放回待处理槽 + 套用它记录下来的参数。 */ + async function reuseRecord(rec, btn) { + const img = rec.image; + try { + const res = await fetch(img.urls.original); + if (!res.ok) throw new Error(`取原图失败(HTTP ${res.status})`); + const blob = await res.blob(); + const name = img.name || `image_${rec.index + 1}.png`; + const file = new File([blob], name, { type: blob.type || 'image/png', lastModified: 0 }); + // 优先用这条记录自己的参数;没有就退回该图最近一次的记录 + const params = rec.options || lookupParams(imageKey(file)) || null; + if (offerImage(file, { params })) { + flash(btn, '已就绪 ✓'); + setNote(`已把「${name}」放回待处理,并套用它上次使用的参数${params ? '' : '(无参数记录)'}`); + } else { + flash(btn, '已在待处理'); + } + } catch (err) { + setError(`加入待处理失败:${err.message}`); + } + } + + function updateStageTools() { + const n = state.records.size; + $('result-count').textContent = String(n); + $('btn-clear').disabled = n === 0; + $('btn-zip').disabled = n === 0; + // 批量清单展开时不再显示「还没有结果」(它自己就是结果) + $('empty').classList.toggle('hidden', n > 0 || state.batch.reportVisible); + } + + /** 清空整个结果列表(显式操作;被清掉的记录不会因轮询复活)。 */ + function clearResults(silent) { + for (const uid of state.records.keys()) state.dismissed.add(uid); + state.records.clear(); + state.selected = null; + $('results').textContent = ''; + updateStageTools(); + if (!silent) setError(''); + } + + function scrollResultsToBottom() { + const grid = $('results'); + if (grid && typeof grid.scrollHeight === 'number') grid.scrollTop = grid.scrollHeight; + } + + /* ----------------------------- 结果渲染 ----------------------------- */ + function buildCompare(rec) { + const { compare } = rec.refs; + const img = rec.image; + const base = el('img', { class: 'before', alt: '原图', src: img.urls.original }); + const after = el('img', { class: 'after', alt: '抠图结果', src: viewUrl(img) }); + const divider = el('div', { class: 'divider' }); + const slider = el('input', { type: 'range', min: '0', max: '100', value: '50', 'aria-label': '对比分割线' }); + slider.addEventListener('input', () => compare.style.setProperty('--p', `${slider.value}%`)); + compare.style.setProperty('--p', '50%'); + compare.append( + base, after, divider, slider, + el('span', { class: 'tag left', text: '原图' }), + el('span', { class: 'tag right', text: '结果' }), + ); + rec.refs.afterImg = after; + } + + function viewUrl(img) { + const urls = img.urls || {}; + return urls[state.view] || urls.cutout || urls.original; + } + + function setView(view) { + state.view = view; + for (const btn of $('view-switch').querySelectorAll('button')) { + btn.classList.toggle('active', btn.dataset.view === view); + } + for (const rec of state.records.values()) { + const img = rec.image; + if (!rec.refs.afterImg || !img) continue; + if (view === 'mask' && !img.urls?.mask) continue; + rec.refs.afterImg.src = viewUrl(img); + } + } + + function stat(label, value) { + return el('span', {}, el('b', { text: `${label} ` }), document.createTextNode(String(value))); + } + + function downloadLink(url, filename, text) { + return el('a', { class: 'btn sm ghost', href: url, download: filename, text }); + } + + async function copyImage(rec, button) { + try { + const res = await fetch(rec.image.urls.cutout); + const blob = await res.blob(); + await navigator.clipboard.write([new ClipboardItem({ 'image/png': blob })]); + flash(button, '已复制'); + } catch (err) { + setError(`复制失败(浏览器需支持 Clipboard API):${err.message}`); + } + } + + /* ----------------------------- 打包下载 ----------------------------- */ + function uniqueName(used, name) { + let candidate = name; + let n = 2; + while (used.has(candidate)) { + const dot = name.lastIndexOf('.'); + candidate = dot > 0 ? `${name.slice(0, dot)}-${n}${name.slice(dot)}` : `${name}-${n}`; + n += 1; + } + used.add(candidate); + return candidate; + } + + function stamp() { + const d = new Date(); + const p = (v) => String(v).padStart(2, '0'); + return `${d.getFullYear()}${p(d.getMonth() + 1)}${p(d.getDate())}-${p(d.getHours())}${p(d.getMinutes())}${p(d.getSeconds())}`; + } + + function saveBlob(blob, filename) { + const url = URL.createObjectURL(blob); + const a = document.createElement('a'); + a.href = url; + a.download = filename; + if (document.body) document.body.append(a); + a.click(); + a.remove(); + setTimeout(() => URL.revokeObjectURL(url), 4000); + } + + /** 把结果列表里所有记录打成一个 ZIP(记录可能来自多个任务)。 */ + async function downloadZip() { + const recs = [...state.records.values()]; + if (!recs.length) return; + if (!window.BiRefNetZip) { setError('打包模块未加载(web/zip.js)'); return; } + + const btn = $('btn-zip'); + const origin = btn.dataset.label || btn.textContent; + btn.dataset.label = origin; + btn.disabled = true; + btn.textContent = '打包中…'; + setError(''); + let okMsg = null; + try { + const kinds = state.view === 'mask' ? ['cutout', 'mask'] : ['cutout']; + const used = new Set(); + const entries = []; + for (const rec of recs) { + const urls = rec.image.urls || {}; + for (const kind of kinds) { + if (!urls[kind]) continue; + const res = await fetch(urls[kind]); + if (!res.ok) continue; + const data = new Uint8Array(await res.arrayBuffer()); + const label = rec.image.label || `image_${rec.index + 1}`; + const name = `${String(rec.index + 1).padStart(3, '0')}_${label}_${kind}.png`; + entries.push({ name: uniqueName(used, name), data }); + } + } + if (!entries.length) throw new Error('没有可打包的结果文件'); + saveBlob(new Blob([window.BiRefNetZip.buildZip(entries)], { type: 'application/zip' }), + `birefnet_results_${stamp()}.zip`); + okMsg = `已打包 ${entries.length} 个文件`; + } catch (err) { + setError(`打包失败:${err.message}`); + } finally { + btn.textContent = origin; + updateStageTools(); + } + if (okMsg) flash(btn, okMsg); + } + + loadParams(); + init(); +})(); diff --git a/web/index.html b/web/index.html new file mode 100644 index 0000000..a76730f --- /dev/null +++ b/web/index.html @@ -0,0 +1,267 @@ + + + + + +BiRefNet WebUI · 智能抠图 + + + + +
    +
    + +

    BiRefNet WebUI

    +
    + +
    +
    设备
    检测中…
    +
    模型
    —
    +
    模型代码
    —
    +
    + +
    + + +
    +
    + +
    + + + + +
    +
    +

    结果 0

    +
    +
    + + + +
    + + +
    +
    + +

    结果会一直保留在列表里(传新图、改参数都不会清空)· 点文件名或「重新处理」会套用它上次使用的参数 · 点卡片选中 · 点 × 移除单条

    + + + +
    + +

    还没有结果

    +

    左侧拖入图片 → 选择模型 → 点击「开始抠图」。处理完成后可拖拽分割线对比原图与抠图效果,记录会累积在右侧供随时重新处理。

    +
    + +
    +
    +
    + + + + + diff --git a/web/style.css b/web/style.css new file mode 100644 index 0000000..e9a7a05 --- /dev/null +++ b/web/style.css @@ -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; } +} diff --git a/web/test_app.js b/web/test_app.js new file mode 100644 index 0000000..26cd4c9 --- /dev/null +++ b/web/test_app.js @@ -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); +} diff --git a/web/zip.js b/web/zip.js new file mode 100644 index 0000000..90cb96b --- /dev/null +++ b/web/zip.js @@ -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);