apolinario commited on
Commit
82e465e
·
1 Parent(s): b26b44f

Krea 2 LoRA trainer (HF Jobs backend)

Browse files
README.md CHANGED
@@ -1,13 +1,56 @@
1
  ---
2
- title: Krea2 Lora Trainer
3
- emoji: 👁
4
  colorFrom: indigo
5
  colorTo: yellow
6
  sdk: gradio
7
  sdk_version: 6.19.0
8
- python_version: '3.13'
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: Krea 2 LoRA Trainer
3
+ emoji: 🎨
4
  colorFrom: indigo
5
  colorTo: yellow
6
  sdk: gradio
7
  sdk_version: 6.19.0
8
+ python_version: '3.12'
9
  app_file: app.py
10
+ hardware: cpu-basic
11
+ pinned: true
12
+ hf_oauth: true
13
+ hf_oauth_scopes:
14
+ - read-repos
15
+ - write-repos
16
+ - manage-repos
17
+ - jobs
18
+ short_description: Train Krea 2 LoRAs on your images via HF Jobs
19
  ---
20
 
21
+ # Krea 2 LoRA Trainer
22
+
23
+ Train a **DreamBooth-LoRA for Krea 2** from your own images, entirely on Hugging Face
24
+ infrastructure:
25
+
26
+ - **Sign in with Hugging Face** — the dataset, the job, and the pushed LoRA all run under
27
+ **your** account and billing (no pasted tokens);
28
+ - the **Space** (this app, `cpu-basic`) collects your images + hyperparameters and submits a job;
29
+ - training runs on **HF Jobs** using the diffusers Krea 2 trainer
30
+ (`examples/dreambooth/train_dreambooth_lora_krea2.py`);
31
+ - the LoRA is **trained on Krea 2 RAW** and **validated / inferred on Krea 2 Turbo**, then pushed
32
+ to the Hub model repo you choose.
33
+
34
+ You only pay for the Job's actual GPU runtime.
35
+
36
+ ## How tokens are used
37
+
38
+ The Krea 2 weights are currently **gated**. The job downloads them with the Space's `KREA_TOKEN`
39
+ secret and passes them to the trainer as **local dirs** — so your own token never needs Krea
40
+ access, and the Krea token never touches your repos. Your token (from sign-in) is used only for
41
+ your dataset and the pushed LoRA.
42
+
43
+ > Set the `KREA_TOKEN` secret to a token with access to `krea/Krea-2-Raw` + `krea/Krea-2-Turbo`.
44
+
45
+ ## diffusers version
46
+
47
+ The trainer lives in diffusers PR #14046 (branch `krea2-lora`). Once it is merged, set the
48
+ `DIFFUSERS_REF` Space **variable** to `main` (or a release tag).
49
+
50
+ ## Usage
51
+
52
+ 1. Sign in with Hugging Face.
53
+ 2. Name your LoRA, set a trigger word / concept, and upload 4–30 images.
54
+ 3. (Optional) caption each image; blanks fall back to the trigger.
55
+ 4. Tweak hyperparameters if you like, pick a GPU flavor, and **Submit training job**.
56
+ 5. Copy the job id into the **Monitor** tab and **Refresh** to stream logs.
__pycache__/app.cpython-312.pyc ADDED
Binary file (16.6 kB). View file
 
__pycache__/jobs.cpython-312.pyc ADDED
Binary file (13.8 kB). View file
 
app.py ADDED
@@ -0,0 +1,270 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Krea 2 LoRA Trainer — HF Space (HF Jobs backend).
2
+
3
+ Sign in with Hugging Face, upload a handful of images (4–30 is ideal), optionally caption them,
4
+ set a trigger word, and submit a DreamBooth-LoRA training job to HF Jobs. The job trains on
5
+ **Krea 2 RAW**, validates / infers on **Krea 2 Turbo**, and pushes the LoRA to your Hub — all
6
+ under your account. The Space runs on `cpu-basic`; the GPU work happens on HF Jobs.
7
+
8
+ The gated Krea 2 weights are downloaded inside the job with the Space's `KREA_TOKEN` secret;
9
+ the user's token is only ever used for their own dataset + the pushed LoRA.
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ import re
15
+
16
+ import gradio as gr
17
+
18
+ import jobs
19
+
20
+ MAX_IMAGES = 40
21
+ MAX_LOG = 60_000
22
+
23
+ LR_SCHEDULERS = ["constant", "cosine", "linear", "constant_with_warmup", "polynomial"]
24
+ OPTIMIZERS = ["adamW", "prodigy"]
25
+ QUANT_CHOICES = [
26
+ ("None — bf16 (best quality, most VRAM)", "none"),
27
+ ("FP8 (faster compute, needs GPU ≥ 8.9)", "fp8"),
28
+ ("4-bit NF4 / QLoRA (lowest VRAM)", "4bit"),
29
+ ]
30
+
31
+ FLAVOR_GUIDE = """**Which GPU?** You're billed per-minute of actual runtime. Krea 2 is a 12B DiT.
32
+
33
+ | Flavor | VRAM | Best for |
34
+ |---|---|---|
35
+ | `l40sx1` | 48 GB | cheapest — pair with **4-bit NF4** quantization |
36
+ | `a100-large` | 80 GB | **recommended (default)** — bf16 with offload + cached latents |
37
+ | `h200` | 141 GB | fastest / highest resolution |
38
+
39
+ First run is slow to start: it downloads the gated Krea 2 RAW + Turbo weights before training.
40
+ """
41
+
42
+ TRIGGER_HELP = (
43
+ "A **trigger** that anchors your concept and is used as the default caption for every image. "
44
+ "For an object/character a rare token like `TOK` works; for a **style**, a descriptive phrase "
45
+ "(e.g. *“hand-drawn children's book illustration”*) works better than a random token."
46
+ )
47
+ CAPTION_HELP = (
48
+ "<details><summary>ℹ️ <b>Captioning tips</b></summary>\n\n"
49
+ "- Captions are <b>optional</b> — blank ones fall back to your trigger.\n"
50
+ "- For a <b>style</b> LoRA, describe what you do <i>not</i> want baked in (subject, scene) and "
51
+ "<i>omit</i> the stylistic parts you want learned, then keep the style trigger phrase.\n"
52
+ "- For an <b>object/character</b>, a trigger word plus the right class noun is enough.\n"
53
+ "</details>"
54
+ )
55
+
56
+
57
+ def _signin_state(profile: gr.OAuthProfile | None):
58
+ if profile is None:
59
+ return (
60
+ "🔒 **You're not signed in.** Use **Sign in with Hugging Face** (top-right) — the dataset, "
61
+ "the training job, and the resulting LoRA all run under **your** account.",
62
+ gr.update(),
63
+ )
64
+ return (
65
+ f"✅ Signed in as **{profile.username}** — the job and the pushed LoRA live under your account.",
66
+ gr.update(placeholder=f"e.g. {profile.username}/my-krea2-lora"),
67
+ )
68
+
69
+
70
+ def load_captioning(images, instance_prompt):
71
+ """Reveal one (image, caption) row per uploaded image; prefill captions with the trigger."""
72
+ n = len(images) if images else 0
73
+ if n > MAX_IMAGES:
74
+ raise gr.Error(f"For now, up to {MAX_IMAGES} images are supported (got {n}).")
75
+ updates = [gr.update(visible=n > 0)] # captioning_area
76
+ for i in range(MAX_IMAGES):
77
+ visible = i < n
78
+ updates.append(gr.update(visible=visible)) # row
79
+ updates.append(gr.update(value=images[i] if visible else None, visible=visible)) # image
80
+ updates.append(gr.update(value=(instance_prompt or "") if visible else None, visible=visible))
81
+ return updates
82
+
83
+
84
+ def gather_dataset(images, *captions):
85
+ """Pair uploaded image paths with their caption textbox values into a list of [path, caption]."""
86
+ images = images or []
87
+ return [[img, (captions[i] if i < len(captions) else "")] for i, img in enumerate(images)]
88
+
89
+
90
+ def start_training(
91
+ dataset_rows, lora_name, instance_prompt, validation_prompt, rank, lora_alpha, max_train_steps,
92
+ learning_rate, lr_scheduler, resolution, repeats, train_batch_size, gradient_accumulation_steps,
93
+ seed, optimizer, use_8bit_adam, cache_latents, gradient_checkpointing, offload, quantization,
94
+ lora_layers, validation_epochs, hub_model_id, flavor, timeout,
95
+ profile: gr.OAuthProfile | None = None, oauth_token: gr.OAuthToken | None = None,
96
+ ):
97
+ if oauth_token is None or profile is None:
98
+ return "❌ Please **sign in with Hugging Face** first (top-right).", "", ""
99
+ if not dataset_rows:
100
+ return "❌ Upload at least one image.", "", ""
101
+ if not (lora_name or "").strip() and not (hub_model_id or "").strip():
102
+ return "❌ Give your LoRA a name (or a Hub model id).", "", ""
103
+
104
+ image_paths = [r[0] for r in dataset_rows]
105
+ captions = [r[1] for r in dataset_rows]
106
+ params = {
107
+ "lora_name": lora_name, "hub_model_id": hub_model_id,
108
+ "instance_prompt": instance_prompt, "validation_prompt": validation_prompt,
109
+ "rank": rank, "lora_alpha": lora_alpha, "max_train_steps": max_train_steps,
110
+ "learning_rate": learning_rate, "lr_scheduler": lr_scheduler, "resolution": resolution,
111
+ "repeats": repeats, "train_batch_size": train_batch_size,
112
+ "gradient_accumulation_steps": gradient_accumulation_steps, "seed": seed,
113
+ "optimizer": optimizer, "use_8bit_adam": bool(use_8bit_adam),
114
+ "cache_latents": bool(cache_latents), "gradient_checkpointing": bool(gradient_checkpointing),
115
+ "offload": bool(offload), "quantization": quantization, "lora_layers": lora_layers,
116
+ "validation_epochs": validation_epochs, "hf_token": oauth_token.token,
117
+ }
118
+ try:
119
+ res = jobs.submit(params, image_paths, captions, flavor=flavor, timeout=timeout)
120
+ except Exception as e: # noqa: BLE001
121
+ return f"❌ Submission failed: {e}", "", ""
122
+
123
+ status = f"✅ Job submitted on **{flavor}**, running as **{profile.username}**."
124
+ link = (
125
+ f"**Job:** [{res['job_id']}]({res['url']}) \n"
126
+ f"**Dataset:** `{res['dataset_repo']}` \n"
127
+ f"**LoRA will be pushed to:** `{res['hub_model_id']}`"
128
+ )
129
+ return status, link, res["job_id"]
130
+
131
+
132
+ def refresh(job_id, oauth_token: gr.OAuthToken | None = None):
133
+ if not (job_id or "").strip():
134
+ return "Enter a job id.", ""
135
+ token = oauth_token.token if oauth_token else ""
136
+ st = jobs.job_status(job_id.strip(), token)
137
+ logs = jobs.job_logs(job_id.strip(), token)
138
+ return f"**Status:** `{st}`", logs[-MAX_LOG:] if len(logs) > MAX_LOG else logs
139
+
140
+
141
+ with gr.Blocks(title="Krea 2 LoRA Trainer") as demo:
142
+ with gr.Row(equal_height=True):
143
+ gr.Markdown(
144
+ "# 🎨 Krea 2 LoRA Trainer\n"
145
+ "Train a LoRA on your own images — trains on **Krea 2 RAW**, validates on **Turbo**, "
146
+ "runs on **HF Jobs**, pushed to your Hub.",
147
+ )
148
+ gr.LoginButton(scale=0, min_width=220)
149
+
150
+ banner = gr.Markdown()
151
+
152
+ with gr.Tabs():
153
+ with gr.Tab("Train"):
154
+ lora_name = gr.Textbox(
155
+ label="LoRA name", placeholder="e.g. my-watercolor-style",
156
+ info="Used for your output model repo (you/<name>) and the dataset repo.",
157
+ )
158
+ instance_prompt = gr.Textbox(
159
+ label="Trigger word / concept", value="TOK", info=TRIGGER_HELP,
160
+ )
161
+ images = gr.File(
162
+ label="Upload your images (4–30 ideal)", file_count="multiple",
163
+ file_types=["image"], height=220,
164
+ )
165
+
166
+ with gr.Column(visible=False) as captioning_area:
167
+ gr.Markdown("**Captions** (optional) — edit per image, or leave as the trigger.")
168
+ gr.Markdown(CAPTION_HELP)
169
+ caption_rows, caption_imgs, caption_txts = [], [], []
170
+ for i in range(MAX_IMAGES):
171
+ with gr.Row(visible=False) as row:
172
+ img = gr.Image(
173
+ show_label=False, interactive=False, height=90, width=90, scale=0,
174
+ )
175
+ cap = gr.Textbox(show_label=False, scale=4, container=False,
176
+ placeholder="caption for this image")
177
+ caption_rows.append(row)
178
+ caption_imgs.append(img)
179
+ caption_txts.append(cap)
180
+
181
+ with gr.Accordion("Advanced options", open=False):
182
+ with gr.Row():
183
+ rank = gr.Number(label="LoRA rank", value=32, precision=0,
184
+ info="Authors recommend 32; raise it for long runs / high-frequency styles.")
185
+ lora_alpha = gr.Number(label="LoRA alpha", value=32, precision=0,
186
+ info="Keep equal to rank (scale 1.0).")
187
+ with gr.Row():
188
+ max_train_steps = gr.Number(label="Training steps", value=1000, precision=0)
189
+ learning_rate = gr.Number(label="Learning rate", value=3e-4,
190
+ info="3e-4 ~ 7e-4 with constant works well.")
191
+ with gr.Row():
192
+ lr_scheduler = gr.Dropdown(LR_SCHEDULERS, value="constant", label="LR scheduler")
193
+ resolution = gr.Number(label="Resolution", value=1024, precision=0)
194
+ with gr.Row():
195
+ repeats = gr.Number(label="Dataset repeats", value=1, precision=0)
196
+ seed = gr.Number(label="Seed", value=0, precision=0)
197
+ lora_layers = gr.Textbox(
198
+ label="Target layers (optional)", placeholder="to_q,to_k,to_v,to_out.0,to_gate",
199
+ info="Comma-separated. Blank = the recommended full set. For long runs, narrow to "
200
+ "attention and raise the rank.",
201
+ )
202
+ with gr.Accordion("Memory / performance", open=False):
203
+ with gr.Row():
204
+ quantization = gr.Dropdown(QUANT_CHOICES, value="none", label="Quantization")
205
+ optimizer = gr.Dropdown(OPTIMIZERS, value="adamW", label="Optimizer")
206
+ with gr.Row():
207
+ use_8bit_adam = gr.Checkbox(label="8-bit Adam", value=True)
208
+ cache_latents = gr.Checkbox(label="Cache latents", value=True)
209
+ with gr.Row():
210
+ gradient_checkpointing = gr.Checkbox(label="Gradient checkpointing", value=True)
211
+ offload = gr.Checkbox(label="CPU offload (VAE + text encoder)", value=False)
212
+ with gr.Row():
213
+ train_batch_size = gr.Number(label="Batch size", value=1, precision=0)
214
+ gradient_accumulation_steps = gr.Number(label="Grad accumulation", value=1, precision=0)
215
+ with gr.Accordion("Validation", open=False):
216
+ validation_prompt = gr.Textbox(
217
+ label="Validation prompt", placeholder="(defaults to your trigger)",
218
+ info="Generated on Turbo every N epochs to preview progress.",
219
+ )
220
+ validation_epochs = gr.Number(label="Validate every N epochs", value=25, precision=0)
221
+
222
+ with gr.Group():
223
+ hub_model_id = gr.Textbox(
224
+ label="Output Hub model id (optional)", placeholder="you/my-krea2-lora",
225
+ info="Where the trained LoRA is pushed. Blank = you/<lora-name>.",
226
+ )
227
+ with gr.Row():
228
+ flavor = gr.Dropdown(jobs.FLAVORS, value=jobs.DEFAULT_FLAVOR, label="GPU flavor")
229
+ timeout = gr.Textbox(label="Timeout", value="3h",
230
+ info="Max job runtime. First run downloads the model first.")
231
+ with gr.Accordion("GPU guide", open=False):
232
+ gr.Markdown(FLAVOR_GUIDE)
233
+
234
+ submit_btn = gr.Button("🚀 Submit training job", variant="primary", size="lg")
235
+ status = gr.Markdown()
236
+ joblink = gr.Markdown()
237
+
238
+ with gr.Tab("Monitor"):
239
+ with gr.Row():
240
+ job_id = gr.Textbox(label="Job id", scale=3)
241
+ refresh_btn = gr.Button("🔄 Refresh", scale=1)
242
+ mon_status = gr.Markdown()
243
+ mon_logs = gr.Textbox(label="Job logs", lines=22, autoscroll=True, max_lines=22)
244
+
245
+ dataset_state = gr.State([])
246
+
247
+ # outputs must match load_captioning's interleaved returns: area, then (row, img, cap) per image
248
+ caption_outputs = [captioning_area]
249
+ for r, im, c in zip(caption_rows, caption_imgs, caption_txts):
250
+ caption_outputs += [r, im, c]
251
+
252
+ demo.load(_signin_state, inputs=None, outputs=[banner, hub_model_id])
253
+ images.change(load_captioning, inputs=[images, instance_prompt], outputs=caption_outputs)
254
+
255
+ submit_btn.click(
256
+ gather_dataset, inputs=[images, *caption_txts], outputs=dataset_state,
257
+ ).then(
258
+ start_training,
259
+ inputs=[dataset_state, lora_name, instance_prompt, validation_prompt, rank, lora_alpha,
260
+ max_train_steps, learning_rate, lr_scheduler, resolution, repeats, train_batch_size,
261
+ gradient_accumulation_steps, seed, optimizer, use_8bit_adam, cache_latents,
262
+ gradient_checkpointing, offload, quantization, lora_layers, validation_epochs,
263
+ hub_model_id, flavor, timeout],
264
+ outputs=[status, joblink, job_id],
265
+ )
266
+ refresh_btn.click(refresh, inputs=[job_id], outputs=[mon_status, mon_logs])
267
+
268
+
269
+ if __name__ == "__main__":
270
+ demo.queue(default_concurrency_limit=4).launch()
jobs.py ADDED
@@ -0,0 +1,265 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """HF Jobs backend for the Krea 2 LoRA trainer Space.
2
+
3
+ Per training request the Space (cpu-basic, no GPU/torch):
4
+ 1. stages the uploaded images + a `metadata.jsonl` (per-image captions),
5
+ 2. pushes them to a private HF **dataset** repo under the signed-in user,
6
+ 3. generates a self-contained UV job script,
7
+ 4. submits it with `HfApi.run_uv_job(... token=<user>)` so the job runs + is billed
8
+ to the signed-in user and the trained LoRA is pushed to their Hub.
9
+
10
+ Token split (important):
11
+ * The **gated Krea 2 weights** (`krea/Krea-2-Raw`, `krea/Krea-2-Turbo`) are not public.
12
+ They are pre-downloaded *inside the job* with the Space's `KREA_TOKEN` secret and passed
13
+ to the trainer as **local dirs**, so `from_pretrained` needs no Krea auth.
14
+ * Everything else (dataset download, `create_repo`/`upload_folder` of the LoRA) uses the
15
+ job's ambient `HF_TOKEN` env = the **signed-in user's** token. The Krea token never
16
+ touches the user's repos and the user's token never needs Krea access.
17
+ """
18
+
19
+ from __future__ import annotations
20
+
21
+ import json
22
+ import os
23
+ import re
24
+ import shutil
25
+ import tempfile
26
+ from pathlib import Path
27
+
28
+ from huggingface_hub import HfApi
29
+
30
+ # The Krea 2 LoRA trainer lives in diffusers PR #14046 (branch `krea2-lora`). Once it is merged,
31
+ # set the `DIFFUSERS_REF` Space variable to `main` (or a release tag) — no code change needed.
32
+ DIFFUSERS_REF = os.environ.get("DIFFUSERS_REF", "krea2-lora")
33
+ BASE_MODEL_RAW = "krea/Krea-2-Raw" # non-distilled base — train LoRA on this
34
+ BASE_MODEL_TURBO = "krea/Krea-2-Turbo" # 8-step distilled — validate / infer on this
35
+
36
+ DEFAULT_FLAVOR = "a100-large"
37
+ FLAVORS = ["l40sx1", "a100-large", "h200"]
38
+ IMAGE_EXTS = {".png", ".jpg", ".jpeg", ".webp", ".bmp"}
39
+
40
+ # Deterministic on-job paths the trainer reads from (baked into the CLI args below).
41
+ JOB_RAW = "/tmp/krea/raw"
42
+ JOB_TURBO = "/tmp/krea/turbo"
43
+ JOB_DATA = "/tmp/data"
44
+ JOB_OUT = "/tmp/out"
45
+ JOB_BNB = "/tmp/bnb.json"
46
+
47
+
48
+ def slug(name: str) -> str:
49
+ s = re.sub(r"[^a-zA-Z0-9-]+", "-", (name or "").strip()).strip("-").lower()
50
+ return s or "krea2-lora"
51
+
52
+
53
+ def _namespace(token: str) -> str:
54
+ from huggingface_hub import whoami # noqa: PLC0415
55
+ return whoami(token=token)["name"]
56
+
57
+
58
+ def build_metadata(image_paths: list[str], captions: list[str], instance_prompt: str) -> list[dict]:
59
+ """One `metadata.jsonl` row per image: {file_name, prompt}. Empty captions fall back to the
60
+ instance prompt (the trigger sentence). Files are renamed to a stable `0000.ext` order."""
61
+ rows = []
62
+ fallback = (instance_prompt or "a photo").strip()
63
+ for i, p in enumerate(image_paths):
64
+ cap = ""
65
+ if i < len(captions) and captions[i]:
66
+ cap = str(captions[i]).strip()
67
+ rows.append({"file_name": f"{i:04d}{Path(p).suffix.lower()}", "prompt": cap or fallback})
68
+ return rows
69
+
70
+
71
+ def build_train_args(params: dict, hub_model_id: str) -> list[str]:
72
+ """Turn UI params into the `train_dreambooth_lora_krea2.py` CLI. Krea weights are passed as
73
+ local dirs (pre-downloaded in the job); the dataset is a local imagefolder (image/prompt cols)."""
74
+ instance_prompt = (params.get("instance_prompt") or "TOK").strip()
75
+ val_prompt = (params.get("validation_prompt") or instance_prompt).strip()
76
+ args = [
77
+ "--pretrained_model_name_or_path", JOB_RAW,
78
+ "--validation_model_path", JOB_TURBO,
79
+ "--dataset_name", JOB_DATA,
80
+ "--image_column", "image",
81
+ "--caption_column", "prompt",
82
+ "--instance_prompt", instance_prompt,
83
+ "--output_dir", JOB_OUT,
84
+ "--mixed_precision", "bf16",
85
+ "--resolution", str(int(params["resolution"])),
86
+ "--train_batch_size", str(int(params["train_batch_size"])),
87
+ "--gradient_accumulation_steps", str(int(params["gradient_accumulation_steps"])),
88
+ "--repeats", str(int(params["repeats"])),
89
+ "--rank", str(int(params["rank"])),
90
+ "--lora_alpha", str(int(params["lora_alpha"])),
91
+ "--learning_rate", str(float(params["learning_rate"])),
92
+ "--lr_scheduler", str(params["lr_scheduler"]),
93
+ "--lr_warmup_steps", "0",
94
+ "--max_train_steps", str(int(params["max_train_steps"])),
95
+ "--optimizer", str(params["optimizer"]),
96
+ "--seed", str(int(params["seed"])),
97
+ "--validation_prompt", val_prompt,
98
+ "--validation_epochs", str(int(params["validation_epochs"])),
99
+ "--num_validation_images", "2",
100
+ "--push_to_hub",
101
+ "--hub_model_id", hub_model_id,
102
+ ]
103
+ if params.get("lora_layers"):
104
+ args += ["--lora_layers", str(params["lora_layers"]).strip()]
105
+ if params.get("gradient_checkpointing", True):
106
+ args += ["--gradient_checkpointing"]
107
+ if params.get("cache_latents", True):
108
+ args += ["--cache_latents"]
109
+ if params.get("offload"):
110
+ args += ["--offload"]
111
+ if params.get("use_8bit_adam") and str(params["optimizer"]).lower() == "adamw":
112
+ args += ["--use_8bit_adam"]
113
+ quant = params.get("quantization", "none")
114
+ if quant == "fp8":
115
+ args += ["--do_fp8_training"]
116
+ elif quant == "4bit":
117
+ args += ["--bnb_quantization_config_path", JOB_BNB]
118
+ return args
119
+
120
+
121
+ # --------------------------------------------------------------------------------------
122
+ # UV job script (runs on HF Jobs GPU hardware)
123
+ # --------------------------------------------------------------------------------------
124
+
125
+ JOB_SCRIPT_TEMPLATE = '''# /// script
126
+ # requires-python = ">=3.10"
127
+ # dependencies = [
128
+ # "git+https://github.com/huggingface/diffusers.git@{ref}",
129
+ # "torch",
130
+ # "torchvision",
131
+ # "transformers>=4.41.2",
132
+ # "accelerate>=0.31.0",
133
+ # "peft>=0.11.1",
134
+ # "datasets",
135
+ # "bitsandbytes",
136
+ # "prodigyopt",
137
+ # "ftfy",
138
+ # "sentencepiece",
139
+ # "hf_transfer",
140
+ # "huggingface_hub[hf-xet]",
141
+ # ]
142
+ # ///
143
+ """Auto-generated Krea 2 DreamBooth-LoRA job. Trains on Krea 2 RAW, validates on Turbo."""
144
+ import json, os, subprocess, sys, urllib.request
145
+ from pathlib import Path
146
+
147
+ os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "1"
148
+
149
+ REF = "{ref}"
150
+ RAW, TURBO, DATA, BNB = "{raw}", "{turbo}", "{data}", "{bnb}"
151
+ DATASET_REPO = {dataset_repo!r}
152
+ QUANT = {quant!r}
153
+ TRAIN_ARGS = {train_args}
154
+ KREA_TOKEN = os.environ["KREA_TOKEN"] # gated Krea weights ONLY
155
+
156
+ SCRIPT_URL = (
157
+ "https://raw.githubusercontent.com/huggingface/diffusers/"
158
+ + REF + "/examples/dreambooth/train_dreambooth_lora_krea2.py"
159
+ )
160
+
161
+
162
+ def main():
163
+ from huggingface_hub import snapshot_download
164
+
165
+ print("=== 1/4 download gated Krea 2 weights (krea token) ===", flush=True)
166
+ snapshot_download("krea/Krea-2-Raw", local_dir=RAW, token=KREA_TOKEN)
167
+ snapshot_download("krea/Krea-2-Turbo", local_dir=TURBO, token=KREA_TOKEN)
168
+
169
+ print("=== 2/4 download dataset (user token / HF_TOKEN env) ===", flush=True)
170
+ snapshot_download(DATASET_REPO, repo_type="dataset", local_dir=DATA)
171
+
172
+ if QUANT == "4bit":
173
+ Path(BNB).write_text(json.dumps({{
174
+ "load_in_4bit": True, "bnb_4bit_quant_type": "nf4",
175
+ "bnb_4bit_compute_dtype": "bfloat16",
176
+ }}))
177
+
178
+ print("=== 3/4 fetch trainer script @ " + REF + " ===", flush=True)
179
+ urllib.request.urlretrieve(SCRIPT_URL, "/tmp/train_dreambooth_lora_krea2.py")
180
+
181
+ print("=== 4/4 accelerate launch (pushes LoRA to the Hub) ===", flush=True)
182
+ cmd = [sys.executable, "-m", "accelerate.commands.launch",
183
+ "/tmp/train_dreambooth_lora_krea2.py", *TRAIN_ARGS]
184
+ print(">>> " + " ".join(cmd), flush=True)
185
+ subprocess.run(cmd, check=True)
186
+ print("=== DONE ===", flush=True)
187
+
188
+
189
+ if __name__ == "__main__":
190
+ main()
191
+ '''
192
+
193
+
194
+ def submit(params: dict, image_paths: list[str], captions: list[str],
195
+ flavor: str, timeout: str) -> dict:
196
+ """Stage dataset → push private dataset repo → generate UV script → submit job.
197
+ Returns {job_id, url, dataset_repo, hub_model_id}."""
198
+ token = (params.get("hf_token") or "").strip()
199
+ if not token:
200
+ raise ValueError("Missing user token (sign in with Hugging Face).")
201
+ if not os.environ.get("KREA_TOKEN"):
202
+ raise RuntimeError("Space is missing the KREA_TOKEN secret (gated Krea 2 access).")
203
+
204
+ ns = _namespace(token)
205
+ name = slug(params.get("lora_name", ""))
206
+ dataset_repo = f"{ns}/{name}-dataset"
207
+ hub_model_id = (params.get("hub_model_id") or "").strip() or f"{ns}/{name}"
208
+
209
+ api = HfApi(token=token)
210
+ tmp = Path(tempfile.mkdtemp(prefix="krea2-"))
211
+ try:
212
+ # stage images under stable names + metadata.jsonl
213
+ rows = build_metadata(image_paths, captions, params.get("instance_prompt", ""))
214
+ for row, src in zip(rows, image_paths):
215
+ shutil.copy(src, tmp / row["file_name"])
216
+ (tmp / "metadata.jsonl").write_text(
217
+ "\n".join(json.dumps(r) for r in rows) + "\n"
218
+ )
219
+
220
+ # push the dataset (private, user namespace)
221
+ api.create_repo(dataset_repo, repo_type="dataset", private=True, exist_ok=True, token=token)
222
+ api.upload_folder(repo_id=dataset_repo, repo_type="dataset", folder_path=str(tmp), token=token)
223
+
224
+ # render the job script
225
+ train_args = build_train_args(params, hub_model_id)
226
+ script = JOB_SCRIPT_TEMPLATE.format(
227
+ ref=DIFFUSERS_REF, raw=JOB_RAW, turbo=JOB_TURBO, data=JOB_DATA, bnb=JOB_BNB,
228
+ dataset_repo=dataset_repo, quant=params.get("quantization", "none"),
229
+ train_args=json.dumps(train_args),
230
+ )
231
+ script_path = tmp / "job_train.py"
232
+ script_path.write_text(script)
233
+
234
+ job = api.run_uv_job(
235
+ str(script_path),
236
+ flavor=flavor,
237
+ timeout=timeout,
238
+ # HF_TOKEN = user (push + dataset); KREA_TOKEN = gated Krea weights only.
239
+ secrets={"HF_TOKEN": token, "KREA_TOKEN": os.environ["KREA_TOKEN"]},
240
+ token=token,
241
+ )
242
+ job_id = getattr(job, "id", "") or ""
243
+ url = getattr(job, "url", "") or (f"https://huggingface.co/jobs/{ns}/{job_id}" if job_id else "")
244
+ return {"job_id": job_id, "url": url, "dataset_repo": dataset_repo, "hub_model_id": hub_model_id}
245
+ finally:
246
+ shutil.rmtree(tmp, ignore_errors=True)
247
+
248
+
249
+ def job_logs(job_id: str, token: str = "") -> str:
250
+ try:
251
+ return "\n".join(HfApi(token=token).fetch_job_logs(job_id=job_id, token=token))
252
+ except Exception as e: # noqa: BLE001
253
+ return f"(could not fetch logs: {e})"
254
+
255
+
256
+ def job_status(job_id: str, token: str = "") -> str:
257
+ try:
258
+ job = HfApi(token=token).inspect_job(job_id=job_id, token=token)
259
+ status = getattr(job, "status", None)
260
+ stage = getattr(status, "stage", None)
261
+ if stage is None and isinstance(status, dict):
262
+ stage = status.get("stage")
263
+ return str(stage or status or "UNKNOWN")
264
+ except Exception as e: # noqa: BLE001
265
+ return f"UNKNOWN ({e})"
requirements.txt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ gradio[oauth]>=6.18
2
+ huggingface_hub[hf-xet]>=1.5
3
+ hf_transfer