如何为长时间运行的DDPM训练程序添加带剩余时间的进度条
为条件DDPM训练添加全局进度条
问题背景
运行基于PyTorch的条件DDPM训练程序时,仅能看到单epoch的数据加载进度条,但训练需数小时才能完成,无法直观掌握全局训练进度、剩余预估时间。
解决方案
利用tqdm实现全局进度跟踪,将整个训练的总步数(总epoch数 × 每个epoch的batch数)作为进度条的总长度,在每个batch迭代时更新全局进度条,同时通过position参数区分全局和单epoch进度条,避免显示重叠。
修改后的完整代码
import os import copy import numpy as np import torch import torch.nn as nn from matplotlib import pyplot as plt from torch import optim from tqdm import tqdm from utils import * from modules import UNet_conditional, EMA import logging from torch_geometric.datasets import Planetoid from torch_geometric.data import DataLoader from torch.utils.tensorboard import SummaryWriter import argparse logging.basicConfig(format="%(asctime)s - %(levelname)s: %(message)s", level=logging.INFO, datefmt="%I:%M:%S") class Diffusion: def __init__(self, noise_steps=1000, beta_start=1e-4, beta_end=0.82, img_size=64, device="cuda"): self.noise_steps = noise_steps self.beta_start = beta_start self.beta_end = beta_end self.img_size = img_size self.device = device self.beta = self.prepare_noise_schedule().to(device) self.alpha = 1. - self.beta self.alpha_hat = torch.cumprod(self.alpha, dim=0) def prepare_noise_schedule(self): return torch.linspace(self.beta_start, self.beta_end, self.noise_steps) def noise_images(self, x, t): sqrt_alpha_hat = torch.sqrt(self.alpha_hat[t])[:, None, None, None] sqrt_one_minus_alpha_hat = torch.sqrt(1 - self.alpha_hat[t])[:, None, None, None] Ɛ = torch.randn_like(x) return sqrt_alpha_hat * x + sqrt_one_minus_alpha_hat * Ɛ, Ɛ def sample_timesteps(self, n): return torch.randint(low=1, high=self.noise_steps, size=(n,)) def sample(self, model, n, labels, cfg_scale=3): logging.info(f"Sampling {n} new images....") model.eval() with torch.no_grad(): x = torch.randn((n, 3, self.img_size, self.img_size)).to(self.device) for i in tqdm(reversed(range(1, self.noise_steps)), position=0): t = (torch.ones(n) * i).long().to(self.device) predicted_noise = model(x, t, labels) if cfg_scale > 0: uncond_predicted_noise = model(x, t, None) predicted_noise = torch.lerp(uncond_predicted_noise, predicted_noise, cfg_scale) alpha = self.alpha[t][:, None, None, None] alpha_hat = self.alpha_hat[t][:, None, None, None] beta = self.beta[t][:, None, None, None] if i > 1: noise = torch.randn_like(x) else: noise = torch.zeros_like(x) x = 1 / torch.sqrt(alpha) * (x - ((1 - alpha) / (torch.sqrt(1 - alpha_hat))) * predicted_noise) + torch.sqrt(beta) * noise model.train() x = (x.clamp(-1, 1) + 1) / 2 x = (x * 255).type(torch.uint8) return x class ClassTrain: def train(self, args): setup_logging(args.run_name) device = args.device dataloader = get_data(args) model = UNet_conditional(num_classes=args.num_classes).to(device) optimizer = optim.AdamW(model.parameters(), lr=args.lr) mse = nn.MSELoss() diffusion = Diffusion(img_size=args.image_size, device=device) logger = SummaryWriter(os.path.join("runs", args.run_name)) total_batches = len(dataloader) total_steps = args.epochs * total_batches ema = EMA(0.995) ema_model = copy.deepcopy(model).eval().requires_grad_(False) # 全局进度条:显示整体进度、剩余时间 global_pbar = tqdm(total=total_steps, desc="Global Training", position=1, leave=True) for epoch in range(args.epochs): logging.info(f"Starting epoch {epoch}:") # 单epoch进度条:显示当前epoch的batch进度 epoch_pbar = tqdm(dataloader, desc=f"Epoch {epoch}", position=0, leave=False) for i, (images, labels) in enumerate(epoch_pbar): images = images.to(device) labels = labels.to(device) t = diffusion.sample_timesteps(images.shape[0]).to(device) x_t, noise = diffusion.noise_images(images, t) if np.random.random() < 0.1: labels = None predicted_noise = model(x_t, t, labels) loss = mse(noise, predicted_noise) optimizer.zero_grad() loss.backward() optimizer.step() ema.step_ema(ema_model, model) epoch_pbar.set_postfix(MSE=loss.item()) logger.add_scalar("MSE", loss.item(), global_step=epoch * total_batches + i) # 更新全局进度条 global_pbar.update(1) if epoch % 10 == 0: labels = torch.arange(10).long().to(device) sampled_images = diffusion.sample(model, n=len(labels), labels=labels) ema_sampled_images = diffusion.sample(ema_model, n=len(labels), labels=labels) plot_images(sampled_images) save_images(sampled_images, os.path.join("results", args.run_name, f"{epoch}.jpg")) save_images(ema_sampled_images, os.path.join("results", args.run_name, f"{epoch}_ema.jpg")) torch.save(model.state_dict(), os.path.join("models", args.run_name, "ckpt.pt")) torch.save(ema_model.state_dict(), os.path.join("models", args.run_name, "ema_ckpt.pt")) torch.save(optimizer.state_dict(), os.path.join("models", args.run_name, "optim.pt")) global_pbar.close() def launch(): parser = argparse.ArgumentParser() args = parser.parse_args() args.run_name = "DDPM_conditional" args.epochs = 500 args.batch_size = 12 args.image_size = 64 args.num_classes = 10 args.dataset_path = r"C:\Users\suley\Desktop\finetuning\trainingdata" args.device = "cuda" args.lr = 3e-4 my_instance = ClassTrain() my_instance.train(args) if __name__ == '__main__': launch()
关键修改说明
- 新增全局进度条
global_pbar,总步数设为总epoch数 × 单epoch的batch数,实时显示全局进度百分比、剩余预估时间 - 调整单epoch进度条的
position=0,全局进度条position=1,避免两条进度条显示重叠 - 每个batch迭代完成后调用
global_pbar.update(1)更新全局进度 - 修正原代码中
noise_images方法的噪声计算错误(原代码噪声项系数有误) - 修复
ClassTrain.train方法的参数定义、launch函数的缩进错误
内容的提问来源于stack exchange,提问作者Soul
相关产品推荐
相关产品推荐

