qwenprep / app.py
Quantumbraid's picture
Persist caption indexes for resumable projects
2dbffef verified
Raw History Blame Contribute Delete
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
@spaces.GPU(size="xlarge", duration=900)
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()
@spaces.GPU(size="large", duration=900)
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)
@spaces.GPU(size="large", duration=900)
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)])