optimize zero3 training and inference

This commit is contained in:
费政聪
2024-07-25 11:40:52 +08:00
committed by GitHub
parent ebe2b1a28a
commit 5fbdafe2ca
2 changed files with 30 additions and 24 deletions
+8 -7
View File
@@ -5,7 +5,7 @@
# LICENSE file in the root directory of this source tree.
"""
Sample new images from a pre-trained DiT.
Sample new images from a pre-trained DiT-MoE.
"""
import torch
torch.backends.cuda.matmul.allow_tf32 = True
@@ -18,6 +18,7 @@ from models import DiT_models
import argparse
def main(args):
# Setup PyTorch:
torch.manual_seed(args.seed)
@@ -55,13 +56,13 @@ def main(args):
if args.ckpt is None:
print('only for testing middle ckpts')
if args.model == "DiT-S/2":
ckpt_path = "results/002-DiT-S-2/checkpoints/ckpt.pt"
ckpt_path = "dit_moe_s_8E2A.pt"
elif args.model == "DiT-B/2":
ckpt_path = "results/003-DiT-B-2/checkpoints/ckpt.pt"
ckpt_path = "dit_moe_b_8E2A.pt"
elif args.model == "DiT-XL/2":
ckpt_path = "results/deepspeed-DiT-XL-2/checkpoints/ckpt.pt"
else:
pass
else:
ckpt_path = "results/deepspeed-DiT-G-2/checkpoints/ckpt.pt"
else:
ckpt_path = args.ckpt
@@ -107,7 +108,7 @@ def main(args):
elif args.model == "DiT-XL/2":
save_image(samples, "sample_xl.png", nrow=4, normalize=True, value_range=(-1, 1))
else:
pass
save_image(samples, "sample_g.png", nrow=4, normalize=True, value_range=(-1, 1))
if __name__ == "__main__":
@@ -120,7 +121,7 @@ if __name__ == "__main__":
parser.add_argument('--num_experts', default=8, type=int,)
parser.add_argument('--num_experts_per_tok', default=2, type=int,)
parser.add_argument("--num-sampling-steps", type=int, default=250)
parser.add_argument("--seed", type=int, default=22)
parser.add_argument("--seed", type=int, default=2024)
parser.add_argument("--ckpt", type=str, default=None, )
args = parser.parse_args()
main(args)
+22 -17
View File
@@ -31,8 +31,8 @@ from models import DiT_models
from diffusion import create_diffusion
from diffusers.models import AutoencoderKL
from download import find_model
import deepspeed
from deepspeed.utils import safe_get_full_fp32_param
#################################################################################
# Training Helper Functions #
@@ -113,13 +113,13 @@ def main(args):
print(f"Starting rank={rank}, seed={seed}, world_size={dist.get_world_size()}.")
# Setup an experiment folder:
# Setup an experiment folder
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
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}")
@@ -223,16 +223,21 @@ def main(args):
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}")
if train_steps % args.ckpt_every == 0 and train_steps > 0:
# zero3 should parameter gathering
if 'zero3' in args.deepspeed_config:
checkpoint_path = f"{checkpoint_dir}/{train_steps:07d}"
model_engine.save_checkpoint(checkpoint_path)
else:
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
@@ -257,7 +262,7 @@ if __name__ == "__main__":
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("--ckpt-every", type=int, default=10_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()