Team Ai
Apppublic

wkea/blockdiffusion-api

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
texture_caption_dataset.py120 linesDownload Raw Back to data
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