Views
No views yet
pip install --upgrade diffusers transformers scipy1import torch
2from diffusers import StableDiffusionPipeline
3model_id = "CompVis/stable-diffusion-v1-4"
4device = "cuda"
5pipe = StableDiffusionPipeline.from_pretrained(model_id, torch_dtype=torch.float16)
6pipe = pipe.to(device)
7prompt = "a photo of an astronaut riding a horse on mars"
8image = pipe(prompt).images[0]
9
10image.save("astronaut_rides_horse.png")1import torch
2pipe = StableDiffusionPipeline.from_pretrained(model_id, torch_dtype=torch.float16)
3pipe = pipe.to(device)
4pipe.enable_attention_slicing()
5prompt = "a photo of an astronaut riding a horse on mars"
6image = pipe(prompt).images[0]
7
8image.save("astronaut_rides_horse.png")from_pretrained:1from diffusers import StableDiffusionPipeline, EulerDiscreteScheduler
2model_id = "CompVis/stable-diffusion-v1-4"
3# Use the Euler scheduler here instead
4scheduler = EulerDiscreteScheduler.from_pretrained(model_id, subfolder="scheduler")
5pipe = StableDiffusionPipeline.from_pretrained(model_id, scheduler=scheduler, torch_dtype=torch.float16)
6pipe = pipe.to("cuda")
7prompt = "a photo of an astronaut riding a horse on mars"
8image = pipe(prompt).images[0]
9
10image.save("astronaut_rides_horse.png")1import jax
2import numpy as np
3from flax.jax_utils import replicate
4from flax.training.common_utils import shard
5from diffusers import FlaxStableDiffusionPipeline
6pipeline, params = FlaxStableDiffusionPipeline.from_pretrained(
7 "CompVis/stable-diffusion-v1-4", revision="flax", dtype=jax.numpy.bfloat16
8)
9prompt = "a photo of an astronaut riding a horse on mars"
10prng_seed = jax.random.PRNGKey(0)
11num_inference_steps = 50
12num_samples = jax.device_count()
13prompt = num_samples * [prompt]
14prompt_ids = pipeline.prepare_inputs(prompt)
15# shard inputs and rng
16params = replicate(params)
17prng_seed = jax.random.split(prng_seed, num_samples)
18prompt_ids = shard(prompt_ids)
19images = pipeline(prompt_ids, params, prng_seed, num_inference_steps, jit=True).images
20images = pipeline.numpy_to_pil(np.asarray(images.reshape((num_samples,) + images.shape[-3:])))FlaxStableDiffusionPipeline in bfloat16 precision instead of the default float32 precision as done above. You can do so by telling diffusers to load the weights from "bf16" branch.1import jax
2import numpy as np
3from flax.jax_utils import replicate
4from flax.training.common_utils import shard
5from diffusers import FlaxStableDiffusionPipeline
6pipeline, params = FlaxStableDiffusionPipeline.from_pretrained(
7 "CompVis/stable-diffusion-v1-4", revision="bf16", dtype=jax.numpy.bfloat16
8)
9prompt = "a photo of an astronaut riding a horse on mars"
10prng_seed = jax.random.PRNGKey(0)
11num_inference_steps = 50
12num_samples = jax.device_count()
13prompt = num_samples * [prompt]
14prompt_ids = pipeline.prepare_inputs(prompt)
15# shard inputs and rng
16params = replicate(params)
17prng_seed = jax.random.split(prng_seed, num_samples)
18prompt_ids = shard(prompt_ids)
19images = pipeline(prompt_ids, params, prng_seed, num_inference_steps, jit=True).images
20images = pipeline.numpy_to_pil(np.asarray(images.reshape((num_samples,) + images.shape[-3:])))