mpasila commited on
Commit
b5e5a50
·
verified ·
1 Parent(s): d73fa99

feat(models): add support for custom Hugging Face base checkpoints

Browse files

Allow users to specify a custom base model from any accessible Hugging Face repository instead of relying solely on the provided catalog.

- Implemented lazy downloading and caching of custom base models into a managed ComfyUI directory using a stable namespace.
- Added UI controls for repository ID, filename, and optional revision.
- Integrated custom base model resolution into the generation and runtime preparation pipelines.
- Updated settings profiles (schema v3) to preserve custom base model configurations.
- Added validation to prevent path traversal and ensure supported file extensions.
- Included unit tests for custom model downloading and validation.

Files changed (5) hide show
  1. README.md +8 -1
  2. app.py +190 -41
  3. settings_utils.py +49 -1
  4. test_app.py +26 -0
  5. test_settings_utils.py +38 -0
README.md CHANGED
@@ -20,6 +20,7 @@ This Space runs Krea 2 Turbo through headless ComfyUI execution. It provides:
20
  - Instruction-based image editing with an optional second reference image.
21
  - Krea 2 Identity Edit v1.2 conditioning for the edit mode.
22
  - Two selectable custom Krea 2 Turbo checkpoints.
 
23
  - A catalog of Krea 2 LoRAs with signed weight controls and search/filtering.
24
  - Custom LoRAs from any accessible Hugging Face repository.
25
  - Seed, sampling, grounding, and reference-fidelity controls.
@@ -37,6 +38,12 @@ The base-model selector provides both checkpoints from [mpasila/Krea-2-Models](h
37
  | `qwen3vl_4b_fp8_scaled.safetensors` | Krea 2 Qwen3-VL text encoder |
38
  | `qwen_image_vae.safetensors` | Krea 2 VAE |
39
 
 
 
 
 
 
 
40
  Edit mode also downloads `krea2_identity_edit_v1_2.safetensors` from
41
  [conradlocke/krea2-identity-edit](https://huggingface.co/conradlocke/krea2-identity-edit).
42
  The edit nodes come from [ComfyUI-Krea2Edit](https://github.com/lbouaraba/comfyui-krea2edit).
@@ -79,4 +86,4 @@ edit instruction more strongly; higher values usually preserve reference identit
79
 
80
  Generated PNGs include a `krea_settings` JSON object and a readable `parameters` field. Source image
81
  bytes are never stored in exported settings profiles. Profiles also preserve the selected checkpoint,
82
- catalog weights, and custom Hugging Face LoRA rows.
 
20
  - Instruction-based image editing with an optional second reference image.
21
  - Krea 2 Identity Edit v1.2 conditioning for the edit mode.
22
  - Two selectable custom Krea 2 Turbo checkpoints.
23
+ - An option to download and use a custom base checkpoint from any accessible Hugging Face repository.
24
  - A catalog of Krea 2 LoRAs with signed weight controls and search/filtering.
25
  - Custom LoRAs from any accessible Hugging Face repository.
26
  - Seed, sampling, grounding, and reference-fidelity controls.
 
38
  | `qwen3vl_4b_fp8_scaled.safetensors` | Krea 2 Qwen3-VL text encoder |
39
  | `qwen_image_vae.safetensors` | Krea 2 VAE |
40
 
41
+ Choose **Custom Hugging Face base model** in the base checkpoint selector to provide a repository,
42
+ relative model filename, and optional revision. Supported files use `.safetensors`, `.ckpt`, `.pt`,
43
+ or `.bin` extensions. The checkpoint is downloaded lazily into ComfyUI's managed diffusion-model
44
+ directory and cached in a stable repository/revision namespace. Private repositories require the
45
+ `HF_TOKEN` or `HUGGINGFACE_HUB_TOKEN` environment variable.
46
+
47
  Edit mode also downloads `krea2_identity_edit_v1_2.safetensors` from
48
  [conradlocke/krea2-identity-edit](https://huggingface.co/conradlocke/krea2-identity-edit).
49
  The edit nodes come from [ComfyUI-Krea2Edit](https://github.com/lbouaraba/comfyui-krea2edit).
 
86
 
87
  Generated PNGs include a `krea_settings` JSON object and a readable `parameters` field. Source image
88
  bytes are never stored in exported settings profiles. Profiles also preserve the selected checkpoint,
89
+ custom base-model repository/file/revision, catalog weights, and custom Hugging Face LoRA rows.
app.py CHANGED
@@ -46,7 +46,9 @@ from settings_utils import (
46
  extract_image_settings,
47
  normalize_custom_loras,
48
  parse_settings_text,
 
49
  stable_custom_lora_namespace,
 
50
  validate_custom_lora,
51
  write_png_metadata,
52
  )
@@ -67,6 +69,7 @@ CUSTOM_KREA_REPO = "mpasila/Krea-2-Models"
67
  CUSTOM_KREA_FILE = "pornmasterKrea2_v2TurboInt8.safetensors"
68
  MUSE_KREA_FILE = "museByStableYogi_v25EXTENDEDTURBO.safetensors"
69
  DEFAULT_BASE_MODEL = CUSTOM_KREA_FILE
 
70
  BASE_MODELS = {
71
  MUSE_KREA_FILE: {
72
  "label": "Muse By Stable Yogi Krea2",
@@ -264,8 +267,8 @@ def _load_lora_catalog() -> None:
264
  print(f"[lora] loaded {len(_lora_catalog)} Krea LoRAs", flush=True)
265
 
266
 
267
- def _ensure_base_model(base_model: str, progress: gr.Progress | None = None) -> None:
268
- """Download the selected approved base checkpoint if it is not cached."""
269
  if base_model not in BASE_MODELS:
270
  raise ValueError("unsupported Krea base model")
271
  item = BASE_MODELS[base_model]
@@ -275,20 +278,82 @@ def _ensure_base_model(base_model: str, progress: gr.Progress | None = None) ->
275
  MODELS / "diffusion_models" / item["filename"],
276
  item["label"],
277
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
278
 
279
 
280
  def _ensure_models(
281
  base_model: str = DEFAULT_BASE_MODEL,
 
282
  progress: gr.Progress | None = None,
283
- ) -> None:
284
  total = len(DOWNLOADS) + 1
285
  for index, (repo, filename, destination, label) in enumerate(DOWNLOADS):
286
  if progress:
287
  progress(index / total, desc=f"downloading {label}")
288
  _download_to_dest(repo, filename, destination, label)
 
 
 
 
289
  if progress:
290
  progress(len(DOWNLOADS) / total, desc="downloading selected Krea checkpoint")
291
- _ensure_base_model(base_model)
292
 
293
 
294
  def _ensure_lora(hf_filename: str) -> str:
@@ -368,9 +433,17 @@ def _ref(node: str, output: int = 0) -> list[Any]:
368
  return [node, output]
369
 
370
 
 
 
 
 
 
 
 
 
 
371
  def _t2i_workflow(base_model: str = DEFAULT_BASE_MODEL) -> dict[str, Any]:
372
- if base_model not in BASE_MODELS:
373
- raise ValueError("unsupported Krea base model")
374
  cache_key = f"text2image:{base_model}"
375
  if cache_key in _workflow_cache:
376
  return json.loads(json.dumps(_workflow_cache[cache_key]))
@@ -395,8 +468,7 @@ def _edit_workflow(
395
  has_second_reference: bool,
396
  base_model: str = DEFAULT_BASE_MODEL,
397
  ) -> dict[str, Any]:
398
- if base_model not in BASE_MODELS:
399
- raise ValueError("unsupported Krea base model")
400
  _read_source_workflow(EDIT_SOURCE)
401
  workflow: dict[str, Any] = {
402
  "1": {"class_type": "LoadImage", "inputs": {"image": ""}},
@@ -604,11 +676,13 @@ def _execute_workflow(workflow: dict[str, Any]) -> list[str]:
604
 
605
  def _prepare_runtime(
606
  base_model: str = DEFAULT_BASE_MODEL,
 
607
  progress: gr.Progress | None = None,
608
- ) -> None:
609
  _ensure_comfy()
610
- _ensure_models(base_model, progress)
611
  _init_comfy_nodes()
 
612
 
613
 
614
  def get_gpu_duration(*args: Any, **kwargs: Any) -> int:
@@ -621,8 +695,8 @@ def get_gpu_duration(*args: Any, **kwargs: Any) -> int:
621
  gen_budget = kwargs.get("gen_budget", args[17] if len(args) > 17 else 0)
622
  if gen_budget and int(gen_budget) > 0:
623
  return max(MIN_GPU_SECONDS, min(MAX_GPU_SECONDS, int(gen_budget)))
624
- lora_weights = kwargs.get("lora_weights", args[19] if len(args) > 19 else {}) or {}
625
- custom_loras = kwargs.get("custom_loras", args[20] if len(args) > 20 else []) or []
626
  lora_count = sum(
627
  1 for value in lora_weights.values()
628
  if value and abs(float(value)) > 1e-6
@@ -659,6 +733,9 @@ def generate(
659
  randomize_seed: bool,
660
  gen_budget: float,
661
  base_model: str = DEFAULT_BASE_MODEL,
 
 
 
662
  lora_weights: dict[str, float] | None = None,
663
  custom_loras: list[dict[str, Any]] | None = None,
664
  progress: gr.Progress = gr.Progress(track_tqdm=True),
@@ -671,9 +748,16 @@ def generate(
671
  effective_edit_prompt = (edit_prompt or prompt or "").strip()
672
  if sampler not in SAMPLERS or scheduler not in SCHEDULERS:
673
  raise ValueError("unsupported sampler or scheduler")
674
- if base_model not in BASE_MODELS:
 
 
 
 
 
 
 
675
  raise ValueError("unsupported Krea base model")
676
- _prepare_runtime(base_model, progress)
677
 
678
  enabled_loras: list[tuple[str, float]] = []
679
  active_catalog: list[dict[str, Any]] = []
@@ -699,14 +783,14 @@ def generate(
699
  if mode == "text2image":
700
  width = max(512, min(MAX_WIDTH, int(width) // 64 * 64))
701
  height = max(512, min(MAX_HEIGHT, int(height) // 64 * 64))
702
- workflow = _t2i_workflow(base_model)
703
  else:
704
  primary_name, width, height = _prepare_edit_image(primary_image, target_megapixels)
705
  staged.append(INPUT / primary_name)
706
  second_name = _stage_image(second_image, "reference") if second_image else None
707
  if second_name:
708
  staged.append(INPUT / second_name)
709
- workflow = _edit_workflow(bool(second_name), base_model)
710
 
711
  if mode == "text2image":
712
  _inject_t2i(
@@ -759,6 +843,7 @@ def generate(
759
  gen_budget=float(gen_budget),
760
  effective_seed=effective_seed,
761
  base_model=base_model,
 
762
  catalog_loras=active_catalog,
763
  custom_loras=active_custom,
764
  )
@@ -800,6 +885,9 @@ def _profile_from_values(
800
  randomize: bool,
801
  budget: float,
802
  base_model: str = DEFAULT_BASE_MODEL,
 
 
 
803
  catalog_loras: list[dict[str, Any]] | None = None,
804
  custom_loras: list[dict[str, Any]] | None = None,
805
  ) -> dict[str, Any]:
@@ -821,6 +909,15 @@ def _profile_from_values(
821
  randomize_seed=bool(randomize),
822
  gen_budget=float(budget),
823
  base_model=base_model,
 
 
 
 
 
 
 
 
 
824
  catalog_loras=catalog_loras,
825
  custom_loras=custom_loras,
826
  )
@@ -833,10 +930,30 @@ def create_ui() -> gr.Blocks:
833
  with gr.Column(scale=1):
834
  mode = gr.Radio(["text2image", "edit"], value="text2image", label="mode")
835
  base_model = gr.Dropdown(
836
- choices=[(item["label"], filename) for filename, item in BASE_MODELS.items()],
 
 
 
837
  value=DEFAULT_BASE_MODEL,
838
  label="base Krea checkpoint",
839
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
840
  with gr.Column(visible=False) as image_inputs:
841
  primary = gr.Image(type="filepath", label="primary image / scene")
842
  second = gr.Image(type="filepath", label="optional second reference")
@@ -930,6 +1047,11 @@ def create_ui() -> gr.Blocks:
930
 
931
  mode.change(on_mode_change, inputs=[mode], outputs=[image_inputs, t2i_resolution, edit_controls, edit_prompt])
932
 
 
 
 
 
 
933
  all_lora_filenames = list(lora_slider_map.keys())
934
  all_lora_sliders = list(lora_slider_map.values())
935
 
@@ -973,11 +1095,17 @@ def create_ui() -> gr.Blocks:
973
  def _generate_wrapper(*values):
974
  base_values = values[:18]
975
  base_model_value = values[18]
976
- custom_rows = values[19]
977
- lora_weights = _catalog_weights(values[20:])
 
 
 
978
  return generate(
979
  *base_values,
980
  base_model=base_model_value,
 
 
 
981
  lora_weights=lora_weights,
982
  custom_loras=custom_rows,
983
  )
@@ -985,14 +1113,16 @@ def create_ui() -> gr.Blocks:
985
  generation_inputs = [
986
  mode, prompt, edit_prompt, primary, second, width, height, target_mp,
987
  grounding, ref_boost, ref_boost_a, steps, cfg, sampler, scheduler,
988
- seed, randomize, gen_budget, base_model, custom_loras, *all_lora_sliders,
 
989
  ]
990
  button.click(_generate_wrapper, inputs=generation_inputs, outputs=[gallery, status, used_seed])
991
 
992
  profile_inputs = [
993
  mode, prompt, edit_prompt, width, height, target_mp, grounding,
994
  ref_boost, ref_boost_a, steps, cfg, sampler, scheduler, seed,
995
- randomize, gen_budget, base_model, custom_loras, *all_lora_sliders,
 
996
  profile_name,
997
  ]
998
 
@@ -1000,16 +1130,19 @@ def create_ui() -> gr.Blocks:
1000
  name_value = str(values[-1] or "krea2")
1001
  catalog = [
1002
  {"hf_filename": filename, "weight": float(weight or 0.0)}
1003
- for filename, weight in zip(all_lora_filenames, values[18:-1])
1004
  if weight and abs(float(weight)) > 1e-6
1005
  ]
1006
- normalized_custom, warnings = normalize_custom_loras(values[17])
1007
  if warnings:
1008
  return None, "settings export failed — " + "; ".join(warnings)
1009
  try:
1010
  data = _profile_from_values(
1011
  *values[:16],
1012
  values[16],
 
 
 
1013
  catalog,
1014
  normalized_custom,
1015
  )
@@ -1024,34 +1157,50 @@ def create_ui() -> gr.Blocks:
1024
  export_button.click(export_profile, inputs=profile_inputs, outputs=[export_file, profile_status]).then(lambda: gr.update(visible=True), outputs=[export_file])
1025
 
1026
  settings_outputs = [
1027
- mode, base_model, prompt, edit_prompt, width, height, target_mp, grounding,
 
1028
  ref_boost, ref_boost_a, steps, cfg, sampler, scheduler, seed,
1029
  randomize, gen_budget, custom_loras, *all_lora_sliders, profile_status, image_inputs,
1030
- t2i_resolution, edit_controls,
1031
  ]
1032
  settings_keys = [
1033
- "mode", "base_model", "prompt", "edit_prompt", "width", "height",
 
1034
  "target_megapixels", "grounding_px", "ref_boost", "ref_boost_a",
1035
  "steps", "cfg", "sampler_name", "scheduler", "effective_seed",
1036
  "randomize_seed", "gen_budget",
1037
  ]
1038
 
1039
  def settings_updates(data: dict[str, Any], warnings: list[str], error: str = ""):
1040
- editing = data.get("mode") == "edit"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1041
  updates = []
1042
  for key in settings_keys:
1043
- if key not in data:
1044
  updates.append(gr.update())
1045
- elif key == "base_model" and data[key] not in BASE_MODELS:
1046
- warnings.append(f"ignored unsupported base model: {data[key]!r}")
1047
  updates.append(gr.update())
1048
- elif key == "edit_prompt" and "mode" in data:
1049
- updates.append(gr.update(value=data[key], visible=editing))
1050
  else:
1051
- updates.append(gr.update(value=data[key]))
1052
  custom_rows: list[list[Any]] = []
1053
- if "custom_loras" in data:
1054
- for row in data.get("custom_loras", []) or []:
1055
  try:
1056
  normalized = validate_custom_lora(row)
1057
  except ValueError as exc:
@@ -1061,9 +1210,9 @@ def create_ui() -> gr.Blocks:
1061
  normalized["repo_id"], normalized["filename"],
1062
  normalized["revision"], normalized["weight"],
1063
  ])
1064
- updates.append(gr.update(value=custom_rows) if "custom_loras" in data else gr.update())
1065
  catalog_values = {filename: 0.0 for filename in all_lora_filenames}
1066
- for row in data.get("catalog_loras", []) or []:
1067
  if not isinstance(row, dict):
1068
  warnings.append("ignored malformed catalog LoRA entry")
1069
  continue
@@ -1086,10 +1235,10 @@ def create_ui() -> gr.Blocks:
1086
  message += f" (+{len(warnings) - 6} more)"
1087
  updates.extend([
1088
  message,
1089
- gr.update(visible=editing) if "mode" in data else gr.update(),
1090
- gr.update(visible=not editing) if "mode" in data else gr.update(),
1091
- gr.update(visible=editing) if "mode" in data else gr.update(),
1092
- gr.update(visible=editing) if "mode" in data else gr.update(),
1093
  ])
1094
  return tuple(updates)
1095
 
 
46
  extract_image_settings,
47
  normalize_custom_loras,
48
  parse_settings_text,
49
+ stable_custom_base_model_namespace,
50
  stable_custom_lora_namespace,
51
+ validate_custom_base_model,
52
  validate_custom_lora,
53
  write_png_metadata,
54
  )
 
69
  CUSTOM_KREA_FILE = "pornmasterKrea2_v2TurboInt8.safetensors"
70
  MUSE_KREA_FILE = "museByStableYogi_v25EXTENDEDTURBO.safetensors"
71
  DEFAULT_BASE_MODEL = CUSTOM_KREA_FILE
72
+ CUSTOM_BASE_MODEL = "__custom_huggingface_base_model__"
73
  BASE_MODELS = {
74
  MUSE_KREA_FILE: {
75
  "label": "Muse By Stable Yogi Krea2",
 
267
  print(f"[lora] loaded {len(_lora_catalog)} Krea LoRAs", flush=True)
268
 
269
 
270
+ def _ensure_base_model(base_model: str, progress: gr.Progress | None = None) -> str:
271
+ """Download an approved base checkpoint and return its ComfyUI name."""
272
  if base_model not in BASE_MODELS:
273
  raise ValueError("unsupported Krea base model")
274
  item = BASE_MODELS[base_model]
 
278
  MODELS / "diffusion_models" / item["filename"],
279
  item["label"],
280
  )
281
+ return base_model
282
+
283
+
284
+ def _ensure_custom_base_model(value: dict[str, Any]) -> str:
285
+ """Download a validated custom HF base model into a managed ComfyUI folder."""
286
+ normalized = validate_custom_base_model(value)
287
+ namespace = stable_custom_base_model_namespace(
288
+ normalized["repo_id"], normalized["filename"], normalized["revision"]
289
+ )
290
+ diffusion_root = MODELS / "diffusion_models"
291
+ destination_dir = diffusion_root / "huggingface" / namespace
292
+ relative_file = pathlib.PurePosixPath(normalized["filename"])
293
+ destination = destination_dir.joinpath(*relative_file.parts)
294
+ if destination.exists():
295
+ relative = destination.resolve().relative_to(diffusion_root.resolve())
296
+ return pathlib.PurePosixPath(*relative.parts).as_posix()
297
+
298
+ token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACE_HUB_TOKEN")
299
+ kwargs: dict[str, Any] = {
300
+ "repo_id": normalized["repo_id"],
301
+ "filename": normalized["filename"],
302
+ "local_dir": str(destination_dir),
303
+ "token": token,
304
+ }
305
+ if normalized["revision"]:
306
+ kwargs["revision"] = normalized["revision"]
307
+ try:
308
+ downloaded = pathlib.Path(hf_hub_download(**kwargs))
309
+ except Exception as exc:
310
+ raise RuntimeError(
311
+ f"failed to download custom base model "
312
+ f"{normalized['repo_id']}/{normalized['filename']}: {exc}"
313
+ ) from exc
314
+ if not downloaded.exists():
315
+ raise RuntimeError(f"HF Hub returned a missing custom base-model path: {downloaded}")
316
+ try:
317
+ relative = downloaded.resolve().relative_to(diffusion_root.resolve())
318
+ except ValueError as exc:
319
+ raise RuntimeError("custom base-model download escaped the managed model directory") from exc
320
+ print(f"[models] ready: custom base model {normalized['repo_id']}/{normalized['filename']}", flush=True)
321
+ return pathlib.PurePosixPath(*relative.parts).as_posix()
322
+
323
+
324
+ def _resolve_base_model(
325
+ base_model: str,
326
+ custom_base_model: dict[str, Any] | None = None,
327
+ ) -> str:
328
+ """Resolve the UI selection to the relative model name expected by ComfyUI."""
329
+ if base_model == CUSTOM_BASE_MODEL:
330
+ if not custom_base_model:
331
+ raise ValueError("enter a custom Hugging Face base-model repository and file")
332
+ return _ensure_custom_base_model(custom_base_model)
333
+ if base_model not in BASE_MODELS:
334
+ raise ValueError("unsupported Krea base model")
335
+ if custom_base_model:
336
+ raise ValueError("custom base-model details require the custom model option")
337
+ return _ensure_base_model(base_model)
338
 
339
 
340
  def _ensure_models(
341
  base_model: str = DEFAULT_BASE_MODEL,
342
+ custom_base_model: dict[str, Any] | None = None,
343
  progress: gr.Progress | None = None,
344
+ ) -> str:
345
  total = len(DOWNLOADS) + 1
346
  for index, (repo, filename, destination, label) in enumerate(DOWNLOADS):
347
  if progress:
348
  progress(index / total, desc=f"downloading {label}")
349
  _download_to_dest(repo, filename, destination, label)
350
+ if base_model == CUSTOM_BASE_MODEL:
351
+ if progress:
352
+ progress(len(DOWNLOADS) / total, desc="downloading custom Hugging Face checkpoint")
353
+ return _resolve_base_model(base_model, custom_base_model)
354
  if progress:
355
  progress(len(DOWNLOADS) / total, desc="downloading selected Krea checkpoint")
356
+ return _resolve_base_model(base_model)
357
 
358
 
359
  def _ensure_lora(hf_filename: str) -> str:
 
433
  return [node, output]
434
 
435
 
436
+ def _validate_model_name(model_name: str) -> str:
437
+ """Accept only a relative model name produced by the managed model loader."""
438
+ normalized = str(model_name).replace("\\", "/")
439
+ path = pathlib.PurePosixPath(normalized)
440
+ if not normalized or path.is_absolute() or any(part in {"", ".", ".."} for part in path.parts):
441
+ raise ValueError("invalid Krea base model path")
442
+ return path.as_posix()
443
+
444
+
445
  def _t2i_workflow(base_model: str = DEFAULT_BASE_MODEL) -> dict[str, Any]:
446
+ base_model = _validate_model_name(base_model)
 
447
  cache_key = f"text2image:{base_model}"
448
  if cache_key in _workflow_cache:
449
  return json.loads(json.dumps(_workflow_cache[cache_key]))
 
468
  has_second_reference: bool,
469
  base_model: str = DEFAULT_BASE_MODEL,
470
  ) -> dict[str, Any]:
471
+ base_model = _validate_model_name(base_model)
 
472
  _read_source_workflow(EDIT_SOURCE)
473
  workflow: dict[str, Any] = {
474
  "1": {"class_type": "LoadImage", "inputs": {"image": ""}},
 
676
 
677
  def _prepare_runtime(
678
  base_model: str = DEFAULT_BASE_MODEL,
679
+ custom_base_model: dict[str, Any] | None = None,
680
  progress: gr.Progress | None = None,
681
+ ) -> str:
682
  _ensure_comfy()
683
+ resolved_base_model = _ensure_models(base_model, custom_base_model, progress)
684
  _init_comfy_nodes()
685
+ return resolved_base_model
686
 
687
 
688
  def get_gpu_duration(*args: Any, **kwargs: Any) -> int:
 
695
  gen_budget = kwargs.get("gen_budget", args[17] if len(args) > 17 else 0)
696
  if gen_budget and int(gen_budget) > 0:
697
  return max(MIN_GPU_SECONDS, min(MAX_GPU_SECONDS, int(gen_budget)))
698
+ lora_weights = kwargs.get("lora_weights", args[22] if len(args) > 22 else {}) or {}
699
+ custom_loras = kwargs.get("custom_loras", args[23] if len(args) > 23 else []) or []
700
  lora_count = sum(
701
  1 for value in lora_weights.values()
702
  if value and abs(float(value)) > 1e-6
 
733
  randomize_seed: bool,
734
  gen_budget: float,
735
  base_model: str = DEFAULT_BASE_MODEL,
736
+ custom_base_repo: str = "",
737
+ custom_base_filename: str = "",
738
+ custom_base_revision: str = "",
739
  lora_weights: dict[str, float] | None = None,
740
  custom_loras: list[dict[str, Any]] | None = None,
741
  progress: gr.Progress = gr.Progress(track_tqdm=True),
 
748
  effective_edit_prompt = (edit_prompt or prompt or "").strip()
749
  if sampler not in SAMPLERS or scheduler not in SCHEDULERS:
750
  raise ValueError("unsupported sampler or scheduler")
751
+ custom_base = None
752
+ if base_model == CUSTOM_BASE_MODEL:
753
+ custom_base = validate_custom_base_model({
754
+ "repo_id": custom_base_repo,
755
+ "filename": custom_base_filename,
756
+ "revision": custom_base_revision,
757
+ })
758
+ elif base_model not in BASE_MODELS:
759
  raise ValueError("unsupported Krea base model")
760
+ resolved_base_model = _prepare_runtime(base_model, custom_base, progress)
761
 
762
  enabled_loras: list[tuple[str, float]] = []
763
  active_catalog: list[dict[str, Any]] = []
 
783
  if mode == "text2image":
784
  width = max(512, min(MAX_WIDTH, int(width) // 64 * 64))
785
  height = max(512, min(MAX_HEIGHT, int(height) // 64 * 64))
786
+ workflow = _t2i_workflow(resolved_base_model)
787
  else:
788
  primary_name, width, height = _prepare_edit_image(primary_image, target_megapixels)
789
  staged.append(INPUT / primary_name)
790
  second_name = _stage_image(second_image, "reference") if second_image else None
791
  if second_name:
792
  staged.append(INPUT / second_name)
793
+ workflow = _edit_workflow(bool(second_name), resolved_base_model)
794
 
795
  if mode == "text2image":
796
  _inject_t2i(
 
843
  gen_budget=float(gen_budget),
844
  effective_seed=effective_seed,
845
  base_model=base_model,
846
+ custom_base_model=custom_base,
847
  catalog_loras=active_catalog,
848
  custom_loras=active_custom,
849
  )
 
885
  randomize: bool,
886
  budget: float,
887
  base_model: str = DEFAULT_BASE_MODEL,
888
+ custom_base_repo: str = "",
889
+ custom_base_filename: str = "",
890
+ custom_base_revision: str = "",
891
  catalog_loras: list[dict[str, Any]] | None = None,
892
  custom_loras: list[dict[str, Any]] | None = None,
893
  ) -> dict[str, Any]:
 
909
  randomize_seed=bool(randomize),
910
  gen_budget=float(budget),
911
  base_model=base_model,
912
+ custom_base_model=(
913
+ validate_custom_base_model({
914
+ "repo_id": custom_base_repo,
915
+ "filename": custom_base_filename,
916
+ "revision": custom_base_revision,
917
+ })
918
+ if base_model == CUSTOM_BASE_MODEL
919
+ else None
920
+ ),
921
  catalog_loras=catalog_loras,
922
  custom_loras=custom_loras,
923
  )
 
930
  with gr.Column(scale=1):
931
  mode = gr.Radio(["text2image", "edit"], value="text2image", label="mode")
932
  base_model = gr.Dropdown(
933
+ choices=[
934
+ *[(item["label"], filename) for filename, item in BASE_MODELS.items()],
935
+ ("Custom Hugging Face base model", CUSTOM_BASE_MODEL),
936
+ ],
937
  value=DEFAULT_BASE_MODEL,
938
  label="base Krea checkpoint",
939
  )
940
+ with gr.Column(visible=False) as custom_base_controls:
941
+ custom_base_repo = gr.Textbox(
942
+ label="custom base-model Hugging Face repository",
943
+ placeholder="namespace/repository",
944
+ )
945
+ custom_base_filename = gr.Textbox(
946
+ label="custom base-model filename",
947
+ placeholder="path/to/model.safetensors",
948
+ )
949
+ custom_base_revision = gr.Textbox(
950
+ label="custom base-model revision (optional)",
951
+ placeholder="main, tag, or commit hash",
952
+ )
953
+ gr.Markdown(
954
+ "Use a relative `.safetensors`, `.ckpt`, `.pt`, or `.bin` file. "
955
+ "Private repositories use `HF_TOKEN` or `HUGGINGFACE_HUB_TOKEN`."
956
+ )
957
  with gr.Column(visible=False) as image_inputs:
958
  primary = gr.Image(type="filepath", label="primary image / scene")
959
  second = gr.Image(type="filepath", label="optional second reference")
 
1047
 
1048
  mode.change(on_mode_change, inputs=[mode], outputs=[image_inputs, t2i_resolution, edit_controls, edit_prompt])
1049
 
1050
+ def on_base_model_change(value: str):
1051
+ return gr.update(visible=value == CUSTOM_BASE_MODEL)
1052
+
1053
+ base_model.change(on_base_model_change, inputs=[base_model], outputs=[custom_base_controls])
1054
+
1055
  all_lora_filenames = list(lora_slider_map.keys())
1056
  all_lora_sliders = list(lora_slider_map.values())
1057
 
 
1095
  def _generate_wrapper(*values):
1096
  base_values = values[:18]
1097
  base_model_value = values[18]
1098
+ custom_repo = values[19]
1099
+ custom_filename = values[20]
1100
+ custom_revision = values[21]
1101
+ custom_rows = values[22]
1102
+ lora_weights = _catalog_weights(values[23:])
1103
  return generate(
1104
  *base_values,
1105
  base_model=base_model_value,
1106
+ custom_base_repo=custom_repo,
1107
+ custom_base_filename=custom_filename,
1108
+ custom_base_revision=custom_revision,
1109
  lora_weights=lora_weights,
1110
  custom_loras=custom_rows,
1111
  )
 
1113
  generation_inputs = [
1114
  mode, prompt, edit_prompt, primary, second, width, height, target_mp,
1115
  grounding, ref_boost, ref_boost_a, steps, cfg, sampler, scheduler,
1116
+ seed, randomize, gen_budget, base_model, custom_base_repo, custom_base_filename,
1117
+ custom_base_revision, custom_loras, *all_lora_sliders,
1118
  ]
1119
  button.click(_generate_wrapper, inputs=generation_inputs, outputs=[gallery, status, used_seed])
1120
 
1121
  profile_inputs = [
1122
  mode, prompt, edit_prompt, width, height, target_mp, grounding,
1123
  ref_boost, ref_boost_a, steps, cfg, sampler, scheduler, seed,
1124
+ randomize, gen_budget, base_model, custom_base_repo, custom_base_filename,
1125
+ custom_base_revision, custom_loras, *all_lora_sliders,
1126
  profile_name,
1127
  ]
1128
 
 
1130
  name_value = str(values[-1] or "krea2")
1131
  catalog = [
1132
  {"hf_filename": filename, "weight": float(weight or 0.0)}
1133
+ for filename, weight in zip(all_lora_filenames, values[21:-1])
1134
  if weight and abs(float(weight)) > 1e-6
1135
  ]
1136
+ normalized_custom, warnings = normalize_custom_loras(values[20])
1137
  if warnings:
1138
  return None, "settings export failed — " + "; ".join(warnings)
1139
  try:
1140
  data = _profile_from_values(
1141
  *values[:16],
1142
  values[16],
1143
+ values[17],
1144
+ values[18],
1145
+ values[19],
1146
  catalog,
1147
  normalized_custom,
1148
  )
 
1157
  export_button.click(export_profile, inputs=profile_inputs, outputs=[export_file, profile_status]).then(lambda: gr.update(visible=True), outputs=[export_file])
1158
 
1159
  settings_outputs = [
1160
+ mode, base_model, custom_base_repo, custom_base_filename, custom_base_revision,
1161
+ prompt, edit_prompt, width, height, target_mp, grounding,
1162
  ref_boost, ref_boost_a, steps, cfg, sampler, scheduler, seed,
1163
  randomize, gen_budget, custom_loras, *all_lora_sliders, profile_status, image_inputs,
1164
+ t2i_resolution, edit_controls, custom_base_controls,
1165
  ]
1166
  settings_keys = [
1167
+ "mode", "base_model", "custom_base_repo", "custom_base_filename", "custom_base_revision",
1168
+ "prompt", "edit_prompt", "width", "height",
1169
  "target_megapixels", "grounding_px", "ref_boost", "ref_boost_a",
1170
  "steps", "cfg", "sampler_name", "scheduler", "effective_seed",
1171
  "randomize_seed", "gen_budget",
1172
  ]
1173
 
1174
  def settings_updates(data: dict[str, Any], warnings: list[str], error: str = ""):
1175
+ view_data = dict(data)
1176
+ custom_base = data.get("custom_base_model")
1177
+ if custom_base:
1178
+ try:
1179
+ normalized_base = validate_custom_base_model(custom_base)
1180
+ except ValueError as exc:
1181
+ warnings.append(f"ignored custom base model: {exc}")
1182
+ else:
1183
+ view_data.update({
1184
+ "custom_base_repo": normalized_base["repo_id"],
1185
+ "custom_base_filename": normalized_base["filename"],
1186
+ "custom_base_revision": normalized_base["revision"],
1187
+ })
1188
+ editing = view_data.get("mode") == "edit"
1189
+ custom_base_visible = view_data.get("base_model") == CUSTOM_BASE_MODEL
1190
  updates = []
1191
  for key in settings_keys:
1192
+ if key not in view_data:
1193
  updates.append(gr.update())
1194
+ elif key == "base_model" and view_data[key] not in (*BASE_MODELS, CUSTOM_BASE_MODEL):
1195
+ warnings.append(f"ignored unsupported base model: {view_data[key]!r}")
1196
  updates.append(gr.update())
1197
+ elif key == "edit_prompt" and "mode" in view_data:
1198
+ updates.append(gr.update(value=view_data[key], visible=editing))
1199
  else:
1200
+ updates.append(gr.update(value=view_data[key]))
1201
  custom_rows: list[list[Any]] = []
1202
+ if "custom_loras" in view_data:
1203
+ for row in view_data.get("custom_loras", []) or []:
1204
  try:
1205
  normalized = validate_custom_lora(row)
1206
  except ValueError as exc:
 
1210
  normalized["repo_id"], normalized["filename"],
1211
  normalized["revision"], normalized["weight"],
1212
  ])
1213
+ updates.append(gr.update(value=custom_rows) if "custom_loras" in view_data else gr.update())
1214
  catalog_values = {filename: 0.0 for filename in all_lora_filenames}
1215
+ for row in view_data.get("catalog_loras", []) or []:
1216
  if not isinstance(row, dict):
1217
  warnings.append("ignored malformed catalog LoRA entry")
1218
  continue
 
1235
  message += f" (+{len(warnings) - 6} more)"
1236
  updates.extend([
1237
  message,
1238
+ gr.update(visible=editing) if "mode" in view_data else gr.update(),
1239
+ gr.update(visible=not editing) if "mode" in view_data else gr.update(),
1240
+ gr.update(visible=editing) if "mode" in view_data else gr.update(),
1241
+ gr.update(visible=custom_base_visible) if "base_model" in view_data else gr.update(),
1242
  ])
1243
  return tuple(updates)
1244
 
settings_utils.py CHANGED
@@ -14,8 +14,9 @@ from PIL import Image, PngImagePlugin
14
 
15
 
16
  APP_ID = "krea-2-turbo-i2i"
17
- PROFILE_SCHEMA_VERSION = 2
18
  CUSTOM_LORA_EXTENSIONS = {".safetensors", ".pt", ".ckpt", ".bin"}
 
19
  _HF_REPO_RE = re.compile(
20
  r"^[A-Za-z0-9][A-Za-z0-9._-]{0,95}/[A-Za-z0-9][A-Za-z0-9._-]{0,95}$"
21
  )
@@ -116,12 +117,53 @@ def normalize_custom_loras(rows: Any) -> tuple[list[dict[str, Any]], list[str]]:
116
  return normalized, warnings
117
 
118
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
119
  def stable_custom_lora_namespace(repo_id: str, filename: str, revision: str = "") -> str:
120
  """Return a filesystem-safe stable namespace for one remote LoRA."""
121
  identity = "\x00".join((repo_id, revision, filename)).encode("utf-8")
122
  return hashlib.sha256(identity).hexdigest()[:20]
123
 
124
 
 
 
 
 
 
 
125
  def build_settings(
126
  *,
127
  mode: str,
@@ -142,6 +184,7 @@ def build_settings(
142
  gen_budget: float,
143
  effective_seed: int | None = None,
144
  base_model: str = "",
 
145
  catalog_loras: list[dict[str, Any]] | None = None,
146
  custom_loras: list[dict[str, Any]] | None = None,
147
  ) -> dict[str, Any]:
@@ -150,6 +193,11 @@ def build_settings(
150
  "app": APP_ID,
151
  "schema_version": PROFILE_SCHEMA_VERSION,
152
  "base_model": _text(base_model, "base model", 256),
 
 
 
 
 
153
  "mode": _text(mode, "mode", 32),
154
  "prompt": _text(prompt, "prompt"),
155
  "edit_prompt": _text(edit_prompt, "edit prompt"),
 
14
 
15
 
16
  APP_ID = "krea-2-turbo-i2i"
17
+ PROFILE_SCHEMA_VERSION = 3
18
  CUSTOM_LORA_EXTENSIONS = {".safetensors", ".pt", ".ckpt", ".bin"}
19
+ CUSTOM_BASE_MODEL_EXTENSIONS = {".safetensors", ".pt", ".ckpt", ".bin"}
20
  _HF_REPO_RE = re.compile(
21
  r"^[A-Za-z0-9][A-Za-z0-9._-]{0,95}/[A-Za-z0-9][A-Za-z0-9._-]{0,95}$"
22
  )
 
117
  return normalized, warnings
118
 
119
 
120
+ def validate_custom_base_model(value: Any) -> dict[str, str]:
121
+ """Validate one custom Hugging Face diffusion-model reference."""
122
+ if isinstance(value, dict):
123
+ repo_id = value.get("repo_id", "")
124
+ filename = value.get("filename", value.get("file", ""))
125
+ revision = value.get("revision", "")
126
+ elif isinstance(value, (list, tuple)):
127
+ values = list(value) + [""] * 3
128
+ repo_id, filename, revision = values[:3]
129
+ else:
130
+ raise ValueError("custom base model must be an object or three-column array")
131
+
132
+ repo_id = _text(repo_id, "custom base-model repository", 193)
133
+ if not _HF_REPO_RE.fullmatch(repo_id):
134
+ raise ValueError(
135
+ f"invalid Hugging Face repository ID {repo_id!r}; expected namespace/name"
136
+ )
137
+
138
+ filename = _text(filename, "custom base-model file", 512).replace("\\", "/")
139
+ path = pathlib.PurePosixPath(filename)
140
+ if (
141
+ not filename
142
+ or path.is_absolute()
143
+ or any(part in {"", ".", ".."} for part in path.parts)
144
+ or path.suffix.lower() not in CUSTOM_BASE_MODEL_EXTENSIONS
145
+ ):
146
+ allowed = ", ".join(sorted(CUSTOM_BASE_MODEL_EXTENSIONS))
147
+ raise ValueError(f"custom base-model file must be relative and end in {allowed}")
148
+
149
+ revision = _text(revision, "custom base-model revision", 256)
150
+ if any(ord(char) < 32 for char in revision):
151
+ raise ValueError("custom base-model revision contains control characters")
152
+ return {"repo_id": repo_id, "filename": filename, "revision": revision}
153
+
154
+
155
  def stable_custom_lora_namespace(repo_id: str, filename: str, revision: str = "") -> str:
156
  """Return a filesystem-safe stable namespace for one remote LoRA."""
157
  identity = "\x00".join((repo_id, revision, filename)).encode("utf-8")
158
  return hashlib.sha256(identity).hexdigest()[:20]
159
 
160
 
161
+ def stable_custom_base_model_namespace(repo_id: str, filename: str, revision: str = "") -> str:
162
+ """Return a filesystem-safe stable namespace for one remote base model."""
163
+ identity = "\x00".join((repo_id, revision, filename)).encode("utf-8")
164
+ return hashlib.sha256(identity).hexdigest()[:20]
165
+
166
+
167
  def build_settings(
168
  *,
169
  mode: str,
 
184
  gen_budget: float,
185
  effective_seed: int | None = None,
186
  base_model: str = "",
187
+ custom_base_model: dict[str, Any] | None = None,
188
  catalog_loras: list[dict[str, Any]] | None = None,
189
  custom_loras: list[dict[str, Any]] | None = None,
190
  ) -> dict[str, Any]:
 
193
  "app": APP_ID,
194
  "schema_version": PROFILE_SCHEMA_VERSION,
195
  "base_model": _text(base_model, "base model", 256),
196
+ "custom_base_model": (
197
+ validate_custom_base_model(custom_base_model)
198
+ if custom_base_model
199
+ else None
200
+ ),
201
  "mode": _text(mode, "mode", 32),
202
  "prompt": _text(prompt, "prompt"),
203
  "edit_prompt": _text(edit_prompt, "edit prompt"),
test_app.py CHANGED
@@ -71,6 +71,32 @@ class WorkflowTests(unittest.TestCase):
71
  self.assertEqual(workflow["4"]["inputs"]["model"], ["user_lora_1", 0])
72
  self.assertEqual(workflow["5"]["inputs"]["clip"], ["user_lora_1", 1])
73
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
74
  def test_edit_user_loras_follow_identity_adapter(self):
75
  workflow = app._edit_workflow(False, app.MUSE_KREA_FILE)
76
  app._inject_edit(
 
71
  self.assertEqual(workflow["4"]["inputs"]["model"], ["user_lora_1", 0])
72
  self.assertEqual(workflow["5"]["inputs"]["clip"], ["user_lora_1", 1])
73
 
74
+ def test_custom_base_model_download_uses_managed_namespace(self):
75
+ original_models = app.MODELS
76
+ with tempfile.TemporaryDirectory() as directory:
77
+ app.MODELS = Path(directory) / "models"
78
+
79
+ def fake_download(**kwargs):
80
+ destination = Path(kwargs["local_dir"]) / kwargs["filename"]
81
+ destination.parent.mkdir(parents=True, exist_ok=True)
82
+ destination.write_bytes(b"test")
83
+ return str(destination)
84
+
85
+ try:
86
+ with patch.object(app, "hf_hub_download", side_effect=fake_download):
87
+ relative = app._ensure_custom_base_model({
88
+ "repo_id": "org/repo",
89
+ "filename": "sub/base.safetensors",
90
+ "revision": "main",
91
+ })
92
+ finally:
93
+ app.MODELS = original_models
94
+
95
+ self.assertTrue(relative.startswith("huggingface/"))
96
+ self.assertTrue(relative.endswith("/sub/base.safetensors"))
97
+ workflow = app._t2i_workflow(relative)
98
+ self.assertEqual(workflow["1"]["inputs"]["unet_name"], relative)
99
+
100
  def test_edit_user_loras_follow_identity_adapter(self):
101
  workflow = app._edit_workflow(False, app.MUSE_KREA_FILE)
102
  app._inject_edit(
test_settings_utils.py CHANGED
@@ -14,6 +14,7 @@ from settings_utils import (
14
  extract_image_settings,
15
  normalize_custom_loras,
16
  parse_settings_text,
 
17
  validate_custom_lora,
18
  write_png_metadata,
19
  )
@@ -49,6 +50,43 @@ class SettingsTests(unittest.TestCase):
49
  }],
50
  )
51
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
52
  def test_png_metadata_round_trip(self):
53
  settings = self.settings()
54
  with tempfile.TemporaryDirectory() as directory:
 
14
  extract_image_settings,
15
  normalize_custom_loras,
16
  parse_settings_text,
17
+ validate_custom_base_model,
18
  validate_custom_lora,
19
  write_png_metadata,
20
  )
 
50
  }],
51
  )
52
 
53
+ def test_custom_base_model_is_preserved_in_settings(self):
54
+ settings = build_settings(
55
+ mode="text2image",
56
+ prompt="a landscape",
57
+ edit_prompt="",
58
+ width=1024,
59
+ height=1024,
60
+ target_megapixels=1.4,
61
+ grounding_px=768,
62
+ ref_boost=1.0,
63
+ ref_boost_a=1.0,
64
+ steps=8,
65
+ cfg=1.0,
66
+ sampler_name="euler",
67
+ scheduler="beta",
68
+ seed=2,
69
+ randomize_seed=False,
70
+ gen_budget=0,
71
+ base_model="__custom_huggingface_base_model__",
72
+ custom_base_model={
73
+ "repo_id": "org/repo",
74
+ "filename": "models/base.safetensors",
75
+ "revision": "main",
76
+ },
77
+ )
78
+ self.assertEqual(settings["custom_base_model"]["repo_id"], "org/repo")
79
+ self.assertEqual(settings["custom_base_model"]["filename"], "models/base.safetensors")
80
+
81
+ def test_custom_base_model_validation_blocks_unsafe_paths(self):
82
+ valid = validate_custom_base_model(["org/repo", "base.safetensors", "main"])
83
+ self.assertEqual(valid["filename"], "base.safetensors")
84
+ with self.assertRaises(ValueError):
85
+ validate_custom_base_model({
86
+ "repo_id": "org/repo",
87
+ "filename": "../escape.safetensors",
88
+ })
89
+
90
  def test_png_metadata_round_trip(self):
91
  settings = self.settings()
92
  with tempfile.TemporaryDirectory() as directory: