multimodalart HF Staff commited on
Commit
95f85fc
·
verified ·
1 Parent(s): 2e16561

Upload folder using huggingface_hub

Browse files
Files changed (3) hide show
  1. README.md +8 -9
  2. app.py +374 -0
  3. requirements.txt +5 -0
README.md CHANGED
@@ -1,13 +1,12 @@
1
  ---
2
- title: Kroma Krea2 Lora Demo
3
- emoji: 🐨
4
- colorFrom: purple
5
- colorTo: yellow
6
  sdk: gradio
7
  sdk_version: 6.22.0
8
- python_version: '3.12'
9
  app_file: app.py
10
- pinned: false
11
- ---
12
-
13
- Check out the configuration reference at https://huggingface.co/docs/hub/spaces-config-reference
 
1
  ---
2
+ title: Kroma Krea 2 LoRA
3
+ emoji: 🎨
4
+ colorFrom: red
5
+ colorTo: indigo
6
  sdk: gradio
7
  sdk_version: 6.22.0
 
8
  app_file: app.py
9
+ short_description: Kroma style LoRA for Krea 2 Turbo text-to-image
10
+ python_version: "3.12"
11
+ startup_duration_timeout: 30m
12
+ ---
app.py ADDED
@@ -0,0 +1,374 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Kroma v0.1 — LoRA for Krea 2 Turbo.
2
+
3
+ Loads the Kroma LoRA (which includes both rank-256 LoRA adapters and fused
4
+ ``.diff`` weight deltas for RMSNorm / modulation tensors) on top of the
5
+ Krea-2-Turbo base pipeline and runs text-to-image inference.
6
+ """
7
+
8
+ import os
9
+
10
+ os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
11
+
12
+ import re
13
+
14
+ import spaces # MUST be before torch / diffusers
15
+ import torch
16
+ import gradio as gr
17
+ import numpy as np
18
+ from safetensors.torch import load_file
19
+ from huggingface_hub import hf_hub_download
20
+ from diffusers import Krea2Pipeline
21
+
22
+ BASE_MODEL_ID = "krea/Krea-2-Turbo"
23
+ LORA_REPO_ID = "lodestones/Kroma"
24
+ LORA_FILENAME = "kroma-v0.1.safetensors"
25
+ MAX_SEED = np.iinfo(np.int32).max
26
+ MAX_IMAGE_SIZE = 1536
27
+
28
+
29
+ # ---------------------------------------------------------------------------
30
+ # Custom LoRA + .diff loading
31
+ # ---------------------------------------------------------------------------
32
+ # The Kroma safetensors contains two kinds of tensors:
33
+ # 1. Standard LoRA adapters (lora_A / lora_B) — ComfyUI key naming, no ".weight" suffix
34
+ # 2. ".diff" weight deltas for RMSNorm scales and modulation layers
35
+ #
36
+ # diffusers' built-in _convert_non_diffusers_krea2_lora_to_diffusers() expects
37
+ # ".lora_A.weight" / ".lora_B.weight" suffixes and raises ValueError on any
38
+ # remaining (non-LoRA) keys. We therefore:
39
+ # - Separate LoRA keys from .diff keys
40
+ # - Convert LoRA keys to diffusers format (add .weight suffix, remap modules)
41
+ # - Load LoRA adapters via pipe.load_lora_weights()
42
+ # - Manually add .diff deltas to the corresponding model parameters
43
+
44
+ _ATTN_MAP = {"wq": "to_q", "wk": "to_k", "wv": "to_v", "wo": "to_out.0", "gate": "to_gate"}
45
+ _FF_MAP = {"gate": "ff.gate", "up": "ff.up", "down": "ff.down"}
46
+ _STANDALONE_MAP = {
47
+ "first": "img_in",
48
+ "last.linear": "final_layer.linear",
49
+ "tmlp.0": "time_embed.linear_1",
50
+ "tmlp.2": "time_embed.linear_2",
51
+ "tproj.1": "time_mod_proj",
52
+ "txtmlp.1": "txt_in.linear_1",
53
+ "txtmlp.3": "txt_in.linear_2",
54
+ "txtfusion.projector": "text_fusion.projector",
55
+ }
56
+
57
+
58
+ def _convert_lora_module(module_path: str) -> str | None:
59
+ """Map a ComfyUI/Krea2 module path to its diffusers equivalent."""
60
+ m = re.match(r"blocks\.(\d+)\.(attn|mlp)\.(\w+)$", module_path)
61
+ if m:
62
+ idx, kind, sub = m.groups()
63
+ if kind == "attn" and sub in _ATTN_MAP:
64
+ return f"transformer_blocks.{idx}.attn.{_ATTN_MAP[sub]}"
65
+ if kind == "mlp" and sub in _FF_MAP:
66
+ return f"transformer_blocks.{idx}.{_FF_MAP[sub]}"
67
+ return None
68
+ m = re.match(r"txtfusion\.(layerwise_blocks|refiner_blocks)\.(\d+)\.(attn|mlp)\.(\w+)$", module_path)
69
+ if m:
70
+ block, idx, kind, sub = m.groups()
71
+ if kind == "attn" and sub in _ATTN_MAP:
72
+ return f"text_fusion.{block}.{idx}.attn.{_ATTN_MAP[sub]}"
73
+ if kind == "mlp" and sub in _FF_MAP:
74
+ return f"text_fusion.{block}.{idx}.{_FF_MAP[sub]}"
75
+ return None
76
+ return _STANDALONE_MAP.get(module_path)
77
+
78
+
79
+ def _convert_diff_module(module_path: str) -> str | None:
80
+ """Map a ComfyUI .diff module path to a dotted attribute path in the transformer."""
81
+ # blocks.N.* -> transformer_blocks.N.*
82
+ m = re.match(r"blocks\.(\d+)\.(.+)$", module_path)
83
+ if m:
84
+ idx, rest = m.groups()
85
+ return f"transformer_blocks.{idx}.{rest}"
86
+
87
+ # txtfusion.* -> text_fusion.*
88
+ m = re.match(r"txtfusion\.(.+)$", module_path)
89
+ if m:
90
+ return f"text_fusion.{m.group(1)}"
91
+
92
+ # last.* -> final_layer.*
93
+ m = re.match(r"last\.(.+)$", module_path)
94
+ if m:
95
+ rest = m.group(1)
96
+ if rest == "norm.scale":
97
+ return "final_layer.norm.weight"
98
+ if rest == "modulation.lin":
99
+ return "final_layer.scale_shift_table"
100
+ return f"final_layer.{rest}"
101
+
102
+ # txtmlp.N -> txt_in.linear_{N//2+1} (0,1->linear_1 2,3->linear_2)
103
+ # Actually txtmlp.0 is txt_in.norm, txtmlp.1 is linear_1, txtmlp.3 is linear_2
104
+ # But .diff for txtmlp is a scale.diff -> txt_in.norm.weight
105
+ m = re.match(r"txtmlp\.(\d+)\.scale$", module_path)
106
+ if m:
107
+ return "txt_in.norm.weight"
108
+
109
+ return None
110
+
111
+
112
+ def load_kroma_lora(pipe: Krea2Pipeline, lora_path: str, strength: float = 1.0):
113
+ """Load Kroma LoRA adapters + .diff deltas into the pipeline."""
114
+ state_dict = load_file(lora_path)
115
+
116
+ # Strip the "diffusion_model." prefix
117
+ state_dict = {
118
+ k.removeprefix("diffusion_model."): v for k, v in state_dict.items()
119
+ }
120
+
121
+ lora_state_dict = {}
122
+ diff_state_dict = {}
123
+
124
+ for key, value in state_dict.items():
125
+ # LoRA keys: end with .lora_A or .lora_B (no .weight suffix)
126
+ m = re.match(r"^(.+)\.(lora_[AB])$", key)
127
+ if m:
128
+ module_path, lora_type = m.groups()
129
+ diffusers_module = _convert_lora_module(module_path)
130
+ if diffusers_module is None:
131
+ print(f" [Kroma LoRA] Skipping unmapped LoRA key: {key}")
132
+ continue
133
+ new_key = f"transformer.{diffusers_module}.{lora_type}.weight"
134
+ lora_state_dict[new_key] = value
135
+ continue
136
+
137
+ # .diff keys
138
+ m = re.match(r"^(.+)\.diff$", key)
139
+ if m:
140
+ module_path = m.group(1)
141
+ diffusers_path = _convert_diff_module(module_path)
142
+ if diffusers_path is None:
143
+ print(f" [Kroma LoRA] Skipping unmapped .diff key: {key}")
144
+ continue
145
+ diff_state_dict[diffusers_path] = value
146
+ continue
147
+
148
+ print(f" [Kroma LoRA] Skipping unknown key: {key}")
149
+
150
+ # Load LoRA adapters via diffusers' standard loader
151
+ print(f" [Kroma LoRA] Loading {len(lora_state_dict)} LoRA tensors")
152
+ pipe.load_lora_weights(lora_state_dict, adapter_name="kroma")
153
+ pipe.set_adapters(["kroma"], adapter_weights=[strength])
154
+
155
+ # Apply .diff deltas manually by adding to existing parameters
156
+ transformer = pipe.transformer
157
+ print(f" [Kroma LoRA] Applying {len(diff_state_dict)} .diff deltas")
158
+ for path, delta in diff_state_dict.items():
159
+ # Navigate the dotted path to the parameter
160
+ obj = transformer
161
+ parts = path.split(".")
162
+ for part in parts[:-1]:
163
+ idx = int(part) if part.isdigit() else None
164
+ obj = obj[idx] if idx is not None else getattr(obj, part)
165
+ param_name = parts[-1]
166
+ # For numeric indices at the end (e.g. to_out.0)
167
+ if param_name.isdigit():
168
+ idx = int(param_name)
169
+ obj = obj[idx]
170
+ param_name = "weight"
171
+
172
+ param = getattr(obj, param_name)
173
+ if param.shape != delta.shape:
174
+ print(
175
+ f" [Kroma LoRA] Shape mismatch for {path}: "
176
+ f"param={param.shape} vs delta={delta.shape}"
177
+ )
178
+ continue
179
+ # Add the delta (scaled by strength) to the parameter
180
+ param.data.add_(delta.to(param.dtype) * strength)
181
+ print(f" [Kroma LoRA] Applied .diff to {path} ({param.shape})")
182
+
183
+
184
+ # ---------------------------------------------------------------------------
185
+ # Model loading
186
+ # ---------------------------------------------------------------------------
187
+ print(f"Loading base pipeline from {BASE_MODEL_ID}...")
188
+ pipe = Krea2Pipeline.from_pretrained(BASE_MODEL_ID, torch_dtype=torch.bfloat16)
189
+ pipe.to("cuda")
190
+
191
+ # Download and load the Kroma LoRA
192
+ lora_path = hf_hub_download(LORA_REPO_ID, LORA_FILENAME)
193
+ print(f"Loading Kroma LoRA from {lora_path}...")
194
+ load_kroma_lora(pipe, lora_path, strength=1.0)
195
+ print("Kroma LoRA loaded successfully.")
196
+
197
+
198
+ # ---------------------------------------------------------------------------
199
+ # Inference
200
+ # ---------------------------------------------------------------------------
201
+ @spaces.GPU(duration=60)
202
+ def generate(
203
+ prompt: str,
204
+ negative_prompt: str = "",
205
+ seed: int = 0,
206
+ randomize_seed: bool = True,
207
+ width: int = 1024,
208
+ height: int = 1024,
209
+ num_inference_steps: int = 8,
210
+ guidance_scale: float = 0.0,
211
+ lora_scale: float = 1.0,
212
+ progress=gr.Progress(track_tqdm=True),
213
+ ):
214
+ """Generate an image from a text prompt using Krea 2 Turbo + Kroma LoRA.
215
+
216
+ Args:
217
+ prompt: What to generate.
218
+ negative_prompt: What to avoid generating.
219
+ seed: RNG seed for reproducibility.
220
+ randomize_seed: If True, pick a random seed each run.
221
+ width: Output image width in pixels.
222
+ height: Output image height in pixels.
223
+ num_inference_steps: Denoising steps (8 for Turbo).
224
+ guidance_scale: Classifier-free guidance (0.0 for Turbo).
225
+ lora_scale: Kroma LoRA strength (1.0 = full effect).
226
+ """
227
+ import random
228
+
229
+ if randomize_seed:
230
+ seed = random.randint(0, MAX_SEED)
231
+ seed = int(seed)
232
+
233
+ generator = torch.Generator(device="cuda").manual_seed(seed)
234
+
235
+ # Adjust LoRA strength at runtime
236
+ pipe.set_adapters(["kroma"], adapter_weights=[lora_scale])
237
+
238
+ image = pipe(
239
+ prompt=prompt,
240
+ negative_prompt=negative_prompt if negative_prompt else None,
241
+ width=width,
242
+ height=height,
243
+ num_inference_steps=num_inference_steps,
244
+ guidance_scale=guidance_scale,
245
+ generator=generator,
246
+ ).images[0]
247
+
248
+ return image, seed
249
+
250
+
251
+ # ---------------------------------------------------------------------------
252
+ # UI
253
+ # ---------------------------------------------------------------------------
254
+ examples = [
255
+ "A weathered fisherman mending nets on a misty dock at dawn, warm lantern light",
256
+ "A close-up portrait of a street violinist in the rain, neon reflections in puddles",
257
+ "Sun-drenched Mediterranean village alley with laundry lines and stray cats",
258
+ ]
259
+
260
+ css = """
261
+ #col-container { max-width: 1100px; margin: 0 auto; }
262
+ .dark .gradio-container { color: var(--body-text-color); }
263
+ """
264
+
265
+ with gr.Blocks(theme=gr.themes.Citrus(), css=css) as demo:
266
+ with gr.Column(elem_id="col-container"):
267
+ gr.Markdown(
268
+ """
269
+ # Kroma v0.1 — LoRA for Krea 2 Turbo
270
+ A style LoRA fine-tune for [Krea 2](https://huggingface.co/krea/Krea-2-Turbo)
271
+ by [lodestones](https://huggingface.co/lodestones/Kroma).
272
+ Runs on ZeroGPU with 8-step Turbo sampling.
273
+ """
274
+ )
275
+
276
+ with gr.Row():
277
+ prompt = gr.Text(
278
+ label="Prompt",
279
+ show_label=False,
280
+ max_lines=1,
281
+ placeholder="Describe the image you want to generate",
282
+ container=False,
283
+ scale=4,
284
+ )
285
+ run_button = gr.Button("Generate", variant="primary", scale=1)
286
+
287
+ result = gr.Image(label="Result", show_label=False)
288
+
289
+ with gr.Accordion("Advanced settings", open=False):
290
+ negative_prompt = gr.Text(
291
+ label="Negative prompt",
292
+ max_lines=1,
293
+ placeholder="Enter a negative prompt (optional)",
294
+ )
295
+
296
+ with gr.Row():
297
+ seed = gr.Slider(
298
+ label="Seed",
299
+ minimum=0,
300
+ maximum=MAX_SEED,
301
+ step=1,
302
+ value=0,
303
+ )
304
+ randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
305
+
306
+ with gr.Row():
307
+ width = gr.Slider(
308
+ label="Width",
309
+ minimum=512,
310
+ maximum=MAX_IMAGE_SIZE,
311
+ step=64,
312
+ value=1024,
313
+ )
314
+ height = gr.Slider(
315
+ label="Height",
316
+ minimum=512,
317
+ maximum=MAX_IMAGE_SIZE,
318
+ step=64,
319
+ value=1024,
320
+ )
321
+
322
+ with gr.Row():
323
+ guidance_scale = gr.Slider(
324
+ label="Guidance scale (CFG)",
325
+ minimum=0.0,
326
+ maximum=10.0,
327
+ step=0.1,
328
+ value=0.0,
329
+ )
330
+ num_inference_steps = gr.Slider(
331
+ label="Inference steps",
332
+ minimum=1,
333
+ maximum=28,
334
+ step=1,
335
+ value=8,
336
+ )
337
+
338
+ lora_scale = gr.Slider(
339
+ label="LoRA strength (Kroma)",
340
+ minimum=0.0,
341
+ maximum=1.5,
342
+ step=0.05,
343
+ value=1.0,
344
+ )
345
+
346
+ gr.Examples(
347
+ examples=examples,
348
+ inputs=[prompt],
349
+ outputs=[result, seed],
350
+ fn=generate,
351
+ cache_examples=True,
352
+ cache_mode="lazy",
353
+ )
354
+
355
+ gr.on(
356
+ triggers=[run_button.click, prompt.submit],
357
+ fn=generate,
358
+ inputs=[
359
+ prompt,
360
+ negative_prompt,
361
+ seed,
362
+ randomize_seed,
363
+ width,
364
+ height,
365
+ num_inference_steps,
366
+ guidance_scale,
367
+ lora_scale,
368
+ ],
369
+ outputs=[result, seed],
370
+ api_name="generate",
371
+ )
372
+
373
+ if __name__ == "__main__":
374
+ demo.launch(mcp_server=True)
requirements.txt ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ diffusers @ git+https://github.com/huggingface/diffusers.git
2
+ transformers
3
+ accelerate
4
+ safetensors
5
+ sentencepiece