https://github.com/lucidrains/imagen-pytorch Skip to content Sign up * Product + Features + Mobile + Actions + Codespaces + Packages + Security + Code review + Issues + Integrations + GitHub Sponsors + Customer stories * Team * Enterprise * Explore + Explore GitHub + Learn and contribute + Topics + Collections + Trending + Learning Lab + Open source guides + Connect with others + The ReadME Project + Events + Community forum + GitHub Education + GitHub Stars program * Marketplace * Pricing + Plans + Compare plans + Contact Sales + Education [ ] * # In this repository All GitHub | Jump to | * No suggested jump to results * # In this repository All GitHub | Jump to | * # In this user All GitHub | Jump to | * # In this repository All GitHub | Jump to | Sign in Sign up {{ message }} lucidrains / imagen-pytorch Public * Notifications * Fork 36 * Star 1.1k Implementation of Imagen, Google's Text-to-Image Neural Network, in Pytorch License MIT license 1.1k stars 36 forks Star Notifications * Code * Issues 1 * Pull requests 0 * Actions * Projects 0 * Wiki * Security * Insights More * Code * Issues * Pull requests * Actions * Projects * Wiki * Security * Insights This commit does not belong to any branch on this repository, and may belong to a fork outside of the repository. main Switch branches/tags [ ] Branches Tags Could not load branches Nothing to show {{ refName }} default View all branches Could not load tags Nothing to show {{ refName }} default View all tags 1 branch 36 tags Code Latest commit @lucidrains lucidrains put efficient grid attention back on the todo ... df33050 May 26, 2022 put efficient grid attention back on the todo df33050 Git stats * 75 commits Files Permalink Failed to load latest commit information. Type Name Latest commit message Commit time .github/workflows scaffold May 24, 2022 imagen_pytorch fix bug with get_times in noise scheduler May 26, 2022 .gitignore Initial commit May 23, 2022 LICENSE Initial commit May 23, 2022 README.md put efficient grid attention back on the todo May 26, 2022 imagen.png Add files via upload May 23, 2022 setup.py fix bug with get_times in noise scheduler May 26, 2022 View code Imagen - Pytorch (wip) Install Usage Shoutouts Todo Citations README.md [imagen] Imagen - Pytorch (wip) 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. Please join Join us on Discord if you are interested in helping out with the replication with the LAION community Install $ pip install imagen-pytorch Usage import torch from imagen_pytorch import Unet, SRUnet, Imagen # unet for imagen unet1 = Unet( dim = 32, cond_dim = 512, channels = 3, dim_mults = (1, 2, 4, 8), num_resnet_blocks = 3, layer_attns = (False, True, True, True), layer_cross_attns = (False, True, True, True) ) unet2 = SRUnet( dim = 32, cond_dim = 512, channels = 3, dim_mults = (1, 2, 4, 8), num_resnet_blocks = (2, 4, 8, 8) ) # imagen, which contains the unets above (base unet and super resoluting ones) imagen = Imagen( unets = (unet1, unet2), image_sizes = (64, 256), beta_schedules = ('cosine', 'linear'), timesteps = 1000, cond_drop_prob = 0.5 ).cuda() # mock images (get a lot of this) and text encodings from large T5 text_embeds = torch.randn(4, 256, 768).cuda() 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, text_embeds = text_embeds, unet_number = i) loss.backward() # do the above for many many many many steps # now you can sample an image based on the text embeddings from the cascading ddpm images = imagen.sample(texts = [ 'a whale breaching from afar', 'young girl blowing out candles on her birthday cake', 'fireworks with blue and green sparkles' ], cond_scale = 2.) images.shape # (3, 3, 256, 256) 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 import torch from imagen_pytorch import Unet, SRUnet, Imagen, ImagenTrainer # unet for imagen unet1 = Unet( dim = 32, cond_dim = 512, channels = 3, dim_mults = (1, 2, 4, 8), num_resnet_blocks = 3, layer_attns = (False, True, True, True), ) unet2 = SRUnet( dim = 32, cond_dim = 512, channels = 3, dim_mults = (1, 2, 4, 8), num_resnet_blocks = (2, 4, 8, 8) ) # imagen, which contains the unets above (base unet and super resoluting ones) imagen = Imagen( unets = (unet1, unet2), text_encoder_name = 't5-large', image_sizes = (64, 256), beta_schedules = ('cosine', 'linear'), timesteps = 1000, cond_drop_prob = 0.5 ).cuda() # wrap imagen with the trainer class trainer = ImagenTrainer(imagen) # mock images (get a lot of this) and text encodings from large T5 text_embeds = torch.randn(4, 256, 1024).cuda() images = torch.randn(4, 3, 256, 256).cuda() # feed images into imagen, training each unet in the cascade for i in (1, 2): loss = trainer(images, text_embeds = text_embeds, unet_number = i) trainer.update(unet_number = i) # do the above for many many many many steps # now you can sample an image based on the text embeddings from the cascading ddpm images = trainer.sample(texts = [ 'a puppy looking anxiously at a giant donut on the table', 'the milky way galaxy in the style of monet' ], cond_scale = 2.) images.shape # (2, 3, 256, 256) Shoutouts * 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 * Jorge Gomes for helping out with the T5 loading code and advice on the correct T5 version * You? It isn't done yet, chip in if you are a researcher or skilled ML engineer Todo * [*] use huggingface transformers for T5-small text embeddings * [*] add dynamic thresholding * [*] add dynamic thresholding DALLE2 and video-diffusion repository as well * [*] allow for one to set T5-large (and perhaps small factory method to take in any huggingface transformer) * [*] add the lowres noise level with the pseudocode in appendix, and figure out what is this sweep they do at inference time * [*] port over some training code from DALLE2 * [*] need to be able to use a different noise schedule per unet (cosine was used for base, but linear for SR) * [*] just make one master-configurable unet * [*] complete resnet block (biggan inspired? but with groupnorm) - complete self attention * [*] complete conditioning embedding block (and make it completely configurable, whether it be attention, film etc) * [ ] add attention pooling option, in addition to cross attention and film * [ ] figure out if learned variance was used at all, and remove it if it was inconsequential * [ ] switch to continuous timesteps instead of discretized, as it seems that is what they used for all stages - first figure out the linear noise schedule case from the variational ddpm paper https://openreview.net/forum?id=2LdBqxc1Yv * [ ] exercise efficient attention expertise + explore skip layer excitation * [ ] consider using perceiver-resampler from https://github.com/ lucidrains/flamingo-pytorch in place of attention pooling * [ ] add optional cosine decay schedule with warmup, for each unet, to trainer * [ ] try out grid attention Citations @inproceedings{Saharia2022PhotorealisticTD, title = {Photorealistic Text-to-Image Diffusion Models with Deep Language Understanding}, author = {Chitwan Saharia and William Chan and Saurabh Saxena and Lala Li and Jay Whang and Emily L. Denton and Seyed Kamyar Seyed Ghasemipour and Burcu Karagol Ayan and Seyedeh Sara Mahdavi and Raphael Gontijo Lopes and Tim Salimans and Jonathan Ho and David Fleet and Mohammad Norouzi}, year = {2022} } @inproceedings{Tu2022MaxViTMV, title = {MaxViT: Multi-Axis Vision Transformer}, author = {Zhengzhong Tu and Hossein Talebi and Han Zhang and Feng Yang and Peyman Milanfar and Alan Conrad Bovik and Yinxiao Li}, year = {2022}, url = {https://arxiv.org/abs/2204.01697} } About Implementation of Imagen, Google's Text-to-Image Neural Network, in Pytorch Topics deep-learning artificial-intelligence text-to-image Resources Readme License MIT license Stars 1.1k stars Watchers 37 watching Forks 36 forks Releases 36 v0.0.37 Latest May 26, 2022 + 35 releases Packages 0 No packages published Contributors 5 * @lucidrains * @jorgemcgomes * @semitrivial * @ProGamerGov * @Wheest Languages * Python 100.0% * (c) 2022 GitHub, Inc. * Terms * Privacy * Security * Status * Docs * Contact GitHub * Pricing * API * Training * Blog * About You can't perform that action at this time. You signed in with another tab or window. Reload to refresh your session. You signed out in another tab or window. Reload to refresh your session.