def ddpm_sampling(model, num_steps=1000):
    """DDPM采样算法"""
    # 1. 从纯噪声开始
    x_t = torch.randn(1, 3, 512, 512)
    
    # 2. 逐步去噪
    for t in reversed(range(num_steps)):
        # 预测噪声
        epsilon_pred = model(x_t, t)
        
        # 计算均值
        alpha_bar_t = get_alpha_bar(t)
        mu_t = (x_t - torch.sqrt(1 - alpha_bar_t) * epsilon_pred) / torch.sqrt(alpha_bar_t)
        
        # 计算方差
        if t > 0:
            beta_t = get_beta(t)
            sigma_t = torch.sqrt(beta_t)
            z = torch.randn_like(x_t)
            x_t = mu_t + sigma_t * z
        else:
            x_t = mu_t  # 最后一步不加噪声
    
    return x_t
