krea2-reference / app.py
OzzyGT's picture
OzzyGT HF Staff
simpler names
db98262
Raw History Blame Contribute Delete
13.7 kB
import gc
import os
import sys
IS_ZERO_GPU = os.environ.get("SPACES_ZERO_GPU") is not None
# `spaces` must be imported before any CUDA-initializing package, so this block stays above the
# torch/diffusers imports below.
if IS_ZERO_GPU:
import spaces
else:
class _NoSpaces:
@staticmethod
def GPU(*args, **kwargs):
if len(args) == 1 and callable(args[0]) and not kwargs:
return args[0]
def deco(fn):
return fn
return deco
spaces = _NoSpaces()
import gradio as gr # noqa: E402
import numpy as np # noqa: E402
import torch # noqa: E402
from PIL import Image, ImageChops # noqa: E402
from diffusers.modular_pipelines import ( # noqa: E402
ComponentsManager,
ModularPipelineBlocks,
)
BLOCK_REPO = os.environ.get("KREA2_BLOCKS", "OzzyGT/krea2_reference_blocks")
MODELS = {
"Krea 2 Turbo — SDNQ int8 (~18GB)": {
"repo": "OzzyGT/Krea_2_Turbo_sdnq_dynamic_8bit",
"distilled": True,
"steps": 8,
},
"Krea 2 Turbo — SDNQ int4 (~10GB)": {
"repo": "OzzyGT/Krea_2_Turbo_sdnq_dynamic_4bit",
"distilled": True,
"steps": 8,
},
}
DEFAULT_MODEL = "Krea 2 Turbo — SDNQ int8 (~18GB)"
MODE_LORAS = {
"off": ("", ""),
"append": (
"ostris/krea2_turbo_style_reference",
"krea2_style_reference.safetensors",
),
"prepend": (
"conradlocke/krea2-identity-edit",
"krea2_identity_edit_v1_2.safetensors",
),
}
COMPONENTS = ["text_encoder", "tokenizer", "transformer", "vae", "scheduler"]
MAX_SEQUENCE_LENGTH = 512
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
# Saturated on purpose: the composite fallback in `_split_editor` diffs strokes against the photo under them.
BRUSH_COLOR = "#ff00ff"
MODE_LABELS = {
"Vision only": "off",
"Style reference": "append",
"Identity edit": "prepend",
}
MASK_MODE_LABELS = {
"Exclude, blank tokens": "exclude_blank",
"Exclude, keep tokens": "exclude",
"De-emphasize only": "deemphasize",
}
# ----------------------------------------------------------------------------- pipeline lifecycle
# One pipeline resident at a time; switching the checkpoint reloads.
_current = {"id": None, "pipe": None, "lora": None, "distilled": False}
def _resolve(choice: str) -> str:
"""Map a dropdown label to its repo id; anything else is used verbatim."""
preset = MODELS.get(choice)
return preset["repo"] if preset else choice
def _load(checkpoint: str, distilled: bool):
import sdnq # noqa: F401 registers the SDNQ backend with transformers so quantized checkpoints load
# `auto_map` names the distilled blockset, so the base one is pulled off the loaded module by name.
blocks = ModularPipelineBlocks.from_pretrained(BLOCK_REPO, trust_remote_code=True)
blocks_name = "Krea2TurboReferenceAutoBlocks" if distilled else "Krea2ReferenceAutoBlocks"
blocks = getattr(sys.modules[type(blocks).__module__], blocks_name)()
manager = ComponentsManager()
# The block repo's modular index supplies the per-component subfolders; `load_components` overrides
# which repo they load from.
pipe = blocks.init_pipeline(BLOCK_REPO, components_manager=manager)
pipe.load_components(names=COMPONENTS, pretrained_model_name_or_path=checkpoint, dtype=torch.bfloat16)
if IS_ZERO_GPU:
# ZeroGPU hands out a large GPU for the duration of the call, so no offloading is needed.
pipe.to(DEVICE)
elif DEVICE == "cuda":
manager.enable_auto_cpu_offload(device=DEVICE)
return pipe
def get_pipe(choice: str):
checkpoint = _resolve(choice)
preset = MODELS.get(choice)
distilled = bool(preset["distilled"]) if preset else _current.get("distilled", False)
if _current["id"] != (checkpoint, distilled):
_current.update(id=None, pipe=None, lora=None)
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
_current.update(
id=(checkpoint, distilled),
distilled=distilled,
pipe=_load(checkpoint, distilled),
)
print(f"loaded {checkpoint}")
return _current["pipe"]
# On ZeroGPU the platform owns the GPU between calls, so preloading costs nothing. Locally nothing loads
# until the Load button, so this Space does not squat on VRAM while another app is running.
if IS_ZERO_GPU:
get_pipe(DEFAULT_MODEL)
@spaces.GPU(duration=120)
def load_model(choice):
get_pipe(choice)
return f"✅ Loaded: {_resolve(choice)}"
def unload_model():
_current.update(id=None, pipe=None, lora=None)
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
return "⚪ No model loaded"
def _sync_lora(pipe, repo: str, weight_name: str):
"""Load, swap or drop the edit LoRA."""
repo, weight_name = (repo or "").strip(), (weight_name or "").strip()
key = (repo, weight_name)
if _current["lora"] == key:
return
try:
pipe.unload_lora_weights()
except Exception:
pass
if repo:
pipe.load_lora_weights(repo, **({"weight_name": weight_name} if weight_name else {}))
_current["lora"] = key
# ----------------------------------------------------------------------------- reference
def _split_editor(value):
"""Pull `(reference, mask)` out of an ImageEditor value.
The painted layer's alpha channel is the mask, in the blocks' polarity: painted is the part of the
reference to use. A mask of `None` means "no mask", not "use nothing".
"""
if not value:
return None, None
background = value.get("background") if isinstance(value, dict) else None
if background is None:
return None, None
mask = None
for layer in value.get("layers") or []:
if layer is None:
continue
alpha = Image.fromarray(np.array(layer.convert("RGBA"))[:, :, 3], mode="L")
if alpha.getextrema()[1] == 0:
continue
mask = alpha if mask is None else ImageChops.lighter(mask, alpha)
# Fallback for an editor build that flattens strokes into the composite instead of keeping a layer:
# the painted region is then wherever the composite differs from the background.
composite = value.get("composite")
if mask is None and composite is not None and composite.size == background.size:
diff = ImageChops.difference(composite.convert("RGB"), background.convert("RGB")).convert("L")
if diff.getextrema()[1] > 8:
mask = diff.point(lambda v: 255 if v > 8 else 0)
return background.convert("RGB"), mask
class _StepProgress:
"""Stands in for the loop block's tqdm so only the denoise steps reach the Gradio bar."""
def __init__(self, progress, total):
self._progress = progress
self._total = max(1, int(total or 1))
self._step = 0
def __enter__(self):
self._progress(0.0, desc=f"Step 0/{self._total}")
return self
def __exit__(self, *exc):
return False
def update(self, n=1):
self._step = min(self._step + int(n), self._total)
self._progress(self._step / self._total, desc=f"Step {self._step}/{self._total}")
def _report_steps(blocks, progress):
"""Swap `progress_bar` on every denoise loop block for one that drives `progress`.
Modular pipelines expose no per-step callback, but `progress_bar` is a plain method on
`LoopSequentialPipelineBlocks`, so an instance attribute shadows it.
"""
for block in (blocks.sub_blocks or {}).values():
if hasattr(type(block), "progress_bar"):
block.progress_bar = lambda iterable=None, total=None, _p=progress: _StepProgress(_p, total)
_report_steps(block, progress)
# ----------------------------------------------------------------------------- generate
@spaces.GPU(duration=180)
def generate(
model_choice,
prompt,
mode_label,
mask_mode_label,
mask_reference_latents,
reference,
subject,
style,
grounding_px,
width,
height,
steps,
seed,
randomize_seed,
progress=gr.Progress(),
):
progress(0.0, desc="Loading model")
pipe = get_pipe(model_choice)
# `pipe.blocks` hands back a deepcopy; the pipeline runs `_blocks`, so the shim has to go there.
_report_steps(pipe._blocks, progress)
mode = MODE_LABELS[mode_label]
# The LoRA follows the reference mode: each latent mode only works with the adapter it was trained for.
_sync_lora(pipe, *MODE_LORAS[mode])
image, mask = _split_editor(reference)
if mode != "off" and image is None:
raise gr.Error(f"`{mode}` needs a reference image; add one or switch to vision-only.")
if randomize_seed:
seed = int(torch.randint(0, 2**31 - 1, (1,)).item())
generator = torch.Generator("cpu").manual_seed(int(seed))
kwargs = {
"prompt": prompt or "",
"height": int(height),
"width": int(width),
"num_inference_steps": int(steps),
"max_sequence_length": MAX_SEQUENCE_LENGTH,
"generator": generator,
"output_type": "pil",
}
if image is not None:
kwargs.update(
reference_images=image,
reference_mode=mode,
reference_subject_strength=float(subject),
reference_style_strength=float(style),
grounding_px=int(grounding_px),
mask_reference_latents=bool(mask_reference_latents),
)
if mask is not None:
kwargs["reference_masks"] = mask
kwargs["reference_mask_mode"] = MASK_MODE_LABELS[mask_mode_label]
images = pipe(**kwargs).values["images"]
return images[0], int(seed)
# ----------------------------------------------------------------------------- UI
with gr.Blocks(title="Krea 2 — reference conditioning") as demo:
gr.Markdown(
"# Krea 2 — reference-image conditioning\n"
"Vision-path references on the stock checkpoint, plus the two clean-latent edit-LoRA modes, via the "
f"[`{BLOCK_REPO}`](https://huggingface.co/{BLOCK_REPO}) modular blocks."
)
with gr.Row():
with gr.Column(scale=1):
prompt = gr.Textbox(label="Prompt", lines=3, placeholder="a photograph of ...")
with gr.Row():
mode = gr.Dropdown(
choices=list(MODE_LABELS),
value="Vision only",
label="Reference mode",
)
mask_mode = gr.Dropdown(
choices=list(MASK_MODE_LABELS),
value="Exclude, blank tokens",
label="Mask behaviour",
)
reference = gr.ImageEditor(
label="Reference — paint the part to use",
type="pil",
layers=False,
brush=gr.Brush(colors=[BRUSH_COLOR], default_color=BRUSH_COLOR, color_mode="fixed"),
height=360,
)
with gr.Row():
subject = gr.Slider(0.0, 1.5, value=1.0, step=0.05, label="Subject")
style = gr.Slider(0.0, 1.5, value=1.0, step=0.05, label="Style")
mask_reference_latents = gr.Checkbox(value=False, label="Mask the reference latents")
with gr.Accordion("Advanced", open=False):
with gr.Row():
width = gr.Slider(512, 2048, value=1024, step=32, label="Width")
height = gr.Slider(512, 2048, value=1024, step=32, label="Height")
with gr.Row():
steps = gr.Slider(
1,
60,
value=MODELS[DEFAULT_MODEL]["steps"],
step=1,
label="Steps",
)
grounding_px = gr.Slider(0, 1536, value=768, step=64, label="Grounding px (prepend)")
with gr.Row():
seed = gr.Number(value=0, precision=0, label="Seed")
randomize_seed = gr.Checkbox(value=True, label="Randomize")
with gr.Column(scale=1):
model_choice = gr.Dropdown(choices=list(MODELS), value=DEFAULT_MODEL, label="Model")
with gr.Row(visible=not IS_ZERO_GPU):
load_btn = gr.Button("Load model")
unload_btn = gr.Button("Unload")
status = gr.Markdown("🟢 Ready (ZeroGPU)" if IS_ZERO_GPU else "⚪ No model loaded")
run = gr.Button("Generate", variant="primary")
gallery = gr.Image(label="Result", type="pil", format="png", interactive=False)
used_seed = gr.Number(label="Seed used", precision=0, interactive=False)
def on_model_change(choice):
"""A preset carries its own schedule; a custom path leaves whatever is set alone."""
preset = MODELS.get(choice)
return gr.update() if preset is None else gr.update(value=preset["steps"])
model_choice.change(on_model_change, inputs=[model_choice], outputs=[steps])
load_btn.click(load_model, inputs=[model_choice], outputs=[status])
unload_btn.click(unload_model, outputs=[status])
run.click(
generate,
inputs=[
model_choice,
prompt,
mode,
mask_mode,
mask_reference_latents,
reference,
subject,
style,
grounding_px,
width,
height,
steps,
seed,
randomize_seed,
],
outputs=[gallery, used_seed],
)
if __name__ == "__main__":
demo.queue().launch()