wkea/blockdiffusion-api
0
1from __future__ import annotations2 3from dataclasses import dataclass4from typing import List, Optional, Tuple5 6from PIL import Image7import torch8from torch.utils.data import Dataset9from torchvision import transforms as T10 11from blockdiffusion.data.captions import CaptionItem12from blockdiffusion.text.simple_tokenizer import SimpleVocab13 14 15class TextureCaptionDataset(Dataset):16 """17 图文配对纹理数据集:18 - 输入:CaptionItem 列表(包含 image_path 与 text)19 - 输出:(image[-1,1], tokens, mask)20 """21 22 def __init__(23 self,24 items: List[CaptionItem],25 vocab: SimpleVocab,26 image_size: int = 16,27 channels: int = 3,28 max_text_len: int = 32,29 random_flip: bool = True,30 pretokenize: bool = False,31 append_alpha_tag: bool = False,32 alpha_tag_transparent: str = "transparent",33 alpha_tag_opaque: str = "opaque",34 verbose: bool = False,35 log_every: int = 50_000,36 ):37 if len(items) == 0:38 raise ValueError("items 不能为空")39 self.items = items40 self.vocab = vocab41 self.image_size = int(image_size)42 self.channels = int(channels)43 if self.channels not in (3, 4):44 raise ValueError("channels 仅支持 3(RGB) 或 4(RGBA)")45 self.max_text_len = int(max_text_len)46 self.pretokenize = bool(pretokenize)47 self.append_alpha_tag = bool(append_alpha_tag)48 self.alpha_tag_transparent = str(alpha_tag_transparent or "").strip()49 self.alpha_tag_opaque = str(alpha_tag_opaque or "").strip()50 self.verbose = bool(verbose)51 self.log_every = int(log_every)52 53 # 透明/不透明标签需要在 __getitem__ 里基于图像 alpha 判定;54 # pretokenize 不读取图像,因此两者不兼容。55 if self.append_alpha_tag and self.pretokenize:56 raise ValueError("append_alpha_tag 与 pretokenize 不兼容:请关闭 pretokenize")57 58 tfms = [59 T.Resize((self.image_size, self.image_size), interpolation=T.InterpolationMode.NEAREST),60 ]61 if random_flip:62 tfms.append(T.RandomHorizontalFlip(p=0.5))63 tfms.extend(64 [65 T.ToTensor(),66 # 用 Normalize 替代 lambda,保证 Windows 多进程 DataLoader 可 pickle67 T.Normalize(mean=[0.5] * self.channels, std=[0.5] * self.channels),68 ]69 )70 self.transform = T.Compose(tfms)71 72 # 可选:预编码 tokens/mask(大数据集会耗时/占内存,默认关闭)73 self._tokens = None74 self._mask = None75 if self.pretokenize:76 tokens_list: List[torch.Tensor] = []77 mask_list: List[torch.Tensor] = []78 for i, it in enumerate(items, start=1):79 ids = vocab.encode(it.text, max_len=self.max_text_len)80 tok = torch.tensor(ids, dtype=torch.long)81 m = (tok != vocab.pad_id)82 tokens_list.append(tok)83 mask_list.append(m)84 if self.verbose and self.log_every > 0 and (i % self.log_every == 0):85 print(f"[pretok] 已预编码 {i} 条文本...")86 self._tokens = tokens_list87 self._mask = mask_list88 if self.verbose:89 print(f"[pretok] 完成:共 {len(items)} 条文本")90 91 def __len__(self) -> int:92 return len(self.items)93 94 def __getitem__(self, idx: int):95 it = self.items[idx]96 with Image.open(it.image_path) as im:97 im = im.convert("RGBA" if self.channels == 4 else "RGB")98 x = self.transform(im)99 if self._tokens is not None and self._mask is not None:100 return x, self._tokens[idx], self._mask[idx]101 # 未预编码:按需编码(更快启动、更省内存)102 text = it.text103 # 可选:自动追加透明/不透明标签,让模型学会“透明度受提示词影响”104 if self.append_alpha_tag and self.channels == 4:105 # im 已被 convert("RGBA")106 alpha = im.getchannel("A")107 mn, mx = alpha.getextrema()108 # 只要出现 <255 就认为“含透明”109 tag = self.alpha_tag_transparent if int(mn) < 255 else self.alpha_tag_opaque110 tag = (tag or "").strip()111 if tag:112 text = f"{text} {tag}".strip() if text else tag113 114 ids = self.vocab.encode(text, max_len=self.max_text_len)115 tok = torch.tensor(ids, dtype=torch.long)116 m = (tok != self.vocab.pad_id)117 return x, tok, m118 119 120 