From f0874b4b8e47ad5141cd6aba89b33b85951a7510 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E8=B4=B9=E6=94=BF=E8=81=AA?= <3300949577@qq.com> Date: Fri, 19 Jul 2024 16:10:36 +0800 Subject: [PATCH] add flash attention, deepspeed --- config/zero2.json | 35 +++++ config/zero3.json | 38 +++++ config/zero3_offload.json | 46 ++++++ download.py | 4 +- models.py | 114 +++++++++++--- sample.py | 8 +- train_deepspeed.py | 303 ++++++++++++++++++++++++++++++++++++++ 7 files changed, 525 insertions(+), 23 deletions(-) create mode 100644 config/zero2.json create mode 100644 config/zero3.json create mode 100644 config/zero3_offload.json create mode 100644 train_deepspeed.py diff --git a/config/zero2.json b/config/zero2.json new file mode 100644 index 0000000..224cf20 --- /dev/null +++ b/config/zero2.json @@ -0,0 +1,35 @@ +{ + "fp16": { + "enabled": true, + "loss_scale": 0, + "loss_scale_window": 1000, + "initial_scale_power": 16, + "hysteresis": 2, + "min_loss_scale": 1 + }, + "optimizer": { + "type": "AdamW", + "params": { + "lr": 0.0001, + "betas": [ + 0.9, + 0.999 + ], + "eps": 1e-8, + "weight_decay": 0 + } + }, + "bf16": { + "enabled": false + }, + "train_micro_batch_size_per_gpu": 32, + "train_batch_size": 256, + "gradient_accumulation_steps": 1, + "zero_optimization": { + "stage": 2, + "overlap_comm": true, + "contiguous_gradients": true, + "sub_group_size": 1e9 + }, + "steps_per_print": 10000, +} diff --git a/config/zero3.json b/config/zero3.json new file mode 100644 index 0000000..b0cd248 --- /dev/null +++ b/config/zero3.json @@ -0,0 +1,38 @@ +{ + "fp16": { + "enabled": true, + "loss_scale": 0, + "loss_scale_window": 1000, + "initial_scale_power": 16, + "hysteresis": 2, + "min_loss_scale": 1 + }, + "optimizer": { + "type": "AdamW", + "params": { + "lr": 0.0001, + "betas": [ + 0.9, + 0.999 + ], + "eps": 1e-8, + "weight_decay": 0 + } + }, + "bf16": { + "enabled": false + }, + "train_micro_batch_size_per_gpu": 2, + "train_batch_size": 16, + "gradient_accumulation_steps": 1, + "zero_optimization": { + "stage": 3, + "overlap_comm": true, + "contiguous_gradients": true, + "sub_group_size": 1e9, + "stage3_max_live_parameters": 1e9, + "stage3_max_reuse_distance": 1e9, + "stage3_gather_16bit_weights_on_model_save": true + }, + "steps_per_print": 10000, +} diff --git a/config/zero3_offload.json b/config/zero3_offload.json new file mode 100644 index 0000000..deba3b8 --- /dev/null +++ b/config/zero3_offload.json @@ -0,0 +1,46 @@ +{ + "fp16": { + "enabled": true, + "loss_scale": 0, + "loss_scale_window": 1000, + "initial_scale_power": 16, + "hysteresis": 2, + "min_loss_scale": 1 + }, + "optimizer": { + "type": "AdamW", + "params": { + "lr": 0.0001, + "betas": [ + 0.9, + 0.999 + ], + "eps": 1e-8, + "weight_decay": 0 + } + }, + "bf16": { + "enabled": false + }, + "train_micro_batch_size_per_gpu": 2, + "train_batch_size": 16, + "gradient_accumulation_steps": 1, + "zero_optimization": { + "stage": 3, + "offload_optimizer": { + "device": "cpu", + "pin_memory": true + }, + "offload_param": { + "device": "cpu", + "pin_memory": true + }, + "overlap_comm": true, + "contiguous_gradients": true, + "sub_group_size": 1e9, + "stage3_max_live_parameters": 1e9, + "stage3_max_reuse_distance": 1e9, + "gather_16bit_weights_on_model_save": true + }, + "steps_per_print": 10000, +} diff --git a/download.py b/download.py index de22d45..b3e8b99 100644 --- a/download.py +++ b/download.py @@ -25,7 +25,9 @@ def find_model(model_name): assert os.path.isfile(model_name), f'Could not find DiT checkpoint at {model_name}' checkpoint = torch.load(model_name, map_location=lambda storage, loc: storage) if "ema" in checkpoint: # supports checkpoints from train.py - checkpoint = checkpoint["ema"] + checkpoint = checkpoint["ema"] + elif "model" in checkpoint: + checkpoint = checkpoint["model"] return checkpoint diff --git a/models.py b/models.py index d0d657d..a196de5 100644 --- a/models.py +++ b/models.py @@ -17,6 +17,18 @@ from timm.models.vision_transformer import PatchEmbed, Attention, Mlp import torch.nn.functional as F +try: + import flash_attn + if hasattr(flash_attn, '__version__') and int(flash_attn.__version__[0]) == 2: + from flash_attn.flash_attn_interface import flash_attn_kvpacked_func + from flash_attn.modules.mha import FlashSelfAttention + else: + from flash_attn.flash_attn_interface import flash_attn_unpadded_kvpacked_func + from flash_attn.modules.mha import FlashSelfAttention +except Exception as e: + print(f'flash_attn import failed: {e}') + + # selected_ids_list = [] @@ -66,8 +78,8 @@ class TimestepEmbedder(nn.Module): return embedding def forward(self, t): - t_freq = self.timestep_embedding(t, self.frequency_embedding_size) - t_emb = self.mlp(t_freq) + t_freq = self.timestep_embedding(t, self.frequency_embedding_size) + t_emb = self.mlp(t_freq)#.half()) return t_emb @@ -106,8 +118,6 @@ class LabelEmbedder(nn.Module): ################################################################################# - - class MoEGate(nn.Module): def __init__(self, embed_dim, num_experts=16, num_experts_per_tok=2, aux_loss_alpha=0.01): super().__init__() @@ -193,7 +203,7 @@ class AddAuxiliaryLoss(torch.autograd.Function): class MoeMLP(nn.Module): - def __init__(self, hidden_size, intermediate_size): + def __init__(self, hidden_size, intermediate_size, pretraining_tp=2): super().__init__() self.hidden_size = hidden_size @@ -202,13 +212,14 @@ class MoeMLP(nn.Module): self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) self.act_fn = nn.SiLU() - self.pretraining_tp = 2 + self.pretraining_tp = pretraining_tp def forward(self, x): if self.pretraining_tp > 1: slice = self.intermediate_size // self.pretraining_tp gate_proj_slices = self.gate_proj.weight.split(slice, dim=0) - up_proj_slices = self.up_proj.weight.split(slice, dim=0) + up_proj_slices = self.up_proj.weight.split(slice, dim=0) + # print(self.up_proj.weight.size(), self.down_proj.weight.size()) down_proj_slices = self.down_proj.weight.split(slice, dim=1) gate_proj = torch.cat( @@ -231,16 +242,16 @@ class SparseMoeBlock(nn.Module): """ A mixed expert module containing shared experts. """ - def __init__(self, embed_dim, mlp_ratio=4, num_experts=16, num_experts_per_tok=2): + def __init__(self, embed_dim, mlp_ratio=4, num_experts=16, num_experts_per_tok=2, pretraining_tp=2): super().__init__() self.num_experts_per_tok = num_experts_per_tok - self.experts = nn.ModuleList([MoeMLP(hidden_size = embed_dim, intermediate_size = mlp_ratio * embed_dim) for i in range(num_experts)]) + self.experts = nn.ModuleList([MoeMLP(hidden_size = embed_dim, intermediate_size = mlp_ratio * embed_dim, pretraining_tp=pretraining_tp) for i in range(num_experts)]) self.gate = MoEGate(embed_dim=embed_dim, num_experts=num_experts, num_experts_per_tok=num_experts_per_tok) self.n_shared_experts = 2 if self.n_shared_experts is not None: intermediate_size = embed_dim * self.n_shared_experts - self.shared_experts = MoeMLP(hidden_size = embed_dim, intermediate_size = intermediate_size) + self.shared_experts = MoeMLP(hidden_size = embed_dim, intermediate_size = intermediate_size, pretraining_tp=pretraining_tp) def forward(self, hidden_states): identity = hidden_states @@ -254,9 +265,9 @@ class SparseMoeBlock(nn.Module): flat_topk_idx = topk_idx.view(-1) if self.training: hidden_states = hidden_states.repeat_interleave(self.num_experts_per_tok, dim=0) - y = torch.empty_like(hidden_states) - for i, expert in enumerate(self.experts): - y[flat_topk_idx == i] = expert(hidden_states[flat_topk_idx == i]) + y = torch.empty_like(hidden_states, dtype=hidden_states.dtype) + for i, expert in enumerate(self.experts): + y[flat_topk_idx == i] = expert(hidden_states[flat_topk_idx == i]).float() y = (y.view(*topk_weight.shape, -1) * topk_weight.unsqueeze(-1)).sum(dim=1) y = y.view(*orig_shape) y = AddAuxiliaryLoss.apply(y, aux_loss) @@ -305,6 +316,64 @@ class RMSNorm(nn.Module): +################################################################################# +# Flash attention Layer. # +################################################################################# + +class FlashSelfMHAModified(nn.Module): + """ + self-attention with flashattention + """ + def __init__(self, + dim, + num_heads, + qkv_bias=True, + qk_norm=False, + attn_drop=0.0, + proj_drop=0.0, + device=None, + dtype=None, + norm_layer=nn.LayerNorm, + ): + factory_kwargs = {'device': device, 'dtype': dtype} + super().__init__() + self.dim = dim + self.num_heads = num_heads + assert self.dim % num_heads == 0, "self.kdim must be divisible by num_heads" + self.head_dim = self.dim // num_heads + assert self.head_dim % 8 == 0 and self.head_dim <= 128, "Only support head_dim <= 128 and divisible by 8" + + self.Wqkv = nn.Linear(dim, 3 * dim, bias=qkv_bias, **factory_kwargs) + # TODO: eps should be 1 / 65530 if using fp16 + self.q_norm = norm_layer(self.head_dim, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity() + self.k_norm = norm_layer(self.head_dim, elementwise_affine=True, eps=1e-6) if qk_norm else nn.Identity() + self.inner_attn = FlashSelfAttention(attention_dropout=attn_drop) + self.out_proj = nn.Linear(dim, dim, bias=qkv_bias, **factory_kwargs) + self.proj_drop = nn.Dropout(proj_drop) + + def forward(self, x,): + """ + Parameters + ---------- + x: torch.Tensor + (batch, seqlen, hidden_dim) (where hidden_dim = num heads * head dim) + """ + b, s, d = x.shape + + qkv = self.Wqkv(x) + qkv = qkv.view(b, s, 3, self.num_heads, self.head_dim) # [b, s, 3, h, d] + q, k, v = qkv.unbind(dim=2) # [b, s, h, d] + q = self.q_norm(q).half() # [b, s, h, d] + k = self.k_norm(k).half() + + qkv = torch.stack([q, k, v], dim=2) # [b, s, 3, h, d] + context = self.inner_attn(qkv) + out = self.out_proj(context.view(b, s, d)) + out = self.proj_drop(out) + + return out + + ################################################################################# # Core DiT Model # ################################################################################# @@ -315,16 +384,20 @@ class DiTBlock(nn.Module): """ def __init__( self, hidden_size, num_heads, mlp_ratio=4, - num_experts=8, num_experts_per_tok=2, **block_kwargs + num_experts=8, num_experts_per_tok=2, pretraining_tp=2, + use_flash_attn=False, **block_kwargs ): super().__init__() self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) - self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=True, **block_kwargs) + if use_flash_attn: + self.attn = FlashSelfMHAModified(hidden_size, num_heads=num_heads, qkv_bias=True, qk_norm=True) + else: + self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=True, **block_kwargs) self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) mlp_hidden_dim = int(hidden_size * mlp_ratio) approx_gelu = lambda: nn.GELU(approximate="tanh") # self.mlp = Mlp(in_features=hidden_size, hidden_features=mlp_hidden_dim, act_layer=approx_gelu, drop=0) - self.moe = SparseMoeBlock(hidden_size, mlp_ratio, num_experts, num_experts_per_tok) + self.moe = SparseMoeBlock(hidden_size, mlp_ratio, num_experts, num_experts_per_tok, pretraining_tp) self.adaLN_modulation = nn.Sequential( nn.SiLU(), @@ -374,7 +447,9 @@ class DiT(nn.Module): class_dropout_prob=0.1, num_classes=1000, num_experts=8, num_experts_per_tok=2, + pretraining_tp=2, learn_sigma=True, + use_flash_attn=False, ): super().__init__() self.learn_sigma = learn_sigma @@ -391,7 +466,7 @@ class DiT(nn.Module): self.pos_embed = nn.Parameter(torch.zeros(1, num_patches, hidden_size), requires_grad=False) self.blocks = nn.ModuleList([ - DiTBlock(hidden_size, num_heads, mlp_ratio, num_experts, num_experts_per_tok, ) for _ in range(depth) + DiTBlock(hidden_size, num_heads, mlp_ratio, num_experts, num_experts_per_tok, pretraining_tp, use_flash_attn) for _ in range(depth) ]) self.final_layer = FinalLayer(hidden_size, patch_size, self.out_channels) self.initialize_weights() @@ -454,6 +529,8 @@ class DiT(nn.Module): t: (N,) tensor of diffusion timesteps y: (N,) tensor of class labels """ + #x = x.half() + # t = t.half() x = self.x_embedder(x) + self.pos_embed # (N, T, D), where T = H * W / patch_size ** 2 t = self.t_embedder(t) # (N, D) y = self.y_embedder(y, self.training) # (N, D) @@ -480,7 +557,8 @@ class DiT(nn.Module): cond_eps, uncond_eps = torch.split(eps, len(eps) // 2, dim=0) half_eps = uncond_eps + cfg_scale * (cond_eps - uncond_eps) eps = torch.cat([half_eps, half_eps], dim=0) - return torch.cat([eps, rest], dim=1) + return torch.cat([eps, rest], dim=1) + ################################################################################# diff --git a/sample.py b/sample.py index 48659e9..2e58b74 100644 --- a/sample.py +++ b/sample.py @@ -40,10 +40,10 @@ def main(args): # Auto-download a pre-trained model or load a custom DiT checkpoint from train.py: if args.model == "DiT-S/2": - ckpt_path = "results/002-DiT-S-2/checkpoints/ckpt.pt" + # ckpt_path = "results/002-DiT-S-2/checkpoints/ckpt.pt" + ckpt_path = "results/deepspeed-DiT-S-2/checkpoints/0000001.pt" else: - # ckpt_path = "results/003-DiT-B-2/checkpoints/0750000.pt" - ckpt_path = "ckpt_clean.pt" + ckpt_path = "results/003-DiT-B-2/checkpoints/0750000.pt" state_dict = find_model(ckpt_path) model.load_state_dict(state_dict) @@ -81,7 +81,7 @@ def main(args): if __name__ == "__main__": parser = argparse.ArgumentParser() - parser.add_argument("--model", type=str, choices=list(DiT_models.keys()), default="DiT-B/2") + parser.add_argument("--model", type=str, choices=list(DiT_models.keys()), default="DiT-S/2") parser.add_argument("--vae", type=str, choices=["ema", "mse"], default="mse") parser.add_argument("--image-size", type=int, choices=[256, 512], default=256) parser.add_argument("--num-classes", type=int, default=1000) diff --git a/train_deepspeed.py b/train_deepspeed.py new file mode 100644 index 0000000..2eff45d --- /dev/null +++ b/train_deepspeed.py @@ -0,0 +1,303 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +""" +A training script for DiT using deepspeed. +""" +import torch +# the first flag below was False when we tested this script but True makes A100 training a lot faster: +torch.backends.cuda.matmul.allow_tf32 = True +torch.backends.cudnn.allow_tf32 = True +import torch.distributed as dist +from torch.nn.parallel import DistributedDataParallel as DDP +from torch.utils.data import DataLoader +from torch.utils.data.distributed import DistributedSampler +from torchvision.datasets import ImageFolder +from torchvision import transforms +import numpy as np +from collections import OrderedDict +from PIL import Image +from copy import deepcopy +from glob import glob +from time import time +import argparse +import logging +import os + +from models import DiT_models +from diffusion import create_diffusion +from diffusers.models import AutoencoderKL +from download import find_model + + +import deepspeed + +################################################################################# +# Training Helper Functions # +################################################################################# + +@torch.no_grad() +def update_ema(ema_model, model, decay=0.9999): + """ + Step the EMA model towards the current model. + """ + ema_params = OrderedDict(ema_model.named_parameters()) + model_params = OrderedDict(model.named_parameters()) + + for name, param in model_params.items(): + # TODO: Consider applying only to params that require_grad to avoid small numerical changes of pos_embed + ema_params[name].mul_(decay).add_(param.data, alpha=1 - decay) + + +def requires_grad(model, flag=True): + """ + Set requires_grad flag for all parameters in a model. + """ + for p in model.parameters(): + p.requires_grad = flag + + +def cleanup(): + """ + End DDP training. + """ + dist.destroy_process_group() + + +def create_logger(logging_dir): + """ + Create a logger that writes to a log file and stdout. + """ + if dist.get_rank() == 0: # real logger + logging.basicConfig( + level=logging.INFO, + format='[\033[34m%(asctime)s\033[0m] %(message)s', + datefmt='%Y-%m-%d %H:%M:%S', + handlers=[logging.StreamHandler(), logging.FileHandler(f"{logging_dir}/log.txt")] + ) + logger = logging.getLogger(__name__) + else: # dummy logger (does nothing) + logger = logging.getLogger(__name__) + logger.addHandler(logging.NullHandler()) + return logger + + +def center_crop_arr(pil_image, image_size): + """ + Center cropping implementation from ADM. + https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126 + """ + while min(*pil_image.size) >= 2 * image_size: + pil_image = pil_image.resize( + tuple(x // 2 for x in pil_image.size), resample=Image.BOX + ) + + scale = image_size / min(*pil_image.size) + pil_image = pil_image.resize( + tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC + ) + + arr = np.array(pil_image) + crop_y = (arr.shape[0] - image_size) // 2 + crop_x = (arr.shape[1] - image_size) // 2 + return Image.fromarray(arr[crop_y: crop_y + image_size, crop_x: crop_x + image_size]) + + +################################################################################# +# Training Loop # +################################################################################# + +def main(args): + """ + Trains a new DiT model. + """ + assert torch.cuda.is_available(), "Training currently requires at least one GPU." + + deepspeed.init_distributed() + + # Setup DDP: + #dist.init_process_group("nccl") + #assert args.global_batch_size % dist.get_world_size() == 0, f"Batch size must be divisible by world size." + rank = args.local_rank + device = rank % torch.cuda.device_count() + seed = args.global_seed * dist.get_world_size() + rank + torch.manual_seed(seed) + torch.cuda.set_device(device) + print(f"Starting rank={rank}, seed={seed}, world_size={dist.get_world_size()}.") + + + # Setup an experiment folder: + if rank == 0: + os.makedirs(args.results_dir, exist_ok=True) # Make results folder (holds all experiment subfolders) + experiment_index = len(glob(f"{args.results_dir}/*")) + model_string_name = args.model.replace("/", "-") # e.g., DiT-XL/2 --> DiT-XL-2 (for naming folders) + experiment_dir = f"{args.results_dir}/deepspeed-{model_string_name}" # Create an experiment folder + checkpoint_dir = f"{experiment_dir}/checkpoints" # Stores saved model checkpoints + os.makedirs(checkpoint_dir, exist_ok=True) + logger = create_logger(experiment_dir) + logger.info(f"Experiment directory created at {experiment_dir}") + else: + logger = create_logger(None) + + # Create model: + assert args.image_size % 8 == 0, "Image size must be divisible by 8 (for the VAE encoder)." + latent_size = args.image_size // 8 + model = DiT_models[args.model]( + input_size=latent_size, + num_classes=args.num_classes, + num_experts=args.num_experts, + num_experts_per_tok=args.num_experts_per_tok, + pretraining_tp=1, + use_flash_attn=True + ) + + if args.resume is not None: + print('load from: ', args.resume) + state_dict = find_model(args.resume) + model.load_state_dict(state_dict) + + + # Note that parameter initialization is done within the DiT constructor + # ema = deepcopy(model).to(device) # Create an EMA of the model for use after training + # requires_grad(ema, False) + # model = DDP(model.to(device), device_ids=[rank]) + # model = DDP(model.to(device), device_ids=[device]) + diffusion = create_diffusion(timestep_respacing="") # default: 1000 steps, linear noise schedule + vae = AutoencoderKL.from_pretrained(args.vae_path).to(device) + logger.info(f"DiT Parameters: {sum(p.numel() for p in model.parameters()):,}") + + # Setup optimizer (we used default Adam betas=(0.9, 0.999) and a constant learning rate of 1e-4 in our paper): + # opt = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0) + + # Setup data: + transform = transforms.Compose([ + transforms.Lambda(lambda pil_image: center_crop_arr(pil_image, args.image_size)), + transforms.RandomHorizontalFlip(), + transforms.ToTensor(), + transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True) + ]) + dataset = ImageFolder(args.data_path, transform=transform) + sampler = DistributedSampler( + dataset, + num_replicas=dist.get_world_size(), + rank=rank, + shuffle=True, + seed=args.global_seed + ) + loader = DataLoader( + dataset, + batch_size=args.train_batch_size, #int(args.global_batch_size // dist.get_world_size()), + shuffle=False, + sampler=sampler, + num_workers=args.num_workers, + pin_memory=True, + drop_last=True + ) + logger.info(f"Dataset contains {len(dataset):,} images ({args.data_path})") + + # Prepare models for training: + # update_ema(ema, model.module, decay=0) # Ensure EMA is initialized with synced weights + # model.train() # important! This enables embedding dropout for classifier-free guidance + # ema.eval() # EMA model should always be in eval mode + + model_engine, opt, _, __ = deepspeed.initialize( + args=args, model=model, model_parameters=model.parameters()) + + # Variables for monitoring/logging purposes: + train_steps = 0 + log_steps = 0 + running_loss = 0 + start_time = time() + + logger.info(f"Training for {args.epochs} epochs...") + for epoch in range(args.epochs): + sampler.set_epoch(epoch) + logger.info(f"Beginning epoch {epoch}...") + data_iter_step = 0 + for x, y in loader: + model_engine.train() + x = x.to(device) + y = y.to(device) + with torch.no_grad(): + # Map input images to latent space + normalize latents: + x = vae.encode(x).latent_dist.sample().mul_(0.18215) + t = torch.randint(0, diffusion.num_timesteps, (x.shape[0],), device=device) + model_kwargs = dict(y=y) + with torch.autocast(device_type='cuda'): + loss_dict = diffusion.training_losses(model, x, t, model_kwargs) + loss = loss_dict["loss"].mean() + #if (data_iter_step + 1) % args.accum_iter == 0: + # opt.zero_grad() + #loss.backward() + model_engine.backward(loss) + model_engine.step() + # opt.step() + # update_ema(ema, model.module) + + data_iter_step += 1 + # Log loss values: + running_loss += loss.item() + log_steps += 1 + train_steps += 1 + if train_steps % args.log_every == 0: + # Measure training speed: + torch.cuda.synchronize() + end_time = time() + steps_per_sec = log_steps / (end_time - start_time) + # Reduce loss history over all processes: + avg_loss = torch.tensor(running_loss / log_steps, device=device) + dist.all_reduce(avg_loss, op=dist.ReduceOp.SUM) + avg_loss = avg_loss.item() / dist.get_world_size() + logger.info(f"(step={train_steps:07d}) Train Loss: {avg_loss:.4f}, Train Steps/Sec: {steps_per_sec:.2f}") + # Reset monitoring variables: + running_loss = 0 + log_steps = 0 + start_time = time() + + # Save DiT checkpoint: + if train_steps % args.ckpt_every == 0 and train_steps > 0: + if rank == 0: + checkpoint = { + "model": model.state_dict(), + "opt": opt.state_dict(), + "args": args + } + checkpoint_path = f"{checkpoint_dir}/{train_steps:07d}.pt" + torch.save(checkpoint, checkpoint_path) + logger.info(f"Saved checkpoint to {checkpoint_path}") + dist.barrier() + + model.eval() # important! This disables randomized embedding dropout + # do any sampling/FID calculation/etc. with ema (or model) in eval mode ... + + logger.info("Done!") + cleanup() + + +if __name__ == "__main__": + # Default args here will train DiT-XL/2 with the hyperparameters we used in our paper (except training iters). + parser = argparse.ArgumentParser() + parser.add_argument("--data-path", type=str, required=True) + parser.add_argument("--results-dir", type=str, default="results") + parser.add_argument("--resume", type=str, default=None) + parser.add_argument("--model", type=str, choices=list(DiT_models.keys()), default="DiT-S/2") + parser.add_argument("--vae-path", type=str, default='/maindata/data/shared/multimodal/zhengcong.fei/ckpts/sd-vae-ft-mse') + parser.add_argument("--image-size", type=int, choices=[256, 512], default=256) + parser.add_argument("--num-classes", type=int, default=1000) + parser.add_argument("--epochs", type=int, default=1400) + parser.add_argument("--train_batch_size", type=int, default=2) + parser.add_argument("--global-seed", type=int, default=1234) + parser.add_argument("--num-workers", type=int, default=4) + parser.add_argument("--log-every", type=int, default=100) + parser.add_argument('--accum_iter', default=8, type=int,) + parser.add_argument('--num_experts', default=8, type=int,) + parser.add_argument('--num_experts_per_tok', default=2, type=int,) + parser.add_argument("--ckpt-every", type=int, default=50_000) + parser.add_argument('--local-rank', type=int, default=-1, help='local rank passed from distributed launcher') + parser = deepspeed.add_config_arguments(parser) + args = parser.parse_args() + print(args) + main(args)