Baike.dev
All toolsAI codingTrendingOpen sourceNewsSubmit
Log in
< Back to tools
D

DALLE-pytorch

> 数据库
Open source

Implementation / replication of DALL-E, OpenAI's Text to Image Transformer, in Pytorch

5.6K stars0 likes0 views
WebsiteGitHub

About

Implementation / replication of DALL-E, OpenAI's Text to Image Transformer, in Pytorch

# DALL-E in Pytorch


Released DALLE Models
Web-Hostable DALLE Checkpoints
Yannic Kilcher's video

Implementation / replication of DALL-E (paper), OpenAI's Text to Image Transformer, in Pytorch. It will also contain CLIP for ranking the generations. --- [Quick Start](https://github.com/lucidrains/DALLE-pytorch/wiki) Deep Daze or Big Sleep are great alternatives! For generating video and audio, please see NÜWA ## Appreciation This library could not have been possible without the contributions of janEbert, Clay, robvanvolt, Romain Beaumont, and Alexander! ## Status

- Hannu has managed to train a small 6 layer DALL-E on a dataset of just 2000 landscape images! (2048 visual tokens) - Kobiso, a research engineer from Naver, has trained on the CUB200 dataset here, using full and deepspeed sparse attention - (3/15/21) afiaka87 has managed one epoch using a reversible DALL-E and the dVaE here - TheodoreGalanos has trained on 150k layouts with the following results

- Rom1504 has trained on 50k fashion images with captions with a really small DALL-E (2 layers) for just 24 hours with the following results

- afiaka87 trained for 6 epochs on the same dataset as before thanks to the efficient 16k VQGAN with the following results

Thanks to the amazing "mega b#6696" you can generate from this checkpoint in colab - - (5/2/21) First 1.3B DALL-E from has been trained and released to the public! - (4/8/22) Moving onwards to DALLE-2! ## Install ```bash $ pip install dalle-pytorch ``` ## Usage Train VAE ``` … ``` Train DALL-E with pretrained VAE from above ``` … ``` To prime with a starting crop of an image, simply pass two more arguments ```python img_prime = torch.randn(4, 3, 256, 256) images = dalle.generate_images( text, img = img_prime, num_init_img_tokens = (14 * 32) # you can set the size of the initial crop, defaults to a little less than ~1/2 of the tokens, as done in the paper ) images.shape # (4, 3, 256, 256) ``` You may also want to generate text using DALL-E. For that call this function: ```python text_tokens, texts = dalle.generate_texts(tokenizer, text) ``` ## OpenAI's Pretrained VAE You can also skip the training of the VAE altogether, using the pretrained model released by OpenAI! The wrapper class should take care of downloading and caching the model for you auto-magically. ``` … ``` ## Taming Transformer's Pretrained VQGAN VAE You can also use the pretrained VAE offered by the authors of Taming Transformers! Currently only the VAE with a codebook size of 1024 is offered, with the hope that it may train a little faster than OpenAI's, which has a size of 8192. In contrast to OpenAI's VAE, it also has an extra layer of downsampling, so the image sequence length is 256 instead of 1024 (this will lead to a 16 reduction in training costs, when you do the math). Whether it will generalize as well as the original DALL-E is up to the citizen scientists out there to discover. Update - it works! ```python from dalle_pytorch import VQGanVAE vae = VQGanVAE() # the rest is the same as the above example ``` The default VQGan is the codebook size 1024 one trained on imagenet. If you wish to use a different one, you can use the `vqgan_model_path` and `vqgan_config_path` to pass the .ckpt file and the .yaml file. These options can be used both in train-dalle script or as argument of VQGanVAE class. Other pretrained VQGAN can be found in [taming transformers readme](https://github.com/CompVis/taming-transformers#overview-of-pretrained-models). If you want to train a custom one you can [follow this guide](https://github.com/CompVis/taming-transformers/pull/54) ## Adjust text conditioning strength Recently there has surfaced a new technique for guiding diffusion models without a classifier. The gist of the technique involves randomly dropping out the text condition during training, and at inference time, deriving the rough direction from unconditional to conditional distributions. Katherine Crowson outlined in a tweet how this could work for autoregressive attention models. I have decided to include her idea in this repository for further exploration. One only has to account for two extra keyword arguments on training (`null_cond_prob`) and generation (`cond_scale`). ``` … ``` That's it! ## Ranking the generations Train CLIP ```python import torch from dalle_pytorch import CLIP clip = CLIP( dim_text = 512, dim_image = 512, dim_latent = 512, num_text_tokens = 10000, text_enc_depth = 6, text_seq_len = 256, text_heads = 8, num_visual_tokens = 512, visual_enc_depth = 6, visual_image_size = 256, visual_patch_size = 32, visual_heads = 8 ) text = torch.randint(0, 10000, (4, 256)) images = torch.randn(4, 3, 256, 256) mask = torch.ones_like(text).bool() loss = clip(text, images, text_mask = mask, return_loss = True) loss.backward() ``` To get the similarity scores from your trained Clipper, just do ```python images, scores = dalle.generate_images(text, mask = mask, clip = clip) scores.shape # (2,) images.shape # (2, 3, 256, 256) # do your topk here, in paper they sampled 512 and chose top 32 ``` Or you can just use the official CLIP model to rank the images from DALL-E ## Scaling depth In the blog post, they used 64 layers to achieve their results. I added reversible networks, from the Reformer paper, in order for users to attempt to scale depth at the cost of compute. Reversible networks allow you to scale to any depth at no memory cost, but a little over 2x compute cost (each layer is rerun on the backward pass). Simply set the `reversible` keyword to `True` for the `DALLE` class ```python dalle = DALLE( dim = 1024, vae = vae, num_text_tokens = 10000, text_seq_len = 256, depth = 64, heads = 16, reversible = True # <-- reversible networks https://arxiv.org/abs/2001.04451 ) ``` ## Sparse Attention The blogpost alluded to a mixture of different types of sparse attention, used mainly on the image (while the text presumably had full causal attention). I have done my best to replicate these types of sparse attention, on the scant details released. Primarily, it seems as though they are doing causal axial row / column attention, combined with a causal convolution-like attention. By default `DALLE` will use full attention for all layers, but you can specify the attention type per layer as follows. - `full` full attention - `axial_row` axial attention, along the rows of the image feature map - `axial_col` axial attention, along the columns of the image feature map - `conv_like` convolution-like attention, for the image feature map The sparse attention only applies to the image. Text will always receive full attention, as said in the blogpost. ```python dalle = DALLE( dim = 1024, vae = vae, num_text_tokens = 10000, text_seq_len = 256, depth = 64, heads = 16, reversible = True, attn_types = ('full', 'axial_row', 'axial_col', 'conv_like') # cycles between these four types of attention ) ``` ## Deepspeed Sparse Attention You can also train with Microsoft Deepspeed's Sparse Attention, with any combination of dense and sparse attention that you'd like. However, you will have to endure the installation process. First, you need to install Deepspeed with Sparse Attention ```bash $ sh install_deepspeed.sh ``` Next, you need to install the pip package `triton`. It will need to be a version `< 1.0` because that's what Microsoft used. ```bash $ pip install triton==0.4.2 ``` If both of the above succeeded, now you can train with Sparse Attention! ```python dalle = DALLE( dim = 512, vae = vae, num_text_tokens = 10000, text_seq_len = 256, depth = 64, heads = 8, attn_types = ('full', 'sparse') # interleave sparse and dense attention for 64 layers ) ``` ## Training This section will outline how to train the discrete variational autoencoder as well as the final multi-modal transformer (DALL-E). We are going to use Weights & Biases for all the experiment tracking. (You can also do everything in this section in a Google Colab, link below) Train in Colab ```bash $ pip install wandb ``` Followed by ```bash $ wandb login ``` ### VAE To train the VAE, you just need to run ```python $ python train_vae.py --image_folder /path/to/your/images ``` If you installed everything correctly, a link to the experiments page should show up in your terminal. You can follow your link there and customize your experiment, like the example layout below. You can of course open up the training script at `./train_vae.py`, where you can modify the constants, what is passed to Weights & Biases, or any other tricks you know to make the VAE learn better. Model will be saved periodically to `./vae.pt` In the experiment tracker, you will have to monitor the hard reconstruction, as we are essentially teaching the network to compress images into discrete visual tokens for use in the transformer as a visual vocabulary. Weights and Biases will allow you to monitor the temperature annealing, image reconstructions (encoder and decoder working properly), as well as to watch out for codebook collapse (where the network decides to only use a few tokens out of what you provide it). Once you have trained a decent VAE to your satisfaction, you can move on to the next step with your model weights at `./vae.pt`. ### DALL-E Training ## Training using an Image-Text-Folder Now you just have to invoke the `./train_dalle.py` script, indicating which VAE model you would like to use, as well as the path to your folder if images and text. The dataset I am currently working with contains a fo

Issues· 0 open

View all issuesOpen on GitHub

No open issues yet, or sync has not completed.

> Tags

Pythonartificial-intelligenceattention-mechanismdeep-learningmulti-modal

No comments yet. Be the first to share.

> Details

PublishedAug 1, 2026
UpdatedSep 17, 2026
Category数据库
PricingOpen source

> Related tools

P
PostgreSQL
功能强大的开源关系型数据库
R
Redis
内存数据结构存储,常用作缓存与队列
M
MySQL
广泛使用的开源关系型数据库