Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
filename_captions.py68 linesDownload Raw Back to data
1from __future__ import annotations
2
3import os
4from typing import List, Optional
5
6from blockdiffusion.data.captions import CaptionItem
7
8
9_IMG_EXTS = (".png", ".jpg", ".jpeg", ".webp", ".bmp")
10
11
12def _caption_from_filename(filename: str) -> str:
13    """
14    从文件名提取“弱标注提示词”。
15
16    适配命名格式:
17    - HASH-tag_tag_tag.png
18    - HASH-side.png / HASH-top.png / HASH-bottom.png
19
20    提取规则:
21    - 去掉扩展名
22    - 以第一个 '-' 作为分隔,取其后部分(若不存在 '-' 则使用全名)
23    - '_' / '-' 统一当作空格
24    """
25    name = os.path.splitext(os.path.basename(filename))[0]
26    if "-" in name:
27        # 只切一次,避免 tag 本身含 '-' 时被过度切分
28        _, tail = name.split("-", 1)
29    else:
30        tail = name
31    tail = tail.replace("_", " ").replace("-", " ").strip().lower()
32    # 多空格归一
33    tail = " ".join([t for t in tail.split(" ") if t])
34    return tail
35
36
37def build_caption_items_from_filenames(
38    root_dir: str,
39    verbose: bool = False,
40    log_every: int = 10_000,
41) -> List[CaptionItem]:
42    """
43    扫描目录下所有图片(含子目录),用文件名自动生成 captions。
44    返回 CaptionItem 列表,可直接用于 TextureCaptionDataset。
45    """
46    root_dir = os.path.abspath(root_dir)
47    items: List[CaptionItem] = []
48    seen = 0
49    for dirpath, _, filenames in os.walk(root_dir):
50        for fn in filenames:
51            if not fn.lower().endswith(_IMG_EXTS):
52                continue
53            path = os.path.join(dirpath, fn)
54            text = _caption_from_filename(fn)
55            # 若提取失败就给空字符串;TextEncoder 会把空句子当作“无条件”
56            items.append(CaptionItem(image_path=path, text=text))
57            seen += 1
58            if verbose and log_every > 0 and (seen % int(log_every) == 0):
59                print(f"[scan] 已发现图片 {seen} 张...")
60    items.sort(key=lambda x: x.image_path)
61    if len(items) == 0:
62        raise ValueError(f"目录下未找到图片:{root_dir}")
63    if verbose:
64        print(f"[scan] 完成:共 {len(items)} 张图片")
65    return items
66
67
68