| import torch | |
| from PIL import Image | |
| from transformers import ( | |
| CLIPTokenizer, | |
| CLIPTextModel | |
| ) | |
| from diffusers import ( | |
| UNet2DConditionModel, | |
| DDPMScheduler | |
| ) | |
| from ema_pytorch import EMA | |
| # ========================== | |
| # CONFIG | |
| # ========================== | |
| CHECKPOINT = "model.pt" | |
| DEVICE = "cuda" if torch.cuda.is_available() else "cpu" | |
| PROMPT = "a pokemon named eevee" | |
| STEPS = 1000 | |
| # ========================== | |
| # TEXT | |
| # ========================== | |
| tokenizer = CLIPTokenizer.from_pretrained( | |
| "openai/clip-vit-base-patch32" | |
| ) | |
| text_encoder = CLIPTextModel.from_pretrained( | |
| "openai/clip-vit-base-patch32" | |
| ) | |
| text_encoder.to(DEVICE) | |
| text_encoder.eval() | |
| # ========================== | |
| # MODEL | |
| # ========================== | |
| unet = UNet2DConditionModel( | |
| sample_size=128, | |
| in_channels=3, | |
| out_channels=3, | |
| layers_per_block=2, | |
| block_out_channels=( | |
| 128, | |
| 256, | |
| 512, | |
| 512, | |
| ), | |
| down_block_types=( | |
| "DownBlock2D", | |
| "DownBlock2D", | |
| "CrossAttnDownBlock2D", | |
| "DownBlock2D" | |
| ), | |
| up_block_types=( | |
| "UpBlock2D", | |
| "CrossAttnUpBlock2D", | |
| "UpBlock2D", | |
| "UpBlock2D" | |
| ), | |
| cross_attention_dim=512 | |
| ) | |
| # ========================== | |
| # LOAD EMA | |
| # ========================== | |
| checkpoint = torch.load( | |
| CHECKPOINT, | |
| map_location=DEVICE | |
| ) | |
| unet.load_state_dict( | |
| checkpoint["model"] | |
| ) | |
| ema = EMA( | |
| unet, | |
| beta=0.9999 | |
| ) | |
| ema.load_state_dict( | |
| checkpoint["ema"] | |
| ) | |
| # use EMA weights | |
| ema.ema_model.to(DEVICE) | |
| unet = ema.ema_model | |
| unet.to(DEVICE) | |
| unet.eval() | |
| # ========================== | |
| # SCHEDULER | |
| # ========================== | |
| scheduler = DDPMScheduler( | |
| num_train_timesteps=1000, | |
| beta_schedule="scaled_linear", | |
| prediction_type="epsilon" | |
| ) | |
| # ========================== | |
| # TEXT EMBEDDING | |
| # ========================== | |
| tokens = tokenizer( | |
| PROMPT, | |
| padding="max_length", | |
| max_length=77, | |
| truncation=True, | |
| return_tensors="pt" | |
| ) | |
| tokens = { | |
| k:v.to(DEVICE) | |
| for k,v in tokens.items() | |
| } | |
| with torch.no_grad(): | |
| text = text_encoder( | |
| **tokens | |
| ).last_hidden_state | |
| # ========================== | |
| # GENERATE | |
| # ========================== | |
| image = torch.randn( | |
| (1,3,128,128), | |
| device=DEVICE | |
| ) | |
| scheduler.set_timesteps( | |
| STEPS | |
| ) | |
| with torch.no_grad(): | |
| for t in scheduler.timesteps: | |
| noise_pred = unet( | |
| image, | |
| t, | |
| encoder_hidden_states=text | |
| ).sample | |
| image = scheduler.step( | |
| noise_pred, | |
| t, | |
| image | |
| ).prev_sample | |
| # ========================== | |
| # SAVE | |
| # ========================== | |
| image = ( | |
| image | |
| .clamp(-1,1) | |
| .add(1) | |
| .div(2) | |
| ) | |
| image = ( | |
| image[0] | |
| .permute(1,2,0) | |
| .cpu() | |
| .numpy() | |
| ) | |
| image = ( | |
| image * 255 | |
| ).astype("uint8") | |
| Image.fromarray( | |
| image | |
| ).save( | |
| "pokemon_test.png" | |
| ) | |
| print( | |
| "Saved pokemon_test.png" | |
| ) | |