add flash attention, deepspeed

This commit is contained in:
费政聪
2024-07-19 16:10:36 +08:00
committed by GitHub
parent 137cc19666
commit f0874b4b8e
7 changed files with 525 additions and 23 deletions
+35
View File
@@ -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,
}
+38
View File
@@ -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,
}
+46
View File
@@ -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,
}
+2
View File
@@ -26,6 +26,8 @@ def find_model(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"]
elif "model" in checkpoint:
checkpoint = checkpoint["model"]
return checkpoint
+91 -13
View File
@@ -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 = []
@@ -67,7 +79,7 @@ class TimestepEmbedder(nn.Module):
def forward(self, t):
t_freq = self.timestep_embedding(t, self.frequency_embedding_size)
t_emb = self.mlp(t_freq)
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)
# 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)
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])
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)
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)
@@ -483,6 +560,7 @@ class DiT(nn.Module):
return torch.cat([eps, rest], dim=1)
#################################################################################
# Sine/Cosine Positional Embedding Functions #
#################################################################################
+4 -4
View File
@@ -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)
+303
View File
@@ -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)