| print("importing gc") |
| import gc |
| print("gc imported") |
| print("importing os") |
| import os |
| print("os imported") |
| print("importing torch") |
| import torch |
| print("torch imported") |
| print("importing pillow") |
| from PIL import Image |
| print("imported pillow") |
| print("importing clip") |
| from transformers import CLIPTokenizer, CLIPTextModel |
| print("clip imported") |
| print("importing diffusers") |
| from diffusers import UNet2DConditionModel, DDPMScheduler |
| print("imported diffusers") |
| print("importing safetensors") |
| from safetensors.torch import load_file |
| print("safetensors imported") |
|
|
| |
| |
| |
|
|
| MODEL_PATH = "model.safetensors" |
| CLIP_NAME = "openai/clip-vit-base-patch32" |
|
|
| DEVICE = torch.device("cpu") |
|
|
| IMAGE_SIZE = 128 |
|
|
|
|
| |
| |
| |
|
|
| def free_ram(): |
| gc.collect() |
|
|
| |
| try: |
| torch._C._malloc_trim(0) |
| except Exception: |
| pass |
|
|
|
|
| |
| |
| |
|
|
| def create_unet(): |
|
|
| return 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 |
| ) |
|
|
|
|
| |
| |
| |
|
|
| def encode_prompt(prompt): |
|
|
| print("Loading CLIP tokenizer...") |
|
|
| tokenizer = CLIPTokenizer.from_pretrained( |
| CLIP_NAME |
| ) |
|
|
| print("Loading CLIP...") |
|
|
| text_encoder = CLIPTextModel.from_pretrained( |
| CLIP_NAME |
| ) |
|
|
| text_encoder.eval() |
|
|
| |
| |
| |
| |
| |
| |
|
|
| print("Quantizing CLIP Linear layers to INT8...") |
|
|
| text_encoder = torch.ao.quantization.quantize_dynamic( |
| text_encoder, |
| { |
| torch.nn.Linear |
| }, |
| dtype=torch.qint8 |
| ) |
|
|
| tokens = tokenizer( |
| prompt, |
| padding="max_length", |
| max_length=77, |
| truncation=True, |
| return_tensors="pt" |
| ) |
|
|
| print("Running CLIP...") |
|
|
| with torch.no_grad(): |
|
|
| text = text_encoder( |
| **tokens |
| ).last_hidden_state |
|
|
| |
| |
| |
|
|
| del text_encoder |
| del tokenizer |
| del tokens |
|
|
| free_ram() |
|
|
| print("CLIP unloaded.") |
|
|
| return text |
|
|
|
|
| |
| |
| |
|
|
| def load_unet(): |
|
|
| print("Creating UNet...") |
|
|
| unet = create_unet() |
|
|
| print("Loading UNet weights...") |
|
|
| state = load_file( |
| MODEL_PATH, |
| device="cpu" |
| ) |
|
|
| unet.load_state_dict( |
| state |
| ) |
|
|
| del state |
|
|
| free_ram() |
|
|
| unet.eval() |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| print("Quantizing UNet Linear layers to INT8...") |
|
|
| unet = torch.ao.quantization.quantize_dynamic( |
| unet, |
| { |
| torch.nn.Linear |
| }, |
| dtype=torch.qint8 |
| ) |
|
|
| free_ram() |
|
|
| print("UNet loaded.") |
|
|
| return unet |
|
|
|
|
| |
| |
| |
|
|
| def generate( |
| prompt, |
| steps=250, |
| seed=0 |
| ): |
|
|
| |
| |
| |
|
|
| if seed == 0: |
|
|
| seed = torch.randint( |
| 0, |
| 2**32 - 1, |
| (1,), |
| dtype=torch.int64 |
| ).item() |
|
|
| seed = int(seed) |
|
|
| print( |
| f"Seed: {seed}" |
| ) |
|
|
| |
| |
| |
| |
|
|
| print() |
| print("========== CLIP ==========") |
|
|
| text = encode_prompt( |
| prompt |
| ) |
|
|
| print( |
| f"Text embedding shape: {tuple(text.shape)}" |
| ) |
|
|
| |
| |
| |
| |
| |
|
|
| free_ram() |
|
|
| |
| |
| |
| |
|
|
| print() |
| print("========== UNET ==========") |
|
|
| unet = load_unet() |
|
|
| scheduler = DDPMScheduler( |
| num_train_timesteps=1000, |
| beta_schedule="scaled_linear", |
| prediction_type="epsilon" |
| ) |
|
|
| generator = torch.Generator( |
| device="cpu" |
| ).manual_seed( |
| seed |
| ) |
|
|
| image = torch.randn( |
| (1, 3, IMAGE_SIZE, IMAGE_SIZE), |
| generator=generator, |
| device="cpu" |
| ) |
|
|
| scheduler.set_timesteps( |
| int(steps) |
| ) |
|
|
| |
| |
| |
|
|
| with torch.no_grad(): |
|
|
| for i, t in enumerate( |
| scheduler.timesteps |
| ): |
|
|
| print( |
| f"\rDiffusion step " |
| f"{i + 1}/{steps}", |
| end="", |
| flush=True |
| ) |
|
|
| noise_pred = unet( |
| image, |
| t, |
| encoder_hidden_states=text |
| ).sample |
|
|
| image = scheduler.step( |
| noise_pred, |
| t, |
| image |
| ).prev_sample |
|
|
| print() |
|
|
| |
| |
| |
|
|
| image = ( |
| image |
| .clamp(-1, 1) |
| .add(1) |
| .div(2) |
| ) |
|
|
| image = ( |
| image[0] |
| .permute(1, 2, 0) |
| .cpu() |
| .numpy() |
| ) |
|
|
| image = ( |
| image * 255 |
| ).astype("uint8") |
|
|
| result = Image.fromarray( |
| image |
| ) |
|
|
| |
| |
| |
|
|
| del unet |
| del scheduler |
| del text |
| del image |
| del noise_pred |
|
|
| free_ram() |
|
|
| print("UNet unloaded.") |
|
|
| return result, seed |
|
|
|
|
| |
| |
| |
|
|
| if __name__ == "__main__": |
|
|
| prompt = input( |
| "Pokemon prompt: " |
| ) |
|
|
| steps_input = input( |
| "Steps [250]: " |
| ) |
|
|
| if steps_input.strip(): |
| steps = int(steps_input) |
| else: |
| steps = 250 |
|
|
| seed_input = input( |
| "Seed [0 = random]: " |
| ) |
|
|
| if seed_input.strip(): |
| seed = int(seed_input) |
| else: |
| seed = 0 |
|
|
| result, used_seed = generate( |
| prompt, |
| steps, |
| seed |
| ) |
|
|
| output_path = "generated.png" |
|
|
| result.save( |
| output_path |
| ) |
|
|
| print() |
| print( |
| f"Saved: {output_path}" |
| ) |
|
|
| print( |
| f"Seed: {used_seed}" |
| ) |
|
|