← Back to list

動手做一個語意圖片生成模型

使用擴散生成模型(Stable Diffusion Model, DDPMs)

KL Wu · 2025-08-29 02:54 · 0 claps · 17.9 min read
#ddpm #stable-diffusion #圖片生成
Open on Medium ↗
Wiki topics: MM · Multimodal & Generative Media

動手做一個語意圖片生成模型

使用擴散生成模型(Stable Diffusion Model, DDPMs)

關於圖片生成的模型,對抗式生成網路(GANs)的概念直觀且易於理解。其核心是同時訓練生成器與鑑別器,通過兩者相互對抗的方式持續優化權重參數,最終實現高品質的圖片生成。如需深入了解,可參考先前相關的「動手做GANs」內容。

[embed]GAN 對抗式生成網路實作 -動漫美少女圖像生成器 GAN(Generative Adversarial Network)對抗式生成網路,由Ian…medium.com

[embed]GAN 生成對抗網路-使用捲積神經網路(DCGAN) 使用 GAN 生成動漫美少女圖片的實驗與改進medium.com

[embed]從風景相片到莫內畫作:探索CycleGAN的模型架構 最近,ChatGPT推出了一項令人興奮的功能:將普通相片轉換為吉卜力動畫風格的圖片!這項技術瞬間引爆熱潮,網友們爭相將自拍、風景照甚至寵物照片轉成宮崎駿筆下的夢幻畫風。然而,這股熱潮也讓OpenAI的伺服器不堪重負,最終不得不暫停服務。這場…medium.com

然而,生成對抗網路(GANs)的訓練過程並不簡單,常因梯度不穩定或模式崩潰(僅能生成近似圖片)而面臨挑戰。因此,穩定擴散模型(Stable Diffusion)應運而生,成功解決了GANs訓練中的諸多問題,並在高品質圖片生成任務上取得了顯著成效。接下來,我們將初步探索擴散生成模型的強大潛力!

關於擴散生成模型的理論基礎,可參考相關論文(https://arxiv.org/abs/2006.11239)。網路上已有大量關於其原理的討論,本文不再贅述理論細節,而是聚焦於以入門方式引導讀者,透過簡單的語意輸入(例如指定生成的手寫數字),打造一個能根據輸入語意生成相應圖片的模型。對於初學者而言,MNIST手寫數字資料集無疑是最佳的練習起點。

本篇目標是構建一個圖片生成模型,根據輸入的數字生成對應的手寫數字圖片。我們將採用Colab環境並基於PyTorch框架,逐步完成模型的搭建。

Step 1. 製作資料集

首先,我們需要安裝製作資料集所需的套件。我們將直接使用PyTorch提供的工具庫,包括:

  • datasets:用於下載和管理資料集。
  • transforms:提供資料轉換和預處理的工具。
  • DataLoader:用於處理小批次訓練資料的工具。

以下先載入需要的函示庫:

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
import math

pytorch 提供CPU或是GPU計算依環境調度的功能,使用起來比Tensorflow更方便設置環境。

首先設置GPU計算環境:

 # Device
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

著,我們將建立資料集調用器。透過下載的資料集存放在硬碟中,並在小批次訓練時利用資料載入器(DataLoader)從硬碟空間動態取得每個批次的訓練資料。這種方式無需將整個資料集預先載入記憶體,從而有效節省記憶體空間。

# DataLoader
    transform = transforms.Compose([
        transforms.ToTensor(),
        transforms.Normalize((0.5,), (0.5,))
    ])
    train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
    train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=2)

Step 2. 雜訊排程器

擴散模型的核心在於逐步為訓練圖片添加雜訊,並訓練模型預測每一階段的雜訊,直至圖片最終轉為隨機高斯雜訊。在生成圖片時,模型則執行反向過程:從隨機高斯雜訊開始,通過預測並移除雜訊,逐步還原生成高品質圖片。

雜訊的添加強度通常由高斯分布的標準差(β值)決定,範圍一般在0.0001至0.02之間,數值越小,離散程度越低。根據所需的步數(例如1000步),雜訊強度會依序遞增。在訓練過程中,模型通過隨機採樣學習預測被添加的雜訊。

為此,我們需要設計一個雜訊排程器,以方便訓練時獲取添加雜訊的圖片及其對應的雜訊資訊,從而支援模型的有效訓練。

# DDPM Scheduler 
class DDPMScheduler:
    def __init__(self, num_timesteps=1000, beta_start=0.0001, beta_end=0.02):
        self.num_timesteps = num_timesteps
        self.beta_start = beta_start
        self.beta_end = beta_end

        self.betas = torch.linspace(beta_start, beta_end, num_timesteps).to(device)
        self.alphas = 1.0 - self.betas
        self.alphas_cumprod = torch.cumprod(self.alphas, dim=0)
        self.sqrt_alphas_cumprod = torch.sqrt(self.alphas_cumprod)
        self.sqrt_one_minus_alphas_cumprod = torch.sqrt(1.0 - self.alphas_cumprod)
        self.sqrt_recip_alphas = torch.sqrt(1.0 / self.alphas)
        self.posterior_variance = self.betas * (1. - self.alphas_cumprod) / (1. - self.alphas_cumprod)

    def add_noise(self, x_0, noise, t):
        t = t.to(self.betas.device)
        sqrt_alpha_prod = self.sqrt_alphas_cumprod[t].view(-1, 1, 1, 1)
        sqrt_one_minus_alpha_prod = self.sqrt_one_minus_alphas_cumprod[t].view(-1, 1, 1, 1)
        noisy_data = sqrt_alpha_prod * x_0 + sqrt_one_minus_alpha_prod * noise
        return noisy_data

Step 3. 位置編碼

為使模型能夠學習任意採樣步數對應的雜訊,我們需要對步數進行編碼。這裡採用Transformer模型中提出的位置編碼方法,通過正弦和餘弦函數,結合權重參數,進行編碼嵌入,以有效捕捉步數資訊。

 import torch.nn

# Sinusoidal Embedding
class SinusoidalEmbedding(nn.Module):
    def __init__(self, dim):
        super().__init__()
        self.dim = dim

    def forward(self, time):
        time = time.float()
        device = time.device
        half_dim = self.dim // 2
        emb = math.log(10000) / (half_dim - 1)
        emb = torch.exp(torch.arange(half_dim, device=device) * -emb)
        emb = time[:, None] * emb[None, :]
        emb = torch.cat((emb.sin(), emb.cos()), dim=-1)
        return emb

Step 4. 雜訊預測模型

https://theaisummer.com/unet-architectures/

https://theaisummer.com/unet-architectures/

擴散模型的訓練層需包含兩個核心架構:首先是辨識帶有雜訊的圖片特徵,其次是預測雜訊的分布。通常採用U-Net架構作為捲積神經網路的核心。U-Net因其結構形似字母「U」而得名,包含向下採樣和向上採樣兩個部分:向下採樣負責解析圖片特徵,向上採樣則負責生成圖片(即預測雜訊)。

MNIST資料集的圖片為單通道黑白圖片(通道數為1)。需要注意的是,PyTorch對圖片張量的軸順序定義與其他框架(如TensorFlow)不同,其格式為(Channel, Width, Height),即通道軸為第0軸。

在模型物件的設計上,為提升靈活性與可擴展性,我們將捲積層、向下採樣、向上採樣及U-Net架構分別封裝為獨立物件。這種模組化設計便於後續適應多通道圖片或更高解析度圖片生成的需求。

在U-Net模型中,除了輸入圖片的資訊外,還需將圖片對應的數字標籤進行編碼嵌入,並與圖片資料一同輸入模型進行訓練。這一設計使模型能夠學習數字的語意特徵,從而具備根據指定數字生成相應圖片的能力。

# DoubleConv with time and condition embedding
class DoubleConv(nn.Module):
    def __init__(self, in_ch, out_ch, emb_dim):
        super().__init__()
        self.conv1 = nn.Conv2d(in_ch, out_ch, 3, padding=1)
        self.conv2 = nn.Conv2d(out_ch, out_ch, 3, padding=1)
        self.relu = nn.ReLU(inplace=True)
        self.emb_proj = nn.Linear(emb_dim, out_ch)

    def forward(self, x, emb):
        x = self.conv1(x)
        emb = self.emb_proj(emb).unsqueeze(-1).unsqueeze(-1)
        x = x + emb
        x = self.relu(x)
        x = self.conv2(x)
        x = self.relu(x)
        return x

# Down block
class Down(nn.Module):
    def __init__(self, in_ch, out_ch, emb_dim):
        super().__init__()
        self.pool = nn.MaxPool2d(2)
        self.conv = DoubleConv(in_ch, out_ch, emb_dim)

    def forward(self, x, emb):
        x = self.pool(x)
        x = self.conv(x, emb)
        return x

# Up block
class Up(nn.Module):
    def __init__(self, in_ch, out_ch, emb_dim):
        super().__init__()
        self.up = nn.ConvTranspose2d(in_ch, out_ch, kernel_size=2, stride=2)
        self.conv = DoubleConv(out_ch * 2, out_ch, emb_dim)

    def forward(self, x, skip, emb):
        x = self.up(x)
        x = torch.cat([skip, x], dim=1)
        x = self.conv(x, emb)
        return x

# Conditional UNet
class ConditionalUNet(nn.Module):
    def __init__(self, in_ch=1, out_ch=1, time_dim=128, num_classes=10):
        super().__init__()
        self.time_dim = time_dim
        self.num_classes = num_classes

        # Time embedding
        self.time_emb = SinusoidalEmbedding(time_dim)
        self.time_mlp = nn.Sequential(
            nn.Linear(time_dim, time_dim),
            nn.ReLU(),
            nn.Linear(time_dim, time_dim)
        )

        # Label embedding (for digit-number prompt)
        self.label_emb = nn.Embedding(num_classes, time_dim)

        # UNet layers
        self.inc = DoubleConv(in_ch, 32, time_dim * 2)  # Concat time and label emb
        self.down1 = Down(32, 64, time_dim * 2)
        self.down2 = Down(64, 128, time_dim * 2)
        self.up1 = Up(128, 64, time_dim * 2)
        self.up2 = Up(64, 32, time_dim * 2)
        self.out = nn.Conv2d(32, out_ch, 1)

    def forward(self, x, t, labels):
        # Time embedding
        t_emb = self.time_emb(t)
        t_emb = self.time_mlp(t_emb)

        # Label embedding
        l_emb = self.label_emb(labels)

        # Concat embeddings
        emb = torch.cat([t_emb, l_emb], dim=-1)

        # Forward pass
        x1 = self.inc(x, emb)
        x2 = self.down1(x1, emb)
        x3 = self.down2(x2, emb)

        x = self.up1(x3, x2, emb)
        x = self.up2(x, x1, emb)
        x = self.out(x)
        return x

Step 5. 模型訓練

上述資料集準備以及模型架構建立完成後,就可以準備開始訓練,這邊先製作模型訓練函式,方便呼叫使用:

# Training function
def train(model, scheduler, train_loader, optimizer, device, epochs=10):
    model.train()
    for epoch in range(epochs):
        total_loss = 0 #reset total_loss
        for batch_idx, (data, labels) in enumerate(train_loader):
            data, labels = data.to(device), labels.to(device)
            optimizer.zero_grad()

            # Sample timesteps and noise
            t = scheduler.get_timesteps(data.size(0), device)
            noise = torch.randn_like(data)

            # Add noise
            noisy_data = scheduler.add_noise(data, noise, t)

            # Predict noise
            predicted_noise = model(noisy_data, t, labels)

            # Loss: MSE between predicted and actual noise
            loss = nn.MSELoss()(predicted_noise, noise)
            loss.backward()
            optimizer.step()

            total_loss += loss.item()

        print(f"Epoch {epoch+1}/{epochs}, Loss: {total_loss / len(train_loader):.4f}")
    #Save trained model
    torch.save(model.state_dict(), f"model_{epoch+1}.pth")

訓練前設定dataset下載資料並設定好相關的資料轉換器(transform)、資料調用器(dataloader)、雜訊排程器(scheduler)、優化器(Adam, 學習率0.0002),便可呼叫訓練函式進行模型訓練:

# Device
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

    # DataLoader
    transform = transforms.Compose([
        transforms.ToTensor(),
        # transforms.Normalize((0.1307,), (0.3081,))
        transforms.Normalize((0.5,), (0.5,))
        ])
    train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
    train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True, num_workers=2)

    # Model and Scheduler
    scheduler = DDPMScheduler(num_timesteps=1000)
    model = ConditionalUNet().to(device)
    optimizer = optim.Adam(model.parameters(), lr=2e-4)

    # Train
    epochs = 50
    # train(model, scheduler, train_loader, optimizer, device, epochs=epochs)  # Adjust epochs as needed

模型訓練過程

模型訓練過程

Step 6. 模型推論驗證

模型訓練完之後,終於來到模型驗證的部分,我們先製作一個圖片生成函式,方便我們呼叫:

# Generation function
def generate(model, save_weights, scheduler, device, num_samples=10, labels=None, steps=1000):
    model.eval()
    model.load_state_dict(save_weights)
    with torch.no_grad():
        # Start from pure noise
        x = torch.randn(num_samples, 1, 28, 28).to(device)
        if labels is None:
            labels = torch.randint(0, 10, (num_samples,)).to(device)
        else:
            labels = torch.tensor([labels] * num_samples).to(device)

        for t in reversed(range(steps)):
            t_tensor = torch.full((num_samples,), t, dtype=torch.long, device=device)
            predicted_noise = model(x, t_tensor, labels)
            x = scheduler.sample_previous_timestep(x, t_tensor, predicted_noise)

        return x.cpu()

我們可以直接載入先前訓練好的模型,並設定想要產生的數字(例如5):

import matplotlib.pyplot as plt
import os
# Generate samples for digit 5
    save_weights = torch.load(f"model_{epochs}.pth")
    labels = 5
    generated = generate(model, save_weights, scheduler, device, num_samples=5, labels=labels)

    # Save or display
    os.makedirs("generated", exist_ok=True)
    for i in range(5):
        plt.imshow(generated[i].squeeze(), cmap="gray")
        plt.axis("off")
        plt.savefig(f"generated/digit_{labels}_{i}.png")
        plt.close()
    print("Generated images saved in 'generated' folder.")

手寫數字5的生成圖片

手寫數字5的生成圖片

擴散模型通過雜訊添加與去雜訊的正反向過程,為圖片生成任務提供了穩定且高效的解決方案。本文以MNIST資料集為例,介紹了從資料準備、雜訊排程到U-Net模型設計的入門實作流程。透過DataLoader的記憶體優化、步數編碼的靈活設計以及U-Net的模組化架構,讀者可快速構建一個能根據數字語意生成手寫圖片的模型。未來可進一步探索進階主題,例如多通道圖片生成或更複雜的語意控制,以拓展擴散模型的應用潛力。


메타데이터
post_id
039fdecdbf8d
slug
動手做一個語意圖片生成模型-039fdecdbf8d
url
https://medium.com/@wkol1126/%E5%8B%95%E6%89%8B%E5%81%9A%E4%B8%80%E5%80%8B%E8%AA%9E%E6%84%8F%E5%9C%96%E7%89%87%E7%94%9F%E6%88%90%E6%A8%A1%E5%9E%8B-039fdecdbf8d
canonical_url
https://medium.com/@wkol1126/%E5%8B%95%E6%89%8B%E5%81%9A%E4%B8%80%E5%80%8B%E8%AA%9E%E6%84%8F%E5%9C%96%E7%89%87%E7%94%9F%E6%88%90%E6%A8%A1%E5%9E%8B-039fdecdbf8d
author_url
https://medium.com/@wkol1126
status
ok
fetched_at
2026-07-17 21:46:37