Spaces:
Running on Zero
Running on Zero
Download app.py from OzzyGT/krea2-reference: direct link, hf CLI and curl.
- Browser
- Download file 13.7 kB
-
https://huggingface.co/spaces/OzzyGT/krea2-reference/resolve/main/app.py
- Command line
-
hf download hf://spaces/OzzyGT/krea2-reference/app.py
-
curl -L -o app.py https://huggingface.co/spaces/OzzyGT/krea2-reference/resolve/main/app.py
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: | |
| 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) | |
| 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 | |
| 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() | |