def ddim_sampling(model, num_steps=50):
    """DDIM采样算法"""
    # 1. 从纯噪声开始
    x_t = torch.randn(1, 3, 512, 512)
    
    # 2. 确定性去噪
    for t in reversed(range(0, 1000, 1000//num_steps)):
        # 预测噪声
        epsilon_pred = model(x_t, t)
        
        # 预测原始图像
        alpha_bar_t = get_alpha_bar(t)
        x_0_pred = (x_t - torch.sqrt(1 - alpha_bar_t) * epsilon_pred) / torch.sqrt(alpha_bar_t)
        
        # 计算下一步
        if t > 0:
            alpha_bar_prev = get_alpha_bar(t - 1000//num_steps)
            x_t = torch.sqrt(alpha_bar_prev) * x_0_pred + torch.sqrt(1 - alpha_bar_prev) * epsilon_pred
        else:
            x_t = x_0_pred
    
    return x_t
