From e83dd07dc0e21dac96bbf43685597f85e12feb39 Mon Sep 17 00:00:00 2001 From: ila Date: Mon, 22 Jun 2026 18:10:54 +0800 Subject: [PATCH] feat: use image concurrency for outfit directories --- src/app/widgets/ai_outfit_panel.py | 3 +- src/core/ai_outfit.py | 81 +++++++++++++++++++++--------- src/core/outfit_batch.py | 23 +++++---- tasks.md | 14 +++--- tests/test_ai_outfit.py | 39 ++++++++++++++ tests/test_ai_outfit_panel.py | 9 ++++ tests/test_outfit_batch.py | 7 ++- 7 files changed, 131 insertions(+), 45 deletions(-) diff --git a/src/app/widgets/ai_outfit_panel.py b/src/app/widgets/ai_outfit_panel.py index 55afd0e..ec6208d 100644 --- a/src/app/widgets/ai_outfit_panel.py +++ b/src/app/widgets/ai_outfit_panel.py @@ -112,6 +112,7 @@ class _OutfitWorker(QObject): task, self._prompt, self._output_dir, self._model_config, quality=self._quality, resolution=self._resolution, request_interval=self._options.request_interval, + image_concurrency=self._options.concurrency, image_log=self.log.emit, ) @@ -337,7 +338,7 @@ class AiOutfitPanel(QWidget): self._quality.addItems(_QUALITIES) pairs = [ - ("并发数", self._concurrency), ("新请求间隔", self._interval), + ("图片并发数", self._concurrency), ("新请求间隔", self._interval), ("单任务冷却", self._cooldown), ("失败重试", self._retry_count), ("分辨率", self._resolution), ("JPG 质量", self._quality), ] diff --git a/src/core/ai_outfit.py b/src/core/ai_outfit.py index 97b00d5..4f8dc9d 100644 --- a/src/core/ai_outfit.py +++ b/src/core/ai_outfit.py @@ -1,7 +1,9 @@ import logging import os import re +import threading import time +from concurrent.futures import ThreadPoolExecutor, as_completed from io import BytesIO from pathlib import Path @@ -144,6 +146,7 @@ def generate_outfit_image( resolution="1K", api_client=None, request_interval=0.0, + image_concurrency=1, image_log=None, ): """Generate outfit image(s) for one Excel row and return OutfitResult. @@ -160,7 +163,8 @@ def generate_outfit_image( return _generate_directory_outfit( task, prompt_template, output_dir, model_config, quality=quality, resolution=resolution, api_client=api_client, - request_interval=request_interval, image_log=image_log, + request_interval=request_interval, image_concurrency=image_concurrency, + image_log=image_log, ) try: @@ -198,6 +202,7 @@ def _generate_directory_outfit( resolution, api_client, request_interval, + image_concurrency, image_log, ): """Fan one directory row out into per-image generations (docs/11 §4.1/§9.1). @@ -228,34 +233,43 @@ def _generate_directory_outfit( subdir_name = os.path.basename(os.path.normpath(directory)) subdir_path = str(Path(output_dir) / safe_product_filename(subdir_name)) prompt = render_prompt(prompt_template, task, resolution=resolution) - client = api_client or ImageApiClient(model_config) interval = float(request_interval or 0.0) + workers = max(1, int(image_concurrency or 1)) + limiter = _RequestStartLimiter(interval) + + def generate_one(index, image_path): + output_path = make_outfit_subdir_path(output_dir, subdir_name, image_path.stem) + if output_path.exists(): + emit("第 {} 行 第 {}/{} 张已存在,跳过:{}".format( + task.row_index, index, total, image_path.name)) + return True, str(output_path), "" + try: + limiter.wait() + client = api_client or ImageApiClient(model_config) + image_bytes = client.generate(prompt, str(image_path), resolution=resolution) + save_jpg_under_limit(image_bytes, output_path, quality=quality) + emit("第 {} 行 第 {}/{} 张完成:{}".format( + task.row_index, index, total, image_path.name)) + return True, str(output_path), "" + except Exception as exc: # noqa: BLE001 - record and continue + emit("第 {} 行 第 {}/{} 张失败:{}({})".format( + task.row_index, index, total, image_path.name, exc)) + return False, "", "{}:{}".format(image_path.name, exc) total = len(images) outputs = [] failures = [] - called = False - for index, image_path in enumerate(images, start=1): - output_path = make_outfit_subdir_path(output_dir, subdir_name, image_path.stem) - if output_path.exists(): - outputs.append(str(output_path)) - emit("第 {} 行 第 {}/{} 张已存在,跳过:{}".format( - task.row_index, index, total, image_path.name)) - continue - try: - if called and interval > 0: - time.sleep(interval) - image_bytes = client.generate(prompt, str(image_path), resolution=resolution) - called = True - save_jpg_under_limit(image_bytes, output_path, quality=quality) - outputs.append(str(output_path)) - emit("第 {} 行 第 {}/{} 张完成:{}".format( - task.row_index, index, total, image_path.name)) - except Exception as exc: # noqa: BLE001 - record and continue - called = True - failures.append("{}:{}".format(image_path.name, exc)) - emit("第 {} 行 第 {}/{} 张失败:{}({})".format( - task.row_index, index, total, image_path.name, exc)) + with ThreadPoolExecutor(max_workers=workers) as executor: + futures = [ + executor.submit(generate_one, index, image_path) + for index, image_path in enumerate(images, start=1) + ] + for future in as_completed(futures): + ok, output, error = future.result() + if ok: + outputs.append(output) + else: + failures.append(error) if failures: error = "{} 张中 {} 张失败:{}".format(total, len(failures), ";".join(failures)) @@ -278,6 +292,25 @@ def _coerce_quality(quality): return QUALITY_BALANCED +class _RequestStartLimiter: + """Thread-safe minimum interval between AI request starts.""" + + def __init__(self, interval_seconds): + self.interval_seconds = max(0.0, float(interval_seconds or 0.0)) + self._lock = threading.Lock() + self._next_time = 0.0 + + def wait(self): + if self.interval_seconds <= 0: + return + with self._lock: + now = time.monotonic() + if now < self._next_time: + time.sleep(self._next_time - now) + now = time.monotonic() + self._next_time = now + self.interval_seconds + + def _crop_square(img): width, height = img.size side = min(width, height) diff --git a/src/core/outfit_batch.py b/src/core/outfit_batch.py index 15e3249..7b36846 100644 --- a/src/core/outfit_batch.py +++ b/src/core/outfit_batch.py @@ -19,6 +19,8 @@ LogCallback = Callable[[str], None] class OutfitBatchOptions: """Batch generation controls for AI outfit tasks.""" + # Kept as "concurrency" for app_config compatibility. Business meaning is + # now image concurrency inside one directory row; Excel row concurrency is 1. concurrency: int = 1 request_interval: float = 0.0 task_cooldown: float = 0.0 @@ -60,7 +62,7 @@ class RateLimiter: class OutfitBatchRunner: - """Run outfit tasks with concurrency, rate limiting, retries and stop.""" + """Run Excel rows sequentially with retries, heartbeat and gentle stop.""" def __init__( self, @@ -80,11 +82,6 @@ class OutfitBatchRunner: self._sleep = sleep_func or time.sleep self._time = time_func or time.monotonic self._stop_event = threading.Event() - self._limiter = RateLimiter( - self.options.request_interval, - now_func=self._time, - sleep_func=self._sleep, - ) def stop(self): """Request a gentle stop: no new tasks, in-flight tasks finish.""" @@ -94,11 +91,16 @@ class OutfitBatchRunner: total = len(self.tasks) results = [] completed = 0 - max_workers = max(1, int(self.options.concurrency or 1)) + image_concurrency = max(1, int(self.options.concurrency or 1)) + row_workers = 1 - self._log("开始 AI 穿搭批量生成:{} 项,并发 {}".format(total, max_workers)) + self._log( + "开始 AI 穿搭批量生成:{} 项,Excel 行并发 1,图片并发 {}".format( + total, image_concurrency + ) + ) - with ThreadPoolExecutor(max_workers=max_workers) as executor: + with ThreadPoolExecutor(max_workers=row_workers) as executor: task_iter = iter(self.tasks) futures = {} submitted_all = False @@ -106,7 +108,7 @@ class OutfitBatchRunner: while futures or not submitted_all: while ( not self._stop_event.is_set() - and len(futures) < max_workers + and len(futures) < row_workers and not submitted_all ): try: @@ -114,7 +116,6 @@ class OutfitBatchRunner: except StopIteration: submitted_all = True break - self._limiter.wait() future = executor.submit(self._run_one_with_retry, task) futures[future] = _RunningTask( task=task, diff --git a/tasks.md b/tasks.md index b8f26c9..020295f 100644 --- a/tasks.md +++ b/tasks.md @@ -1225,10 +1225,10 @@ 设计取舍:外层 `OutfitBatchRunner` 固定 Excel 行并发为 1;界面文案改为「图片并发数」。配置字段可先兼容复用现有 `outfit_concurrency`,但业务含义改为目录内图片并发;若后续改字段名,需要迁移旧配置。 - [x] 文档更新:`docs/11-ai-outfit.md` 明确 Excel 行并发固定 1、图片并发数只作用于目录内部;`docs/ui-ai-outfit.html` 示例文案改为「图片并发数」 -- [ ] `ai_outfit_panel.py`:右侧生成设置 label 从「并发数」改为「图片并发数」,日志启动行同时显示「Excel 行并发 1 / 图片并发 N」 -- [ ] `_OutfitWorker` / `OutfitBatchRunner`:外层 Excel 行任务固定顺序处理,不再用界面并发值作为行 `max_workers` -- [ ] `generate_outfit_image`:增加或接入 `image_concurrency` 参数;单文件行保持顺序单张处理 -- [ ] `core/ai_outfit.py` 目录分支:使用 `ThreadPoolExecutor(max_workers=图片并发数)` 并发生成目录内图片;已存在输出仍跳过,部分失败仍聚合到整行结果 -- [ ] 保留「新请求间隔」启动节流;图片并发数 >1 时日志/缩略图完成顺序允许与文件名排序不同 -- [ ] 补测试:Excel 行顺序处理、目录内图片并发、单文件行不并发、已存在跳过、部分失败聚合、UI label/default/config 兼容 -- [ ] 验证:相关单测通过;离屏启动 AI 穿搭页确认 label 和日志文案;如工作区清理后再跑全套测试 +- [x] `ai_outfit_panel.py`:右侧生成设置 label 从「并发数」改为「图片并发数」,日志启动行同时显示「Excel 行并发 1 / 图片并发 N」 +- [x] `_OutfitWorker` / `OutfitBatchRunner`:外层 Excel 行任务固定顺序处理,不再用界面并发值作为行 `max_workers` +- [x] `generate_outfit_image`:增加或接入 `image_concurrency` 参数;单文件行保持顺序单张处理 +- [x] `core/ai_outfit.py` 目录分支:使用 `ThreadPoolExecutor(max_workers=图片并发数)` 并发生成目录内图片;已存在输出仍跳过,部分失败仍聚合到整行结果 +- [x] 保留「新请求间隔」启动节流;图片并发数 >1 时日志/缩略图完成顺序允许与文件名排序不同 +- [x] 补测试:Excel 行顺序处理、目录内图片并发、单文件行不并发、已存在跳过、部分失败聚合、UI label/default/config 兼容 +- [~] 验证:语法检查、`test_outfit_batch.py`、`test_ai_outfit.py`、`test_ai_outfit_panel.py` 通过;全套 `python -m unittest discover -s tests` 当前被工作区未提交的 `packaging/default_config/ai_models.json` 改动阻塞(模型顺序与出厂模板规范不一致),待清理该文件后重跑 diff --git a/tests/test_ai_outfit.py b/tests/test_ai_outfit.py index e2cdeaa..7c1d176 100644 --- a/tests/test_ai_outfit.py +++ b/tests/test_ai_outfit.py @@ -3,6 +3,8 @@ import os import shutil import sys import tempfile +import threading +import time import unittest from io import BytesIO from pathlib import Path @@ -262,6 +264,22 @@ class TestAiOutfitCore(unittest.TestCase): self.assertEqual(len(result.output_paths), 2) self.assertFalse((out / "a" / "img2.jpg").exists()) + def test_generate_directory_uses_image_concurrency(self): + from core.ai_outfit import generate_outfit_image + + d = self._make_dir_with_images("a", files=("img1.png", "img2.png", "img3.png", "img4.png")) + out = self.tmp / "out" + client = _ConcurrentRecordingClient(self._image_bytes()) + + result = generate_outfit_image( + self._dir_task(d), "x", out, model_config={}, api_client=client, + image_concurrency=2) + + self.assertTrue(result.success, result.error) + self.assertEqual(client.calls, 4) + self.assertGreaterEqual(client.max_active, 2) + self.assertEqual(len(result.output_paths), 4) + class _FakeClient: def __init__(self, image_bytes): @@ -309,5 +327,26 @@ class _FailOnClient: return self._image_bytes +class _ConcurrentRecordingClient: + """Thread-safe fake client that records concurrent generate() calls.""" + + def __init__(self, image_bytes): + self._image_bytes = image_bytes + self.calls = 0 + self.active = 0 + self.max_active = 0 + self._lock = threading.Lock() + + def generate(self, prompt, image_path, resolution="1K"): + with self._lock: + self.calls += 1 + self.active += 1 + self.max_active = max(self.max_active, self.active) + time.sleep(0.05) + with self._lock: + self.active -= 1 + return self._image_bytes + + if __name__ == "__main__": unittest.main() diff --git a/tests/test_ai_outfit_panel.py b/tests/test_ai_outfit_panel.py index 64bfa36..1a8321b 100644 --- a/tests/test_ai_outfit_panel.py +++ b/tests/test_ai_outfit_panel.py @@ -64,6 +64,15 @@ class TestAiOutfitPanelDefaults(unittest.TestCase): self.assertEqual(Path(panel._output_edit.text()), saved) + def test_generation_setting_label_is_image_concurrency(self): + from PySide6.QtWidgets import QLabel + + panel = self._panel() + + labels = [label.text() for label in panel.findChildren(QLabel)] + + self.assertIn("图片并发数", labels) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_outfit_batch.py b/tests/test_outfit_batch.py index a4bbdee..049d301 100644 --- a/tests/test_outfit_batch.py +++ b/tests/test_outfit_batch.py @@ -91,15 +91,17 @@ class TestOutfitBatchRunner(unittest.TestCase): self.assertEqual(summary.success_count, 1) self.assertEqual(summary.results[0].attempts, 2) - def test_concurrency_runs_more_than_one_task_at_once(self): + def test_concurrency_option_does_not_parallelize_excel_rows(self): from core.outfit_batch import OutfitBatchOptions, OutfitBatchRunner active = [0] max_active = [0] + started = [] lock = threading.Lock() def generate(task): with lock: + started.append(task.row_index) active[0] += 1 max_active[0] = max(max_active[0], active[0]) time.sleep(0.05) @@ -116,7 +118,8 @@ class TestOutfitBatchRunner(unittest.TestCase): summary = runner.run() self.assertEqual(summary.success_count, 4) - self.assertGreaterEqual(max_active[0], 2) + self.assertEqual(max_active[0], 1) + self.assertEqual(started, [2, 3, 4, 5]) def test_task_cooldown_is_applied(self): from core.outfit_batch import OutfitBatchOptions, OutfitBatchRunner