mirror of
https://github.com/zenlm/enso.git
synced 2026-07-26 22:30:28 +00:00
debug
This commit is contained in:
Binary file not shown.
@@ -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
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user