wkea/blockdiffusion-api
0
1from __future__ import annotations2 3import os4from typing import List, Sequence, Tuple5 6from PIL import Image7import torch8from torch.utils.data import Dataset9from torchvision import transforms as T10 11 12_IMG_EXTS = (".png", ".jpg", ".jpeg", ".webp", ".bmp")13 14 15def _list_images(root: str) -> List[str]:16 paths: List[str] = []17 for dirpath, _, filenames in os.walk(root):18 for fn in filenames:19 if fn.lower().endswith(_IMG_EXTS):20 paths.append(os.path.join(dirpath, fn))21 paths.sort()22 return paths23 24 25class TextureFolderDataset(Dataset):26 """27 纹理图片数据集:28 - 读取目录下所有图片(含子目录)29 - 统一转换为 RGB30 - 缩放到 16×1631 - 输出范围为 [-1, 1] 的张量(C,H,W)32 """33 34 def __init__(35 self,36 root_dir: str,37 image_size: int = 16,38 channels: int = 3,39 random_flip: bool = True,40 ):41 self.root_dir = root_dir42 self.image_size = int(image_size)43 self.channels = int(channels)44 if self.channels not in (3, 4):45 raise ValueError("channels 仅支持 3(RGB) 或 4(RGBA)")46 47 self.paths = _list_images(root_dir)48 if len(self.paths) == 0:49 raise ValueError(f"数据集目录为空或未找到图片:{root_dir}")50 51 tfms = [52 T.Resize((self.image_size, self.image_size), interpolation=T.InterpolationMode.NEAREST),53 ]54 if random_flip:55 tfms.append(T.RandomHorizontalFlip(p=0.5))56 tfms.extend(57 [58 T.ToTensor(), # [0,1]59 # 用 Normalize 替代 lambda,保证 Windows 多进程 DataLoader 可 pickle60 T.Normalize(mean=[0.5] * self.channels, std=[0.5] * self.channels), # [-1,1]61 ]62 )63 self.transform = T.Compose(tfms)64 65 def __len__(self) -> int:66 return len(self.paths)67 68 def __getitem__(self, idx: int) -> torch.Tensor:69 path = self.paths[idx]70 # 使用 nearest resize,尽量保留像素风格边缘71 with Image.open(path) as im:72 im = im.convert("RGBA" if self.channels == 4 else "RGB")73 x = self.transform(im)74 return x75 76 77 