Diffusers
Safetensors
clamp-checkpoints / diffusion /inference.py
nexh98's picture
Add CLAMP classification and diffusion checkpoints
5746967
Raw History Blame Contribute Delete
1.74 kB
"""
Minimal inference example for the CLAMP diffusion checkpoint (clamp_ESD).
Requires: diffusers, transformers, torch, accelerate, safetensors
pip install diffusers transformers torch accelerate safetensors
Note: pipe.load_lora_weights() and unet.load_lora_adapter() do NOT work with this
checkpoint (it was saved via the older UNet-level attn-procs format, not the newer
pipeline-level LoRA format, so the key prefixes don't match what those methods
expect - they will silently attach zero LoRA parameters). Use
unet.load_attn_procs() instead, as below - it is deprecated in newer diffusers
versions but is the one that actually works for this file.
"""
import torch
from diffusers import StableDiffusionPipeline
from huggingface_hub import snapshot_download
REPO_ID = "nexh98/clamp-checkpoints"
# Base model: Stable Diffusion V1-4 with the paper's ESD erasure already applied
# to the sub-concepts. Using this (rather than vanilla SD v1.4) is required to
# reproduce the paper's reported harmful/benign adaptation numbers.
base_model_dir = snapshot_download(REPO_ID, allow_patterns="diffusion/base_model/*")
lora_dir = snapshot_download(REPO_ID, allow_patterns="diffusion/clamp_ESD/*")
base_model_path = f"{base_model_dir}/diffusion/base_model"
lora_path = f"{lora_dir}/diffusion/clamp_ESD"
pipe = StableDiffusionPipeline.from_pretrained(
base_model_path, torch_dtype=torch.float16, safety_checker=None
)
pipe.unet.load_attn_procs(lora_path, weight_name="pytorch_lora_weights.safetensors")
pipe = pipe.to("cuda")
image = pipe(
"A picture of a Kavri", # one of the paper's merged/harmful concepts (elephant+zebra)
num_inference_steps=25,
guidance_scale=7.5,
).images[0]
image.save("clamp_diffusion_sample.png")