This commit is contained in:
费政聪
2024-08-14 14:22:50 +08:00
committed by GitHub
parent 936fb27b55
commit 5d9775735d
3 changed files with 34 additions and 15 deletions
Binary file not shown.
+11 -6
View File
@@ -1,11 +1,16 @@
import torch import torch
from torch.nn.parallel import DistributedDataParallel as DDP
class RectifiedFlow(torch.nn.Module): class RectifiedFlow(torch.nn.Module):
def __init__(self, model, ln=True): def __init__(self, model, ln=True):
super().__init__() super().__init__()
self.model = model self.model = model
self.ln = ln self.ln = ln
self.stratified = False self.stratified = False
if isinstance(model, DDP):
self.learn_sigma = model.module.learn_sigma
else:
self.learn_sigma = model.learn_sigma
def forward(self, x, cond): def forward(self, x, cond):
@@ -31,7 +36,7 @@ class RectifiedFlow(torch.nn.Module):
# make t, zt into same dtype as x # make t, zt into same dtype as x
zt, t = zt.to(x.dtype), t.to(x.dtype) zt, t = zt.to(x.dtype), t.to(x.dtype)
vtheta = self.model(zt, t, cond) vtheta = self.model(zt, t, cond)
if self.model.learn_sigma == True: if self.learn_sigma == True:
vtheta, _ = vtheta.chunk(2, dim=1) vtheta, _ = vtheta.chunk(2, dim=1)
batchwise_mse = ((z1 - x - vtheta) ** 2).mean(dim=list(range(1, len(x.shape)))) batchwise_mse = ((z1 - x - vtheta) ** 2).mean(dim=list(range(1, len(x.shape))))
tlist = batchwise_mse.detach().cpu().reshape(-1).tolist() tlist = batchwise_mse.detach().cpu().reshape(-1).tolist()
@@ -49,11 +54,11 @@ class RectifiedFlow(torch.nn.Module):
t = torch.tensor([t] * b).to(z.device) t = torch.tensor([t] * b).to(z.device)
vc = self.model(z, t, cond) vc = self.model(z, t, cond)
if self.model.learn_sigma == True: if self.learn_sigma == True:
vc, _ = vc.chunk(2, dim=1) vc, _ = vc.chunk(2, dim=1)
if null_cond is not None: if null_cond is not None:
vu = self.model(z, t, null_cond) vu = self.model(z, t, null_cond)
if self.model.learn_sigma == True: if self.learn_sigma == True:
vu, _ = vu.chunk(2, dim=1) vu, _ = vu.chunk(2, dim=1)
vc = vu + cfg * (vc - vu) vc = vu + cfg * (vc - vu)
@@ -72,11 +77,11 @@ class RectifiedFlow(torch.nn.Module):
t = torch.tensor([t] * b).to(z.device) t = torch.tensor([t] * b).to(z.device)
vc = self.model(z, t, cond) vc = self.model(z, t, cond)
if self.model.learn_sigma == True: if self.learn_sigma == True:
vc, _ = vc.chunk(2, dim=1) vc, _ = vc.chunk(2, dim=1)
if null_cond is not None: if null_cond is not None:
vu = self.model(z, t, null_cond) vu = self.model(z, t, null_cond)
if self.model.learn_sigma == True: if self.learn_sigma == True:
vu, _ = vu.chunk(2, dim=1) vu, _ = vu.chunk(2, dim=1)
vc = vu + cfg * (vc - vu) vc = vu + cfg * (vc - vu)
x = z - i * dt * vc x = z - i * dt * vc
+23 -9
View File
@@ -28,7 +28,8 @@ import logging
import os import os
from models import DiT_models from models import DiT_models
from diffusion import create_diffusion from diffusion import create_diffusion
from diffusion.rectified_flow import RectifiedFlow
from diffusers.models import AutoencoderKL from diffusers.models import AutoencoderKL
from download import find_model from download import find_model
@@ -128,7 +129,7 @@ def main(args):
os.makedirs(args.results_dir, exist_ok=True) # Make results folder (holds all experiment subfolders) os.makedirs(args.results_dir, exist_ok=True) # Make results folder (holds all experiment subfolders)
experiment_index = len(glob(f"{args.results_dir}/*")) 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) model_string_name = args.model.replace("/", "-") # e.g., DiT-XL/2 --> DiT-XL-2 (for naming folders)
experiment_dir = f"{args.results_dir}/{experiment_index:03d}-{model_string_name}" # Create an experiment folder experiment_dir = f"{args.results_dir}/{model_string_name}-{args.num_experts}E{args.num_experts_per_tok}A{args.pretraining_tp}TP" # Create an experiment folder
checkpoint_dir = f"{experiment_dir}/checkpoints" # Stores saved model checkpoints checkpoint_dir = f"{experiment_dir}/checkpoints" # Stores saved model checkpoints
os.makedirs(checkpoint_dir, exist_ok=True) os.makedirs(checkpoint_dir, exist_ok=True)
logger = create_logger(experiment_dir) logger = create_logger(experiment_dir)
@@ -157,7 +158,13 @@ def main(args):
requires_grad(ema, False) requires_grad(ema, False)
# model = DDP(model.to(device), device_ids=[rank]) # model = DDP(model.to(device), device_ids=[rank])
model = DDP(model.to(device), device_ids=[device]) model = DDP(model.to(device), device_ids=[device])
diffusion = create_diffusion(timestep_respacing="") # default: 1000 steps, linear noise schedule
if args.rf:
logger.info("train with rectified flow")
diffusion = RectifiedFlow(model)
else:
diffusion = create_diffusion(timestep_respacing="") # default: 1000 steps, linear noise schedule
vae = AutoencoderKL.from_pretrained(args.vae_path).to(device) vae = AutoencoderKL.from_pretrained(args.vae_path).to(device)
logger.info(f"DiT Parameters: {sum(p.numel() for p in model.parameters()):,}") logger.info(f"DiT Parameters: {sum(p.numel() for p in model.parameters()):,}")
@@ -212,10 +219,15 @@ def main(args):
with torch.no_grad(): with torch.no_grad():
# Map input images to latent space + normalize latents: # Map input images to latent space + normalize latents:
x = vae.encode(x).latent_dist.sample().mul_(0.18215) 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) if args.rf:
loss_dict = diffusion.training_losses(model, x, t, model_kwargs) loss, _ = diffusion.forward(x, y)
loss = loss_dict["loss"].mean() else:
t = torch.randint(0, diffusion.num_timesteps, (x.shape[0],), device=device)
model_kwargs = dict(y=y)
loss_dict = diffusion.training_losses(model, x, t, model_kwargs)
loss = loss_dict["loss"].mean()
if (data_iter_step + 1) % args.accum_iter == 0: if (data_iter_step + 1) % args.accum_iter == 0:
opt.zero_grad() opt.zero_grad()
loss.backward() loss.backward()
@@ -278,9 +290,11 @@ if __name__ == "__main__":
parser.add_argument("--global-seed", type=int, default=2024) parser.add_argument("--global-seed", type=int, default=2024)
parser.add_argument("--num-workers", type=int, default=4) parser.add_argument("--num-workers", type=int, default=4)
parser.add_argument("--log-every", type=int, default=100) parser.add_argument("--log-every", type=int, default=100)
parser.add_argument('--accum_iter', default=8, type=int,) parser.add_argument('--accum_iter', default=1, type=int,)
parser.add_argument('--num_experts', 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('--num_experts_per_tok', default=2, type=int,)
parser.add_argument("--ckpt-every", type=int, default=50_000) parser.add_argument('--pretraining_tp', default=2, type=int,)
parser.add_argument("--ckpt-every", type=int, default=50_000)
parser.add_argument("--rf", type=bool, default=False)
args = parser.parse_args() args = parser.parse_args()
main(args) main(args)