CoolFace
Apppublic

jax-diffusers-event/canny_coyo1m

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
2likes
app.py60 linesDownload Raw Back to root
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