import os import json import time import uuid import socket import traceback from datetime import datetime, timezone from pathlib import Path import torch import spaces import gradio as gr from transformers import AutoProcessor, PreTrainedTokenizerBase import torch.backends.cuda torch.backends.cuda.enable_flash_sdp(False) torch.backends.cuda.enable_mem_efficient_sdp(False) from models.qwen3_vl_transformers import Qwen3VLForConditionalGeneration from models.pipeline import generate_image, DEFAULT_TIMESTEPS from prompt_agent_v2 import refine_prompt MODEL_REPO_ID = "HiDream-ai/HiDream-O1-Image-Dev-2604" MODEL_DIR = (Path(__file__).resolve().parent / "models" / "gen_model").resolve() MODEL_PATH = ( (os.environ.get("LOCAL_MODEL_PATH") or "").strip() or (os.environ.get("MODEL_PATH") or "").strip() or (str(MODEL_DIR) if MODEL_DIR.exists() else MODEL_REPO_ID) ) HEIGHT = int(os.environ.get("IMAGE_HEIGHT", "1024")) WIDTH = int(os.environ.get("IMAGE_WIDTH", "1024")) DEFAULT_SEED = 42 NUM_INFERENCE_STEPS = int(os.environ.get("NUM_INFERENCE_STEPS", "20")) GUIDANCE_SCALE = 0.0 SHIFT = 1.0 SCHEDULER_NAME = os.environ.get("SCHEDULER_NAME", "flow_match") NOISE_SCALE_START = float(os.environ.get("NOISE_SCALE_START", "8.0")) NOISE_SCALE_END = float(os.environ.get("NOISE_SCALE_END", "8.0")) NOISE_CLIP_STD = float(os.environ.get("NOISE_CLIP_STD", "8.0")) MODEL_DTYPE_NAME = os.environ.get("MODEL_DTYPE", "float16").lower().strip() _DTYPE_MAP = { "float16": torch.float16, "fp16": torch.float16, "bfloat16": torch.bfloat16, "bf16": torch.bfloat16, "float32": torch.float32, "fp32": torch.float32, } MODEL_DTYPE = _DTYPE_MAP.get(MODEL_DTYPE_NAME, torch.float16) # Directory to persist every generation (image + metadata sidecar). # Falls back to ./data when /data is not writable (e.g. local dev). _DEFAULT_OUTPUT_DIR = "/data" try: Path(_DEFAULT_OUTPUT_DIR).mkdir(parents=True, exist_ok=True) _probe = Path(_DEFAULT_OUTPUT_DIR) / ".write_test" _probe.touch() _probe.unlink() OUTPUT_DIR = Path(_DEFAULT_OUTPUT_DIR) except Exception: OUTPUT_DIR = Path(os.environ.get("OUTPUT_DIR", "./data")).resolve() OUTPUT_DIR.mkdir(parents=True, exist_ok=True) print(f"[app] Saving generations to {OUTPUT_DIR}") REFINE_BASE_URL = os.environ.get("PROMPT_REFINE_BASE_URL", "") REFINE_API_KEY = os.environ.get("PROMPT_REFINE_API_KEY", "") REFINE_MODEL_NAME = os.environ.get("PROMPT_REFINE_MODEL", "") def _get_tokenizer(processor): if isinstance(processor, PreTrainedTokenizerBase): return processor return processor.tokenizer def _add_special_tokens(tokenizer): tokenizer.boi_token = "<|boi_token|>" tokenizer.bor_token = "<|bor_token|>" tokenizer.eor_token = "<|eor_token|>" tokenizer.bot_token = "<|bot_token|>" tokenizer.tms_token = "<|tms_token|>" print(f"[app] Loading processor and model from {MODEL_PATH}") processor = AutoProcessor.from_pretrained(MODEL_PATH) if not torch.cuda.is_available(): raise RuntimeError( "A CUDA GPU is required for independent local inference. " "Upgrade this Space hardware (recommended: zero-a10g or a10g-large)." ) model = ( Qwen3VLForConditionalGeneration.from_pretrained( MODEL_PATH, torch_dtype=MODEL_DTYPE ) .eval() .to("cuda") ) tokenizer = _get_tokenizer(processor) _add_special_tokens(tokenizer) print("[app] Model loaded.") @spaces.GPU(duration=300) def _run_generation(prompt: str, seed: int): return generate_image( model=model, processor=processor, prompt=prompt, ref_image_paths=[], height=HEIGHT, width=WIDTH, num_inference_steps=NUM_INFERENCE_STEPS, guidance_scale=GUIDANCE_SCALE, shift=SHIFT, timesteps_list=DEFAULT_TIMESTEPS, scheduler_name=SCHEDULER_NAME, seed=int(seed), keep_original_aspect=False, layout_bboxes=None, noise_scale_start=NOISE_SCALE_START, noise_scale_end=NOISE_SCALE_END, noise_clip_std=NOISE_CLIP_STD, ) def _rewrite_prompt(prompt: str) -> str: if not REFINE_BASE_URL or not REFINE_MODEL_NAME: return prompt return refine_prompt( prompt, model_id=REFINE_MODEL_NAME, base_url=REFINE_BASE_URL, api_key=REFINE_API_KEY or "EMPTY", ) def _make_record_id() -> str: # Sortable timestamp + short random suffix; safe on every filesystem. now = datetime.now(timezone.utc) return f"{now.strftime('%Y%m%dT%H%M%S')}_{uuid.uuid4().hex[:8]}" def _save_record( record_id: str, image, *, original_prompt: str, final_prompt: str, seed: int, use_rewrite: bool, rewrite_error: str | None, started_at: float, rewrite_started_at: float | None, rewrite_ended_at: float | None, generation_started_at: float, generation_ended_at: float, ended_at: float, status: str, error: str | None = None, ) -> tuple[Path, Path]: image_path = OUTPUT_DIR / f"{record_id}.png" meta_path = OUTPUT_DIR / f"{record_id}.json" if image is not None: try: image.save(image_path, format="PNG") except Exception as exc: error = error or f"failed to save image: {exc}" status = "image_save_failed" def _iso(ts: float | None) -> str | None: if ts is None: return None return datetime.fromtimestamp(ts, tz=timezone.utc).isoformat() metadata = { "id": record_id, "status": status, "error": error, "host": socket.gethostname(), "image_file": image_path.name if image is not None else None, "metadata_file": meta_path.name, "request": { "original_prompt": original_prompt, "seed": int(seed), "use_rewrite": bool(use_rewrite), }, "prompt": { "original": original_prompt, "final": final_prompt, "rewritten": use_rewrite and final_prompt != original_prompt, "rewrite_error": rewrite_error, }, "model": { "image_model_path": MODEL_PATH, "rewrite_model": REFINE_MODEL_NAME or None, "rewrite_base_url": REFINE_BASE_URL or None, }, "generation_params": { "height": HEIGHT, "width": WIDTH, "num_inference_steps": NUM_INFERENCE_STEPS, "guidance_scale": GUIDANCE_SCALE, "shift": SHIFT, "scheduler_name": SCHEDULER_NAME, "noise_scale_start": NOISE_SCALE_START, "noise_scale_end": NOISE_SCALE_END, "noise_clip_std": NOISE_CLIP_STD, }, "timing": { "started_at": _iso(started_at), "ended_at": _iso(ended_at), "rewrite_started_at": _iso(rewrite_started_at), "rewrite_ended_at": _iso(rewrite_ended_at), "generation_started_at": _iso(generation_started_at), "generation_ended_at": _iso(generation_ended_at), "rewrite_seconds": ( rewrite_ended_at - rewrite_started_at if rewrite_started_at and rewrite_ended_at else None ), "generation_seconds": generation_ended_at - generation_started_at if generation_ended_at and generation_started_at else None, "total_seconds": ended_at - started_at, }, } if image is not None: try: metadata["image"] = { "format": "PNG", "size": list(image.size), "mode": image.mode, } except Exception: pass try: with open(meta_path, "w", encoding="utf-8") as fh: json.dump(metadata, fh, ensure_ascii=False, indent=2) except Exception as exc: print(f"[app] Failed to write metadata for {record_id}: {exc}") return image_path, meta_path def text_to_image(prompt: str, seed: int, use_rewrite: bool): if not prompt or not prompt.strip(): raise gr.Error("Please enter a prompt.") record_id = _make_record_id() started_at = time.time() rewrite_started_at = None rewrite_ended_at = None rewrite_error = None final_prompt = prompt image = None try: if use_rewrite: rewrite_started_at = time.time() try: final_prompt = _rewrite_prompt(prompt) except Exception as exc: rewrite_error = f"{type(exc).__name__}: {exc}" final_prompt = prompt finally: rewrite_ended_at = time.time() generation_started_at = time.time() image = _run_generation(final_prompt, seed) generation_ended_at = time.time() ended_at = time.time() _save_record( record_id, image, original_prompt=prompt, final_prompt=final_prompt, seed=seed, use_rewrite=use_rewrite, rewrite_error=rewrite_error, started_at=started_at, rewrite_started_at=rewrite_started_at, rewrite_ended_at=rewrite_ended_at, generation_started_at=generation_started_at, generation_ended_at=generation_ended_at, ended_at=ended_at, status="ok", ) return image, final_prompt except Exception as exc: ended_at = time.time() try: _save_record( record_id, image, original_prompt=prompt, final_prompt=final_prompt, seed=seed, use_rewrite=use_rewrite, rewrite_error=rewrite_error, started_at=started_at, rewrite_started_at=rewrite_started_at, rewrite_ended_at=rewrite_ended_at, generation_started_at=started_at, generation_ended_at=ended_at, ended_at=ended_at, status="error", error=f"{type(exc).__name__}: {exc}\n{traceback.format_exc()}", ) except Exception as save_exc: print(f"[app] Failed to persist error record {record_id}: {save_exc}") raise CUSTOM_CSS = """ .page-footer { margin-top: 32px; padding: 20px 0 8px 0; border-top: 1px solid var(--border-color-primary, #e5e7eb); text-align: center; } .page-footer .footer-links a { margin: 0 12px; text-decoration: none; font-weight: 500; } .page-footer .tagline { margin-top: 8px; font-size: 0.9em; opacity: 0.75; } """ with gr.Blocks(title="HiDream-O1-Image-Dev-2604", css=CUSTOM_CSS) as demo: gr.Markdown("# HiDream-O1-Image-Dev-2604\nA minimal text-to-image demo.") with gr.Row(): with gr.Column(): prompt = gr.Textbox( label="Prompt", lines=6, placeholder="Describe the image you want to generate...", ) seed = gr.Number( label="Seed", value=DEFAULT_SEED, precision=0, ) use_rewrite = gr.Checkbox( label="Rewrite prompt before generation", value=False, ) run_btn = gr.Button("Generate", variant="primary") with gr.Column(): output_image = gr.Image(label="Output", type="pil") final_prompt = gr.Textbox( label="Final prompt used", lines=6, interactive=False, ) gr.HTML( """ """ ) run_btn.click( fn=text_to_image, inputs=[prompt, seed, use_rewrite], outputs=[output_image, final_prompt], ) prompt.submit( fn=text_to_image, inputs=[prompt, seed, use_rewrite], outputs=[output_image, final_prompt], ) if __name__ == "__main__": demo.launch()