Views
No views yet
1git clone https://github.com/entrpn/diffusers
2cd diffusers
3git checkout lcm_flax
4pip install jax[tpu] -f https://storage.googleapis.com/jax-releases/libtpu_releases.html
5pip install transformers flax torch torchvision
6pip install .1import os
2from diffusers import FlaxStableDiffusionXLPipeline
3import torch
4import time
5import jax
6import jax.numpy as jnp
7from flax.jax_utils import replicate
8import numpy as np
9from jax.experimental.compilation_cache import compilation_cache as cc
10cc.initialize_cache(os.path.expanduser("~/jax_cache"))
11
12from diffusers import (
13 FlaxUNet2DConditionModel,
14 FlaxLCMScheduler
15)
16
17base_model = "stabilityai/stable-diffusion-xl-base-1.0"
18weight_dtype = jnp.bfloat16
19revision= 'refs/pr/95'
20
21pipeline, params = FlaxStableDiffusionXLPipeline.from_pretrained(
22 base_model, revision=revision, dtype=weight_dtype
23 )
24
25del params["unet"]
26
27unet, unet_params = FlaxUNet2DConditionModel.from_pretrained(
28 "jffacevedo/flax_lcm_unet",
29 dtype=weight_dtype,
30)
31
32scheduler, scheduler_state = FlaxLCMScheduler.from_pretrained(
33 base_model,
34 subfolder="scheduler",
35 revision=revision,
36 dtype=jnp.float32
37)
38
39params["unet"] = unet_params
40pipeline.unet = unet
41
42pipeline.scheduler = scheduler
43
44params = jax.tree_util.tree_map(lambda x: x.astype(weight_dtype), params)
45params["scheduler"] = scheduler_state
46
47default_prompt = "high-quality photo of a baby dolphin playing in a pool and wearing a party hat"
48default_neg_prompt = ""
49default_seed = 42
50default_guidance_scale = 1.0
51default_num_steps = 4
52
53def tokenize_prompt(prompt, neg_prompt):
54 prompt_ids = pipeline.prepare_inputs(prompt)
55 neg_prompt_ids = pipeline.prepare_inputs(neg_prompt)
56 return prompt_ids, neg_prompt_ids
57
58NUM_DEVICES = jax.device_count()
59
60p_params = replicate(params)
61
62def replicate_all(prompt_ids, neg_prompt_ids, seed):
63 p_prompt_ids = replicate(prompt_ids)
64 p_neg_prompt_ids = replicate(neg_prompt_ids)
65 rng = jax.random.PRNGKey(seed)
66 rng = jax.random.split(rng, NUM_DEVICES)
67 return p_prompt_ids, p_neg_prompt_ids, rng
68
69def generate(
70 prompt,
71 negative_prompt,
72 seed=default_seed,
73 guidance_scale=default_guidance_scale,
74 num_inference_steps=default_num_steps,
75):
76 prompt_ids, neg_prompt_ids = tokenize_prompt(prompt, negative_prompt)
77 prompt_ids, neg_prompt_ids, rng = replicate_all(prompt_ids, neg_prompt_ids, seed)
78 images = pipeline(
79 prompt_ids,
80 p_params,
81 rng,
82 num_inference_steps=num_inference_steps,
83 guidance_scale=guidance_scale,
84 do_classifier_free_guidance=False,
85 jit=True,
86 ).images
87 print("images.shape: ", images.shape)
88 # convert the images to PIL
89 images = images.reshape((images.shape[0] * images.shape[1], ) + images.shape[-3:])
90 return pipeline.numpy_to_pil(np.array(images))
91
92start = time.time()
93print(f"Compiling ...")
94generate(default_prompt, default_neg_prompt)
95print(f"Compiled in {time.time() - start}")
96
97dts = []
98i = 0
99for x in range(2):
100
101 start = time.time()
102 prompt = "Self-portrait oil painting, a beautiful cyborg with golden hair, 8k"
103 neg_prompt = ""
104
105 print(f"Prompt: {prompt}")
106 images = generate(prompt, neg_prompt)
107 t = time.time() - start
108 print(f"Inference in {t}")
109
110 dts.append(t)
111 for img in images:
112 img.save(f'{i:06d}.jpg')
113 i += 1
114
115mean = np.mean(dts)
116stdev = np.std(dts)
117print(f"batches: {i}, Mean {mean:.2f} sec/batch± {stdev * 1.96 / np.sqrt(len(dts)):.2f} (95%)")