在 Pytorch 中实现 Google 的文字到图像转换模型 Imagen
Implementation of Imagen, Google's Text-to-Image Neural Network that beats DALL-E2, in Pytorch. It is the new SOTA for text-to-image synthesis.
Architecturally, it is actually much simpler than DALL-E2. It consists of a cascading DDPM conditioned on text embeddings from a large pretrained T5 model (attention network). It also contains dynamic clipping for improved classifier free guidance, noise level conditioning, and a memory efficient unet design.
It appears neither CLIP nor prior network is needed after all. And so research continues.
AI Coffee Break with Letitia | Assembly AI | Yannic Kilcher
Please join if you are interested in helping out with the replication with the LAION community
StabilityAI for the generous sponsorship, as well as my other sponsors out there
Huggingface for their amazing transformers library. The text encoder portion is pretty much taken care of because of them
Jonathan Ho for bringing about a revolution in generative artificial intelligence through his seminal paper
Sylvain and Zachary for the Accelerate library, which this repository uses for distributed training
Jorge Gomes for helping out with the T5 loading code and advice on the correct T5 version
Katherine Crowson, for her beautiful code, which helped me understand the continuous time version of gaussian diffusion
Marunine and Netruk44, for reviewing code, sharing experimental results, and help with debugging
Marunine for providing a potential solution for a color shifting issue in the memory efficient u-nets. Thanks to Jacob for sharing experimental comparisons between the base and memory-efficient unets
Marunine for finding numerous bugs, resolving an issue with resize right, and for sharing his experimental configurations and results
MalumaDev for proposing the use of pixel shuffle upsampler to fix checkboard artifacts
Valentin for pointing out insufficient skip connections in the unet, as well as the specific method of attention conditioning in the base-unet in the appendix
BIGJUN for catching a big bug with continuous time gaussian diffusion noise level conditioning at inference time
Bingbing for identifying a bug with sampling and order of normalizing and noising with low resolution conditioning image
Kay for contributing one line command training of Imagen!
Hadrien Reynaud for testing out text-to-video on a medical dataset, sharing his results, and identifying issues!
$ pip install imagen-pytorch
…
For simpler training, you can directly supply text strings instead of precomputing text encodings. (Although for scaling purposes, you will definitely want to precompute the textual embeddings + mask)
The number of textual captions must match the batch size of the images if you go this route.
# mock images and text (get a lot of this)
texts = [
'a child screaming at finding a worm within a half-eaten apple',
'lizard running across the desert on two feet',
'waking up to a psychedelic landscape',
'seashells sparkling in the shallow waters'
]
images = torch.randn(4, 3, 256, 256).cuda()
# feed images into imagen, training each unet in the cascade
for i in (1, 2):
loss = imagen(images, texts = texts, unet_number = i)
loss.backward()
With the ImagenTrainer wrapper class, the exponential moving averages for all of the U-nets in the cascading DDPM will be automatically taken care of when calling update
…
You can also train Imagen without text (unconditional image generation) as follows
…
Or train only super-resoluting unets
…
At any time you can save and load the trainer and all associated states with the save and load methods. It is recommended you use these methods instead of manually saving with a state_dict call, as there are some device memory management being done underneath the hood within the trainer.
ex.
trainer.save('./path/to/checkpoint.pt')
trainer.load('./path/to/checkpoint.pt')
trainer.steps # (2,) step number for each of the unets, in this case 2
You can also rely on the ImagenTrainer to automatically train off DataLoader instances. You simply have to craft your DataLoader to return either images (for unconditional case), or of ('images', 'text_embeds') for text-guided generation.
ex. unconditional training
…
Thanks to Accelerate, you can do multi GPU training easily with two steps.
First you need to invoke accelerate config in the same directory as your training script (say it is named train.py)
$ accelerate config
Next, instead of calling python train.py as you would for single GPU, you would use the accelerate CLI as so
$ accelerate launch train.py
That's it!
Imagen can also be used via CLI directly.
ex.
$ imagen config
or
$ imagen config --path ./configs/config.json
In the config you are able to change settings for the trainer, dataset and the imagen config.
The Imagen config parameters can be found here
The Elucidated Imagen config parameters can be found here
The Imagen Trainer config parameters can be found here
For the dataset parameters all dataloader parameters can be used.
This command allows you to train or resume training your model
ex.
$ imagen train
or
$ imagen train --unet 2 --epoches 10
You can pass following arguments to the training command.
--config specify the config file to use for training [default: ./imagen_config.json]--unet the index of the unet to train [default: 1]--epoches how many epoches to train for [default: 50]Be aware when sampling your checkpoint should have trained all unets to get a usable result.
ex.
$ imagen sample --model ./path/to/model/checkpoint.pt "a squirrel raiding the birdfeeder"
# image is saved to ./a_squirrel_raiding_the_birdfeeder.png
You can pass following arguments to the sample command.
--model specify the model file to use for sampling--cond_scale conditioning scale (classifier free guidance) in decoder--load_ema load EMA version of unets if availableIn order to use a saved checkpoint with this feature, you either must instantiate your Imagen instance using the config classes, ImagenConfig and ElucidatedImagenConfig or create a checkpoint via the CLI directly
For proper training, you'll likely want to setup config-driven training anyways.
ex.
import torch
from imagen_pytorch import ImagenConfig, ElucidatedImagenConfig, ImagenTrainer
# in this example, using elucidated imagen
imagen = ElucidatedImagenConfig(
unets = [
dict(dim = 32, dim_mults = (1, 2, 4, 8)),
dict(dim = 32, dim_mults = (1, 2, 4, 8))
],
image_sizes = (64, 128),
cond_drop_prob = 0.5,
num_sample_steps = 32
).create()
trainer = ImagenTrainer(imagen)
# do your training ...
# then save it
trainer.save('./checkpoint.pt')
# you should see a message informing you that ./checkpoint.pt is commandable from the terminal
It really should be as simple as that
You can also pass this checkpoint file around, and anyone can continue finetune on their own data
from imagen_pytorch import load_imagen_from_checkpoint, ImagenTrainer
imagen = load_imagen_from_checkpoint('./checkpoint.pt')
trainer = ImagenTrainer(imagen)
# continue training / fine-tuning
Inpainting follows the formulation laid out by the recent Repaint paper. Simply pass in inpaint_images and inpaint_masks to the sample function on either Imagen or ElucidatedImagen
inpaint_images = torch.randn(4, 3, 512, 512).cuda() # (batch, channels, height, width)
inpaint_masks = torch.ones((4, 512, 512)).bool().cuda() # (batch, height, width)
inpainted_images = trainer.sample(texts = [
'a whale breaching from afar',
'young girl blowing out candles on her birthday cake',
'fireworks with blue and green sparkles',
'dust motes swirling in the morning sunshine on the windowsill'
], inpaint_images = inpaint_images, inpaint_masks = inpaint_masks, cond_scale = 5.)
inpainted_images # (4, 3, 512, 512)
For video, similarly pass in your videos to inpaint_videos keyword on .sample. Inpainting mask can either be the same across all frames (batch, height, width) or different (batch, frames, height, width)
inpaint_videos = torch.randn(4, 3, 8, 512, 512).cuda() # (batch, channels, frames, height, width)
inpaint_masks = torch.ones((4, 8, 512, 512)).bool().cuda() # (batch, frames, height, width)
inpainted_videos = trainer.sample(texts = [
'a whale breaching from afar',
'young girl blowing out candles on her birthday cake',
'fireworks with blue and green sparkles',
'dust motes swirling in the morning sunshine on the windowsill'
], inpaint_videos = inpaint_videos, inpaint_masks = inpaint_masks, cond_scale = 5.)
inpainted_videos # (4, 3, 8, 512, 512)
Tero Karras of StyleGAN fame has written a new paper with results that have been corroborated by a number of independent researchers as well as on my own machine. I have decided to create a version of Imagen, the ElucidatedImagen, so that one can use the new elucidated DDPM for text-guided cascading generation.
Simply import ElucidatedImagen, and then instantiate the instance as you did before. The hyperparameters are different than the usual ones for discrete and continuous time gaussian diffusion, and can be individualized for each unet in the cascade.
Ex.
…
This repository will also start accumulating new research around text guided video synthesis. For st
暂无开放 Issues,或尚未同步最近议题。