Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
captions.py72 linesDownload Raw Back to data
1from __future__ import annotations
2
3import csv
4import json
5import os
6from dataclasses import dataclass
7from typing import Iterable, List, Sequence, Tuple
8
9
10@dataclass(frozen=True)
11class CaptionItem:
12    """
13    一条图文配对样本:
14    - image_path:图片路径(绝对路径)
15    - text:文本描述/提示词
16    """
17
18    image_path: str
19    text: str
20
21
22def load_captions(captions_path: str, data_root: str) -> List[CaptionItem]:
23    """
24    读取 caption 文件,支持:
25    - .jsonl:每行 {"file": "...", "text": "..."},file 可相对 data_root 或绝对路径
26    - .csv:两列 file,text(含表头),file 同上
27    """
28    captions_path = os.path.abspath(captions_path)
29    data_root = os.path.abspath(data_root)
30
31    if not os.path.exists(captions_path):
32        raise FileNotFoundError(f"未找到 captions 文件:{captions_path}")
33
34    items: List[CaptionItem] = []
35    ext = os.path.splitext(captions_path)[1].lower()
36
37    def resolve_path(p: str) -> str:
38        p = (p or "").strip()
39        if not p:
40            return ""
41        if os.path.isabs(p):
42            return p
43        return os.path.join(data_root, p)
44
45    if ext == ".jsonl":
46        with open(captions_path, "r", encoding="utf-8") as f:
47            for line in f:
48                line = line.strip()
49                if not line:
50                    continue
51                obj = json.loads(line)
52                img = resolve_path(str(obj.get("file", "")))
53                txt = str(obj.get("text", ""))
54                if img and os.path.exists(img):
55                    items.append(CaptionItem(image_path=img, text=txt))
56    elif ext == ".csv":
57        with open(captions_path, "r", encoding="utf-8") as f:
58            reader = csv.DictReader(f)
59            for row in reader:
60                img = resolve_path(str(row.get("file", "")))
61                txt = str(row.get("text", ""))
62                if img and os.path.exists(img):
63                    items.append(CaptionItem(image_path=img, text=txt))
64    else:
65        raise ValueError("captions 文件仅支持 .jsonl 或 .csv")
66
67    if len(items) == 0:
68        raise ValueError("captions 文件未读取到有效样本(请检查 file 路径与图片是否存在)")
69    return items
70
71
72