返回 VideoClaw
reference_agent.py
根目录 / video-claw / video-claw / backend / core / agents / reference_agent.py
1 # -*- coding: utf-8 -*-
2 """
3 阶段4: 参考图生成智能体
4 - 基于阶段3分镜(shots),为每个分镜生成「首帧图像提示词」,再据此生成参考图
5 - 首帧提示词由 LLM 根据 shot 的 plot、visual_prompt、duration 生成
6 - 阶段5生视频时使用阶段3的原始分镜描述,而非首帧提示词
7 - 支持逐项实时预览、重新生成、多版本管理
8 """
9
10 import os
11 import re
12 import glob
13 import json
14 import asyncio
15 import logging
16 from typing import Any, Optional, Dict, List
17 from concurrent.futures import ThreadPoolExecutor, as_completed
18
19 from .base_agent import AgentInterface
20 from prompts.loader import load_prompt
21
22 logger = logging.getLogger(__name__)
23
24
25 def ratio_to_size(ratio: str) -> str:
26 """将视频比例转换为图像尺寸"""
27 size_map = {
28 "16:9": "1920*1080",
29 "9:16": "1080*1920",
30 "1:1": "1024*1024",
31 "4:3": "1024*768",
32 "3:4": "768*1024",
33 "21:9": "2560*1080",
34 }
35 return size_map.get(ratio, "1920*1080")
36
37
38 class ReferenceGeneratorAgent(AgentInterface):
39 """参考图生成:分镜(阶段3) → 首帧提示词(LLM) → 参考图(图像模型)"""
40
41 def __init__(self):
42 super().__init__(name="ReferenceGenerator")
43
44 # ─── 版本管理 ───
45
46 @staticmethod
47 def _scenes_base(sid: str) -> str:
48 return os.path.join('code/result/image', str(sid), 'Scenes')
49
50 def _list_versions(self, sid: str, shot_id: str) -> List[str]:
51 """列出某个分镜的所有历史版本
52 命名: shot_001_01.jpg, shot_001_01_v2.jpg, ...
53 """
54 return self._list_versions_static(sid, shot_id)
55
56 @staticmethod
57 def _list_versions_static(sid: str, shot_id: str) -> List[str]:
58 """列出某个分镜的所有历史版本(静态方法,供外部调用)"""
59 scenes_dir = os.path.join('code/result/image', str(sid), 'Scenes')
60 files = []
61 for ext in ("jpg", "jpeg", "png", "webp", "bmp"):
62 pattern = os.path.join(scenes_dir, f"{shot_id}*.{ext}")
63 files.extend(glob.glob(pattern))
64 files = sorted(set(files), key=os.path.getmtime)
65 return files
66
67 def _next_version_path(self, sid: str, shot_id: str) -> str:
68 """获取下一个版本路径"""
69 scenes_dir = self._scenes_base(sid)
70 os.makedirs(scenes_dir, exist_ok=True)
71
72 existing = self._list_versions(sid, shot_id)
73 if not existing:
74 return os.path.join(scenes_dir, f"{shot_id}.jpg")
75
76 max_v = 1
77 for fp in existing:
78 bn = os.path.splitext(os.path.basename(fp))[0]
79 m = re.search(r'_v(\d+)$', bn)
80 if m:
81 max_v = max(max_v, int(m.group(1)))
82
83 return os.path.join(scenes_dir, f"{shot_id}_v{max_v + 1}.jpg")
84
85 # ─── 素材匹配 ───
86
87 @staticmethod
88 def _build_asset_map(character_design: Dict[str, Any]) -> Dict[str, Dict[str, str]]:
89 """从阶段2生成的素材数据中构建映射,不再直接扫描磁盘"""
90 am: Dict[str, Dict[str, str]] = {'characters': {}, 'settings': {}}
91
92 # 处理角色
93 for char in character_design.get('characters', []):
94 cid = char.get('id') or char.get('character_id')
95 selected = char.get('selected')
96 if cid and selected:
97 am['characters'][cid] = selected
98
99 # 处理场景
100 for setting in character_design.get('settings', []):
101 sid = setting.get('id') or setting.get('setting_id')
102 selected = setting.get('selected')
103 if sid and selected:
104 am['settings'][sid] = selected
105
106 return am
107
108 def _collect_refs(self, segment: dict, asset_map: dict,
109 char_id_map: dict, setting_id_map: dict) -> List[str]:
110 """为一个片段(Segment)收集参考原图路径(角色 + 场景素材)"""
111 refs = []
112 # 1. 角色匹配
113 for cn in segment.get('characters', []):
114 cid = char_id_map.get(cn)
115 # 如果名称不直接匹配,尝试模糊匹配(部分包含)
116 if not cid:
117 for name, _id in char_id_map.items():
118 if name in cn or cn in name:
119 cid = _id
120 break
121
122 if cid and cid in asset_map['characters']:
123 refs.append(os.path.abspath(asset_map['characters'][cid]))
124 logger.info(f"[{segment.get('segment_id', '')}] 添加角色参考图: {cn} -> {cid}")
125
126 # 2. 场景匹配
127 loc = segment.get('location', '')
128 set_id = setting_id_map.get(loc)
129 # 如果名称不直接匹配,尝试模糊匹配
130 if not set_id and loc:
131 for name, _id in setting_id_map.items():
132 if name in loc or loc in name:
133 set_id = _id
134 logger.info(f"[{segment.get('segment_id', '')}] 场景模糊匹配成功: {loc} -> {name} ({set_id})")
135 break
136
137 if set_id and set_id in asset_map['settings']:
138 refs.append(os.path.abspath(asset_map['settings'][set_id]))
139 logger.info(f"[{segment.get('segment_id', '')}] 添加场景参考图: {loc} -> {set_id}")
140 else:
141 logger.warning(f"[{segment.get('segment_id', '')}] 未找到场景参考图: location={loc}, set_id={set_id}, available_settings={list(asset_map['settings'].keys())}")
142
143 logger.info(f"[{segment.get('segment_id', '')}] 共收集 {len(refs)} 张参考图")
144 return refs[:10]
145
146 def _get_descriptions(self, segment: dict, char_id_map: dict, setting_id_map: dict,
147 character_json: dict) -> tuple:
148 """获取片段中涉及的角色和场景描述
149
150 Returns:
151 (character_description, setting_description)
152 """
153 # 角色描述
154 char_descs = []
155 for cn in segment.get('characters', []):
156 cid = char_id_map.get(cn, '')
157 if cid:
158 for c in character_json.get('characters', []):
159 if (c.get('id') or c.get('character_id')) == cid:
160 desc = c.get('description', '')
161 if desc:
162 char_descs.append(f"{cn}: {desc}")
163 break
164
165 # 场景描述
166 loc = segment.get('location', '')
167 set_id = setting_id_map.get(loc)
168 setting_desc = ""
169 if set_id:
170 for s in character_json.get('settings', []):
171 if (s.get('id') or s.get('setting_id')) == set_id:
172 setting_desc = s.get('description', '')
173 break
174
175 return "; ".join(char_descs), setting_desc
176
177 # ─── 首帧提示词生成 ───
178
179 # ─── 预览构建 ───
180
181 def _build_preview(self, sid: str, segments: list, session_data: dict = None) -> list:
182 """构建片段预览列表(含当前状态)"""
183 preview = []
184
185 # 建立 scene_id 到 selected 路径的映射
186 selected_map = {}
187 if session_data and "artifacts" in session_data:
188 ref_gen = session_data["artifacts"].get("reference_generation", {})
189 for scene in ref_gen.get("scenes", []):
190 sid_in_json = scene.get("id")
191 if sid_in_json:
192 selected_map[sid_in_json] = scene.get("selected", "")
193
194 for idx, seg in enumerate(segments, 1):
195 segment_id = seg.get('segment_id', f'seg_unk_{idx}')
196 versions = self._list_versions(sid, segment_id)
197
198 # 优先从 artifacts 中读取 selected 字段,如果没有则回退到最后一个版本
199 selected_path = selected_map.get(segment_id)
200 if not selected_path:
201 selected_path = versions[-1] if versions else ""
202
203 # 获取该段下第一个镜头的 content 作为片段描述
204 plot = seg.get('shots', [])[0].get('content', '') if seg.get('shots') else ""
205 ep_n = seg.get('episode_number', 1)
206 seg_n = seg.get('segment_number', idx)
207
208 preview.append({
209 "id": segment_id,
210 "name": f"第{ep_n}集-片段{seg_n}",
211 "episode": ep_n,
212 "index": seg_n,
213 "description": plot,
214 "selected": selected_path,
215 "versions": versions,
216 "status": "done" if versions else "pending",
217 })
218 return preview
219
220 # ─── 单张生成 ───
221
222 def _generate_image_with_doctor(self, img_client, *, prompt: str, model: str,
223 llm_model: str, context: dict, **kwargs) -> tuple:
224 """Generate once; if doctor rewrites a prompt-related failure, retry once."""
225 try:
226 paths = img_client.generate_image(prompt=prompt, model=model, **kwargs)
227 if not paths:
228 raise RuntimeError("Image generation returned no output")
229 return paths, prompt, None
230 except Exception as exc:
231 from .doctor_agent import DoctorAgent
232
233 doctor = DoctorAgent(llm_model=llm_model)
234 rewrite, diagnosis, rewrite_result = doctor.maybe_rewrite_prompt(
235 stage="reference_generation",
236 model=model,
237 prompt=prompt,
238 error=str(exc),
239 context=context,
240 )
241 if not rewrite:
242 logger.info("Doctor skipped reference prompt rewrite: %s", diagnosis.get("reason"))
243 raise
244
245 logger.info("Doctor rewrote reference prompt: reason_type=%s reason=%s",
246 diagnosis.get("reason_type"), diagnosis.get("reason"))
247 paths = img_client.generate_image(prompt=rewrite, model=model, **kwargs)
248 if not paths:
249 raise RuntimeError("Image generation returned no output after doctor rewrite")
250 return paths, rewrite, rewrite_result
251
252 @staticmethod
253 def _apply_eval_feedback_to_visual_prompt(current_prompt: str, eval_result: dict, version: int) -> str:
254 suggested_prompt = (eval_result.get('suggested_prompt') or '').strip()
255 if suggested_prompt:
256 return suggested_prompt
257
258 hard_failures = eval_result.get('hard_failures') or []
259 soft_issues = eval_result.get('soft_issues') or []
260 issues = eval_result.get('issues') or []
261 suggestion = (eval_result.get('suggestion') or '').strip()
262
263 feedback_lines = []
264 if hard_failures:
265 feedback_lines.append("硬性失败项:" + ";".join(map(str, hard_failures)))
266 if issues:
267 feedback_lines.append("主要问题:" + ";".join(map(str, issues)))
268 if soft_issues:
269 feedback_lines.append("软性问题:" + ";".join(map(str, soft_issues)))
270 if suggestion:
271 feedback_lines.append("修改建议:" + suggestion)
272 if not feedback_lines:
273 return current_prompt
274
275 return (
276 f"{current_prompt}\n\n"
277 f"【第{version + 1}轮VLM评估反馈】\n"
278 f"上一轮参考图未通过评估,请在下一轮生成时优先修正以下问题;不要改变角色核心外貌和场景核心设定:\n"
279 + "\n".join(f"- {line}" for line in feedback_lines)
280 )
281
282 def _generate_one(self, img_client, sid: str, segment: dict,
283 first_frame_prompt: str, refs: List[str],
284 style: str, it2i_model: str, t2i_model: str,
285 video_ratio: str = "16:9", resolution: str = "1080P", vlm_model: str = "qwen3.5-plus",
286 character_description: str = "", setting_description: str = "",
287 llm_model: str = "",
288 max_versions: int = 3) -> tuple:
289 """生成单个片段参考图,返回 (segment_id, path_or_None, eval_result)
290
291 最多生成 max_versions 个版本,如果所有版本都没有达到硬性合格标准,
292 使用 VLM 选择最好的一张作为最终参考图。
293 """
294 segment_id = segment.get('segment_id', '')
295
296 # 仅提取第一个镜头的描述作为当前 Plot
297 plot = segment.get('shots', [])[0].get('content', '') if segment.get('shots') else ""
298 current_visual_prompt = first_frame_prompt
299
300 # 取消时直接跳过,不抛异常,以保留已生成的部分结果
301 if self.cancellation_check and self.cancellation_check():
302 logger.info(f"ReferenceGeneratorAgent: {segment_id} 跳过(用户取消)")
303 return segment_id, None, None, None
304
305 model = it2i_model if refs else t2i_model
306 logger.info(f"[{segment_id}] 使用模型: {model}, 参考图数量: {len(refs) if refs else 0}")
307 if refs:
308 for i, r in enumerate(refs):
309 logger.info(f"[{segment_id}] 参考图[{i}]: {r}")
310
311 # 收集所有生成的版本
312 all_versions = []
313 all_eval_results = []
314
315 for version in range(max_versions):
316 self._check_cancel()
317
318 style_prompt = self._get_style_prompt(style)
319 full_prompt = f"{style_prompt}, {current_visual_prompt}"
320
321 save_path = self._next_version_path(sid, segment_id)
322 save_dir = os.path.dirname(save_path)
323
324 try:
325 paths, used_prompt, rewrite_result = self._generate_image_with_doctor(
326 img_client,
327 prompt=full_prompt,
328 model=model,
329 llm_model=llm_model,
330 context={
331 "segment_id": segment_id,
332 "location": segment.get("location", ""),
333 "characters": segment.get("characters", []),
334 "plot": plot,
335 },
336 image_paths=refs if refs else None,
337 session_id=str(sid),
338 save_dir=save_dir,
339 video_ratio=video_ratio,
340 resolution=resolution,
341 )
342
343 gen = paths[0]
344 if gen != save_path:
345 if os.path.exists(save_path):
346 os.remove(save_path)
347 os.rename(gen, save_path)
348
349 # VLM 评估
350 eval_result = self._evaluate_with_vlm(save_path, segment, plot, current_visual_prompt,
351 character_description=character_description,
352 setting_description=setting_description,
353 vlm_model=vlm_model)
354
355 score = eval_result.get('score', 0)
356 hard_failures = eval_result.get('hard_failures') or []
357 if 'is_acceptable' in eval_result:
358 is_acceptable = bool(eval_result.get('is_acceptable')) and not hard_failures
359 else:
360 is_acceptable = not hard_failures and score >= 7
361
362 logger.info(f"[{segment_id}] 版本{version + 1}: 评分 {score}/10, {'✓通过' if is_acceptable else '✗不通过'}")
363 if hard_failures:
364 logger.warning(f"[{segment_id}] 硬性失败项: {hard_failures}")
365
366 # 记录版本信息
367 eval_result["final_visual_prompt"] = current_visual_prompt
368 if used_prompt != full_prompt:
369 eval_result["doctor_rewrite_prompt"] = used_prompt
370 if rewrite_result:
371 eval_result["rewrite_result"] = rewrite_result
372 all_versions.append(save_path)
373 all_eval_results.append(eval_result)
374
375 # 如果 VLM 判定达到硬性标准,立即返回
376 if is_acceptable:
377 return segment_id, save_path, eval_result, rewrite_result
378
379 # 报告进度
380 if version < max_versions - 1:
381 current_visual_prompt = self._apply_eval_feedback_to_visual_prompt(current_visual_prompt, eval_result, version)
382 logger.info(f"[{segment_id}] 下一轮将使用VLM反馈优化首帧提示词")
383 self._report_progress("参考图", f"重新生成中 ({version + 2}/{max_versions}): {segment_id}", 0)
384
385 except Exception as e:
386 logger.error(f"Segment {segment_id} image generation failed: {e}")
387
388 # 所有版本都没有达到硬性标准,使用 VLM 选择最好的
389 if all_versions:
390 logger.warning(f"[{segment_id}] 所有版本都未达到硬性合格标准,使用VLM选择最佳...")
391 best_path, best_eval = self._select_best_with_vlm(
392 all_versions, segment, plot, current_visual_prompt,
393 character_description=character_description,
394 setting_description=setting_description,
395 vlm_model=vlm_model
396 )
397 if best_path and isinstance(best_eval, dict):
398 best_eval["final_visual_prompt"] = current_visual_prompt
399 hard_failures = best_eval.get("hard_failures") or []
400 is_acceptable = bool(best_eval.get("is_acceptable")) and not hard_failures
401 if is_acceptable:
402 return segment_id, best_path, best_eval, best_eval.get("rewrite_result") if isinstance(best_eval, dict) else None
403 if hard_failures:
404 logger.warning(
405 "[%s] VLM选择的最佳图仍有硬性失败项,保留为候选图供人工确认: %s",
406 segment_id,
407 hard_failures,
408 )
409 return segment_id, best_path, best_eval, best_eval.get("rewrite_result") if isinstance(best_eval, dict) else None
410 # 低分但已有候选图属于可人工确认的软失败,保留最佳图作为候选,避免把阶段误标为 failed。
411 logger.warning(
412 "[%s] VLM选择的最佳图未达分数阈值,仍作为低分候选图保留: %s",
413 segment_id,
414 best_eval.get("issues", []),
415 )
416 return segment_id, best_path, best_eval, best_eval.get("rewrite_result") if isinstance(best_eval, dict) else None
417
418 # 如果没有任何生成成功
419 logger.warning(f"[{segment_id}] 没有成功生成任何图片")
420 if all_eval_results and isinstance(all_eval_results[-1], dict):
421 all_eval_results[-1]["final_visual_prompt"] = current_visual_prompt
422 return segment_id, None, None, None
423
424 def _select_best_with_vlm(self, image_paths: List[str], segment: dict, plot: str, visual_prompt: str,
425 character_description: str = "", setting_description: str = "",
426 vlm_model: str = "qwen3.5-plus") -> tuple:
427 """使用 VLM 从多个版本中选择最好的一张"""
428 from models.vlm_client import VLM
429
430 if not image_paths:
431 return None, None
432
433 segment_id = segment.get('segment_id', '')
434
435 # 加载评估提示词
436 select_prompt = load_prompt('reference', 'eval_select_best', 'zh').format(
437 num_images=len(image_paths),
438 num_images_minus_1=len(image_paths) - 1,
439 plot=plot,
440 visual_prompt=visual_prompt,
441 character_description=character_description,
442 setting_description=setting_description,
443 images_list="\n".join([f"图片{i}: {p}" for i, p in enumerate(image_paths)])
444 )
445
446 try:
447 vlm = VLM()
448 result = vlm.query(select_prompt, image_paths=image_paths, model=vlm_model)
449 logger.info(f"[{segment_id}] VLM选择结果: {result}")
450
451 # 解析 JSON 结果
452 import re
453 json_match = re.search(r'\{[^{}]*\}', result, re.DOTALL)
454 if json_match:
455 selected = json.loads(json_match.group())
456 selected_idx = selected.get('selected_index', 0)
457 if 0 <= selected_idx < len(image_paths):
458 best_path = image_paths[selected_idx]
459 logger.info(f"[{segment_id}] VLM选择第{selected_idx + 1}张作为最佳图片")
460 hard_failures = selected.get('hard_failures') or []
461 score = selected.get('score', 5)
462 # 构建评估结果
463 best_eval = {
464 "score": score,
465 "hard_failures": hard_failures,
466 "soft_issues": selected.get('soft_issues', []),
467 "issues": selected.get('issues', []),
468 "is_acceptable": not hard_failures and score >= 7,
469 "selected_by_vlm": True,
470 "reason": selected.get('reason', '')
471 }
472 return best_path, best_eval
473
474 except Exception as e:
475 logger.error(f"[{segment_id}] VLM选择最佳图片失败: {e}")
476
477 # 如果 VLM 选择失败,仍保留第一张候选图,避免“已有参考图”被误标成失败。
478 return image_paths[0], {
479 "score": 0,
480 "hard_failures": ["VLM选择最佳图片失败,无法确认硬性标准"],
481 "issues": ["VLM选择最佳图片失败"],
482 "is_acceptable": False,
483 "selected_by_vlm": False,
484 }
485
486 def _evaluate_with_vlm(self, image_path: str, segment: dict, plot: str, visual_prompt: str,
487 character_description: str = "", setting_description: str = "",
488 vlm_model: str = "qwen3.5-plus") -> dict:
489 """使用 VLM 评估首帧参考图"""
490 try:
491 from models.vlm_client import VLM
492 vlm = VLM()
493
494 eval_prompt = load_prompt('reference', 'eval_first_frame', 'zh').format(
495 plot=plot,
496 visual_prompt=visual_prompt,
497 character_description=character_description,
498 setting_description=setting_description
499 )
500
501 result = vlm.query(
502 prompt=eval_prompt,
503 image_paths=[image_path],
504 model=vlm_model
505 )
506
507 if result and isinstance(result, list):
508 result_text = result[0] if result else ""
509 elif isinstance(result, str):
510 result_text = result
511 else:
512 result_text = str(result)
513
514 import json
515 try:
516 import re
517 json_match = re.search(r'\{[^{}]*\}', result_text, re.DOTALL)
518 if json_match:
519 eval_result = json.loads(json_match.group())
520 return eval_result
521 except:
522 pass
523
524 return {
525 "score": 0,
526 "hard_failures": ["VLM评估解析失败,无法确认硬性标准"],
527 "issues": ["评估解析失败"],
528 "is_acceptable": False,
529 }
530
531 except Exception as e:
532 logger.warning(f"VLM evaluation failed: {e}")
533 return {
534 "score": 0,
535 "hard_failures": ["VLM评估失败,无法确认硬性标准"],
536 "issues": [str(e)],
537 "is_acceptable": False,
538 }
539
540 # ─── 构建最终 payload ───
541
542 def _build_payload(
543 self,
544 sid: str,
545 segments: list,
546 session_data: dict = None,
547 prompts_map: dict = None,
548 selected_images: dict = None,
549 rewrite_results: dict = None,
550 ) -> dict:
551 """构建最终 payload"""
552 scenes = []
553 if prompts_map is None:
554 prompts_map = {}
555 if selected_images is None:
556 selected_images = {}
557 rewrite_results = rewrite_results or {}
558
559 # 建立 scene_id 到 selected 路径的映射
560 selected_map = {}
561 existing_prompts = {}
562 existing_rewrite_results = {}
563 if session_data and "artifacts" in session_data:
564 ref_gen = session_data["artifacts"].get("reference_generation", {})
565 for scene in ref_gen.get("scenes", []):
566 sid_in_json = scene.get("id")
567 if sid_in_json:
568 selected_map[sid_in_json] = scene.get("selected", "")
569 existing_prompts[sid_in_json] = scene.get("visual_prompt", "")
570 if scene.get("rewrite_result"):
571 existing_rewrite_results[sid_in_json] = scene.get("rewrite_result")
572
573 for idx, seg in enumerate(segments, 1):
574 segment_id = seg.get('segment_id', f'seg_unk_{idx}')
575 versions = self._list_versions(sid, segment_id)
576
577 # 优先从本轮生成的 selected_images 中读取,如果找不到,再从 session 的 artifacts 中读取 selected 字段
578 selected_path = selected_images.get(segment_id)
579 if not selected_path:
580 selected_path = selected_map.get(segment_id)
581 if not selected_path:
582 selected_path = versions[-1] if versions else ""
583
584 # 获取该段下第一个镜头的 content 作为片段描述
585 shots_summary = seg.get('shots', [])[0].get('content', '') if seg.get('shots') else ""
586
587 # 取提示词
588 visual_prompt = prompts_map.get(segment_id) or existing_prompts.get(segment_id) or ""
589
590 # 最终 payload 表示阶段已跑完;仍没有图片的片段应标记为 failed,避免覆盖实时失败状态。
591 status = "done" if selected_path or versions else "failed"
592 item = {
593 "id": segment_id,
594 "name": f"第{seg.get('episode_number', 1)}集-片段{seg.get('segment_number', idx)}",
595 "index": idx,
596 "description": shots_summary,
597 "visual_prompt": visual_prompt,
598 "selected": selected_path,
599 "versions": versions,
600 "status": status,
601 }
602 rewrite_result = rewrite_results.get(segment_id) or existing_rewrite_results.get(segment_id)
603 if rewrite_result:
604 item["rewrite_result"] = rewrite_result
605 scenes.append(item)
606 return {
607 "payload": {
608 "session_id": sid,
609 "scenes": scenes,
610 },
611 "stage_completed": True,
612 }
613
614 # ─── 核心流程 ───
615
616 async def process(self, input_data: Any, intervention: Optional[Dict] = None) -> Dict:
617 from config import settings
618 from models.image_client import ImageClient
619 from models.llm_client import LLM
620
621 # 从编排器注入的 session 快照补齐缺失参数,避免各阶段直接读取 session JSON。
622 input_data = self._merge_session_params(input_data)
623
624 sid = input_data["session_id"]
625
626 style = input_data.get("style", "anime")
627 video_ratio = input_data.get("video_ratio", "16:9")
628 resolution = input_data.get("resolution", "2K")
629 llm_model = self._require_input(input_data, "llm_model")
630 t2i = self._require_input(input_data, "image_t2i_model")
631 it2i = self._require_input(input_data, "image_it2i_model")
632 vlm_model = self._require_input(input_data, "vlm_model")
633 # 根据 enable_concurrency 决定并发数
634 enable_concurrency = input_data.get("enable_concurrency", True)
635 logger.info(f"[ReferenceAgent] enable_concurrency={enable_concurrency}")
636 # 取 t2i 和 it2i 中的最大并发数
637 from models.config_model import get_max_concurrency
638 max_t2i = get_max_concurrency(t2i, enable_concurrency)
639 max_it2i = get_max_concurrency(it2i, enable_concurrency)
640 concurrency = max(max_t2i, max_it2i)
641 logger.info(f"[ReferenceAgent] 使用并发数={concurrency}")
642
643 artifacts = self._session_artifacts(input_data)
644 session_data = {
645 "meta": self._session_meta(input_data),
646 "artifacts": artifacts,
647 }
648
649 # 提取已经存在于 session 中的 visual_prompts
650 session_visual_prompts = {}
651 ref_gen = artifacts.get("reference_generation", {})
652 for scene in ref_gen.get("scenes", []):
653 sid_in_json = scene.get("id")
654 vp = scene.get("visual_prompt")
655 if sid_in_json and vp:
656 session_visual_prompts[sid_in_json] = vp
657
658 img_client = ImageClient(
659 dashscope_api_key=settings.DASHSCOPE_API_KEY,
660 dashscope_base_url=settings.DASHSCOPE_BASE_URL,
661 gpt_api_key=settings.OPENAI_API_KEY,
662 gpt_base_url=settings.OPENAI_BASE_URL,
663 proxy=settings.provider_proxy("openai"),
664 ark_api_key=settings.ARK_API_KEY,
665 ark_base_url=settings.ARK_BASE_URL,
666 )
667
668 episodes = artifacts.get('storyboard', {}).get('episodes', [])
669 if not episodes:
670 raise Exception("未找到分镜剧集数据,请先完成阶段3")
671
672 segments = []
673 for ep in episodes:
674 for seg in ep.get("segments", []):
675 segments.append(seg)
676
677 if not segments:
678 raise Exception("未找到分镜片段数据,请先完成阶段3")
679
680 logger.info(f"[ReferenceAgent] 解析到 {len(segments)} 个拍摄片段")
681
682 script_json = artifacts.get('script_generation', {})
683 character_json = artifacts.get('character_design', {})
684
685 # 判断中英文
686 is_zh = any('\u4e00' <= c <= '\u9fff' for c in script_json.get("title", ""))
687
688 # 构建 name → id 映射(用于素材匹配)
689 char_id_map = {}
690 for c in character_json.get('characters', []):
691 chara_id = c.get('id') or c.get('character_id') or ''
692 char_id_map[c['name']] = chara_id
693
694 setting_id_map = {}
695 for s in character_json.get('settings', []):
696 set_id = s.get('id') or s.get('setting_id') or ''
697 setting_id_map[s['name']] = set_id
698
699 asset_map = self._build_asset_map(character_json)
700
701 # ═══ 介入:重新生成指定分段 ═══
702 if intervention:
703 regen_scenes = intervention.get("regenerate_scenes", [])
704
705 if regen_scenes:
706 self._report_progress("参考图", "重新生成中...", 2)
707
708 fresh_episodes = artifacts.get('storyboard', {}).get('episodes', [])
709
710 fresh_segments = []
711 for ep in fresh_episodes:
712 fresh_segments.extend(ep.get("segments", []))
713
714 fresh_segment_map = {s['segment_id']: s for s in fresh_segments}
715
716 selected_images = {}
717 prompt_map = {} # segment_id → first_frame_prompt
718 rewrite_results_map = {}
719
720 def regen_run():
721 total = len(regen_scenes)
722 done = 0
723 nonlocal selected_images
724 nonlocal prompt_map
725
726 def calc_pct_regen(completed: int) -> int:
727 return min(95, 5 + int(90 * completed / max(total, 1)))
728
729 def regen_segment_run(segment_id: str, index: int):
730 existing_versions = self._list_versions(sid, segment_id)
731 self._report_progress("参考图", f"正在生成: {segment_id}", 5, data={
732 "asset_complete": {
733 "type": "scenes",
734 "id": segment_id,
735 "status": "running",
736 "versions": existing_versions,
737 }
738 })
739 seg = fresh_segment_map.get(segment_id, {})
740 first_shot = seg.get('shots', [])[0] if seg.get('shots') else {}
741 plot = first_shot.get('content', '')
742 char_desc, set_desc = self._get_descriptions(seg, char_id_map, setting_id_map, character_json)
743
744 existing_vp = session_visual_prompts.get(segment_id)
745 if existing_vp:
746 ff_prompt = existing_vp
747 logger.info(f"[{segment_id}] 重新生成时命中已有提示词,复用原提示词...")
748 else:
749 ff_prompt_tpl = load_prompt('reference', 'first_frame', 'zh' if is_zh else 'en')
750 try:
751 local_llm = LLM()
752 ff_prompt_resp = self._cancellable_query(
753 local_llm,
754 prompt=ff_prompt_tpl.format(
755 original_text=script_json.get("original_text", ""),
756 plot=plot,
757 character_description=char_desc,
758 setting_description=set_desc
759 ),
760 model=llm_model
761 )
762 if hasattr(ff_prompt_resp, 'content'):
763 ff_prompt = ff_prompt_resp.content.strip()
764 else:
765 ff_prompt = str(ff_prompt_resp).strip()
766 except Exception as e:
767 logger.error(f"Error generating first-frame prompt for {segment_id}: {e}")
768 ff_prompt = plot[:200]
769
770 logger.info(f"[{segment_id}] first-frame prompt: {ff_prompt}...")
771
772 refs = self._collect_refs(seg, asset_map, char_id_map, setting_id_map)
773 char_desc, set_desc = self._get_descriptions(
774 seg, char_id_map, setting_id_map, character_json
775 )
776 result_segment_id, result_path, eval_result, rewrite_result = self._generate_one(
777 img_client, sid,
778 seg, ff_prompt, refs,
779 style, it2i, t2i, video_ratio, resolution, vlm_model,
780 character_description=char_desc, setting_description=set_desc,
781 llm_model=llm_model,
782 )
783 final_prompt = ff_prompt
784 if isinstance(eval_result, dict):
785 final_prompt = eval_result.get("final_visual_prompt") or ff_prompt
786 return result_segment_id, result_path, eval_result, final_prompt, rewrite_result
787
788 # 并发生成提示词与图像
789 self._report_progress("参考图", f"生成参考图... 0/{total}", 5)
790 with ThreadPoolExecutor(max_workers=concurrency) as executor:
791 futs = {}
792 for i, segment_id in enumerate(regen_scenes):
793 fut = executor.submit(regen_segment_run, segment_id, i)
794 futs[fut] = segment_id
795 for fut in as_completed(futs):
796 segment_id_done = futs[fut]
797 try:
798 _, result_path, eval_result, ff_prompt, rewrite_result = fut.result()
799 prompt_map[segment_id_done] = ff_prompt
800 if rewrite_result:
801 rewrite_results_map[segment_id_done] = rewrite_result
802 except Exception as e:
803 logger.error(f"Regen future error for {segment_id_done}: {e}")
804 result_path = None
805 done += 1
806 pct = calc_pct_regen(done)
807 if result_path:
808 selected_images[segment_id_done] = result_path
809 versions = self._list_versions(sid, segment_id_done)
810 asset_complete = {
811 "type": "scenes", "id": segment_id_done,
812 "status": "done",
813 "selected": result_path,
814 "versions": versions,
815 }
816 if rewrite_results_map.get(segment_id_done):
817 asset_complete["rewrite_result"] = rewrite_results_map[segment_id_done]
818 self._report_progress("参考图", f"完成: {segment_id_done}", pct, data={
819 "asset_complete": asset_complete
820 })
821 else:
822 versions = self._list_versions(sid, segment_id_done)
823 fallback_path = versions[-1] if versions else ""
824 if fallback_path:
825 selected_images[segment_id_done] = fallback_path
826 asset_complete = {
827 "type": "scenes", "id": segment_id_done,
828 "status": "done" if fallback_path else "failed",
829 "selected": fallback_path,
830 "versions": versions,
831 }
832 if rewrite_results_map.get(segment_id_done):
833 asset_complete["rewrite_result"] = rewrite_results_map[segment_id_done]
834 message = f"完成: {segment_id_done}" if fallback_path else f"失败: {segment_id_done}"
835 self._report_progress("参考图", message, pct, data={
836 "asset_complete": asset_complete
837 })
838 # 检查取消
839 if self.cancellation_check and self.cancellation_check():
840 logger.info("ReferenceGeneratorAgent: 用户取消重新生成,停止等待剩余任务")
841 for f in futs:
842 if not f.done():
843 f.cancel()
844 break
845
846 loop = asyncio.get_running_loop()
847 await loop.run_in_executor(None, regen_run)
848
849 self._report_progress("参考图", "完成", 100)
850 return self._build_payload(sid, fresh_segments, session_data, prompt_map, selected_images, rewrite_results_map)
851
852 # ═══ 正常流程:全量生成 ═══
853 self._report_progress("参考图", "加载分镜数据...", 5)
854
855 # 发送预览列表
856 preview = self._build_preview(sid, segments, session_data)
857 self._report_progress("参考图", "加载分镜列表", 8, data={"assets_preview": {"scenes": preview}})
858
859 first_frame_prompts = {} # 提升作用域,用于最后写回结果文件
860 selected_images_map = {} # 提升作用域,记录本轮新生成且 VLM 挑选出来的图片路径
861 rewrite_results_map = {}
862
863 def run():
864 nonlocal first_frame_prompts
865 nonlocal selected_images_map
866 nonlocal rewrite_results_map
867 # 筛选需要生成的(跳过已有图的)
868 pending_segments = []
869 for seg in segments:
870 segment_id = seg['segment_id']
871 existing = self._list_versions(sid, segment_id)
872 if existing:
873 continue
874 pending_segments.append(seg)
875
876 if not pending_segments:
877 self._report_progress("参考图", "所有分镜图已存在", 95)
878 return
879
880 total = len(pending_segments)
881
882 def calc_pct(completed: int) -> int:
883 """并发阶段只按完成数量推进,避免提交任务时进度虚高。"""
884 return min(95, 10 + int(85 * completed / max(total, 1)))
885
886 done = 0
887
888 # 步骤2-6(每片段):流式生成提示词并立即开始图像生成
889 self._report_progress("参考图", f"开始生成... 0/{total}", calc_pct(0))
890
891 with ThreadPoolExecutor(max_workers=concurrency) as executor:
892 futs = {}
893 done = 0
894
895 def segment_run(seg: dict, index: int):
896 segment_id = seg['segment_id']
897
898 self._report_progress("参考图", f"正在生成: {segment_id}", calc_pct(done), data={
899 "asset_complete": {
900 "type": "scenes", "id": segment_id,
901 "status": "running"
902 }
903 })
904
905 first_shot = seg.get('shots', [])[0] if seg.get('shots') else {}
906 plot = first_shot.get('content', '')
907 char_desc, set_desc = self._get_descriptions(seg, char_id_map, setting_id_map, character_json)
908
909 existing_vp = session_visual_prompts.get(segment_id)
910 if existing_vp:
911 ff_prompt = existing_vp
912 logger.info(f"[{segment_id}] 命中已有提示词,复用原提示词...")
913 else:
914 ff_prompt_tpl = load_prompt('reference', "first_frame", 'zh' if is_zh else 'en')
915 try:
916 local_llm = LLM()
917 ff_prompt_resp = self._cancellable_query(
918 local_llm,
919 prompt=ff_prompt_tpl.format(
920 original_text=script_json.get("original_text", ""),
921 plot=plot,
922 character_description=char_desc,
923 setting_description=set_desc
924 ),
925 model=llm_model
926 )
927 if hasattr(ff_prompt_resp, 'content'):
928 ff_prompt = ff_prompt_resp.content.strip()
929 else:
930 ff_prompt = str(ff_prompt_resp).strip()
931 except Exception as e:
932 logger.error(f"Error generating first-frame prompt for {segment_id}: {e}")
933 ff_prompt = plot[:200]
934
935 logger.info(f"[{segment_id}] Prompt ready, starting image generation...")
936
937 refs = self._collect_refs(seg, asset_map, char_id_map, setting_id_map)
938 char_desc, set_desc = self._get_descriptions(seg, char_id_map, setting_id_map, character_json)
939 result_segment_id, result_path, eval_result, rewrite_result = self._generate_one(
940 img_client, sid,
941 seg, ff_prompt, refs,
942 style, it2i, t2i, video_ratio, resolution, vlm_model,
943 character_description=char_desc, setting_description=set_desc,
944 llm_model=llm_model,
945 )
946 final_prompt = ff_prompt
947 if isinstance(eval_result, dict):
948 final_prompt = eval_result.get("final_visual_prompt") or ff_prompt
949 return result_segment_id, result_path, eval_result, final_prompt, rewrite_result
950
951 for i, seg in enumerate(pending_segments):
952 segment_id = seg['segment_id']
953 fut = executor.submit(segment_run, seg, i)
954 futs[fut] = segment_id
955
956 # 4. 等待所有任务完成
957 cancelled = False
958 for fut in as_completed(futs):
959 segment_id_done = futs[fut]
960 try:
961 _, result_path, eval_result, ff_prompt, rewrite_result = fut.result()
962 first_frame_prompts[segment_id_done] = ff_prompt
963 if rewrite_result:
964 rewrite_results_map[segment_id_done] = rewrite_result
965 except Exception as e:
966 logger.error(f"Image future error for {segment_id_done}: {e}")
967 result_path = None
968
969 done += 1
970 pct = calc_pct(done)
971
972 if result_path:
973 selected_images_map[segment_id_done] = result_path
974 versions = self._list_versions(sid, segment_id_done)
975 asset_complete = {
976 "type": "scenes", "id": segment_id_done,
977 "status": "done",
978 "selected": result_path,
979 "versions": versions,
980 }
981 if rewrite_results_map.get(segment_id_done):
982 asset_complete["rewrite_result"] = rewrite_results_map[segment_id_done]
983 self._report_progress("参考图", f"完成: {segment_id_done}", pct, data={
984 "asset_complete": asset_complete
985 })
986 else:
987 versions = self._list_versions(sid, segment_id_done)
988 fallback_path = versions[-1] if versions else ""
989 if fallback_path:
990 selected_images_map[segment_id_done] = fallback_path
991 asset_complete = {
992 "type": "scenes", "id": segment_id_done,
993 "status": "done" if fallback_path else "failed",
994 "selected": fallback_path,
995 "versions": versions,
996 }
997 if rewrite_results_map.get(segment_id_done):
998 asset_complete["rewrite_result"] = rewrite_results_map[segment_id_done]
999 message = f"完成: {segment_id_done}" if fallback_path else f"失败: {segment_id_done}"
1000 self._report_progress("参考图", message, pct, data={
1001 "asset_complete": asset_complete
1002 })
1003
1004 # 检查取消
1005 if self.cancellation_check and self.cancellation_check():
1006 logger.info("ReferenceGeneratorAgent: 用户取消,停止等待剩余任务")
1007 for f in futs:
1008 if not f.done():
1009 f.cancel()
1010 cancelled = True
1011 break
1012
1013 if cancelled:
1014 self._report_progress("参考图", "已取消(保留已完成图片)", 96)
1015 else:
1016 self._report_progress("参考图", "保存结果...", 96)
1017
1018 loop = asyncio.get_running_loop()
1019 try:
1020 await loop.run_in_executor(None, run)
1021 except Exception as e:
1022 if "cancel" in str(e).lower():
1023 logger.info("ReferenceGeneratorAgent: 用户取消,返回已完成部分结果")
1024 self._report_progress("参考图", "已取消(保留已完成图片)", 100)
1025 return self._build_payload(sid, segments, session_data, first_frame_prompts, selected_images_map, rewrite_results_map)
1026 raise
1027
1028 self._report_progress("参考图", "完成", 100)
1029 return self._build_payload(sid, segments, session_data, first_frame_prompts, selected_images_map, rewrite_results_map)
1030
1030 lines PYTHON