Files
BiRefNet_WebUI/birefnet_web/fsutil.py
T
2026-10-08 09:46:47 +08:00

68 lines
2.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- 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)))