jax-diffusers-event/canny_coyo1m
2
1import gradio as gr2import jax3import numpy as np4import jax.numpy as jnp5from flax.jax_utils import replicate6from flax.training.common_utils import shard7from PIL import Image8from diffusers import FlaxStableDiffusionControlNetPipeline, FlaxControlNetModel9import cv210 11def create_key(seed=0):12 return jax.random.PRNGKey(seed)13 14def canny_filter(image):15 gray_image = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY)16 blurred_image = cv2.GaussianBlur(gray_image, (5, 5), 0)17 edges_image = cv2.Canny(blurred_image, 50, 150)18 return edges_image19 20# load control net and stable diffusion v1-521controlnet, controlnet_params = FlaxControlNetModel.from_pretrained(22 "jax-diffusers-event/canny-coyo1m", dtype=jnp.bfloat1623)24pipe, params = FlaxStableDiffusionControlNetPipeline.from_pretrained(25 "runwayml/stable-diffusion-v1-5", controlnet=controlnet, revision="flax", dtype=jnp.bfloat1626)27 28def infer(prompts, negative_prompts, image):29 params["controlnet"] = controlnet_params30 31 num_samples = 1 #jax.device_count()32 rng = create_key(0)33 rng = jax.random.split(rng, jax.device_count())34 im = canny_filter(image)35 canny_image = Image.fromarray(im)36 37 prompt_ids = pipe.prepare_text_inputs([prompts] * num_samples)38 negative_prompt_ids = pipe.prepare_text_inputs([negative_prompts] * num_samples)39 processed_image = pipe.prepare_image_inputs([canny_image] * num_samples)40 41 p_params = replicate(params)42 prompt_ids = shard(prompt_ids)43 negative_prompt_ids = shard(negative_prompt_ids)44 processed_image = shard(processed_image)45 46 output = pipe(47 prompt_ids=prompt_ids,48 image=processed_image,49 params=p_params,50 prng_seed=rng,51 num_inference_steps=50,52 neg_prompt_ids=negative_prompt_ids,53 jit=True,54 ).images55 56 output_images = pipe.numpy_to_pil(np.asarray(output.reshape((num_samples,) + output.shape[-3:])))57 return output_images58 59gr.Interface(infer, inputs=["text", "text", "image"], outputs="gallery").launch()60 