mirror of
https://github.com/hpcaitech/Open-Sora.git
synced 2026-05-21 11:59:01 +02:00
[feat] unify dataloader api
This commit is contained in:
parent
5a1cf0c718
commit
5b9e753039
|
|
@ -100,7 +100,7 @@ outputs = "outputs"
|
|||
wandb = False
|
||||
epochs = 1000
|
||||
log_every = 10
|
||||
ckpt_every = 500
|
||||
ckpt_every = 1
|
||||
|
||||
# optimization settings
|
||||
load = None
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)}")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue