Views
No views yet
1import torch
2from diffusers import UNet2DConditionModel, AutoencoderKL, DDIMScheduler
3from transformers import CLIPTokenizer, CLIPTextModel
4from safetensors.torch import load_file
5
6# Load base models
7vae = AutoencoderKL.from_pretrained("runwayml/stable-diffusion-v1-5", subfolder="vae")
8tokenizer = CLIPTokenizer.from_pretrained("runwayml/stable-diffusion-v1-5", subfolder="tokenizer")
9text_encoder = CLIPTextModel.from_pretrained("runwayml/stable-diffusion-v1-5", subfolder="text_encoder")
10scheduler = DDIMScheduler.from_pretrained("runwayml/stable-diffusion-v1-5", subfolder="scheduler")
11
12# Load fine-tuned UNet
13unet = UNet2DConditionModel.from_pretrained("runwayml/stable-diffusion-v1-5", subfolder="unet")
14unet_state = load_file("unet.safetensors")
15unet.load_state_dict(unet_state)
16
17# Load custom sketch components
18# Note: You'll need the custom SketchCrossAttentionEncoder and SketchTextCombiner classes
19sketch_encoder_state = load_file("sketch_encoder.safetensors")
20sketch_combiner_state = load_file("sketch_text_combiner.safetensors")1# Your sketch should be a 1-channel grayscale image (edges/contours)
2sketch = load_sketch_image("your_sketch.png") # 256x256 grayscale
3prompt = "a red apple"
4
5# Process sketch and text
6sketch_embeddings = sketch_encoder(sketch)
7text_embeddings = text_encoder(tokenize(prompt))
8combined_embeddings = sketch_text_combiner(text_embeddings, sketch_embeddings)
9
10# Generate image
11with torch.no_grad():
12 latents = torch.randn((1, 4, 32, 32)) # 256x256 -> 32x32 latents
13 for t in scheduler.timesteps:
14 noise_pred = unet(latents, t, encoder_hidden_states=combined_embeddings).sample
15 latents = scheduler.step(noise_pred, t, latents).prev_sample
16
17 # Decode to image
18 image = vae.decode(latents / vae.config.scaling_factor).sampleunet.safetensors (3.3GB): Fine-tuned UNet model weightssketch_encoder.safetensors (24MB): Sketch encoder weightssketch_text_combiner.safetensors (16 bytes): Sketch-text combiner weightstraining_info.json: Training metadata1@misc{scribblediffusion-fruit-2024,
2 title={ScribbleDiffusion: Fruit Dataset Fine-tuned Model},
3 author={Your Name},
4 year={2024},
5 howpublished={\\url{https://huggingface.co/your-username/scribblediffusion-fruit}}
6}