wikeeyang's picture
Update app.py
060f27c verified
Raw
History Blame Contribute Delete
11.9 kB
import os
import random
import spaces
import torch
import gradio as gr
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file as load_safetensors
from diffusers import Krea2Pipeline
TURBO_REPO = "unsloth/Krea-2-Turbo"
HD_VAE_REPO = "wikeeyang/Krea2-Turbo-HD-V1"
MAX_SEED = 2**31 - 1
def _remap_vae_keys(state_dict):
"""Remap ComfyUI-format VAE state dict keys to diffusers AutoencoderKLQwenImage keys.
The ComfyUI checkpoint uses a flat sequential naming convention (residual.0,
residual.2, etc.) while diffusers uses semantic names (norm1, conv1, norm2,
conv2). This function translates between the two.
"""
# Mapping from ComfyUI decoder upsamples index to diffusers up_blocks path.
# Derived from the AutoencoderKLQwenImage architecture with dim_mult=[1,2,4,4]
# and num_res_blocks=2: each level has 2 resnets, with upsamplers between levels.
up_map = {
0: "up_blocks.0.resnets.0",
1: "up_blocks.0.resnets.1",
2: "up_blocks.0.resnets.2",
3: "up_blocks.0.upsamplers.0",
4: "up_blocks.1.resnets.0",
5: "up_blocks.1.resnets.1",
6: "up_blocks.1.resnets.2",
7: "up_blocks.1.upsamplers.0",
8: "up_blocks.2.resnets.0",
9: "up_blocks.2.resnets.1",
10: "up_blocks.2.resnets.2",
11: "up_blocks.2.upsamplers.0",
12: "up_blocks.3.resnets.0",
13: "up_blocks.3.resnets.1",
14: "up_blocks.3.resnets.2",
}
def _fix_resnet(rest):
"""Map sequential residual indices to semantic names."""
rest = rest.replace("residual.0.", "norm1.")
rest = rest.replace("residual.2.", "conv1.")
rest = rest.replace("residual.3.", "norm2.")
rest = rest.replace("residual.6.", "conv2.")
rest = rest.replace("shortcut.", "conv_shortcut.")
return rest
def _fix_middle(key, side):
"""Map encoder/decoder middle.X to mid_block structure."""
parts = key.split(".")
idx = int(parts[2])
rest = ".".join(parts[3:])
if idx == 0: # first resnet
return f"{side}.mid_block.resnets.0.{_fix_resnet(rest)}"
elif idx == 1: # attention
return f"{side}.mid_block.attentions.0.{rest}"
elif idx == 2: # second resnet
return f"{side}.mid_block.resnets.1.{_fix_resnet(rest)}"
new_state = {}
for key, val in state_dict.items():
new_key = key
# Top-level convs: quant_conv and post_quant_conv
if key.startswith("conv1."):
new_key = "quant_conv." + key[len("conv1."):]
elif key.startswith("conv2."):
new_key = "post_quant_conv." + key[len("conv2."):]
# Encoder mapping
elif key.startswith("encoder.conv1."):
new_key = "encoder.conv_in." + key[len("encoder.conv1."):]
elif key.startswith("encoder.head.0."):
new_key = "encoder.norm_out." + key[len("encoder.head.0."):]
elif key.startswith("encoder.head.2."):
new_key = "encoder.conv_out." + key[len("encoder.head.2."):]
elif key.startswith("encoder.downsamples."):
parts = key.split(".")
idx = int(parts[2])
rest = ".".join(parts[3:])
rest = _fix_resnet(rest)
new_key = f"encoder.down_blocks.{idx}.{rest}"
elif key.startswith("encoder.middle."):
new_key = _fix_middle(key, "encoder")
# Decoder mapping
elif key.startswith("decoder.conv1."):
new_key = "decoder.conv_in." + key[len("decoder.conv1."):]
elif key.startswith("decoder.head.0."):
new_key = "decoder.norm_out." + key[len("decoder.head.0."):]
elif key.startswith("decoder.head.2."):
new_key = "decoder.conv_out." + key[len("decoder.head.2."):]
elif key.startswith("decoder.upsamples."):
parts = key.split(".")
idx = int(parts[2])
rest = ".".join(parts[3:])
rest = _fix_resnet(rest)
new_key = f"decoder.{up_map[idx]}.{rest}"
elif key.startswith("decoder.middle."):
new_key = _fix_middle(key, "decoder")
new_state[new_key] = val
return new_state
# Load the Krea-2-Turbo pipeline (full diffusers format) then swap in the
# HD-optimized VAE from wikeeyang/Krea2-Turbo-HD-V1 for enhanced detail,
# clarity, and contrast. The HD VAE is a ComfyUI-format single safetensors
# checkpoint; we remap its keys to the diffusers AutoencoderKLQwenImage layout.
pipe = Krea2Pipeline.from_pretrained(TURBO_REPO, torch_dtype=torch.bfloat16)
# Load the HD VAE weights, remap keys from ComfyUI format, and replace the VAE
hd_vae_path = hf_hub_download(HD_VAE_REPO, "Krea2-HD-vae.safetensors")
hd_vae_state = load_safetensors(hd_vae_path)
remapped = _remap_vae_keys(hd_vae_state)
pipe.vae.load_state_dict(remapped, strict=True)
print("HD VAE weights loaded and remapped successfully.")
pipe.to("cuda")
PROMPT_TIPS = """\
Krea 2 Turbo HD is tuned for natural language prompts. Describe the image as you would to a person.
- Write in full sentences or rich phrases. Longer, more specific prompts give the best results.
- Name the things that matter: subject, setting, lighting, color, framing, medium, and mood.
- To render text in the image, wrap the words in quotes, e.g. a neon sign that reads "open late".
- The HD-optimized VAE improves detail rendering, clarity, and contrast over the original.
"""
EXAMPLE_PROMPTS = [
[
"A mesmerizing painting that captures the essence of a dreamy, ethereal punk woman "
"with delicate features. Her short hair frames her introspective face, reflecting the "
"depth of emotion through her soft gaze. The harmonious background, a mix of soft, "
"muted colors, creates a warm and intimate atmosphere. The delicate brushstrokes and "
"vibrant color palette demonstrate the minimalist sophistication, subtle realism, "
"and emotional depth of the artist., Mysterious, vibrant, painting, fashion."
],
[
"This close-up shot captures an Eastern Gray Squirrel perched on a tree branch, "
"its inquisitive gaze directed straight at the viewer. The squirrel's fur is a blend "
"of gray and brown, with hints of reddish-brown on its bushy tail and flanks. Its "
"underparts are a lighter, almost white color, providing a subtle contrast. The fur "
"appears thick and soft, perfectly adapted for the colder climate it inhabits. The "
"squirrel's eyes are large and dark, giving it an alert and curious expression."
],
[
"Create a photorealistic concept poster featuring an elegant timepiece. The watch "
"should be the focal point, displayed in a close-up view. The watch face is a "
"high-end, luxury design, featuring intricate gold hands and markers that are "
"sharply detailed. The reflection on the watch face should depict a gentle, golden "
"sunset over a navy blue ocean. The ocean's surface should have subtle ripples, "
"catching the last light of the day."
],
[
"Super detailed 32K A majestic and ferocious tiger captured mid-leap, its powerful "
"body cutting through a cascade of splashing water, embodying raw energy and primal "
"strength. The tiger's face is at the center of the composition, its intense golden "
"eyes fixed forward, radiating focus and dominance. The intricate black stripes on "
"its fiery orange fur form mesmerizing patterns that emphasize its wild beauty and "
"feral grace. Each strand of fur is meticulously detailed, shimmering with droplets "
"of water that catch the light."
],
]
PLACEHOLDER = (
"Describe your image in natural language. e.g. a russet harvest mouse clinging to a "
'branch, macro photograph, shallow depth of field, creamy green bokeh, soft natural light. '
'Wrap words in "quotes" to render them as text.'
)
CSS = """
#col-container { max-width: 1100px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""
def _duration(prompt, width, height, steps, seed, randomize):
"""Estimate GPU time based on steps and pixel area."""
megapixels = max(1.0, (int(width) * int(height)) / (1024 * 1024))
return int(int(steps) * 2 * megapixels + 25)
@spaces.GPU(duration=_duration, size="xlarge")
def generate(
prompt: str,
width: int = 1024,
height: int = 1024,
steps: int = 8,
seed: int = 42,
randomize: bool = True,
):
"""Generate an image from a text prompt using Krea-2-Turbo with HD VAE.
Args:
prompt: Natural language description of the image to generate.
width: Output image width in pixels.
height: Output image height in pixels.
steps: Number of inference steps (8 for Turbo).
seed: RNG seed for reproducibility.
randomize: If True, pick a random seed.
"""
if not prompt or not prompt.strip():
raise gr.Error("Enter a prompt to generate an image.")
if randomize:
seed = random.randint(0, MAX_SEED)
seed = int(seed)
generator = torch.Generator("cuda").manual_seed(seed)
try:
image = pipe(
prompt=prompt,
height=int(height),
width=int(width),
num_inference_steps=int(steps),
guidance_scale=0.0,
generator=generator,
).images[0]
except RuntimeError as exc:
torch.cuda.empty_cache()
raise gr.Error(
f"Generation failed at {int(width)}x{int(height)}. This is usually GPU "
"out of memory. Try 1024x1024 or a smaller size."
) from exc
return image, seed
with gr.Blocks() as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(
"# Krea 2 Turbo HD\n"
"Text-to-image generation with the HD-optimized VAE from "
"[wikeeyang/Krea2-Turbo-HD-V1](https://huggingface.co/wikeeyang/Krea2-Turbo-HD-V1), "
"built on [krea/Krea-2-Turbo](https://huggingface.co/krea/Krea-2-Turbo). "
"The fine-tuned VAE enhances detail rendering, clarity, and contrast."
)
with gr.Row():
prompt = gr.Textbox(
show_label=False,
placeholder=PLACEHOLDER,
container=False,
scale=4,
lines=3,
autofocus=True,
)
run = gr.Button("Generate", variant="primary", scale=1)
output = gr.Image(label="Result", format="png")
with gr.Accordion("Advanced settings", open=False):
with gr.Row():
width = gr.Slider(512, 2048, value=1024, step=16, label="Width")
height = gr.Slider(512, 2048, value=1024, step=16, label="Height")
steps = gr.Slider(1, 50, value=8, step=1, label="Steps")
with gr.Row():
seed = gr.Slider(0, MAX_SEED, value=42, step=1, label="Seed")
randomize = gr.Checkbox(value=True, label="Randomize seed")
with gr.Accordion("Prompting tips", open=False):
gr.Markdown(PROMPT_TIPS)
gr.Examples(
examples=EXAMPLE_PROMPTS,
inputs=[prompt],
outputs=[output, seed],
fn=generate,
cache_examples=True,
cache_mode="lazy",
label="Example prompts",
)
run.click(
generate,
inputs=[prompt, width, height, steps, seed, randomize],
outputs=[output, seed],
api_name="generate",
)
prompt.submit(
generate,
inputs=[prompt, width, height, steps, seed, randomize],
outputs=[output, seed],
)
demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS)