naykun multimodalart HF Staff commited on
Commit
19e3cdc
·
1 Parent(s): e4eefea

[Admin maintenance] Migrate to ZeroGPU (#2)

Browse files

- [Admin maintenance] Migrate to ZeroGPU (970bb8eb02fb83fe1a4ea8742d765f77d68f7ed8)
- Run on a large ZeroGPU slice (f85981c3fc5e98fcb96ca7b75ab395a22fb137c0)
- Update app.py (91199a503ce1f7aafafff19b54a9a835c7af3b22)
- Add speed (1024px) / quality (2048px) toggle; decode 2K outputs with 1024px VAE tiles (e33ac7d1546107887d5e97b8d7972b79744e05c0)
- Show diffusion progress with gr.Progress(track_tqdm=True) (50995d8da0647cd2e0de08b4824ee16a06ba5ea1)


Co-authored-by: Apolinário from multimodal AI art <multimodalart@users.noreply.huggingface.co>

Files changed (4) hide show
  1. README.md +2 -1
  2. app.py +263 -316
  3. ncii_guard.py +110 -0
  4. requirements.txt +6 -0
README.md CHANGED
@@ -4,9 +4,10 @@ emoji: 🎨
4
  colorFrom: blue
5
  colorTo: purple
6
  sdk: gradio
7
- sdk_version: 5.33.0
8
  python_version: '3.10'
9
  app_file: app.py
 
10
  pinned: true
11
  license: other
12
  license_name: qwen-research
 
4
  colorFrom: blue
5
  colorTo: purple
6
  sdk: gradio
7
+ sdk_version: 6.28.0
8
  python_version: '3.10'
9
  app_file: app.py
10
+ startup_duration_timeout: 1h
11
  pinned: true
12
  license: other
13
  license_name: qwen-research
app.py CHANGED
@@ -1,31 +1,63 @@
 
 
 
 
 
1
  import gradio as gr
2
  import numpy as np
3
  import random
4
- import requests
5
  import io
6
  import time
7
- # import spaces
8
  import uuid
9
  from datetime import datetime
10
 
11
  from PIL import Image
12
 
13
- import os
14
  import base64
15
  import json
 
 
 
 
 
 
16
 
17
  # ============== 配置参数 ==============
18
  # 日志目录,可通过环境变量 LOG_DIR 自定义
19
  LOG_DIR = os.environ.get("LOG_DIR", "./generation_logs_paper_case")
20
 
21
- # ============== API 配置 ==============
22
- API_ENDPOINT = "https://poc-dashscope.aliyuncs.com/api/v1/services/aigc/multimodal-generation/generation"
23
- EDIT_MODEL = "pre-qwen-image-2.1-pro-yunqi"
24
- T2I_MODEL = "pre-qwen-image-2.1-pro-yunqi"
 
 
 
 
 
 
 
 
 
 
25
 
26
  MAX_INPUT_IMAGES = 10
27
  DEFAULT_LANGUAGE = "en"
28
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
29
  # ============== 预设分辨率 ==============
30
  SIZE_PRESETS_2K = {
31
  "2688x1536 (16:9)": (2688, 1536),
@@ -140,295 +172,205 @@ def validate_image_count(images):
140
  raise gr.Error("最多支持 10 张输入图片。 / Up to 10 input images are supported.")
141
 
142
 
143
- def encode_image(pil_image):
144
- buffered = io.BytesIO()
145
- pil_image.save(buffered, format="PNG")
146
- return base64.b64encode(buffered.getvalue()).decode("utf-8")
 
147
 
148
 
149
- def call_image_api(prompt, mode="t2i", images=None, size="1024*1024", negative_prompt=" ", seed=None, prompt_extend=True):
150
- """
151
- 调用 POC API 生成图片
152
-
153
- Args:
154
- prompt: 文本提示词
155
- mode: "t2i" 或 "edit"
156
- images: PIL Image 列表 (edit 模式需要)
157
- size: 尺寸字符串如 "1024*1024"
158
- negative_prompt: 负向提示词
159
- seed: 随机种子
160
- prompt_extend: 是否开启API侧提示词智能改写
161
-
162
- Returns:
163
- tuple: (PIL.Image, dict) 生成的图片和完整API响应
164
- """
165
- validate_image_count(images)
166
 
167
- # 使用独立的 IMAGE_API_KEY 环境变量
168
- api_key = os.environ.get('IMAGE_API_KEY')
169
- if not api_key:
170
- raise EnvironmentError("IMAGE_API_KEY environment variable is not set")
171
 
172
- headers = {
173
- "Authorization": f"Bearer {api_key}",
174
- "Content-Type": "application/json"
175
- }
176
 
177
- # 构造 messages
178
- content = [{"text": prompt}]
179
-
180
- if mode == "edit":
181
- model = EDIT_MODEL
182
- # 添加图片到 content
183
- if images and len(images) > 0:
184
- # 图片放在文本之前
185
- content = []
186
- content.append({"text": prompt})
187
- for img in images:
188
- img_base64 = encode_image(img)
189
- content.append({"image": f"data:image/png;base64,{img_base64}"})
190
- else:
191
- model = T2I_MODEL
192
-
193
- # 构造请求体
194
- payload = {
195
- "model": model,
196
- "input": {
197
- "messages": [
198
- {
199
- "role": "user",
200
- "content": content
201
- }
202
- ]
203
- },
204
- "parameters": {
205
- "watermark": False,
206
- "negative_prompt": negative_prompt,
207
- "prompt_extend": prompt_extend,
208
- "debug": True
209
- }
210
- }
211
-
212
- # 传入 seed 参数
213
- if seed is not None:
214
- payload["parameters"]["seed"] = int(seed)
215
-
216
- # 传入 size 参数(T2I 始终传,Edit 仅在指定时传)
217
- if size:
218
- payload["parameters"]["size"] = size
219
-
220
- print(f"[API] Calling {model} API...")
221
- print(f"[API] Prompt: {prompt[:100]}..." if len(prompt) > 100 else f"[API] Prompt: {prompt}")
222
- if mode == "edit" and images:
223
- print(f"[API] Input images count: {len(images)}")
224
 
225
- # 发送请求
226
- response = requests.post(API_ENDPOINT, headers=headers, json=payload)
227
- print(response)
228
 
229
- if response.status_code != 200:
230
- raise Exception(f"API request failed with status {response.status_code}: {response.text}")
231
 
232
- result = response.json()
233
- print(f"[API] Response: {json.dumps(result, ensure_ascii=False)[:500]}...")
 
 
 
234
 
235
- # 解析响应 - 处理同步和异步两种情况
236
- if "output" in result:
237
- output = result["output"]
238
-
239
- # 检查是否是异步任务
240
- if "task_id" in output and "task_status" in output:
241
- task_id = output["task_id"]
242
- task_status = output["task_status"]
243
- print(f"[API] Async task created: {task_id}, status: {task_status}")
244
-
245
- # 轮询等待任务完成
246
- image, async_result = poll_task_result(task_id, api_key)
247
- return image, async_result
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
248
 
249
- # 同步响应 - 直接返回结果
250
- if "choices" in output and len(output["choices"]) > 0:
251
- choice = output["choices"][0]
252
- if "message" in choice and "content" in choice["message"]:
253
- for item in choice["message"]["content"]:
254
- if "image" in item:
255
- image_url = item["image"]
256
- return download_image(image_url), result
257
 
258
- # 另一种响应格式
259
- if "results" in output and len(output["results"]) > 0:
260
- image_url = output["results"][0].get("url")
261
- if image_url:
262
- return download_image(image_url), result
263
 
264
- raise Exception(f"Failed to parse API response: {result}")
265
 
 
 
266
 
267
- def poll_task_result(task_id, api_key, max_retries=120, interval=2):
 
268
  """
269
- 轮询异步任务结果
270
-
271
- Args:
272
- task_id: 任务ID
273
- api_key: API密钥
274
- max_retries: 最大重试次数
275
- interval: 轮询间隔(秒)
276
-
277
- Returns:
278
- tuple: (PIL.Image, dict) 生成的图片和完整API响应
279
  """
280
- task_url = f"https://poc-dashscope.aliyuncs.com/api/v1/tasks/{task_id}"
281
- headers = {
282
- "Authorization": f"Bearer {api_key}"
283
- }
284
-
285
- for i in range(max_retries):
286
- response = requests.get(task_url, headers=headers)
287
- if response.status_code != 200:
288
- print(f"[API] Task polling failed: {response.status_code}")
289
- time.sleep(interval)
290
- continue
291
-
292
- result = response.json()
293
- output = result.get("output", {})
294
- task_status = output.get("task_status", "UNKNOWN")
295
-
296
- print(f"[API] Task {task_id} status: {task_status} (attempt {i+1}/{max_retries})")
297
-
298
- if task_status == "SUCCEEDED":
299
- # 获取生成的图片
300
- if "results" in output and len(output["results"]) > 0:
301
- image_url = output["results"][0].get("url")
302
- if image_url:
303
- return download_image(image_url), result
304
- raise Exception(f"Task succeeded but no image found in response: {result}")
305
 
306
- elif task_status == "FAILED":
307
- error_msg = output.get("message", "Unknown error")
308
- raise Exception(f"Task failed: {error_msg}")
309
 
310
- elif task_status in ["PENDING", "RUNNING"]:
311
- time.sleep(interval)
312
- continue
313
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
314
  else:
315
- print(f"[API] Unknown task status: {task_status}")
316
- time.sleep(interval)
317
-
318
- raise Exception(f"Task {task_id} timed out after {max_retries * interval} seconds")
319
-
320
-
321
- def download_image(url):
322
- """
323
- 从URL下载图片
324
-
325
- Args:
326
- url: 图片URL
327
-
328
- Returns:
329
- PIL.Image: 下载的图片
330
- """
331
- print(f"[API] Downloading image from: {url[:100]}...")
332
- response = requests.get(url)
333
- if response.status_code != 200:
334
- raise Exception(f"Failed to download image: {response.status_code}")
335
 
336
- with Image.open(io.BytesIO(response.content)) as image:
337
- return image.convert("RGBA")
338
-
339
-
340
-
341
-
342
- # 初始化日志记录器
343
- print(f"Initializing logger with directory: {LOG_DIR}")
344
- logger = GenerationLogger(log_dir=LOG_DIR)
345
-
346
-
347
- # --- UI Constants and Helpers ---
348
- MAX_SEED = np.iinfo(np.int32).max
349
 
350
  # --- Stage 2: Image Generation Function ---
351
  def generate_image_stage(
352
- input_images,
353
  original_prompt,
 
354
  custom_size,
355
  log_dir,
356
- seed=42,
357
- randomize_seed=False,
358
- height=1024,
359
- width=1024,
 
360
  negative_prompt=" ",
361
  prompt_extend=True,
362
  username=None,
363
- progress=gr.Progress(track_tqdm=True),
364
  ):
365
  """
366
- 使用 POC API 生成图片。
367
  根据是否有输入图片自动判断模式:有图片 -> Edit,无图片 -> T2I。
368
- 改写由 API 侧 prompt_extend 参数控制。
369
  """
370
- validate_image_count(input_images)
371
-
372
  # 更新日志目录(如果用户修改了)
373
  if log_dir and log_dir.strip():
374
  logger.set_log_dir(log_dir.strip())
375
 
376
- # 使用传入的 negative_prompt(默认为空格)
377
- if not negative_prompt or not negative_prompt.strip():
378
- negative_prompt = " "
379
-
380
- if randomize_seed:
381
- seed = random.randint(0, MAX_SEED)
382
-
383
- # Load input images into PIL Images from gallery
384
- pil_images = []
385
- if input_images is not None and len(input_images) > 0:
386
- for item in input_images:
387
- try:
388
- if isinstance(item, Image.Image):
389
- pil_images.append(item.convert("RGB"))
390
- elif isinstance(item, tuple) and len(item) > 0 and isinstance(item[0], Image.Image):
391
- pil_images.append(item[0].convert("RGB"))
392
- elif isinstance(item, str):
393
- pil_images.append(Image.open(item).convert("RGB"))
394
- elif hasattr(item, "name"):
395
- pil_images.append(Image.open(item.name).convert("RGB"))
396
- except Exception as e:
397
- print(f"[Warning] Failed to load input image: {e}")
398
- continue
399
- print(f"[API] Loaded {len(pil_images)} input images for editing")
400
-
401
- # 根据是否有输入图片自动判断模式
402
- is_edit_mode = len(pil_images) > 0
403
  mode = "edit" if is_edit_mode else "t2i"
404
 
405
- # 构造 size:仅在用户开启自定义尺寸时传,否则由 API 自动决定
406
  if custom_size:
407
- size = f"{width}*{height}"
408
  else:
409
- size = None
410
-
411
- print(f"[API] Mode: {mode} (auto-detected: {'has input images' if is_edit_mode else 'no input images'})")
412
- print(f"[API] Prompt: '{original_prompt[:100]}...'" if len(original_prompt) > 100 else f"[API] Prompt: '{original_prompt}'")
413
- print(f"[API] Negative Prompt: '{negative_prompt}'")
414
- print(f"[API] Prompt Extend: {prompt_extend}")
415
- print(f"[API] Input images count: {len(pil_images)}")
416
- print(f"[API] Seed: {seed}, Size: {size if size else 'auto (not specified)'}")
417
-
418
- # 调用 API 生成图片
419
- image, api_response = call_image_api(
420
- prompt=original_prompt,
421
- mode=mode,
422
- images=pil_images if len(pil_images) > 0 else None,
423
- size=size,
424
- negative_prompt=negative_prompt,
425
- seed=seed,
426
- prompt_extend=prompt_extend
427
- )
428
 
429
- # Only show actual API rewrite text; never substitute the original prompt.
430
- rewritten_prompt = extract_rewritten_prompt(api_response) if prompt_extend else ""
431
- enhanced_prompt = rewritten_prompt or original_prompt
432
 
433
  # 记录生成日志
434
  params = {
@@ -437,50 +379,25 @@ def generate_image_stage(
437
  "width": width,
438
  "negative_prompt": negative_prompt,
439
  "prompt_extend": prompt_extend,
440
- "input_images_count": len(pil_images),
441
- "api_model": EDIT_MODEL if is_edit_mode else T2I_MODEL
 
442
  }
443
 
444
  log_name = logger.log_generation(
445
  original_prompt=original_prompt,
446
- enhanced_prompt=enhanced_prompt,
447
  image=image,
448
  seed=seed,
449
  params=params,
450
  gpu_id=None,
451
- input_images=pil_images if len(pil_images) > 0 else None,
452
  username=username,
453
- api_response=api_response
454
  )
455
 
456
- print(f"[API] Generation complete, logged as: {log_name}")
457
-
458
- return image, seed, rewritten_prompt
459
-
460
-
461
- def extract_rewritten_prompt(api_response):
462
- """Read the text returned alongside the image when parameters.debug is true."""
463
- if not isinstance(api_response, dict):
464
- return ""
465
- output = api_response.get("output") or {}
466
- debug_info = output.get("debug_info") or {}
467
- rewrite_info = debug_info.get("rewrite_debug_info") or {}
468
- rewritten = rewrite_info.get("rewritten_prompt")
469
- if isinstance(rewritten, str) and rewritten.strip():
470
- return rewritten
471
- # Some responses expose only the prompt actually passed to generation.
472
- actual_prompt = debug_info.get("actual_prompt")
473
- if output.get("rewrite_status") == "success" and isinstance(actual_prompt, str) and actual_prompt.strip():
474
- return actual_prompt
475
- # Compatibility with responses that include rewrite text beside the image.
476
- choices = output.get("choices") or []
477
- if not choices:
478
- return ""
479
- content = (choices[0].get("message") or {}).get("content") or []
480
- return "\n\n".join(
481
- item["text"] for item in content
482
- if isinstance(item, dict) and isinstance(item.get("text"), str) and item["text"].strip()
483
- )
484
 
485
 
486
  def make_placeholder_image(text, width=512, height=320, bg_color=(30, 30, 30), text_color=(200, 200, 200)):
@@ -500,37 +417,48 @@ def make_placeholder_image(text, width=512, height=320, bg_color=(30, 30, 30), t
500
  return img
501
 
502
 
503
- # --- Unified generate_with_enhance generator ---
504
- def generate_with_enhance(
505
- input_images,
 
 
 
 
 
 
 
 
 
 
 
 
 
506
  original_prompt,
507
  enable_extend,
508
  custom_size,
509
  log_dir,
510
  seed,
511
- randomize_seed,
512
  height,
513
  width,
514
  negative_prompt,
515
  request: gr.Request,
 
516
  ):
517
  """
518
- 生成图片,prompt_extend 由 API 侧处理。
519
- Yields intermediate results so the UI updates progressively.
520
  """
 
 
521
  username = request.username if request else None
522
- print(f"[API] Request from user: {username}")
523
-
524
- yield make_placeholder_image("Generating image..."), seed, ""
525
-
526
- image, seed, rewritten_prompt = generate_image_stage(
527
- input_images, original_prompt,
528
- custom_size, log_dir, seed, randomize_seed, height, width,
529
  negative_prompt=negative_prompt,
530
  prompt_extend=enable_extend,
531
  username=username,
532
  )
533
- yield image, seed, rewritten_prompt
534
 
535
 
536
  def make_example_loader(images, text, extend):
@@ -627,14 +555,14 @@ css = """
627
  #edit_text{margin-top: -62px !important}
628
  """
629
 
630
- with gr.Blocks(title="Qwen Image 2.1 Demo", css=css) as demo:
631
  with gr.Column(elem_id="col-container"):
632
  gr.HTML('<a href="https://huggingface.co/Qwen/Qwen-Image-2.1" target="_blank"><img src="https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-Image/image2.1/logo.png" alt="Qwen-Image Logo" width="400" style="display: block; margin: 0 auto;"></a>')
633
  language = gr.Radio(choices=[("中文", "zh"), ("English", "en")], value=DEFAULT_LANGUAGE, label="语言 / Language")
634
  instructions = gr.Markdown("""
635
  ## Qwen Image 2.1 Demo 使用说明
636
  1. 如果不输入图片,默认进入文生图模式;如果输入图片,则进入图像编辑模式(支持1-10张图片)。
637
- 2. 默认"Enable Prompt Extend"为开启状态,API会自动进行提示词智能改写。如果希望直接使用原始提示词,可以关闭该选项。
638
  3. 下方提供了包括文生图,图���图的测试样例,可以作为模型基础能力的参考。
639
  4. 生成透明图时,建议提示词遵循以下格式,将中间的省略号替换为具体画面描述:
640
 
@@ -665,16 +593,21 @@ with gr.Blocks(title="Qwen Image 2.1 Demo", css=css) as demo:
665
  )
666
  with gr.Row():
667
  enable_extend = gr.Checkbox(
668
- label="Enable Prompt Extend (API侧提示词智能改写)",
669
  value=True,
670
  )
 
 
 
 
 
671
  generate_button = gr.Button("Generate Image", variant="primary")
672
 
673
  rewritten_prompt_output = localize(
674
  gr.Textbox(value="", lines=4, max_lines=20, interactive=False),
675
  label=("改写结果", "Rewritten prompt"),
676
- placeholder=("开启智能改写并生成图片后,API 返回的改写提示词会显示在这里;未返回时留空。",
677
- "After generation with prompt enhancement enabled, the API rewrite appears here. Blank if none is returned."),
678
  )
679
  enable_extend.change(fn=lambda: "", inputs=[], outputs=[rewritten_prompt_output], queue=False)
680
 
@@ -691,8 +624,6 @@ with gr.Blocks(title="Qwen Image 2.1 Demo", css=css) as demo:
691
  visible=False,
692
  )
693
 
694
- # gr.Markdown(f"**当前使用 API 模式** (Edit: {EDIT_MODEL}, T2I: {T2I_MODEL})")
695
-
696
  seed = gr.Slider(
697
  label="Seed",
698
  minimum=0,
@@ -715,7 +646,7 @@ with gr.Blocks(title="Qwen Image 2.1 Demo", css=css) as demo:
715
  custom_size = gr.Checkbox(
716
  label="自定义输出尺寸 (Customize output size)",
717
  value=False,
718
- info="关闭时由 API/模型自动决定输出尺寸;开启后使用下方分辨率设置",
719
  )
720
 
721
  # 分辨率预设选择
@@ -814,6 +745,9 @@ with gr.Blocks(title="Qwen Image 2.1 Demo", css=css) as demo:
814
  localize(result, label=("生成结果", "Result"))
815
  localize(prompt, label=("提示词", "Prompt"), placeholder=("描述想生成或编辑的内容,可按上传顺序引用第 1–10 张图…", "Describe what to generate or edit; refer to images 1–10 in upload order…"))
816
  localize(enable_extend, label=("智能改写提示词", "Enhance prompt"))
 
 
 
817
  localize(generate_button, value=("生成图片", "Generate image"))
818
  localize(advanced, label=("高级设置", "Advanced settings"))
819
  localize(log_dir_input, label=("日志保存目录", "Log directory"), placeholder=("输入日志保存目录", "Enter a log directory"))
@@ -834,27 +768,39 @@ with gr.Blocks(title="Qwen Image 2.1 Demo", css=css) as demo:
834
  language.change(fn=switch_language, inputs=[language], outputs=localized_components, queue=False)
835
  input_images.upload(fn=validate_image_count, inputs=[input_images], outputs=[], queue=False)
836
 
837
- # Generate Image button event (with optional prompt enhancement)
 
 
838
  generate_button.click(
839
- fn=generate_with_enhance,
840
  inputs=[
841
  input_images, # input_images (有图片->Edit模式, 无图片->T2I模式)
842
  prompt, # original_prompt
843
- enable_extend, # enable_extend (是否开启API侧改写)
844
  custom_size, # custom_size (是否自定义输出尺寸)
845
- log_dir_input, # log_dir
846
  seed,
847
  randomize_seed,
 
 
 
 
 
 
 
 
 
 
 
 
848
  height,
849
  width,
850
  negative_prompt_input, # negative_prompt (负向提示词)
851
  ],
852
- outputs=[result, seed, rewritten_prompt_output],
853
- concurrency_limit=2
854
  )
855
 
856
- demo.queue(default_concurrency_limit=10, max_size=20)
857
-
858
  if __name__ == "__main__":
859
  # 使用 os.path.realpath 解析真实路径,避免 NAS 软链接导致 Gradio 文件权限检查失败(404)
860
  _script_dir = os.path.realpath(os.path.dirname(os.path.abspath(__file__)))
@@ -865,4 +811,5 @@ if __name__ == "__main__":
865
  demo.launch(
866
  server_name="0.0.0.0",
867
  allowed_paths=_allowed,
 
868
  )
 
1
+ import os
2
+
3
+ os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
4
+
5
+ import spaces
6
  import gradio as gr
7
  import numpy as np
8
  import random
 
9
  import io
10
  import time
 
11
  import uuid
12
  from datetime import datetime
13
 
14
  from PIL import Image
15
 
 
16
  import base64
17
  import json
18
+ import re
19
+
20
+ import torch
21
+ from diffusers import QwenImage21Pipeline
22
+
23
+ import ncii_guard
24
 
25
  # ============== 配置参数 ==============
26
  # 日志目录,可通过环境变量 LOG_DIR 自定义
27
  LOG_DIR = os.environ.get("LOG_DIR", "./generation_logs_paper_case")
28
 
29
+ # ============== 模型配置 ==============
30
+ MODEL_ID = os.environ.get("QWEN_IMAGE_MODEL", "Qwen/Qwen-Image-2.1")
31
+ # Prompt rewriting (Qwen-Image-2.1-PE-T2I / PE-I2I) runs in a companion Space so this
32
+ # Space only holds the diffusion pipeline and fits a `large` ZeroGPU slice.
33
+ PE_SPACE_ID = os.environ.get("PE_SPACE_ID", "hugging-apps/qwen-image-2-1-prompt-enhancer")
34
+ PE_MAX_NEW_TOKENS = {"t2i": 1536, "i2i": 2048}
35
+ GUARD_THRESHOLD = 0.5
36
+ NUM_INFERENCE_STEPS = 28
37
+ QUALITY_RESOLUTIONS = {"speed": 1024, "quality": 2048}
38
+ TRUE_CFG_SCALE = 4.0
39
+ # The prefix KV cache costs ~2 GB per 1K condition image; above this budget it is
40
+ # switched off so many-image edits still fit next to the weights.
41
+ KV_CACHE_BYTES_PER_TOKEN = 32 * 2 * 4096 * 2
42
+ KV_CACHE_BUDGET_GB = 10.0
43
 
44
  MAX_INPUT_IMAGES = 10
45
  DEFAULT_LANGUAGE = "en"
46
 
47
+ pipe = QwenImage21Pipeline.from_pretrained(MODEL_ID, dtype=torch.bfloat16)
48
+ pipe.to("cuda")
49
+ # Decode large outputs in tiles so a 2K VAE decode fits next to the weights.
50
+ pipe.vae.enable_tiling(
51
+ tile_sample_min_height=1536,
52
+ tile_sample_min_width=1536,
53
+ tile_sample_stride_height=1152,
54
+ tile_sample_stride_width=1152,
55
+ )
56
+
57
+ # The NCII classifier runs in a CPU subprocess: a transformers forward in the main
58
+ # process breaks the ZeroGPU worker fork.
59
+ ncii_guard.start()
60
+
61
  # ============== 预设分辨率 ==============
62
  SIZE_PRESETS_2K = {
63
  "2688x1536 (16:9)": (2688, 1536),
 
172
  raise gr.Error("最多支持 10 张输入图片。 / Up to 10 input images are supported.")
173
 
174
 
175
+ def check_prompt_guard(prompt):
176
+ """Reject image-editing prompts the NCII classifier flags. The error is
177
+ deliberately generic and does not say which classifier fired."""
178
+ if ncii_guard.score(prompt or "") >= GUARD_THRESHOLD:
179
+ raise gr.Error("prompt invalid based on our classifiers, try again")
180
 
181
 
182
+ def ratio_to_size(ratio, area=1024 * 1024, multiple=32):
183
+ """Turn a rewriter aspect ratio like "16:9" into a height/width pair of about `area` pixels."""
184
+ match = re.fullmatch(r"\s*(\d+(?:\.\d+)?)\s*:\s*(\d+(?:\.\d+)?)\s*", str(ratio or ""))
185
+ if not match or float(match.group(2)) == 0:
186
+ return None, None
187
+ aspect = float(match.group(1)) / float(match.group(2))
188
+ if not 1 / 4 <= aspect <= 4:
189
+ return None, None
190
+ return aspect_to_size(aspect, area, multiple)
 
 
 
 
 
 
 
 
191
 
 
 
 
 
192
 
193
+ def aspect_to_size(aspect, area, multiple=32):
194
+ width = round((area * aspect) ** 0.5 / multiple) * multiple
195
+ height = round((area / aspect) ** 0.5 / multiple) * multiple
196
+ return height, width
197
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
198
 
199
+ _pe_client = None
 
 
200
 
 
 
201
 
202
+ def enhance_prompt(prompt, image_paths):
203
+ """Rewrite the prompt with the official PE models in the companion Space.
204
+ Returns (rewritten_prompt, wh_ratio)."""
205
+ global _pe_client
206
+ from gradio_client import Client, handle_file
207
 
208
+ if _pe_client is None:
209
+ _pe_client = Client(PE_SPACE_ID, token=os.environ.get("HF_TOKEN"), httpx_kwargs={"timeout": 900}, verbose=False)
210
+ rewritten, wh_ratio, *_ = _pe_client.predict(
211
+ prompt=prompt,
212
+ image_paths=[handle_file(p) for p in image_paths],
213
+ max_new_tokens=PE_MAX_NEW_TOKENS["i2i" if image_paths else "t2i"],
214
+ enable_thinking=False,
215
+ seed=0,
216
+ randomize_seed=True,
217
+ api_name="/enhance",
218
+ )
219
+ return str(rewritten or "").strip(), str(wh_ratio or "").strip()
220
+
221
+
222
+ def kv_cache_fits(n_images):
223
+ return n_images * (1024 // 16) ** 2 * KV_CACHE_BYTES_PER_TOKEN <= KV_CACHE_BUDGET_GB * 1e9
224
+
225
+
226
+ def generation_duration(prompt, image_paths, height, width, negative_prompt, seed):
227
+ """GPU seconds, fitted on this pipeline (eager, large ZeroGPU): step cost scales as
228
+ latent_tokens ** 1.363 and each cached condition image adds a fifth of its tokens."""
229
+ n_images = len(image_paths or [])
230
+ pixels = (int(width) * int(height)) if (width and height) else 1024 * 1024
231
+ prefix = n_images * (1024 // 16) ** 2
232
+ cached = kv_cache_fits(n_images)
233
+ tokens = pixels / 256 + (0.2 * prefix if cached else prefix)
234
+ per_step = 3.4e-6 * tokens ** 1.363 * 1.3
235
+ if negative_prompt:
236
+ per_step *= 2
237
+ fixed = 4 + 2 * n_images + 1.5e-6 * pixels
238
+ return int(min(300, (fixed + NUM_INFERENCE_STEPS * per_step) * 1.25))
239
+
240
+
241
+ @spaces.GPU(duration=generation_duration)
242
+ def run_pipeline(prompt, image_paths, height, width, negative_prompt, seed):
243
+ images = [Image.open(p) for p in image_paths] or None
244
+ kwargs = {"use_kv_cache": kv_cache_fits(len(image_paths))}
245
+ tile, stride = (1024, 768) if max(height or 0, width or 0) > 1536 else (1536, 1152)
246
+ pipe.vae.enable_tiling(
247
+ tile_sample_min_height=tile,
248
+ tile_sample_min_width=tile,
249
+ tile_sample_stride_height=stride,
250
+ tile_sample_stride_width=stride,
251
+ )
252
+ if negative_prompt:
253
+ kwargs.update(negative_prompt=negative_prompt, true_cfg_scale=TRUE_CFG_SCALE)
254
+ return pipe(
255
+ prompt,
256
+ image=images,
257
+ height=height,
258
+ width=width,
259
+ num_inference_steps=NUM_INFERENCE_STEPS,
260
+ generator=torch.Generator("cuda").manual_seed(int(seed)),
261
+ **kwargs,
262
+ ).images[0]
263
+
264
+
265
+ def gallery_paths(input_images):
266
+ """Save gallery images to PNG files so they can be sent to the rewriter Space and
267
+ reopened in the GPU worker without loss (RGBA inputs keep their alpha)."""
268
+ import tempfile
269
+
270
+ paths = []
271
+ for item in input_images or []:
272
+ try:
273
+ if isinstance(item, (tuple, list)) and item:
274
+ item = item[0]
275
+ if isinstance(item, str):
276
+ item = Image.open(item)
277
+ elif hasattr(item, "name") and not isinstance(item, Image.Image):
278
+ item = Image.open(item.name)
279
+ if not isinstance(item, Image.Image):
280
+ continue
281
+ if item.mode not in ("RGB", "RGBA"):
282
+ item = item.convert("RGBA")
283
+ path = tempfile.NamedTemporaryFile(suffix=".png", delete=False).name
284
+ item.save(path)
285
+ paths.append(path)
286
+ except Exception as e:
287
+ print(f"[Warning] Failed to load input image: {e}")
288
+ return paths
289
 
 
 
 
 
 
 
 
 
290
 
291
+ # 初始化日志记录器
292
+ print(f"Initializing logger with directory: {LOG_DIR}")
293
+ logger = GenerationLogger(log_dir=LOG_DIR)
 
 
294
 
 
295
 
296
+ # --- UI Constants and Helpers ---
297
+ MAX_SEED = np.iinfo(np.int32).max
298
 
299
+ # --- Stage 1: screen and rewrite the prompt (CPU, no GPU quota) ---
300
+ def prepare_stage(input_images, original_prompt, enable_extend, custom_size, quality, seed, randomize_seed):
301
  """
302
+ 校验输入、NCII 检查(仅在有输入图片时)、可选的提示词改写,全部在 GPU 之外完成。
303
+ Returns (image_paths, final_prompt, rewritten_prompt, seed, auto_height, auto_width).
 
 
 
 
 
 
 
 
304
  """
305
+ validate_image_count(input_images)
306
+ if not original_prompt or not original_prompt.strip():
307
+ raise gr.Error("请输入提示词。 / Please enter a prompt.")
308
+ image_paths = gallery_paths(input_images)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
309
 
310
+ if image_paths:
311
+ check_prompt_guard(original_prompt)
 
312
 
313
+ if randomize_seed:
314
+ seed = random.randint(0, MAX_SEED)
 
315
 
316
+ rewritten_prompt, wh_ratio = "", ""
317
+ if enable_extend:
318
+ try:
319
+ rewritten_prompt, wh_ratio = enhance_prompt(original_prompt, image_paths)
320
+ except Exception as e:
321
+ print(f"[Warning] Prompt enhancement failed: {e!r}")
322
+ gr.Warning("提示词改写暂不可用,已使用原始提示词。 / Prompt enhancement is unavailable; using the original prompt.")
323
+
324
+ auto_height, auto_width = (None, None)
325
+ if not custom_size:
326
+ area = QUALITY_RESOLUTIONS.get(quality, 1024) ** 2
327
+ if image_paths:
328
+ last_width, last_height = Image.open(image_paths[-1]).size
329
+ auto_height, auto_width = aspect_to_size(last_width / last_height, area)
330
  else:
331
+ auto_height, auto_width = ratio_to_size(wh_ratio, area)
332
+ if auto_height is None:
333
+ auto_height, auto_width = aspect_to_size(1.0, area)
334
+ return image_paths, rewritten_prompt or original_prompt, rewritten_prompt, seed, auto_height, auto_width
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
335
 
 
 
 
 
 
 
 
 
 
 
 
 
 
336
 
337
  # --- Stage 2: Image Generation Function ---
338
  def generate_image_stage(
339
+ image_paths,
340
  original_prompt,
341
+ final_prompt,
342
  custom_size,
343
  log_dir,
344
+ seed,
345
+ height,
346
+ width,
347
+ auto_height,
348
+ auto_width,
349
  negative_prompt=" ",
350
  prompt_extend=True,
351
  username=None,
 
352
  ):
353
  """
 
354
  根据是否有输入图片自动判断模式:有图片 -> Edit,无图片 -> T2I。
 
355
  """
 
 
356
  # 更新日志目录(如果用户修改了)
357
  if log_dir and log_dir.strip():
358
  logger.set_log_dir(log_dir.strip())
359
 
360
+ negative_prompt = (negative_prompt or "").strip()
361
+ is_edit_mode = len(image_paths) > 0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
362
  mode = "edit" if is_edit_mode else "t2i"
363
 
364
+ # 构造 size:开启自定义尺寸时使用设置值,否则由改写模型推荐的比例或模型自动决定
365
  if custom_size:
366
+ out_height, out_width = int(height), int(width)
367
  else:
368
+ out_height, out_width = auto_height, auto_width
369
+
370
+ print(f"Mode: {mode}, prompt extend: {prompt_extend}, input images: {len(image_paths)}, "
371
+ f"seed: {seed}, size: {f'{out_width}x{out_height}' if out_width else 'auto'}")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
372
 
373
+ image = run_pipeline(final_prompt, image_paths, out_height, out_width, negative_prompt, seed)
 
 
374
 
375
  # 记录生成日志
376
  params = {
 
379
  "width": width,
380
  "negative_prompt": negative_prompt,
381
  "prompt_extend": prompt_extend,
382
+ "input_images_count": len(image_paths),
383
+ "model": MODEL_ID,
384
+ "num_inference_steps": NUM_INFERENCE_STEPS,
385
  }
386
 
387
  log_name = logger.log_generation(
388
  original_prompt=original_prompt,
389
+ enhanced_prompt=final_prompt,
390
  image=image,
391
  seed=seed,
392
  params=params,
393
  gpu_id=None,
394
+ input_images=[Image.open(p) for p in image_paths] or None,
395
  username=username,
 
396
  )
397
 
398
+ print(f"Generation complete, logged as: {log_name}")
399
+
400
+ return image
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
401
 
402
 
403
  def make_placeholder_image(text, width=512, height=320, bg_color=(30, 30, 30), text_color=(200, 200, 200)):
 
417
  return img
418
 
419
 
420
+ # --- Two-step flow: prepare (CPU) then generate (GPU) ---
421
+ def prepare_request(input_images, original_prompt, enable_extend, custom_size, quality, seed, randomize_seed):
422
+ image_paths, final_prompt, rewritten_prompt, seed, auto_height, auto_width = prepare_stage(
423
+ input_images, original_prompt, enable_extend, custom_size, quality, seed, randomize_seed,
424
+ )
425
+ request_state = {
426
+ "image_paths": image_paths,
427
+ "final_prompt": final_prompt,
428
+ "auto_height": auto_height,
429
+ "auto_width": auto_width,
430
+ }
431
+ return make_placeholder_image("Generating image..."), seed, rewritten_prompt, request_state
432
+
433
+
434
+ def generate_request(
435
+ request_state,
436
  original_prompt,
437
  enable_extend,
438
  custom_size,
439
  log_dir,
440
  seed,
 
441
  height,
442
  width,
443
  negative_prompt,
444
  request: gr.Request,
445
+ progress=gr.Progress(track_tqdm=True),
446
  ):
447
  """
448
+ 生成图片。提示词改写已在上一步完成(开启时)。
 
449
  """
450
+ if not request_state:
451
+ raise gr.Error("请重新点击生成。 / Please click generate again.")
452
  username = request.username if request else None
453
+ print(f"Request from user: {username}")
454
+ return generate_image_stage(
455
+ request_state["image_paths"], original_prompt, request_state["final_prompt"],
456
+ custom_size, log_dir, seed, height, width,
457
+ request_state["auto_height"], request_state["auto_width"],
 
 
458
  negative_prompt=negative_prompt,
459
  prompt_extend=enable_extend,
460
  username=username,
461
  )
 
462
 
463
 
464
  def make_example_loader(images, text, extend):
 
555
  #edit_text{margin-top: -62px !important}
556
  """
557
 
558
+ with gr.Blocks(title="Qwen Image 2.1 Demo") as demo:
559
  with gr.Column(elem_id="col-container"):
560
  gr.HTML('<a href="https://huggingface.co/Qwen/Qwen-Image-2.1" target="_blank"><img src="https://qianwen-res.oss-cn-beijing.aliyuncs.com/Qwen-Image/image2.1/logo.png" alt="Qwen-Image Logo" width="400" style="display: block; margin: 0 auto;"></a>')
561
  language = gr.Radio(choices=[("中文", "zh"), ("English", "en")], value=DEFAULT_LANGUAGE, label="语言 / Language")
562
  instructions = gr.Markdown("""
563
  ## Qwen Image 2.1 Demo 使用说明
564
  1. 如果不输入图片,默认进入文生图模式;如果输入图片,则进入图像编辑模式(支持1-10张图片)。
565
+ 2. 默认"Enable Prompt Extend"为开启状态,会自动进行提示词智能改写。如果希望直接使用原始提示词,可以关闭该选项。
566
  3. 下方提供了包括文生图,图���图的测试样例,可以作为模型基础能力的参考。
567
  4. 生成透明图时,建议提示词遵循以下格式,将中间的省略号替换为具体画面描述:
568
 
 
593
  )
594
  with gr.Row():
595
  enable_extend = gr.Checkbox(
596
+ label="Enable Prompt Extend (提示词智能改写)",
597
  value=True,
598
  )
599
+ quality = gr.Radio(
600
+ choices=[("Speed (1024px)", "speed"), ("Quality (2048px)", "quality")],
601
+ value="speed",
602
+ show_label=False,
603
+ )
604
  generate_button = gr.Button("Generate Image", variant="primary")
605
 
606
  rewritten_prompt_output = localize(
607
  gr.Textbox(value="", lines=4, max_lines=20, interactive=False),
608
  label=("改写结果", "Rewritten prompt"),
609
+ placeholder=("开启智能改写并生成图片后,改写后的提示词会显示在这里。",
610
+ "After generation with prompt enhancement enabled, the rewritten prompt appears here."),
611
  )
612
  enable_extend.change(fn=lambda: "", inputs=[], outputs=[rewritten_prompt_output], queue=False)
613
 
 
624
  visible=False,
625
  )
626
 
 
 
627
  seed = gr.Slider(
628
  label="Seed",
629
  minimum=0,
 
646
  custom_size = gr.Checkbox(
647
  label="自定义输出尺寸 (Customize output size)",
648
  value=False,
649
+ info="关闭时由模型自动决定输出尺寸;开启后使用下方分辨率设置",
650
  )
651
 
652
  # 分辨率预设选择
 
745
  localize(result, label=("生成结果", "Result"))
746
  localize(prompt, label=("提示词", "Prompt"), placeholder=("描述想生成或编辑的内容,可按上传顺序引用第 1–10 张图…", "Describe what to generate or edit; refer to images 1–10 in upload order…"))
747
  localize(enable_extend, label=("智能改写提示词", "Enhance prompt"))
748
+ localize(quality, choices=(
749
+ [("速度 (1024px)", "speed"), ("质量 (2048px)", "quality")],
750
+ [("Speed (1024px)", "speed"), ("Quality (2048px)", "quality")]))
751
  localize(generate_button, value=("生成图片", "Generate image"))
752
  localize(advanced, label=("高级设置", "Advanced settings"))
753
  localize(log_dir_input, label=("日志保存目录", "Log directory"), placeholder=("输入日志保存目录", "Enter a log directory"))
 
768
  language.change(fn=switch_language, inputs=[language], outputs=localized_components, queue=False)
769
  input_images.upload(fn=validate_image_count, inputs=[input_images], outputs=[], queue=False)
770
 
771
+ # Generate Image button event: screen/rewrite the prompt off-GPU, then generate.
772
+ # Two events so the GPU step is scheduled with a fresh ZeroGPU token after a long rewrite.
773
+ request_state = gr.State(None)
774
  generate_button.click(
775
+ fn=prepare_request,
776
  inputs=[
777
  input_images, # input_images (有图片->Edit模式, 无图片->T2I模式)
778
  prompt, # original_prompt
779
+ enable_extend, # enable_extend (是否开启提示词改写)
780
  custom_size, # custom_size (是否自定义输出尺寸)
781
+ quality,
782
  seed,
783
  randomize_seed,
784
+ ],
785
+ outputs=[result, seed, rewritten_prompt_output, request_state],
786
+ concurrency_limit=2,
787
+ ).success(
788
+ fn=generate_request,
789
+ inputs=[
790
+ request_state,
791
+ prompt,
792
+ enable_extend,
793
+ custom_size,
794
+ log_dir_input, # log_dir
795
+ seed,
796
  height,
797
  width,
798
  negative_prompt_input, # negative_prompt (负向提示词)
799
  ],
800
+ outputs=[result],
801
+ concurrency_limit=2,
802
  )
803
 
 
 
804
  if __name__ == "__main__":
805
  # 使用 os.path.realpath 解析真实路径,避免 NAS 软链接导致 Gradio 文件权限检查失败(404)
806
  _script_dir = os.path.realpath(os.path.dirname(os.path.abspath(__file__)))
 
811
  demo.launch(
812
  server_name="0.0.0.0",
813
  allowed_paths=_allowed,
814
+ css=css,
815
  )
ncii_guard.py ADDED
@@ -0,0 +1,110 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ import os
5
+ import select
6
+ import subprocess
7
+ import sys
8
+ import tempfile
9
+ import threading
10
+
11
+ GUARD_ID = "hfmlsoc/ncii-guard-v02"
12
+
13
+ _lock = threading.Lock()
14
+ _process: subprocess.Popen | None = None
15
+
16
+
17
+ def _read(timeout: float) -> dict:
18
+ readable, _, _ = select.select([_process.stdout], [], [], timeout)
19
+ if not readable:
20
+ raise TimeoutError(f"the guard did not answer within {timeout}s")
21
+ line = _process.stdout.readline()
22
+ if not line:
23
+ raise EOFError("the guard process died")
24
+ return json.loads(line)
25
+
26
+
27
+ def _spawn() -> None:
28
+ global _process
29
+ _process = subprocess.Popen(
30
+ [sys.executable, os.path.abspath(__file__)],
31
+ stdin=subprocess.PIPE,
32
+ stdout=subprocess.PIPE,
33
+ text=True,
34
+ bufsize=1,
35
+ )
36
+ ready = _read(600.0)
37
+ if ready.get("status") != "ready":
38
+ raise RuntimeError(f"guard failed to start: {ready}")
39
+
40
+
41
+ def start() -> None:
42
+ with _lock:
43
+ _spawn()
44
+
45
+
46
+ def score(prompt: str, timeout: float = 60.0) -> float:
47
+ with _lock:
48
+ for attempt in (0, 1):
49
+ try:
50
+ if _process is None or _process.poll() is not None:
51
+ _spawn()
52
+ _process.stdin.write(json.dumps({"prompt": prompt or ""}) + "\n")
53
+ _process.stdin.flush()
54
+ return float(_read(timeout)["score"])
55
+ except Exception:
56
+ if attempt:
57
+ raise
58
+ if _process is not None and _process.poll() is None:
59
+ _process.kill()
60
+
61
+
62
+ def _snapshot() -> str:
63
+ from huggingface_hub import snapshot_download
64
+
65
+ src = snapshot_download(
66
+ GUARD_ID,
67
+ allow_patterns=["config.json", "tokenizer.json", "tokenizer_config.json", "model.safetensors"],
68
+ )
69
+ dst = os.path.join(tempfile.gettempdir(), "ncii-guard-v02-normalised")
70
+ os.makedirs(dst, exist_ok=True)
71
+ for name in os.listdir(src):
72
+ link = os.path.join(dst, name)
73
+ if not os.path.exists(link):
74
+ os.symlink(os.path.realpath(os.path.join(src, name)), link)
75
+
76
+ config = json.load(open(os.path.join(src, "config.json")))
77
+ rope = config.get("rope_parameters")
78
+ if isinstance(rope, dict):
79
+ per_layer = {k: v for k, v in rope.items() if isinstance(v, dict)}
80
+ if per_layer:
81
+ config["rope_parameters"] = per_layer
82
+ config_path = os.path.join(dst, "config.json")
83
+ if os.path.islink(config_path):
84
+ os.remove(config_path)
85
+ with open(config_path, "w") as fh:
86
+ json.dump(config, fh)
87
+ return dst
88
+
89
+
90
+ def _serve() -> None:
91
+ protocol = os.fdopen(os.dup(1), "w", buffering=1)
92
+ os.dup2(2, 1)
93
+
94
+ import torch
95
+ from transformers import AutoModelForSequenceClassification, AutoTokenizer
96
+
97
+ guard_dir = _snapshot()
98
+ tokenizer = AutoTokenizer.from_pretrained(guard_dir)
99
+ model = AutoModelForSequenceClassification.from_pretrained(guard_dir, dtype=torch.float32).eval()
100
+ protocol.write(json.dumps({"status": "ready", "id2label": model.config.id2label}) + "\n")
101
+ for line in sys.stdin:
102
+ text = json.loads(line)["prompt"]
103
+ batch = tokenizer(text, truncation=True, max_length=256, padding=True, return_tensors="pt")
104
+ with torch.no_grad():
105
+ logits = model(**batch).logits.float()
106
+ protocol.write(json.dumps({"score": torch.softmax(logits, dim=-1)[0, 1].item()}) + "\n")
107
+
108
+
109
+ if __name__ == "__main__":
110
+ _serve()
requirements.txt ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ torchvision
2
+ diffusers @ git+https://github.com/huggingface/diffusers.git
3
+ transformers @ git+https://github.com/huggingface/transformers.git
4
+ accelerate
5
+ safetensors
6
+ pillow