Team Ai
Apppublic

wkea/blockdiffusion-api

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