Master-krea2 / app.py
2i's picture
Use bf16 text encoder + bf16 stock fallback: fp8 weights bypass ZeroGPU packing (comfy_kitchen QuantizedTensor) and produce NaN
48fdf5f
Raw
History Blame Contribute Delete
20.5 kB
"""Krea 2 Turbo text-to-image on Gradio, powered by the ComfyUI backend.
Deploys to Hugging Face Spaces (ZeroGPU). Follows the pattern from:
https://huggingface.co/blog/run-comfyui-workflows-on-spaces
Workflow source: Comfy-Org/workflow_templates image_krea2_turbo_t2i.json
UNet: CivitAI PornMaster-Krea2 (see CIVIT_* env vars below)
Text encoder / VAE / LoRA: Comfy-Org/Krea-2 (not gated)
"""
import os
import random
import subprocess
import sys
from typing import Any, Mapping, Sequence, Union
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import spaces # noqa: E402 # MUST precede torch/comfy imports (no-op off ZeroGPU)
# --------------------------------------------------------------------------
# ComfyUI backend
# --------------------------------------------------------------------------
COMFYUI_PATH = os.environ.get("COMFYUI_PATH", os.path.join(os.getcwd(), "ComfyUI"))
def ensure_comfyui() -> None:
if os.path.isfile(os.path.join(COMFYUI_PATH, "nodes.py")):
return
print("Cloning ComfyUI backend (once)...")
subprocess.run(
["git", "clone", "--depth", "1",
"https://github.com/comfyanonymous/ComfyUI.git", COMFYUI_PATH],
check=True,
)
ensure_comfyui()
if COMFYUI_PATH not in sys.path:
sys.path.insert(0, COMFYUI_PATH)
import comfy.options # noqa: E402
comfy.options.enable_args_parsing()
import numpy as np # noqa: E402
import requests # noqa: E402
import torch # noqa: E402
# Inference-only: kill autograd overhead without making tensors "inference"
# tensors (which ComfyUI's fp8 quantized-path cannot re-register on device moves).
torch.set_grad_enabled(False)
from huggingface_hub import hf_hub_download # noqa: E402
from comfy import model_management # noqa: E402
from nodes import ( # noqa: E402
CLIPLoader,
CLIPTextEncode,
ConditioningZeroOut,
EmptyLatentImage,
KSampler,
LoraLoaderModelOnly,
UNETLoader,
VAEDecode,
VAELoader,
)
# Optional LLM prompt enhancement: reuses the qwen3vl text encoder as an LLM,
# no extra model download needed.
try:
from comfy_extras.nodes_textgen import TextGenerate # noqa: F401
HAS_LLM = True
except Exception as exc: # pragma: no cover
HAS_LLM = False
print(f"LLM prompt enhancement unavailable: {exc}")
# --------------------------------------------------------------------------
# Models
# --------------------------------------------------------------------------
# Companion models always pulled from Comfy-Org/Krea-2 (the CivitAI checkpoint
# is UNet-only): (repo_id, subfolder, filename). local_dir = models/ root, so
# the repo subfolder (text_encoders/vae/loras) is replicated under models/.
#
# IMPORTANT: bf16 files only. fp8 models are comfy_kitchen QuantizedTensor
# objects whose .to("cuda") bypasses ZeroGPU's torch patch, so they are never
# packed/streamed to VRAM and inference produces NaN. bf16 packs fine.
COMPANION_MODELS = [
("Comfy-Org/Krea-2", "text_encoders", "qwen3vl_4b_bf16.safetensors"),
("Comfy-Org/Krea-2", "vae", "qwen_image_vae.safetensors"),
("Comfy-Org/Krea-2", "loras", "krea2_darkbrush.safetensors"),
]
STOCK_UNET = ("Comfy-Org/Krea-2", "diffusion_models", "krea2_turbo_bf16.safetensors")
# CivitAI checkpoint (the diffusion model):
# CIVIT_API_KEY secret on the Space (required)
# CIVIT_MODEL_VERSION model version id, default 3171380 = PornMaster V2.5 Turbo fp8
# CIVIT_MODEL_FP fp8 | bf16 | int8
# V2.5 (3171380) is Early Access on CivitAI (needs Buzz to unlock). If it is not
# unlocked yet, the app automatically falls back to 3112108 (Turbo V2 FP8).
CIVIT_API_KEY = os.environ.get("CIVIT_API_KEY", "")
CIVIT_MODEL_VERSION = os.environ.get("CIVIT_MODEL_VERSION", "3171380") # V2.5 Turbo
CIVIT_FALLBACK_VERSION = os.environ.get("CIVIT_FALLBACK_VERSION", "3112108") # Turbo V2 FP8
CIVIT_MODEL_FP = os.environ.get("CIVIT_MODEL_FP", "fp8")
LORA_TRIGGERS = {
"krea2_darkbrush.safetensors": "monochrome ink wash style",
"krea2_dotmatrix.safetensors": "monochrome stippling style",
"krea2_kidsdrawing.safetensors": "naive expressive sketch style",
"krea2_neondrip.safetensors": "textured abstract style",
"krea2_rainywindow.safetensors": "rainy window style",
"krea2_retroanime.safetensors": "purple retro anime style",
"krea2_softwatercolor.safetensors": "art deco watercolor style",
"krea2_sunsetblur.safetensors": "ethereal motion blur style",
"krea2_vintagetarot.safetensors": "vintage tarot style",
}
# System prompt for LLM prompt enhancement (copied from the official template).
LLM_SYSTEM_PROMPT = (
"You are an expert prompt engineer for text-to-image models. Your task is to expand the user's prompt into a "
"highly effective image-generation prompt.\n\n"
"Think step by step about the request before writing the answer:\n"
"- What is the subject and mood?\n"
"- What visual styles, mediums, and lighting options would fit? Consider two or three alternatives and pick the "
"one that best serves the caption.\n"
"- What composition, framing, and grounded details will help the text-to-image model?\n\n"
"Then output a single expanded prompt paragraph.\n\n"
"Follow these rules strictly:\n"
"1. **Faithfulness First:** Preserve all original subjects, actions, colors, and spatial relationships. Do not "
"add new objects, props, characters, or animals unless the user clearly implies them.\n"
"2. **Practical T2I Structure:** Write a prompt that a text-to-image model can parse cleanly. Group subjects with "
"their own attributes and actions. Use grounded phrasing for poses, interactions, and spatial layout.\n"
"3. **Style Planning Stays Internal:** Use your internal reasoning to choose style, medium, framing, and "
"lighting. Do not emit planning tags or wrappers in the visible answer body.\n"
"4. **Text Rendering:** If the user requests visible text, quotes, labels, or typography, specify the exact text "
"clearly and wrap requested words in quotes.\n"
"5. **Avoid Over-Specification:** Do not invent highly specific clothing, colors, materials, or scene details "
"unless the input supports them.\n"
"6. **Structure:** Write one cohesive paragraph after the thinking block. No bullets, JSON, or markdown.\n"
"7. **Respect Existing Detail:** If the user's prompt is already detailed, lightly polish and finalize rather "
"than heavily expanding, preserve their phrasing and direction.\n"
"8. **Respect the Human Form:** Treat depictions of people with dignity. Assume clothing covers genitals and "
"intimate anatomy.\n"
"9. **Preserve User Medium:** When the user explicitly requests a medium (e.g. \"photo of\", \"photograph of\", "
"\"illustration of\", \"painting of\", \"sketch of\", \"3D render of\"), honor it. Do not pivot to a different "
"medium to avoid difficulty, match the user's stated intent.\n\n"
"User's Input:\n\n"
)
def download_civitai_unet(version_id: str) -> str:
"""Download a CivitAI Krea 2 UNet checkpoint.
Returns the filename placed in ComfyUI/models/diffusion_models/.
"""
dest_dir = os.path.join(COMFYUI_PATH, "models", "diffusion_models")
os.makedirs(dest_dir, exist_ok=True)
headers = {"Authorization": f"Bearer {CIVIT_API_KEY}"} if CIVIT_API_KEY else {}
info = requests.get(
f"https://civitai.com/api/v1/model-versions/{version_id}",
headers=headers, timeout=30,
).json()
files = info.get("files", [])
target = next(
(
f for f in files
if f.get("metadata", {}).get("fp") == CIVIT_MODEL_FP
and f.get("metadata", {}).get("format") == "SafeTensor"
),
None,
)
if target is None:
target = next((f for f in files if f.get("metadata", {}).get("format") == "SafeTensor"), None)
if target is None and files:
target = files[0]
if target is None:
raise RuntimeError(f"CivitAI version {version_id} has no files")
filename = target["name"]
out_path = os.path.join(dest_dir, filename)
if os.path.isfile(out_path) and os.path.getsize(out_path) > 1e9:
print(f"CivitAI UNet already present: {filename}")
return filename
# Multi-file versions need type/format/fp params to pick a variant; try a
# ladder in case a param combo is rejected.
base_url = f"https://civitai.com/api/download/models/{version_id}"
param_ladder = [
{"type": target["type"], "format": "SafeTensor", "fp": CIVIT_MODEL_FP},
{"format": "SafeTensor", "fp": CIVIT_MODEL_FP},
{},
]
tmp = out_path + ".part"
for params in param_ladder:
with requests.get(base_url, params=params, headers=headers, stream=True, timeout=(30, 300)) as resp:
if not resp.ok:
print(f"civitai download attempt {resp.status_code}: {resp.text[:120]}")
continue
total = int(resp.headers.get("content-length", 0))
print(f"Downloading {filename} ({total / 1e9:.2f} GB) from CivitAI...")
with open(tmp, "wb") as fh:
for chunk in resp.iter_content(1 << 20):
fh.write(chunk)
os.replace(tmp, out_path)
return filename
raise RuntimeError(f"CivitAI version {version_id} download failed for all URL variants")
def ensure_models() -> str:
"""Download companion models + the UNet. Returns the UNet filename to load.
CivitAI version ladder: configured version -> fallback version -> stock
Comfy-Org fp8 (so the Space still boots even if CivitAI gates the model).
"""
for repo_id, subfolder, filename in COMPANION_MODELS:
hf_hub_download(
repo_id=repo_id,
subfolder=subfolder,
filename=filename,
local_dir=os.path.join(COMFYUI_PATH, "models"),
)
tried = []
for version_id in (CIVIT_MODEL_VERSION, CIVIT_FALLBACK_VERSION):
try:
return download_civitai_unet(version_id)
except Exception as exc:
tried.append(f"{version_id} ({exc})")
print(f"WARNING: CivitAI version {version_id} unavailable: {exc}")
print(f"WARNING: all CivitAI versions failed {tried}; "
"falling back to stock Krea 2 Turbo fp8 from Comfy-Org.")
hf_hub_download(
repo_id=STOCK_UNET[0], subfolder=STOCK_UNET[1], filename=STOCK_UNET[2],
local_dir=os.path.join(COMFYUI_PATH, "models"),
)
return STOCK_UNET[2]
UNET_NAME = ensure_models()
def get_value_at_index(obj: Union[Sequence, Mapping], index: int) -> Any:
try:
return obj[index]
except KeyError:
return obj["result"][index]
def list_loras() -> list[str]:
lora_dir = os.path.join(COMFYUI_PATH, "models", "loras")
return sorted(f for f in os.listdir(lora_dir) if f.endswith(".safetensors")) if os.path.isdir(lora_dir) else []
# --------------------------------------------------------------------------
# Load models at module scope. On ZeroGPU these weights are packed to disk at
# startup and streamed into VRAM per request, so first call after idle is the
# only slow one.
# --------------------------------------------------------------------------
unet_loader = UNETLoader()
UNET = unet_loader.load_unet(unet_name=UNET_NAME, weight_dtype="default")
clip_loader = CLIPLoader()
CLIP = clip_loader.load_clip(clip_name="qwen3vl_4b_bf16.safetensors", type="krea2")
vae_loader = VAELoader()
VAE = vae_loader.load_vae(vae_name="qwen_image_vae.safetensors")
lora_loader = LoraLoaderModelOnly()
text_encode = CLIPTextEncode()
zero_out = ConditioningZeroOut()
empty_latent = EmptyLatentImage()
sampler = KSampler()
vae_decode = VAEDecode()
model_management.load_models_gpu(
[
getattr(get_value_at_index(UNET, 0), "patcher", get_value_at_index(UNET, 0)),
getattr(get_value_at_index(CLIP, 0), "patcher", get_value_at_index(CLIP, 0)),
getattr(get_value_at_index(VAE, 0), "patcher", get_value_at_index(VAE, 0)),
]
)
# --------------------------------------------------------------------------
# Inference
# --------------------------------------------------------------------------
@spaces.GPU(duration=120) # tune: measure worst-case and multiply by ~1.4
def generate_image(
prompt: str,
width: int,
height: int,
seed: int,
steps: int,
cfg: float,
enable_lora: bool,
lora_name: str,
lora_strength: float,
trigger_word: str,
prompt_enhance: bool,
thinking: bool,
max_tokens: int,
) -> np.ndarray:
"""Generate one Krea 2 Turbo image from a text prompt."""
width = max(256, int(width) // 16 * 16)
height = max(256, int(height) // 16 * 16)
seed = int(seed) if int(seed) >= 0 else random.randint(1, 2**63)
steps = max(1, int(steps))
lora_strength = float(lora_strength)
max_tokens = max(16, int(max_tokens))
# NOTE: no torch.inference_mode() here. ComfyUI's fp8-quantized model
# weights are inference tensors, and _quantized_apply() cannot clone them
# while inference mode is active, which crashes device moves during sampling
# and yields NaN latents (black output).
model = get_value_at_index(UNET, 0)
if enable_lora and lora_name:
model = get_value_at_index(
lora_loader.load_lora_model_only(
model=model, lora_name=lora_name, strength_model=lora_strength
),
0,
)
# Optional LLM prompt enhancement (reuses the qwen3vl text encoder).
final_prompt = prompt
if prompt_enhance and HAS_LLM:
sampling_mode = {
"sampling_mode": "on",
"temperature": 0.7,
"top_k": 64,
"top_p": 0.95,
"min_p": 0.05,
"repetition_penalty": 1.05,
"seed": 0,
"presence_penalty": 0.0,
}
enhanced = TextGenerate.execute(
clip=get_value_at_index(CLIP, 0),
prompt=LLM_SYSTEM_PROMPT + prompt,
max_length=max_tokens,
sampling_mode=sampling_mode,
thinking=thinking,
use_default_template=True,
)
final_prompt = str(enhanced[0]).strip()
if enable_lora and trigger_word:
final_prompt = f"{final_prompt}, {trigger_word}"
# Conditioning (krea2 turbo uses cfg=1, so the negative is zeroed out).
positive = text_encode.encode(text=final_prompt, clip=get_value_at_index(CLIP, 0))
cond_t = get_value_at_index(positive, 0)[0][0]
cond_f = cond_t.float()
print(
f"[debug] cond shape={tuple(cond_t.shape)} mean={cond_f.mean().item():.4f} "
f"abs_mean={cond_f.abs().mean().item():.4f} nan={torch.isnan(cond_f).sum().item()}"
)
negative = zero_out.zero_out(conditioning=get_value_at_index(positive, 0))
latent = empty_latent.generate(width=width, height=height, batch_size=1)
sampled = sampler.sample(
model=model,
seed=seed,
steps=steps,
cfg=cfg,
sampler_name="euler",
scheduler="simple",
positive=get_value_at_index(positive, 0),
negative=get_value_at_index(negative, 0),
latent_image=get_value_at_index(latent, 0),
denoise=1.0,
)
lat = get_value_at_index(sampled, 0)["samples"].float()
print(
f"[debug] latent mean={lat.mean().item():.4f} abs_mean={lat.abs().mean().item():.4f} "
f"min={lat.min().item():.4f} max={lat.max().item():.4f} nan={torch.isnan(lat).sum().item()}"
)
decoded = vae_decode.decode(samples=get_value_at_index(sampled, 0), vae=get_value_at_index(VAE, 0))
image = get_value_at_index(decoded, 0)[0]
print(
f"[debug] image mean={image.float().mean().item():.4f} "
f"min={image.float().min().item():.4f} max={image.float().max().item():.4f} "
f"nan={torch.isnan(image.float()).sum().item()}"
)
img_np = (
torch.nan_to_num(image, nan=0.0, posinf=1.0, neginf=0.0)
.mul(255)
.clamp_(0, 255)
.byte()
.cpu()
.numpy()
)
return img_np
# --------------------------------------------------------------------------
# Gradio UI
# --------------------------------------------------------------------------
import gradio as gr # noqa: E402
RESOLUTIONS = {
"1:1 (1024x1024)": (1024, 1024),
"2:3 (832x1216)": (832, 1216),
"3:2 (1216x832)": (1216, 832),
"3:4 (896x1152)": (896, 1152),
"4:3 (1152x896)": (1152, 896),
"9:16 (768x1344)": (768, 1344),
"16:9 (1344x768)": (1344, 768),
}
LORA_CHOICES = list_loras() or ["krea2_darkbrush.safetensors"]
output_image = gr.Image(label="Generated Image")
with gr.Blocks(title="Krea 2 Turbo") as app:
gr.Markdown("# Krea 2 Turbo")
gr.Markdown(
"Krea 2 Turbo text-to-image running on a Gradio app over the ComfyUI backend "
"(workflow: `image_krea2_turbo_t2i.json`). Turbo: 8 steps, CFG 1."
)
with gr.Row():
with gr.Column(scale=1):
prompt_input = gr.Textbox(label="Prompt", lines=3, placeholder="Describe an image...")
resolution = gr.Dropdown(
label="Resolution preset", choices=list(RESOLUTIONS), value="1:1 (1024x1024)"
)
with gr.Row():
width_input = gr.Number(label="Width", value=1024, precision=0)
height_input = gr.Number(label="Height", value=1024, precision=0)
with gr.Row():
seed_input = gr.Number(label="Seed (-1 = random)", value=-1, precision=0)
steps_input = gr.Slider(label="Steps", minimum=1, maximum=20, value=8, step=1)
cfg_input = gr.Slider(label="CFG", minimum=0.0, maximum=10.0, value=1.0, step=0.1)
with gr.Accordion("Style LoRA", open=False):
lora_enable = gr.Checkbox(label="Enable LoRA", value=False)
lora_name_input = gr.Dropdown(
label="LoRA file", choices=LORA_CHOICES, value=LORA_CHOICES[0]
)
lora_strength_input = gr.Slider(label="LoRA strength", minimum=0.0, maximum=2.0, value=0.8, step=0.05)
trigger_input = gr.Textbox(label="Trigger word (auto-appended)", value=LORA_TRIGGERS.get(LORA_CHOICES[0], ""))
with gr.Accordion("Prompt enhancement (LLM)", open=False):
enhance_enable = gr.Checkbox(
label="Enhance prompt with LLM (uses the qwen3vl text encoder)",
value=False,
interactive=HAS_LLM,
)
thinking_input = gr.Checkbox(label="Thinking mode", value=False)
max_tokens_input = gr.Slider(label="Max tokens", minimum=64, maximum=2048, value=512, step=64)
generate_btn = gr.Button("Generate", variant="primary")
gr.Examples(
examples=[
["a cozy cabin in snowy mountains at dusk, warm window light", 1024, 1024, -1, 8, 1.0, False, LORA_CHOICES[0], 0.8, "monochrome ink wash style", False, False, 512],
["a sleek cyberpunk street in the rain, neon signs", 1024, 1024, -1, 8, 1.0, True, LORA_CHOICES[0], 0.8, "monochrome ink wash style", False, False, 512],
],
inputs=[
prompt_input, width_input, height_input, seed_input, steps_input,
cfg_input, lora_enable, lora_name_input, lora_strength_input, trigger_input,
enhance_enable, thinking_input, max_tokens_input,
],
outputs=[output_image],
fn=generate_image,
cache_examples=True,
cache_mode="lazy",
)
with gr.Column(scale=1):
output_image.render()
resolution.change(
lambda name: list(RESOLUTIONS[name]),
inputs=[resolution],
outputs=[width_input, height_input],
)
lora_name_input.change(
lambda name: LORA_TRIGGERS.get(name, ""),
inputs=[lora_name_input],
outputs=[trigger_input],
)
generate_btn.click(
fn=generate_image,
inputs=[
prompt_input, width_input, height_input, seed_input, steps_input, cfg_input,
lora_enable, lora_name_input, lora_strength_input, trigger_input,
enhance_enable, thinking_input, max_tokens_input,
],
outputs=[output_image],
)
if __name__ == "__main__":
app.launch()