Spaces:
Running on Zero
Running on Zero
feat(models): add support for custom Hugging Face base checkpoints
Browse filesAllow 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.
- README.md +8 -1
- app.py +190 -41
- settings_utils.py +49 -1
- test_app.py +26 -0
- 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) ->
|
| 268 |
-
"""Download
|
| 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 |
-
) ->
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
) ->
|
| 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[
|
| 625 |
-
custom_loras = kwargs.get("custom_loras", args[
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|
| 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),
|
| 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=[
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 977 |
-
|
|
|
|
|
|
|
|
|
|
| 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,
|
|
|
|
| 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,
|
|
|
|
| 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[
|
| 1004 |
if weight and abs(float(weight)) > 1e-6
|
| 1005 |
]
|
| 1006 |
-
normalized_custom, warnings = normalize_custom_loras(values[
|
| 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,
|
|
|
|
| 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", "
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1041 |
updates = []
|
| 1042 |
for key in settings_keys:
|
| 1043 |
-
if key not in
|
| 1044 |
updates.append(gr.update())
|
| 1045 |
-
elif key == "base_model" and
|
| 1046 |
-
warnings.append(f"ignored unsupported base model: {
|
| 1047 |
updates.append(gr.update())
|
| 1048 |
-
elif key == "edit_prompt" and "mode" in
|
| 1049 |
-
updates.append(gr.update(value=
|
| 1050 |
else:
|
| 1051 |
-
updates.append(gr.update(value=
|
| 1052 |
custom_rows: list[list[Any]] = []
|
| 1053 |
-
if "custom_loras" in
|
| 1054 |
-
for row in
|
| 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
|
| 1065 |
catalog_values = {filename: 0.0 for filename in all_lora_filenames}
|
| 1066 |
-
for row in
|
| 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
|
| 1090 |
-
gr.update(visible=not editing) if "mode" in
|
| 1091 |
-
gr.update(visible=editing) if "mode" in
|
| 1092 |
-
gr.update(visible=
|
| 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 =
|
| 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:
|