[feat] unify dataloader api

This commit is contained in:
zhengzangw 2024-05-20 08:40:45 +00:00
parent 5a1cf0c718
commit 5b9e753039
9 changed files with 151 additions and 236 deletions

View file

@ -100,7 +100,7 @@ outputs = "outputs"
wandb = False
epochs = 1000
log_every = 10
ckpt_every = 500
ckpt_every = 1
# optimization settings
load = None

View file

@ -1,3 +1,2 @@
from .dataloader import prepare_dataloader, prepare_variable_dataloader
from .datasets import IMG_FPS, VariableVideoTextDataset, VideoTextDataset
from .utils import get_transforms_image, get_transforms_video, is_img, is_vid, save_sample

View file

@ -7,44 +7,59 @@ from torch.distributed import ProcessGroup
from torch.distributed.distributed_c10d import _get_default_group
from torch.utils.data import DataLoader
from .datasets import VariableVideoTextDataset, VideoTextDataset
from .sampler import StatefulDistributedSampler, VariableVideoBatchSampler
# Deterministic dataloader
def get_seed_worker(seed):
def seed_worker(worker_id):
worker_seed = seed
np.random.seed(worker_seed)
torch.manual_seed(worker_seed)
random.seed(worker_seed)
return seed_worker
def prepare_dataloader(
dataset,
batch_size,
batch_size=None,
shuffle=False,
seed=1024,
drop_last=False,
pin_memory=False,
num_workers=0,
process_group: Optional[ProcessGroup] = None,
distributed=True,
bucket_config=None,
num_bucket_build_workers=1,
**kwargs,
):
r"""
Prepare a dataloader for distributed training. The dataloader will be wrapped by
`torch.utils.data.DataLoader` and `StatefulDistributedSampler`.
Args:
dataset (`torch.utils.data.Dataset`): The dataset to be loaded.
shuffle (bool, optional): Whether to shuffle the dataset. Defaults to False.
seed (int, optional): Random worker seed for sampling, defaults to 1024.
add_sampler: Whether to add ``DistributedDataParallelSampler`` to the dataset. Defaults to True.
drop_last (bool, optional): Set to True to drop the last incomplete batch, if the dataset size
is not divisible by the batch size. If False and the size of dataset is not divisible by
the batch size, then the last batch will be smaller, defaults to False.
pin_memory (bool, optional): Whether to pin memory address in CPU memory. Defaults to False.
num_workers (int, optional): Number of worker threads for this dataloader. Defaults to 0.
kwargs (dict): optional parameters for ``torch.utils.data.DataLoader``, more details could be found in
`DataLoader <https://pytorch.org/docs/stable/_modules/torch/utils/data/dataloader.html#DataLoader>`_.
Returns:
:class:`torch.utils.data.DataLoader`: A DataLoader used for training or testing.
"""
_kwargs = kwargs.copy()
if distributed:
if isinstance(dataset, VariableVideoTextDataset):
batch_sampler = VariableVideoBatchSampler(
dataset,
bucket_config,
num_replicas=process_group.size(),
rank=process_group.rank(),
shuffle=shuffle,
seed=seed,
drop_last=drop_last,
verbose=True,
num_bucket_build_workers=num_bucket_build_workers,
)
return (
DataLoader(
dataset,
batch_sampler=batch_sampler,
worker_init_fn=get_seed_worker(seed),
pin_memory=pin_memory,
num_workers=num_workers,
**_kwargs,
),
batch_sampler,
)
elif isinstance(dataset, VideoTextDataset):
process_group = process_group or _get_default_group()
sampler = StatefulDistributedSampler(
dataset,
@ -52,67 +67,18 @@ def prepare_dataloader(
rank=process_group.rank(),
shuffle=shuffle,
)
return (
DataLoader(
dataset,
batch_size=batch_size,
sampler=sampler,
worker_init_fn=get_seed_worker(seed),
drop_last=drop_last,
pin_memory=pin_memory,
num_workers=num_workers,
**_kwargs,
),
sampler,
)
else:
sampler = None
# Deterministic dataloader
def seed_worker(worker_id):
worker_seed = seed
np.random.seed(worker_seed)
torch.manual_seed(worker_seed)
random.seed(worker_seed)
return DataLoader(
dataset,
batch_size=batch_size,
sampler=sampler,
worker_init_fn=seed_worker,
drop_last=drop_last,
pin_memory=pin_memory,
num_workers=num_workers,
**_kwargs,
)
def prepare_variable_dataloader(
dataset,
batch_size,
bucket_config,
shuffle=False,
seed=1024,
drop_last=False,
pin_memory=False,
num_workers=0,
process_group=None,
num_bucket_build_workers=1,
**kwargs,
):
_kwargs = kwargs.copy()
process_group = process_group or _get_default_group()
batch_sampler = VariableVideoBatchSampler(
dataset,
bucket_config,
num_replicas=process_group.size(),
rank=process_group.rank(),
shuffle=shuffle,
seed=seed,
drop_last=drop_last,
verbose=True,
num_bucket_build_workers=num_bucket_build_workers,
)
# Deterministic dataloader
def seed_worker(worker_id):
worker_seed = seed
np.random.seed(worker_seed)
torch.manual_seed(worker_seed)
random.seed(worker_seed)
return torch.utils.data.DataLoader(
dataset,
batch_sampler=batch_sampler,
worker_init_fn=seed_worker,
pin_memory=pin_memory,
num_workers=num_workers,
**_kwargs,
)
raise ValueError(f"Unsupported dataset type: {type(dataset)}")

View file

@ -1,4 +1,3 @@
import warnings
from collections import OrderedDict, defaultdict
from pprint import pformat
from typing import Iterator, List, Optional
@ -48,8 +47,14 @@ class StatefulDistributedSampler(DistributedSampler):
def __len__(self) -> int:
return self.num_samples - self.start_index
def set_start_index(self, start_index: int) -> None:
self.start_index = start_index
def reset(self) -> None:
self.start_index = 0
def state_dict(self) -> dict:
return {"start_index": self.start_index}
def load_state_dict(self, state_dict: dict) -> None:
self.__dict__.update(state_dict)
class VariableVideoBatchSampler(DistributedSampler):
@ -77,42 +82,6 @@ class VariableVideoBatchSampler(DistributedSampler):
self._get_num_batch_cached_bucket_sample_dict = None
self.num_bucket_build_workers = num_bucket_build_workers
def group_by_bucket(self) -> dict:
bucket_sample_dict = OrderedDict()
from pandarallel import pandarallel
pandarallel.initialize(nb_workers=self.num_bucket_build_workers, progress_bar=False)
get_logger().info("Building buckets...")
bucket_ids = self.dataset.data.parallel_apply(
apply,
axis=1,
method=self.bucket.get_bucket_id,
frame_interval=self.dataset.frame_interval,
seed=self.seed + self.epoch,
num_bucket=self.bucket.num_bucket,
)
# group by bucket
# each data sample is put into a bucket with a similar image/video size
for i in range(len(self.dataset)):
bucket_id = bucket_ids[i]
if bucket_id is None:
continue
if bucket_id not in bucket_sample_dict:
bucket_sample_dict[bucket_id] = []
bucket_sample_dict[bucket_id].append(i)
return bucket_sample_dict
def get_num_batch(self) -> int:
bucket_sample_dict = self.group_by_bucket()
self._get_num_batch_cached_bucket_sample_dict = bucket_sample_dict
# calculate the number of batches
if self.verbose:
self._print_bucket_info(bucket_sample_dict)
return self.approximate_num_batch
def __iter__(self) -> Iterator[List[int]]:
if self._get_num_batch_cached_bucket_sample_dict is not None:
bucket_sample_dict = self._get_num_batch_cached_bucket_sample_dict
@ -215,20 +184,46 @@ class VariableVideoBatchSampler(DistributedSampler):
cur_micro_batch = [f"{idx}-{real_t}-{real_h}-{real_w}" for idx in cur_micro_batch]
yield cur_micro_batch
self._reset()
self.reset()
def _reset(self):
self.last_micro_batch_access_index = 0
def __len__(self) -> int:
return self.get_num_batch() // dist.get_world_size()
def state_dict(self, num_steps: int) -> dict:
# the last_micro_batch_access_index in the __iter__ is often
# not accurate during multi-workers and data prefetching
# thus, we need the user to pass the actual steps which have been executed
# to calculate the correct last_micro_batch_access_index
return {"seed": self.seed, "epoch": self.epoch, "last_micro_batch_access_index": num_steps * self.num_replicas}
def group_by_bucket(self) -> dict:
bucket_sample_dict = OrderedDict()
def load_state_dict(self, state_dict: dict) -> None:
self.__dict__.update(state_dict)
from pandarallel import pandarallel
pandarallel.initialize(nb_workers=self.num_bucket_build_workers, progress_bar=False)
get_logger().info("Building buckets...")
bucket_ids = self.dataset.data.parallel_apply(
apply,
axis=1,
method=self.bucket.get_bucket_id,
frame_interval=self.dataset.frame_interval,
seed=self.seed + self.epoch,
num_bucket=self.bucket.num_bucket,
)
# group by bucket
# each data sample is put into a bucket with a similar image/video size
for i in range(len(self.dataset)):
bucket_id = bucket_ids[i]
if bucket_id is None:
continue
if bucket_id not in bucket_sample_dict:
bucket_sample_dict[bucket_id] = []
bucket_sample_dict[bucket_id].append(i)
return bucket_sample_dict
def get_num_batch(self) -> int:
bucket_sample_dict = self.group_by_bucket()
self._get_num_batch_cached_bucket_sample_dict = bucket_sample_dict
# calculate the number of batches
if self.verbose:
self._print_bucket_info(bucket_sample_dict)
return self.approximate_num_batch
def _print_bucket_info(self, bucket_sample_dict: dict) -> None:
# collect statistics
@ -276,19 +271,15 @@ class VariableVideoBatchSampler(DistributedSampler):
)
self.approximate_num_batch = total_batch
def set_epoch(self, epoch: int) -> None:
super().set_epoch(epoch)
def reset(self):
self.last_micro_batch_access_index = 0
def __len__(self) -> int:
warnings.warn(
"The length of VariableVideoBatchSampler is dynamic and may not be accurate. Return the max value."
)
min_batch_size = None
for v in self.bucket.bucket_bs.values():
for bs in v.values():
if bs is not None and (min_batch_size is None or bs < min_batch_size):
min_batch_size = bs
if self.drop_last:
return len(self.dataset) // min_batch_size
else:
return (len(self.dataset) + min_batch_size - 1) // min_batch_size
def state_dict(self, num_steps: int) -> dict:
# the last_micro_batch_access_index in the __iter__ is often
# not accurate during multi-workers and data prefetching
# thus, we need the user to pass the actual steps which have been executed
# to calculate the correct last_micro_batch_access_index
return {"seed": self.seed, "epoch": self.epoch, "last_micro_batch_access_index": num_steps * self.num_replicas}
def load_state_dict(self, state_dict: dict) -> None:
self.__dict__.update(state_dict)

View file

@ -239,12 +239,11 @@ def save(
if lr_scheduler is not None:
booster.save_lr_scheduler(lr_scheduler, os.path.join(save_dir, "lr_scheduler"))
if dist.get_rank() == 0:
sampler_start_idx = step * batch_size if batch_size is not None else None
running_states = {
"epoch": epoch,
"step": step,
"global_step": global_step,
"sample_start_index": sampler_start_idx,
"batch_size": batch_size,
}
save_json(running_states, os.path.join(save_dir, "running_states.json"))
@ -255,6 +254,7 @@ def save(
# only for VariableVideoBatchSampler
torch.save(sampler.state_dict(step), os.path.join(save_dir, "sampler"))
dist.barrier()
return save_dir
def load(
@ -288,5 +288,4 @@ def load(
return (
running_states["epoch"],
running_states["step"],
running_states["sample_start_index"],
)

View file

@ -46,7 +46,7 @@ def main():
logger.info("Building reconstruction dataset...")
dataset = build_module(cfg.dataset, DATASETS)
batch_size = cfg.get("batch_size", 1)
dataloader = prepare_dataloader(
dataloader, _ = prepare_dataloader(
dataset,
batch_size=batch_size,
num_workers=cfg.get("num_workers", 4),

View file

@ -14,7 +14,7 @@ from tqdm import tqdm
from opensora.acceleration.checkpoint import set_grad_checkpoint
from opensora.acceleration.parallel_states import get_data_parallel_group
from opensora.datasets import prepare_dataloader, prepare_variable_dataloader
from opensora.datasets import prepare_dataloader
from opensora.datasets.utils import collate_fn_ignore_none
from opensora.registry import DATASETS, MODELS, SCHEDULERS, build_module
from opensora.utils.ckpt_utils import load, model_gathering, model_sharding, record_model_param_shape, save
@ -101,20 +101,12 @@ def main():
process_group=get_data_parallel_group(),
collate_fn=collate_fn_ignore_none,
)
if cfg.dataset.type == DEFAULT_DATASET_NAME:
dataloader = prepare_dataloader(**dataloader_args)
total_batch_size = cfg.batch_size * dist.get_world_size() // cfg.get("sp_size", 1)
logger.info("Total batch size: %s", total_batch_size)
num_steps_per_epoch = len(dataloader)
sampler_to_io = None
else:
dataloader = prepare_variable_dataloader(
bucket_config=cfg.get("bucket_config", None),
num_bucket_build_workers=cfg.get("num_bucket_build_workers", 1),
**dataloader_args,
)
num_steps_per_epoch = dataloader.batch_sampler.get_num_batch() // dist.get_world_size()
sampler_to_io = None if cfg.get("start_from_scratch ", False) else dataloader.batch_sampler
dataloader, sampler = prepare_dataloader(
bucket_config=cfg.get("bucket_config", None),
num_bucket_build_workers=cfg.get("num_bucket_build_workers", 1),
**dataloader_args,
)
num_steps_per_epoch = len(dataloader)
# ======================================================
# 3. build model
@ -190,7 +182,7 @@ def main():
# == global variables ==
cfg_epochs = cfg.get("epochs", 1000)
start_epoch = start_step = log_step = sampler_start_idx = acc_step = 0
start_epoch = start_step = log_step = acc_step = 0
running_loss = 0.0
logger.info("Training for %s epochs with %s steps per epoch", cfg_epochs, num_steps_per_epoch)
@ -204,13 +196,11 @@ def main():
ema=ema,
optimizer=optimizer,
lr_scheduler=lr_scheduler,
sampler=sampler_to_io,
sampler=None if cfg.get("start_from_scratch ", False) else sampler,
)
if not cfg.get("start_from_scratch ", False):
start_epoch, start_step, sampler_start_idx = ret
start_epoch, start_step = ret
logger.info("Loaded checkpoint %s at epoch %s step %s", cfg.load, start_epoch, start_step)
if cfg.dataset.type == DEFAULT_DATASET_NAME:
dataloader.sampler.set_start_index(sampler_start_idx)
model_sharding(ema)
@ -220,8 +210,7 @@ def main():
dist.barrier()
for epoch in range(start_epoch, cfg_epochs):
# == set dataloader to new epoch ==
if cfg.dataset.type == DEFAULT_DATASET_NAME:
dataloader.sampler.set_epoch(epoch)
sampler.set_epoch(epoch)
dataloader_iter = iter(dataloader)
logger.info("Beginning epoch %s...", epoch)
@ -304,14 +293,14 @@ def main():
ckpt_every = cfg.get("ckpt_every", 0)
if ckpt_every > 0 and (global_step + 1) % ckpt_every == 0:
model_gathering(ema, ema_shape_dict)
save(
save_dir = save(
booster,
exp_dir,
model=model,
ema=ema,
optimizer=optimizer,
lr_scheduler=lr_scheduler,
sampler=sampler_to_io,
sampler=sampler,
epoch=epoch,
step=step + 1,
global_step=global_step + 1,
@ -320,19 +309,13 @@ def main():
if dist.get_rank() == 0:
model_sharding(ema)
logger.info(
"Saved checkpoint at epoch %s step %s global_step %s to %s",
"Saved checkpoint at epoch %s, step %s, global_step %s to %s",
epoch,
step + 1,
global_step + 1,
exp_dir,
save_dir,
)
# NOTE: the continue epochs are not resumed, so we need to reset the sampler start index and start step
if cfg.dataset.type == DEFAULT_DATASET_NAME:
dataloader.sampler.set_start_index(0)
else:
dataloader.batch_sampler.set_epoch(epoch + 1)
logger.info("Epoch done, recomputing batch sampler")
sampler.reset()
start_step = 0

View file

@ -14,7 +14,7 @@ from tqdm import tqdm
from opensora.acceleration.checkpoint import set_grad_checkpoint
from opensora.acceleration.parallel_states import get_data_parallel_group
from opensora.datasets import prepare_dataloader, prepare_variable_dataloader
from opensora.datasets.dataloader import prepare_dataloader
from opensora.datasets.utils import collate_fn_ignore_none
from opensora.registry import DATASETS, MODELS, SCHEDULERS, build_module
from opensora.utils.ckpt_utils import load, model_gathering, model_sharding, record_model_param_shape, save
@ -30,8 +30,6 @@ from opensora.utils.misc import (
)
from opensora.utils.train_utils import MaskGenerator, create_colossalai_plugin, update_ema
DEFAULT_DATASET_NAME = "VideoTextDataset"
def main():
# ======================================================
@ -100,20 +98,12 @@ def main():
process_group=get_data_parallel_group(),
collate_fn=collate_fn_ignore_none,
)
if cfg.dataset.type == DEFAULT_DATASET_NAME:
dataloader = prepare_dataloader(**dataloader_args)
total_batch_size = cfg.batch_size * dist.get_world_size() // cfg.get("sp_size", 1)
logger.info("Total batch size: %s", total_batch_size)
num_steps_per_epoch = len(dataloader)
sampler_to_io = None
else:
dataloader = prepare_variable_dataloader(
bucket_config=cfg.get("bucket_config", None),
num_bucket_build_workers=cfg.get("num_bucket_build_workers", 1),
**dataloader_args,
)
num_steps_per_epoch = dataloader.batch_sampler.get_num_batch() // dist.get_world_size()
sampler_to_io = None if cfg.get("start_from_scratch ", False) else dataloader.batch_sampler
dataloader, sampler = prepare_dataloader(
bucket_config=cfg.get("bucket_config", None),
num_bucket_build_workers=cfg.get("num_bucket_build_workers", 1),
**dataloader_args,
)
num_steps_per_epoch = len(dataloader)
# ======================================================
# 3. build model
@ -189,7 +179,7 @@ def main():
# == global variables ==
cfg_epochs = cfg.get("epochs", 1000)
start_epoch = start_step = log_step = sampler_start_idx = acc_step = 0
start_epoch = start_step = log_step = acc_step = 0
running_loss = 0.0
logger.info("Training for %s epochs with %s steps per epoch", cfg_epochs, num_steps_per_epoch)
@ -203,13 +193,11 @@ def main():
ema=ema,
optimizer=optimizer,
lr_scheduler=lr_scheduler,
sampler=sampler_to_io,
sampler=None if cfg.get("start_from_scratch ", False) else sampler,
)
if not cfg.get("start_from_scratch ", False):
start_epoch, start_step, sampler_start_idx = ret
start_epoch, start_step = ret
logger.info("Loaded checkpoint %s at epoch %s step %s", cfg.load, start_epoch, start_step)
if cfg.dataset.type == DEFAULT_DATASET_NAME:
dataloader.sampler.set_start_index(sampler_start_idx)
model_sharding(ema)
@ -219,8 +207,7 @@ def main():
dist.barrier()
for epoch in range(start_epoch, cfg_epochs):
# == set dataloader to new epoch ==
if cfg.dataset.type == DEFAULT_DATASET_NAME:
dataloader.sampler.set_epoch(epoch)
sampler.set_epoch(epoch)
dataloader_iter = iter(dataloader)
logger.info("Beginning epoch %s...", epoch)
@ -299,14 +286,14 @@ def main():
ckpt_every = cfg.get("ckpt_every", 0)
if ckpt_every > 0 and (global_step + 1) % ckpt_every == 0:
model_gathering(ema, ema_shape_dict)
save(
save_dir = save(
booster,
exp_dir,
model=model,
ema=ema,
optimizer=optimizer,
lr_scheduler=lr_scheduler,
sampler=sampler_to_io,
sampler=sampler,
epoch=epoch,
step=step + 1,
global_step=global_step + 1,
@ -315,19 +302,13 @@ def main():
if dist.get_rank() == 0:
model_sharding(ema)
logger.info(
"Saved checkpoint at epoch %s step %s global_step %s to %s",
"Saved checkpoint at epoch %s, step %s, global_step %s to %s",
epoch,
step + 1,
global_step + 1,
exp_dir,
save_dir,
)
# NOTE: the continue epochs are not resumed, so we need to reset the sampler start index and start step
if cfg.dataset.type == DEFAULT_DATASET_NAME:
dataloader.sampler.set_start_index(0)
else:
dataloader.batch_sampler.set_epoch(epoch + 1)
logger.info("Epoch done, recomputing batch sampler")
sampler.reset()
start_step = 0

View file

@ -30,8 +30,6 @@ from opensora.utils.misc import (
)
from opensora.utils.train_utils import create_colossalai_plugin
DEFAULT_DATASET_NAME = "VideoTextDataset"
def main():
# ======================================================
@ -85,7 +83,7 @@ def main():
# ======================================================
logger.info("Building dataset...")
# == build dataset ==
assert cfg.dataset.type == DEFAULT_DATASET_NAME, "Only support VideoTextDataset for vae training"
assert cfg.dataset.type == "VideoTextDataset", "Only support VideoTextDataset for vae training"
dataset = build_module(cfg.dataset, DATASETS)
logger.info("Dataset contains %s samples.", len(dataset))
@ -100,7 +98,7 @@ def main():
pin_memory=True,
process_group=get_data_parallel_group(),
)
dataloader = prepare_dataloader(**dataloader_args)
dataloader, sampler = prepare_dataloader(**dataloader_args)
total_batch_size = cfg.batch_size * dist.get_world_size() // cfg.get("sp_size", 1)
logger.info("Total batch size: %s", total_batch_size)
num_steps_per_epoch = len(dataloader)
@ -208,12 +206,13 @@ def main():
# == resume ==
if cfg.get("load", None) is not None:
logger.info("Loading checkpoint")
start_epoch, start_step, sampler_start_idx = load(
start_epoch, start_step = load(
booster,
cfg.load,
model=model,
optimizer=optimizer,
lr_scheduler=lr_scheduler,
sampler=sampler,
)
if use_discriminator and os.path.exists(os.path.join(cfg.load, "discriminator")):
booster.load_model(discriminator, os.path.join(cfg.load, "discriminator"))
@ -221,15 +220,13 @@ def main():
dist.barrier()
logger.info("Loaded checkpoint %s at epoch %s step %s", cfg.load, start_epoch, start_step)
dataloader.sampler.set_start_index(sampler_start_idx)
# =======================================================
# 5. training loop
# =======================================================
dist.barrier()
for epoch in range(start_epoch, cfg_epochs):
# == set dataloader to new epoch ==
dataloader.sampler.set_epoch(epoch)
sampler.set_epoch(epoch)
dataiter = iter(dataloader)
logger.info("Beginning epoch %s...", epoch)
@ -385,8 +382,7 @@ def main():
exp_dir,
)
# NOTE: the continue epochs are not resumed, so we need to reset the sampler start index and start step
dataloader.sampler.set_start_index(0)
sampler.reset()
start_step = 0