You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为长时间运行的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.16 20:16:57