A repository where I implement Rectified Flow model from the rectified flow paper.
Тут я реализовываю rectified flow модель на pytorch с удобным интерфейсом для экспериментов. Обучаю модель на CIFAR10. В будущем я планирую добавить и другие небольшие датасеты.
На вход rectified flow получает (X, t) и выдает векторное поле той же размерностью, что и X для момента времени
Лосс модели следующий:
где
Под капотом я использую DiT (Diffusion Transformer). Для входа t в модель используется AdaLN (Adaptive Layer Norm).
Для генерации изображений я реализовал 3 ODE-Solver'а:
- Euler
- Heun
- odeint (torchdiffeq lib)
Формула для Euler ODE солвера:
где
само
Причем генерация происходит не на основных весах модели, а на EMA (Exponential Moving Average) весах. Таким образом генерации должны быть лучше, нежели если бы они были на весах основной модели.
Формула обновления EMA весов:
decay обычно использут от 0.9 до 0.9999
Изображения находятся в results. Надо отметить, что часть генераций на данный момент каша по нескольким причинам:
- Маленькое разрешение (32x32 у CIFAR10) на таком разрешении даже человеку иногда трудно понять, что изображено
- Отсутствие привязки к текстовым эмбеддингам. Планируется добавить в будущем
- configs
- src
- dataset
- modules
- modules.py
- rectified_flow.py
- utils
- data_utils.py
- initialization.py
- utils.py
- train.py
- eval.py
- Запуск обучения:
python train.py
| Flag | Type | Default | Description |
|---|---|---|---|
--device |
str |
cuda |
Какой device использовать (cpu, cuda). |
--config |
str |
.configs/default_config.yaml |
Путь к YAML конфиг файлу. |
--mode |
str |
train |
Режим обучения (train, overfit, debug, train_c (продолжение обучения)), overfit и debug режимы используют их собственные конфиги (overfit_config, debug_config). train_c дает выбор эксперимента, какой вы хотите продолжить. |
--experiment |
str |
None | Для явного указания эксперимента для продолжения обучения |
--wandb |
bool |
из конфига | Подключать wandb или нет |
--batch_size |
int |
из конфига | Количество изображений в батче, переписывает значение из конфига |
--epochs |
int |
из конфига | Количество эпох, переписывает значение из конфига |
--num_training |
int |
из конфига | Количество изображений взятых из датасета для обучения, переписывает значение из конфига |
--decay |
float |
из конфига | Decay для EMA модели, переписывает значение из конфига |
--warmup_epochs |
int |
из конфига | Количество эпох для разогрева, переписывает значение из конфига |
Все, что запущено в debug моде, сохраняется в папку results/debug (она перезаписывается). Все остальное сохраняется в свои отдельные папки экспериментов.
- Тестирование модели:
python eval.py --device --solver['euler', 'heun', default='odeint'] --mode=['grid', 'process']
In this repo, I am implementing a Rectified Flow model using PyTorch, featuring a user-friendly interface for experimentation. The model trained on CIFAR10. I plan to add other small-size datasets in the feature.
Rectified Flow takes (X, t) as input and predicts a vector field of the same dimension as X for a time step
The model's loss function is as follows:
where:
Under the hood, I use a DiT (Diffusion Transformer) architecture. The time step is incorporated into the model using AdaLN (Adaptive Layer Norm).
For image generation, I have implemented three ODE Solvers:
- Euler
- Heun
- odeint (torchdiffeq lib)
The Euler ODE solver formula:
where:
and
Furthermore, generation is performed using EMA (Exponential Moving Average) weights rather than the main model weights. This typically results in higher quality samples compared to using the raw model weights.
The EMA weight update formula:
The decay value usually ranges from 0.9 to 0.9999
The images can be found in the results section. It should be noted that some of the generations currently look like "mush" for several reasons:
- Low resolution: CIFAR-10 uses 32x32 images. At this resolution, even for a human, it can sometimes be difficult to tell what is being depicted.
- Lack of text embeddings: The model is currently unconditional. I plan to add text-conditioning in the future.
- Start training:
python train.py
| Flag | Type | Default | Description |
|---|---|---|---|
--device |
str |
cuda |
Device to use (cuda, cpu). |
--config |
str |
.configs/default_config.yaml |
Path to the YAML configuration file. |
--mode |
str |
train |
Training mode (train, overfit, debug, train_c (continue training)), overfit and debug modes uses their own configs (overfit_config, debug_config). train_c gives you choose an experiment which you want to continue to train. |
--experiment |
str |
None | specify a specific experiment to continue learning |
--wandb |
bool |
from config | Is wandb linked or not |
--batch_size |
int |
from config | Number of samples per training step. Changes config number |
--epochs |
int |
from config | Number of epochs. Changes config number |
--num_training |
int |
from config | Number of training samples from the dataset (now only CIFAR10). Changes config number |
--decay |
float |
from config | Decay for EMA weights. Changes config number |
--warmup_epochs |
int |
from config | Number of warmup epochs. Changes config number |
Everything run in debug mode is saved to the results/debug folder (it will be overwritten). All other runs are saved in their respective experiment folders.
- Model Testing:
python eval.py --device --solver['euler', 'heun', default='odeint'] --mode=['grid', 'process']
Samples with odeint solver:
Generation process with Heun solver (T = 50):
Generation process with Heun solver (T = 25):