mirror of
https://github.com/lllyasviel/Fooocus.git
synced 2026-08-16 13:13:16 +02:00
i
This commit is contained in:
@@ -0,0 +1 @@
|
||||
from .dataset import StableDataModuleFromConfig
|
||||
@@ -0,0 +1,67 @@
|
||||
import pytorch_lightning as pl
|
||||
import torchvision
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
from torchvision import transforms
|
||||
|
||||
|
||||
class CIFAR10DataDictWrapper(Dataset):
|
||||
def __init__(self, dset):
|
||||
super().__init__()
|
||||
self.dset = dset
|
||||
|
||||
def __getitem__(self, i):
|
||||
x, y = self.dset[i]
|
||||
return {"jpg": x, "cls": y}
|
||||
|
||||
def __len__(self):
|
||||
return len(self.dset)
|
||||
|
||||
|
||||
class CIFAR10Loader(pl.LightningDataModule):
|
||||
def __init__(self, batch_size, num_workers=0, shuffle=True):
|
||||
super().__init__()
|
||||
|
||||
transform = transforms.Compose(
|
||||
[transforms.ToTensor(), transforms.Lambda(lambda x: x * 2.0 - 1.0)]
|
||||
)
|
||||
|
||||
self.batch_size = batch_size
|
||||
self.num_workers = num_workers
|
||||
self.shuffle = shuffle
|
||||
self.train_dataset = CIFAR10DataDictWrapper(
|
||||
torchvision.datasets.CIFAR10(
|
||||
root=".data/", train=True, download=True, transform=transform
|
||||
)
|
||||
)
|
||||
self.test_dataset = CIFAR10DataDictWrapper(
|
||||
torchvision.datasets.CIFAR10(
|
||||
root=".data/", train=False, download=True, transform=transform
|
||||
)
|
||||
)
|
||||
|
||||
def prepare_data(self):
|
||||
pass
|
||||
|
||||
def train_dataloader(self):
|
||||
return DataLoader(
|
||||
self.train_dataset,
|
||||
batch_size=self.batch_size,
|
||||
shuffle=self.shuffle,
|
||||
num_workers=self.num_workers,
|
||||
)
|
||||
|
||||
def test_dataloader(self):
|
||||
return DataLoader(
|
||||
self.test_dataset,
|
||||
batch_size=self.batch_size,
|
||||
shuffle=self.shuffle,
|
||||
num_workers=self.num_workers,
|
||||
)
|
||||
|
||||
def val_dataloader(self):
|
||||
return DataLoader(
|
||||
self.test_dataset,
|
||||
batch_size=self.batch_size,
|
||||
shuffle=self.shuffle,
|
||||
num_workers=self.num_workers,
|
||||
)
|
||||
@@ -0,0 +1,80 @@
|
||||
from typing import Optional
|
||||
|
||||
import torchdata.datapipes.iter
|
||||
import webdataset as wds
|
||||
from omegaconf import DictConfig
|
||||
from pytorch_lightning import LightningDataModule
|
||||
|
||||
try:
|
||||
from sdata import create_dataset, create_dummy_dataset, create_loader
|
||||
except ImportError as e:
|
||||
print("#" * 100)
|
||||
print("Datasets not yet available")
|
||||
print("to enable, we need to add stable-datasets as a submodule")
|
||||
print("please use ``git submodule update --init --recursive``")
|
||||
print("and do ``pip install -e stable-datasets/`` from the root of this repo")
|
||||
print("#" * 100)
|
||||
exit(1)
|
||||
|
||||
|
||||
class StableDataModuleFromConfig(LightningDataModule):
|
||||
def __init__(
|
||||
self,
|
||||
train: DictConfig,
|
||||
validation: Optional[DictConfig] = None,
|
||||
test: Optional[DictConfig] = None,
|
||||
skip_val_loader: bool = False,
|
||||
dummy: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.train_config = train
|
||||
assert (
|
||||
"datapipeline" in self.train_config and "loader" in self.train_config
|
||||
), "train config requires the fields `datapipeline` and `loader`"
|
||||
|
||||
self.val_config = validation
|
||||
if not skip_val_loader:
|
||||
if self.val_config is not None:
|
||||
assert (
|
||||
"datapipeline" in self.val_config and "loader" in self.val_config
|
||||
), "validation config requires the fields `datapipeline` and `loader`"
|
||||
else:
|
||||
print(
|
||||
"Warning: No Validation datapipeline defined, using that one from training"
|
||||
)
|
||||
self.val_config = train
|
||||
|
||||
self.test_config = test
|
||||
if self.test_config is not None:
|
||||
assert (
|
||||
"datapipeline" in self.test_config and "loader" in self.test_config
|
||||
), "test config requires the fields `datapipeline` and `loader`"
|
||||
|
||||
self.dummy = dummy
|
||||
if self.dummy:
|
||||
print("#" * 100)
|
||||
print("USING DUMMY DATASET: HOPE YOU'RE DEBUGGING ;)")
|
||||
print("#" * 100)
|
||||
|
||||
def setup(self, stage: str) -> None:
|
||||
print("Preparing datasets")
|
||||
if self.dummy:
|
||||
data_fn = create_dummy_dataset
|
||||
else:
|
||||
data_fn = create_dataset
|
||||
|
||||
self.train_datapipeline = data_fn(**self.train_config.datapipeline)
|
||||
if self.val_config:
|
||||
self.val_datapipeline = data_fn(**self.val_config.datapipeline)
|
||||
if self.test_config:
|
||||
self.test_datapipeline = data_fn(**self.test_config.datapipeline)
|
||||
|
||||
def train_dataloader(self) -> torchdata.datapipes.iter.IterDataPipe:
|
||||
loader = create_loader(self.train_datapipeline, **self.train_config.loader)
|
||||
return loader
|
||||
|
||||
def val_dataloader(self) -> wds.DataPipeline:
|
||||
return create_loader(self.val_datapipeline, **self.val_config.loader)
|
||||
|
||||
def test_dataloader(self) -> wds.DataPipeline:
|
||||
return create_loader(self.test_datapipeline, **self.test_config.loader)
|
||||
@@ -0,0 +1,85 @@
|
||||
import pytorch_lightning as pl
|
||||
import torchvision
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
from torchvision import transforms
|
||||
|
||||
|
||||
class MNISTDataDictWrapper(Dataset):
|
||||
def __init__(self, dset):
|
||||
super().__init__()
|
||||
self.dset = dset
|
||||
|
||||
def __getitem__(self, i):
|
||||
x, y = self.dset[i]
|
||||
return {"jpg": x, "cls": y}
|
||||
|
||||
def __len__(self):
|
||||
return len(self.dset)
|
||||
|
||||
|
||||
class MNISTLoader(pl.LightningDataModule):
|
||||
def __init__(self, batch_size, num_workers=0, prefetch_factor=2, shuffle=True):
|
||||
super().__init__()
|
||||
|
||||
transform = transforms.Compose(
|
||||
[transforms.ToTensor(), transforms.Lambda(lambda x: x * 2.0 - 1.0)]
|
||||
)
|
||||
|
||||
self.batch_size = batch_size
|
||||
self.num_workers = num_workers
|
||||
self.prefetch_factor = prefetch_factor if num_workers > 0 else 0
|
||||
self.shuffle = shuffle
|
||||
self.train_dataset = MNISTDataDictWrapper(
|
||||
torchvision.datasets.MNIST(
|
||||
root=".data/", train=True, download=True, transform=transform
|
||||
)
|
||||
)
|
||||
self.test_dataset = MNISTDataDictWrapper(
|
||||
torchvision.datasets.MNIST(
|
||||
root=".data/", train=False, download=True, transform=transform
|
||||
)
|
||||
)
|
||||
|
||||
def prepare_data(self):
|
||||
pass
|
||||
|
||||
def train_dataloader(self):
|
||||
return DataLoader(
|
||||
self.train_dataset,
|
||||
batch_size=self.batch_size,
|
||||
shuffle=self.shuffle,
|
||||
num_workers=self.num_workers,
|
||||
prefetch_factor=self.prefetch_factor,
|
||||
)
|
||||
|
||||
def test_dataloader(self):
|
||||
return DataLoader(
|
||||
self.test_dataset,
|
||||
batch_size=self.batch_size,
|
||||
shuffle=self.shuffle,
|
||||
num_workers=self.num_workers,
|
||||
prefetch_factor=self.prefetch_factor,
|
||||
)
|
||||
|
||||
def val_dataloader(self):
|
||||
return DataLoader(
|
||||
self.test_dataset,
|
||||
batch_size=self.batch_size,
|
||||
shuffle=self.shuffle,
|
||||
num_workers=self.num_workers,
|
||||
prefetch_factor=self.prefetch_factor,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
dset = MNISTDataDictWrapper(
|
||||
torchvision.datasets.MNIST(
|
||||
root=".data/",
|
||||
train=False,
|
||||
download=True,
|
||||
transform=transforms.Compose(
|
||||
[transforms.ToTensor(), transforms.Lambda(lambda x: x * 2.0 - 1.0)]
|
||||
),
|
||||
)
|
||||
)
|
||||
ex = dset[0]
|
||||
Reference in New Issue
Block a user