Spaces:
Sleeping
Sleeping
Download app.py from Quantumbraid/qwenprep: direct link, hf CLI and curl.
- Browser
- Download file 20 kB
-
https://huggingface.co/spaces/Quantumbraid/qwenprep/resolve/main/app.py
- Command line
-
hf download hf://spaces/Quantumbraid/qwenprep/app.py
-
curl -L -o app.py https://huggingface.co/spaces/Quantumbraid/qwenprep/resolve/main/app.py
20 kB
| import gc | |
| import json | |
| import os | |
| import re | |
| import shutil | |
| import subprocess | |
| import time | |
| from pathlib import Path | |
| import gradio as gr | |
| import spaces | |
| import torch | |
| from PIL import Image | |
| DATA = Path(os.environ.get("DINING_CART_DATA", "/data")) | |
| PROJECTS = DATA / "projects" | |
| LEGACY_RUNS = DATA / "runs" | |
| PERSISTENT_MODEL_CACHE = DATA / "models" / "hf-cache" | |
| LOCAL_MODEL_CACHE = Path(os.environ.get("DINING_CART_LOCAL_CACHE", "/tmp/dining-cart-hf-cache")) | |
| IMAGE_EXTS = {".png", ".jpg", ".jpeg", ".webp"} | |
| RAPID_REPO = "prithivMLmods/Qwen-Image-Edit-Rapid-AIO-V19" | |
| BASE_REPO = "Qwen/Qwen-Image-Edit-2511" | |
| CAPTION_REPO = "Qwen/Qwen2.5-VL-7B-Instruct" | |
| os.environ.setdefault("HF_HOME", str(LOCAL_MODEL_CACHE)) | |
| os.environ.setdefault("HF_HUB_CACHE", str(LOCAL_MODEL_CACHE / "hub")) | |
| DEFAULT_EDIT_PROMPT = "Create a polished, photorealistic training example that preserves the subject's exact identity and distinguishing details while naturally varying pose, viewpoint, expression, clothing, lighting, and setting. Correct blur, compression, anatomy, hands, eyes, and facial detail." | |
| QWEN_CAPTION_PROMPT = "Describe this image as a precise, information-dense image-generation prompt for Qwen Image training. Include subject appearance, pose, expression, clothing, composition, camera, lighting, environment, materials, and fine details. Output only the prompt." | |
| SD_CAPTION_PROMPT = "Write a concise Stable Diffusion training caption. Use direct comma-separated visual tags and short natural phrases. Describe the visible subject, clothing, pose, composition, camera, lighting, background, style, and important objects. Output only the caption." | |
| _pipe = None | |
| _pipe_fast = None | |
| _caption_model = None | |
| _caption_processor = None | |
| def sync_model_cache(source, destination, skip_blobs=False): | |
| if not source.exists(): | |
| return | |
| for source_path in source.rglob("*"): | |
| if not source_path.is_file() or ".locks" in source_path.parts: | |
| continue | |
| relative = source_path.relative_to(source) | |
| if skip_blobs and "blobs" in relative.parts: | |
| continue | |
| destination_path = destination / relative | |
| source_size = source_path.stat().st_size | |
| if destination_path.exists() and destination_path.stat().st_size == source_size: | |
| continue | |
| destination_path.parent.mkdir(parents=True, exist_ok=True) | |
| shutil.copy2(source_path, destination_path) | |
| def stage_model_cache(): | |
| LOCAL_MODEL_CACHE.mkdir(parents=True, exist_ok=True) | |
| PERSISTENT_MODEL_CACHE.mkdir(parents=True, exist_ok=True) | |
| sync_model_cache(PERSISTENT_MODEL_CACHE, LOCAL_MODEL_CACHE, skip_blobs=True) | |
| def persist_model_cache(): | |
| sync_model_cache(LOCAL_MODEL_CACHE, PERSISTENT_MODEL_CACHE) | |
| def safe_name(value): | |
| value = re.sub(r"[^a-zA-Z0-9._-]+", "-", (value or "").strip()).strip("-._") | |
| if not value: | |
| raise gr.Error("Give this project a name.") | |
| return value[:80].lower() | |
| def project_path(project_id): | |
| return PROJECTS / safe_name(project_id) | |
| def ensure_project(project_id): | |
| project = project_path(project_id) | |
| for name in ("source", "target", "captions", "videos"): | |
| (project / name).mkdir(parents=True, exist_ok=True) | |
| metadata = project / "project.json" | |
| if not metadata.exists(): | |
| metadata.write_text(json.dumps({"project": project.name, "created": int(time.time()), "format": 1}, indent=2), encoding="utf-8") | |
| return project | |
| def unique_destination(folder, filename): | |
| stem = safe_name(Path(filename).stem) | |
| suffix = Path(filename).suffix.lower() or ".png" | |
| candidate = folder / f"{stem}{suffix}" | |
| index = 2 | |
| while candidate.exists(): | |
| candidate = folder / f"{stem}-{index:03d}{suffix}" | |
| index += 1 | |
| return candidate | |
| def image_files(folder): | |
| if not folder.exists(): | |
| return [] | |
| return sorted(p for p in folder.iterdir() if p.is_file() and p.suffix.lower() in IMAGE_EXTS) | |
| def write_caption_index(folder): | |
| """Keep a resumable caption manifest beside the images as well as sidecars.""" | |
| captions = {} | |
| for image_path in image_files(folder): | |
| sidecar = image_path.with_suffix(".txt") | |
| if sidecar.exists(): | |
| text = sidecar.read_text(encoding="utf-8").strip() | |
| if text: | |
| captions[image_path.name] = text | |
| (folder / "captions.json").write_text(json.dumps(captions, indent=2, ensure_ascii=False), encoding="utf-8") | |
| def project_summary(project): | |
| return { | |
| "project": project.name, | |
| "starting_images": len(image_files(project / "source")), | |
| "finished_examples": len(image_files(project / "target")), | |
| "captions": len(list((project / "captions").glob("*.txt"))), | |
| } | |
| def list_projects(): | |
| PROJECTS.mkdir(parents=True, exist_ok=True) | |
| return sorted((p.name for p in PROJECTS.iterdir() if p.is_dir()), reverse=True) | |
| def refresh_projects(current=""): | |
| choices = list_projects() | |
| value = current if current in choices else (choices[0] if choices else None) | |
| return gr.update(choices=choices, value=value), (project_summary(project_path(value)) if value else {}) | |
| def create_project(name): | |
| project = ensure_project(name) | |
| choices = list_projects() | |
| return gr.update(choices=choices, value=project.name), project_summary(project) | |
| def inspect_project(project_id): | |
| if not project_id: | |
| return {}, [], [] | |
| project = project_path(project_id) | |
| return project_summary(project), [str(p) for p in image_files(project / "source")], [str(p) for p in image_files(project / "target")] | |
| def add_images(project_id, uploads, destination): | |
| if not uploads: | |
| raise gr.Error("Choose one or more images.") | |
| project = ensure_project(project_id) | |
| folder = project / ("source" if destination == "Starting images" else "target") | |
| added = [] | |
| for upload in uploads: | |
| source = Path(upload) | |
| target = unique_destination(folder, source.name) | |
| with Image.open(source) as image: | |
| image.convert("RGB").save(target, quality=96) | |
| added.append(str(target)) | |
| return f"Added {len(added)} image(s) to {destination.lower()}.", project_summary(project), [str(p) for p in image_files(folder)] | |
| def extract_video(project_id, video, interval, max_frames): | |
| if not video: | |
| raise gr.Error("Choose a video.") | |
| ffmpeg = shutil.which("ffmpeg") | |
| if not ffmpeg: | |
| raise gr.Error("FFmpeg is unavailable in this Space image.") | |
| project = ensure_project(project_id) | |
| video_source = Path(video) | |
| stored_video = unique_destination(project / "videos", video_source.name) | |
| shutil.copy2(video_source, stored_video) | |
| prefix = safe_name(video_source.stem) | |
| pattern = project / "source" / f"{prefix}-frame-%05d.png" | |
| command = [ffmpeg, "-hide_banner", "-loglevel", "error", "-i", str(stored_video), | |
| "-vf", f"fps=1/{max(0.1, float(interval))}", "-frames:v", str(int(max_frames)), | |
| "-vsync", "vfr", str(pattern)] | |
| subprocess.run(command, check=True, timeout=1200) | |
| frames = sorted((project / "source").glob(f"{prefix}-frame-*.png")) | |
| if not frames: | |
| raise gr.Error("FFmpeg did not extract any frames.") | |
| return f"Extracted {len(frames)} starting images.", project_summary(project), [str(p) for p in image_files(project / "source")] | |
| def release_models(): | |
| global _pipe, _pipe_fast, _caption_model, _caption_processor | |
| _pipe = None | |
| _pipe_fast = None | |
| _caption_model = None | |
| _caption_processor = None | |
| gc.collect() | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| def get_pipe(fast): | |
| global _pipe, _pipe_fast, _caption_model, _caption_processor | |
| if _pipe is not None and _pipe_fast == fast: | |
| return _pipe | |
| release_models() | |
| stage_model_cache() | |
| from qwenimage.pipeline_qwenimage_edit_plus import QwenImageEditPlusPipeline | |
| from qwenimage.transformer_qwenimage import QwenImageTransformer2DModel | |
| transformer = QwenImageTransformer2DModel.from_pretrained( | |
| RAPID_REPO if fast else BASE_REPO, | |
| subfolder=None if fast else "transformer", | |
| torch_dtype=torch.bfloat16, | |
| device_map="cuda", | |
| cache_dir=str(LOCAL_MODEL_CACHE), | |
| ) | |
| _pipe = QwenImageEditPlusPipeline.from_pretrained( | |
| BASE_REPO, transformer=transformer, torch_dtype=torch.bfloat16, | |
| cache_dir=str(LOCAL_MODEL_CACHE), | |
| ).to("cuda") | |
| persist_model_cache() | |
| _pipe.set_progress_bar_config(disable=False) | |
| _pipe_fast = fast | |
| return _pipe | |
| def generate_batch(project_id, selected_sources, prompt, count_per_image, fast, seed, group_controls): | |
| if not selected_sources: | |
| raise gr.Error("Select at least one starting image from the gallery.") | |
| project = ensure_project(project_id) | |
| pipe = get_pipe(bool(fast)) | |
| outputs = [] | |
| prompt = (prompt or DEFAULT_EDIT_PROMPT).strip() | |
| steps = 4 if fast else 25 | |
| guidance = 1.0 if fast else 4.0 | |
| pairings_path = project / "pairings.json" | |
| try: | |
| pairings = json.loads(pairings_path.read_text(encoding="utf-8")) if pairings_path.exists() else {} | |
| except json.JSONDecodeError: | |
| pairings = {} | |
| source_paths = [Path(value[0] if isinstance(value, (list, tuple)) else value) for value in selected_sources] | |
| batches = [source_paths] if group_controls else [[path] for path in source_paths] | |
| for batch_index, batch_sources in enumerate(batches): | |
| controls = [] | |
| for source_path in batch_sources: | |
| with Image.open(source_path) as image: | |
| controls.append(image.convert("RGB")) | |
| label = "-".join(path.stem for path in batch_sources[:2]) | |
| for variation in range(int(count_per_image)): | |
| current_seed = int(seed) + batch_index * 1000 + variation | |
| generator = torch.Generator(device="cuda").manual_seed(current_seed) | |
| result = pipe( | |
| image=controls, prompt=prompt, negative_prompt=" ", | |
| num_inference_steps=steps, true_cfg_scale=guidance, | |
| generator=generator, | |
| ).images[0] | |
| destination = unique_destination(project / "target", f"{label}-generated-{variation + 1:02d}.png") | |
| result.save(destination) | |
| outputs.append(str(destination)) | |
| pairings[destination.name] = [path.name for path in batch_sources] | |
| pairings_path.write_text(json.dumps(pairings, indent=2), encoding="utf-8") | |
| return outputs, f"Generated {len(outputs)} finished example(s) with {'Rapid AIO V19' if fast else 'standard Qwen Edit 2511'}.", project_summary(project) | |
| def get_captioner(): | |
| global _pipe, _pipe_fast, _caption_model, _caption_processor | |
| if _caption_model is not None: | |
| return _caption_model, _caption_processor | |
| release_models() | |
| stage_model_cache() | |
| from transformers import AutoProcessor, Qwen2_5_VLForConditionalGeneration | |
| _caption_processor = AutoProcessor.from_pretrained(CAPTION_REPO, cache_dir=str(LOCAL_MODEL_CACHE)) | |
| _caption_model = Qwen2_5_VLForConditionalGeneration.from_pretrained( | |
| CAPTION_REPO, torch_dtype=torch.bfloat16, device_map="cuda", attn_implementation="sdpa", | |
| cache_dir=str(LOCAL_MODEL_CACHE), | |
| ).eval() | |
| persist_model_cache() | |
| return _caption_model, _caption_processor | |
| def caption_image(image, instruction): | |
| model, processor = get_captioner() | |
| messages = [{"role": "user", "content": [{"type": "image", "image": image.convert("RGB")}, {"type": "text", "text": instruction}]}] | |
| inputs = processor.apply_chat_template(messages, tokenize=True, add_generation_prompt=True, return_dict=True, return_tensors="pt") | |
| inputs = {key: value.to(model.device) if hasattr(value, "to") else value for key, value in inputs.items()} | |
| with torch.no_grad(): | |
| generated = model.generate(**inputs, max_new_tokens=384, do_sample=False) | |
| trimmed = generated[:, inputs["input_ids"].shape[1]:] | |
| return processor.batch_decode(trimmed, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0].strip() | |
| def caption_project(project_id, image_set, caption_style, custom_instruction, overwrite): | |
| project = ensure_project(project_id) | |
| folder = project / ("target" if image_set == "Finished examples" else "source") | |
| images = image_files(folder) | |
| if not images: | |
| raise gr.Error(f"This project has no {image_set.lower()}.") | |
| instruction = (custom_instruction or "").strip() | |
| if not instruction: | |
| instruction = QWEN_CAPTION_PROMPT if caption_style == "Qwen / WAN native prompt" else SD_CAPTION_PROMPT | |
| written = 0 | |
| skipped = 0 | |
| for image_path in images: | |
| caption_path = project / "captions" / f"{image_path.stem}.txt" | |
| sidecar = image_path.with_suffix(".txt") | |
| if not overwrite and caption_path.exists() and caption_path.read_text(encoding="utf-8").strip(): | |
| skipped += 1 | |
| continue | |
| with Image.open(image_path) as image: | |
| caption = caption_image(image, instruction) | |
| if not caption: | |
| raise gr.Error(f"Qwen returned an empty caption for {image_path.name}.") | |
| caption_path.write_text(caption, encoding="utf-8") | |
| sidecar.write_text(caption, encoding="utf-8") | |
| written += 1 | |
| write_caption_index(folder) | |
| return {"project": project.name, "captioned": written, "skipped": skipped, "style": caption_style}, project_summary(project) | |
| def caption_run(run_id, instruction=SD_CAPTION_PROMPT, overwrite=False): | |
| """Compatibility endpoint for the existing Stable Diffusion Caboose.""" | |
| run = LEGACY_RUNS / safe_name(run_id) | |
| image_dir = run / "images" | |
| generated_dir = run / "captions" / "generated" | |
| final_dir = run / "captions" / "final" | |
| generated_dir.mkdir(parents=True, exist_ok=True) | |
| final_dir.mkdir(parents=True, exist_ok=True) | |
| images = image_files(image_dir) | |
| if not images: | |
| raise gr.Error("No images found for this Caboose run.") | |
| written = 0 | |
| skipped = 0 | |
| for image_path in images: | |
| output = generated_dir / f"{image_path.stem}.txt" | |
| final = final_dir / f"{image_path.stem}.txt" | |
| if not overwrite and ((output.exists() and output.read_text(encoding="utf-8").strip()) or (final.exists() and final.read_text(encoding="utf-8").strip())): | |
| skipped += 1 | |
| continue | |
| with Image.open(image_path) as image: | |
| output.write_text(caption_image(image, instruction or SD_CAPTION_PROMPT), encoding="utf-8") | |
| written += 1 | |
| status = {"status": "captions_ready", "total": len(images), "newly_captioned": written, "skipped": skipped} | |
| (run / "status.json").write_text(json.dumps(status), encoding="utf-8") | |
| return status | |
| def select_gallery(evt: gr.SelectData, current): | |
| current = list(current or []) | |
| value = evt.value | |
| path = value.get("image", {}).get("path") if isinstance(value, dict) else value | |
| if path and path not in current: | |
| current.append(path) | |
| return current, f"Selected {len(current)} starting image(s)." | |
| def clear_selection(): | |
| return [], "No starting images selected." | |
| with gr.Blocks(title="Dining Cart") as demo: | |
| selected_sources = gr.State([]) | |
| gr.Markdown("# 🚃 Dining Cart\nReusable dataset preparation for Caboose, Qwen, WAN, and Stable Diffusion. Everything is private in the shared `/data` workspace; nothing is published to the community.") | |
| with gr.Row(): | |
| project = gr.Dropdown(list_projects(), label="Project", allow_custom_value=True) | |
| refresh_button = gr.Button("Refresh projects") | |
| new_name = gr.Textbox(label="New project name", placeholder="amber-portraits") | |
| create_button = gr.Button("Create project", variant="primary") | |
| summary = gr.JSON(label="Project contents") | |
| with gr.Tab("Bring in material"): | |
| with gr.Accordion("Add images", open=True): | |
| uploads = gr.File(file_count="multiple", file_types=["image"], type="filepath", label="Images") | |
| destination = gr.Radio(["Starting images", "Finished examples"], value="Starting images", label="Put them in") | |
| add_button = gr.Button("Add images") | |
| add_status = gr.Markdown() | |
| with gr.Accordion("Extract a video", open=False): | |
| video = gr.Video(label="Video") | |
| interval = gr.Number(value=2.0, minimum=0.1, label="Seconds between frames") | |
| max_frames = gr.Slider(1, 500, value=100, step=1, label="Maximum frames") | |
| extract_button = gr.Button("Extract starting images") | |
| extract_status = gr.Markdown() | |
| source_gallery = gr.Gallery(label="Starting images", columns=5, height="auto", object_fit="contain") | |
| with gr.Tab("Create finished examples"): | |
| gr.Markdown("Click starting images in the gallery to select them for batch generation.") | |
| generation_sources = gr.Gallery(label="Choose starting images", columns=5, height="auto", object_fit="contain") | |
| selection_status = gr.Markdown("No starting images selected.") | |
| clear_button = gr.Button("Clear selection") | |
| edit_prompt = gr.Textbox(value=DEFAULT_EDIT_PROMPT, lines=6, label="Qwen editing instruction") | |
| with gr.Row(): | |
| count_per_image = gr.Slider(1, 8, value=1, step=1, label="Finished examples per starting image") | |
| fast_generation = gr.Checkbox(value=True, label="Fast 4-step generation (Rapid AIO V19)") | |
| group_controls = gr.Checkbox(value=True, label="Use selected images together as controls", info="Passes multiple photographs at once so Qwen can separate identity from clothing, light, and background.") | |
| seed = gr.Number(value=42, precision=0, label="Seed") | |
| generate_button = gr.Button("Generate finished examples", variant="primary") | |
| generate_status = gr.Markdown() | |
| target_gallery = gr.Gallery(label="Finished examples", columns=5, height="auto", object_fit="contain") | |
| with gr.Tab("Caption"): | |
| with gr.Row(): | |
| caption_set = gr.Radio(["Finished examples", "Starting images"], value="Finished examples", label="Images to caption") | |
| caption_style = gr.Radio(["Qwen / WAN native prompt", "Stable Diffusion caption / tags"], value="Qwen / WAN native prompt", label="Caption language") | |
| overwrite = gr.Checkbox(value=False, label="Replace existing captions") | |
| custom_instruction = gr.Textbox(label="Optional custom instruction", lines=5, placeholder="Leave blank to use the selected caption language.") | |
| caption_button = gr.Button("Caption project", variant="primary") | |
| caption_result = gr.JSON(label="Caption result") | |
| refresh_button.click(refresh_projects, [project], [project, summary]) | |
| create_button.click(create_project, [new_name], [project, summary]) | |
| project.change(inspect_project, [project], [summary, source_gallery, target_gallery]).then(inspect_project, [project], [summary, generation_sources, target_gallery]) | |
| add_button.click(add_images, [project, uploads, destination], [add_status, summary, source_gallery]).then(inspect_project, [project], [summary, generation_sources, target_gallery]) | |
| extract_button.click(extract_video, [project, video, interval, max_frames], [extract_status, summary, source_gallery]).then(inspect_project, [project], [summary, generation_sources, target_gallery]) | |
| generation_sources.select(select_gallery, [selected_sources], [selected_sources, selection_status]) | |
| clear_button.click(clear_selection, None, [selected_sources, selection_status]) | |
| generate_button.click(generate_batch, [project, selected_sources, edit_prompt, count_per_image, fast_generation, seed, group_controls], [target_gallery, generate_status, summary]) | |
| caption_button.click(caption_project, [project, caption_set, caption_style, custom_instruction, overwrite], [caption_result, summary]) | |
| demo.launch(allowed_paths=[str(DATA)]) | |