Giới thiệu
v-diffusion là phương pháp rút ngắn bước lấy mẫu của mô hình khuếch tán nhờ chưng cất tiến trình. Repository mã nguồn mở này cung cấp bản tái hiện chính xác thuật toán trong bài báo gốc, đồng thời bổ sung nhiều tiện ích giúp huấn luyện và triển khai dễ dàng hơn.
Cấu trúc và nguyên lý hoạt động
Thuật toán chính gồm hai giai đoạn:
- Huấn luyện mô hình gốc: sử dụng DDPM truyền thống với 1000 bước khuếch tán.
- Chưng cất tiến trình: tạo chuỗi mô hình "con", mỗi lần giảm một nửa số bước so với mô hình trước đó (1000 → 500 → 250 → 125 → 64 → 32 → 16 → 8 → 4).
Quá trình này không can thiệp vào kiến trúc mạng; thay vào đó, mô hình mới học cách bắt chước phân phối của mô hình cũ ở mức nhiễu thấp. Bộ lấy mẫu từ rosinality/denoising-diffusion-pytorch được tích hợp để đảm bảo ổn định số học và chất lượng ảnh đầu ra.
Kịch bản ứng dụng
- Tạo ảnh thời gian thực: chỉ cần 4–8 bước để sinh ảnh 512×512, phù hợp demo trực tuyến hoặc ứng dụng AR/VR.
- Thiết bị biên: chạy được trên GPU tầm trung hoặc thậm chí CPU nhờ giảm đáng kể số phép toán.
- Nghiên cứu học thuật: mã nguồn rõ ràng giúp sinh viên và nhà nghiên cứu dễ dàng thử nghiệm ý tưởng mới.
Điểm nổi bật của dự án
| Tính năng | Lợi ích |
|---|---|
| Lấy mẫu siêu nhanh | Giảm 1000 bước xuống còn 4–8 bước, tốc độ tăng ~100× |
| Trọng số có sẵn | Tải về và chạy ngay lập tức, không cần huấn luyện lại |
| API mô-đun | Thay thế mô hình backbone chỉ bằng một dòng import |
| TensorBoard tích hợp | Theo dõi loss, ảnh mẫu và gradient theo thời gian thực |
Cài đặt và chạy thử
# 1. Clone repository
git clone https://github.com/xxx/v-diffusion-pytorch.git
cd v-diffusion-pytorch
# 2. Cài đặt môi trường
pip install -r requirements.txt
# 3. Tải trọng số mẫu
wget https://huggingface.co/xxx/celeba_8step.pt -P checkpoints/celeba/
Lấy mẫu với trọng số có sẵn
python generate.py \
--config configs/celeba_8step.yml \
--ckpt checkpoints/celeba/celeba_8step.pt \
--out_dir samples/celeba/ \
--n_samples 16 \
--steps 8
Chưng cất mô hình của bạn
python distill.py \
--base_cfg configs/ddpm_ffhq.yml \
--base_ckpt logs/ffhq_base/checkpoint_ema.pt \
--target_steps 8 \
--max_epochs 200 \
--batch_size 32 \
--lr 2e-4 \
--log_dir logs/ffhq_distill_8step
Sau khi huấn luyện, mô hình 8 bước sẽ được lưu tại logs/ffhq_distill_8step/checkpoint_best.pt.
Điều chỉnh siêu tham số
--target_steps: chọn 16, 8, 4 tùy nhu cầu tốc độ/chất lượng.--loss_type:msehoặcvlb(variational lower bound).--ema_decay: giá trị 0.9999 thường giúp ảnh mịn hơn.
Thử nghiệm các giá trị trên tập validation nhỏ để tìm ra cấu hình tối ưu cho bài toán của bạn.