wkea/blockdiffusion-api
0
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 