返回 JoyAI-Echo
generation_settings.py
根目录 / echo_longvideo / Director_Agent / nanobot / session / generation_settings.py
1 """Director generation settings: duration, shot count, canvas, and language."""
2
3 from __future__ import annotations
4
5 from typing import Any
6
7 _UNSET = object()
8
9 SESSION_NSHOT_KEY = "n_shots"
10 SESSION_DURATION_KEY = "duration_sec"
11 SESSION_VIDEO_WIDTH_KEY = "video_width"
12 SESSION_VIDEO_HEIGHT_KEY = "video_height"
13 SESSION_LANGUAGE_KEY = "language"
14 SESSION_LLM_TEMPERATURE_KEY = "llm_temperature"
15 SESSION_LLM_TOP_P_KEY = "llm_top_p"
16 SESSION_LLM_TOP_K_KEY = "llm_top_k"
17
18 DEFAULT_NSHOT = 1
19 DEFAULT_DURATION_SEC = 10
20 DEFAULT_VIDEO_WIDTH = 1280
21 DEFAULT_VIDEO_HEIGHT = 736
22 DEFAULT_LANGUAGE = "zh"
23
24 # UI / story_profile.language values
25 LANGUAGE_ZH = "zh"
26 LANGUAGE_EN = "en"
27 VALID_LANGUAGES = frozenset({LANGUAGE_ZH, LANGUAGE_EN})
28
29 # story_profile.dialogue_language values used by PE / shot prompts
30 DIALOGUE_LANGUAGE_BY_LANGUAGE: dict[str, str] = {
31 LANGUAGE_ZH: "Mandarin Chinese",
32 LANGUAGE_EN: "English",
33 }
34 CAPTION_LANGUAGE_BY_LANGUAGE: dict[str, str] = {
35 LANGUAGE_ZH: "Simplified Chinese",
36 LANGUAGE_EN: "English",
37 }
38
39 # duration_sec → n_shots
40 DURATION_TO_NSHOT: dict[int, int] = {
41 10: 1,
42 20: 2,
43 30: 3,
44 60: 6,
45 90: 9,
46 120: 12,
47 150: 15,
48 180: 18,
49 }
50 NSHOT_TO_DURATION: dict[int, int] = {n: d for d, n in DURATION_TO_NSHOT.items()}
51 VALID_DURATIONS = frozenset(DURATION_TO_NSHOT)
52 VALID_NSHOTS = frozenset(NSHOT_TO_DURATION)
53
54
55 def duration_to_n_shots(duration_sec: int) -> int | None:
56 return DURATION_TO_NSHOT.get(int(duration_sec))
57
58
59 def n_shots_to_duration(n_shots: int) -> int | None:
60 return NSHOT_TO_DURATION.get(int(n_shots))
61
62
63 def normalize_duration_sec(value: Any) -> int | None:
64 try:
65 parsed = int(value)
66 except (TypeError, ValueError):
67 return None
68 return parsed if parsed in VALID_DURATIONS else None
69
70
71 def normalize_n_shots(value: Any) -> int | None:
72 try:
73 parsed = int(value)
74 except (TypeError, ValueError):
75 return None
76 return parsed if parsed in VALID_NSHOTS else None
77
78
79 def normalize_language(value: Any) -> str | None:
80 if not isinstance(value, str):
81 return None
82 cleaned = value.strip()
83 if cleaned in VALID_LANGUAGES:
84 return cleaned
85 lowered = cleaned.lower()
86 if lowered in {
87 "zh",
88 "zh-cn",
89 "zh_cn",
90 "chinese",
91 "mandarin",
92 "mandarin chinese",
93 "中文",
94 }:
95 return LANGUAGE_ZH
96 if lowered in {"en", "en-us", "en_us", "english"}:
97 return LANGUAGE_EN
98 return None
99
100
101 def language_to_dialogue_language(language: str | None) -> str | None:
102 if not language:
103 return None
104 return DIALOGUE_LANGUAGE_BY_LANGUAGE.get(language)
105
106
107 def language_to_caption_language(language: str | None) -> str | None:
108 if not language:
109 return None
110 return CAPTION_LANGUAGE_BY_LANGUAGE.get(language)
111
112
113 def normalize_llm_temperature(value: Any) -> float | None:
114 try:
115 parsed = float(value)
116 except (TypeError, ValueError):
117 return None
118 return parsed if 0 <= parsed < 2 else None
119
120
121 def normalize_llm_top_p(value: Any) -> float | None:
122 try:
123 parsed = float(value)
124 except (TypeError, ValueError):
125 return None
126 return parsed if 0 <= parsed <= 1 else None
127
128
129 def normalize_llm_top_k(value: Any) -> int | None:
130 try:
131 parsed = int(value)
132 except (TypeError, ValueError):
133 return None
134 return parsed if 1 <= parsed <= 64 else None
135
136
137 def default_settings() -> dict[str, int | str]:
138 return {
139 "n_shots": DEFAULT_NSHOT,
140 "duration_sec": DEFAULT_DURATION_SEC,
141 "width": DEFAULT_VIDEO_WIDTH,
142 "height": DEFAULT_VIDEO_HEIGHT,
143 "language": DEFAULT_LANGUAGE,
144 }
145
146
147 def get_generation_settings(metadata: dict[str, Any] | None) -> dict[str, int | str]:
148 """Return persisted generation settings, falling back to defaults."""
149 base = default_settings()
150 if not isinstance(metadata, dict):
151 return base
152
153 n_shots = normalize_n_shots(metadata.get(SESSION_NSHOT_KEY))
154 duration_sec = normalize_duration_sec(metadata.get(SESSION_DURATION_KEY))
155 try:
156 width = int(metadata.get(SESSION_VIDEO_WIDTH_KEY))
157 height = int(metadata.get(SESSION_VIDEO_HEIGHT_KEY))
158 except (TypeError, ValueError):
159 width = height = 0
160 if width > 0 and height > 0:
161 base["width"] = width
162 base["height"] = height
163
164 language = normalize_language(metadata.get(SESSION_LANGUAGE_KEY))
165 if language is not None:
166 base["language"] = language
167
168 if n_shots is not None:
169 base["n_shots"] = n_shots
170 mapped = n_shots_to_duration(n_shots)
171 if mapped is not None:
172 base["duration_sec"] = mapped
173 elif duration_sec is not None:
174 base["duration_sec"] = duration_sec
175 mapped = duration_to_n_shots(duration_sec)
176 if mapped is not None:
177 base["n_shots"] = mapped
178 return base
179
180
181 def get_llm_sampling_settings(metadata: dict[str, Any] | None) -> dict[str, float | int]:
182 """Return only explicitly set LLM sampling params (empty dict = gateway defaults)."""
183 if not isinstance(metadata, dict):
184 return {}
185 out: dict[str, float | int] = {}
186 temp = normalize_llm_temperature(metadata.get(SESSION_LLM_TEMPERATURE_KEY))
187 if temp is not None:
188 out["temperature"] = temp
189 top_p = normalize_llm_top_p(metadata.get(SESSION_LLM_TOP_P_KEY))
190 if top_p is not None:
191 out["top_p"] = top_p
192 top_k = normalize_llm_top_k(metadata.get(SESSION_LLM_TOP_K_KEY))
193 if top_k is not None:
194 out["top_k"] = top_k
195 return out
196
197
198 def get_llm_sampling_for_api(metadata: dict[str, Any] | None) -> dict[str, float | int | None]:
199 """API-facing view: keys always present, null when unset."""
200 sampling = get_llm_sampling_settings(metadata)
201 return {
202 "temperature": sampling.get("temperature"),
203 "top_p": sampling.get("top_p"),
204 "top_k": sampling.get("top_k"),
205 }
206
207
208 def apply_llm_sampling_settings(
209 metadata: dict[str, Any],
210 *,
211 temperature: Any = _UNSET,
212 top_p: Any = _UNSET,
213 top_k: Any = _UNSET,
214 ) -> dict[str, float | int | None]:
215 """Persist optional LLM sampling params. Pass ``None`` to clear a field."""
216 resolved = get_llm_sampling_for_api(metadata)
217
218 if temperature is not _UNSET:
219 if temperature is None or temperature == "":
220 metadata.pop(SESSION_LLM_TEMPERATURE_KEY, None)
221 resolved["temperature"] = None
222 else:
223 normalized = normalize_llm_temperature(temperature)
224 if normalized is None:
225 raise ValueError(f"invalid temperature: {temperature}")
226 metadata[SESSION_LLM_TEMPERATURE_KEY] = normalized
227 resolved["temperature"] = normalized
228
229 if top_p is not _UNSET:
230 if top_p is None or top_p == "":
231 metadata.pop(SESSION_LLM_TOP_P_KEY, None)
232 resolved["top_p"] = None
233 else:
234 normalized = normalize_llm_top_p(top_p)
235 if normalized is None:
236 raise ValueError(f"invalid top_p: {top_p}")
237 metadata[SESSION_LLM_TOP_P_KEY] = normalized
238 resolved["top_p"] = normalized
239
240 if top_k is not _UNSET:
241 if top_k is None or top_k == "":
242 metadata.pop(SESSION_LLM_TOP_K_KEY, None)
243 resolved["top_k"] = None
244 else:
245 normalized = normalize_llm_top_k(top_k)
246 if normalized is None:
247 raise ValueError(f"invalid top_k: {top_k}")
248 metadata[SESSION_LLM_TOP_K_KEY] = normalized
249 resolved["top_k"] = normalized
250
251 return resolved
252
253
254 def apply_llm_sampling_from_wire(metadata: dict[str, Any], wire: dict[str, Any] | None) -> bool:
255 """Apply LLM sampling overrides from a WS message envelope. Returns True if updated."""
256 if not isinstance(wire, dict):
257 return False
258 updates: dict[str, Any] = {}
259 for src, dst in (
260 ("temperature", "temperature"),
261 ("topP", "top_p"),
262 ("top_p", "top_p"),
263 ("topK", "top_k"),
264 ("top_k", "top_k"),
265 ):
266 if src in wire:
267 updates[dst] = wire[src]
268 if not updates:
269 return False
270 apply_llm_sampling_settings(metadata, **updates)
271 return True
272
273
274 def apply_generation_settings(
275 metadata: dict[str, Any],
276 *,
277 n_shots: int | None = None,
278 duration_sec: int | None = None,
279 width: int | None = None,
280 height: int | None = None,
281 language: str | None = None,
282 ) -> dict[str, int | str]:
283 """Persist Director generation settings and return the resolved values."""
284 resolved = get_generation_settings(metadata)
285
286 if duration_sec is not None:
287 normalized_duration = normalize_duration_sec(duration_sec)
288 if normalized_duration is None:
289 raise ValueError(f"invalid duration_sec: {duration_sec}")
290 resolved["duration_sec"] = normalized_duration
291 resolved["n_shots"] = duration_to_n_shots(normalized_duration) or DEFAULT_NSHOT
292 elif n_shots is not None:
293 normalized_n = normalize_n_shots(n_shots)
294 if normalized_n is None:
295 raise ValueError(f"invalid n_shots: {n_shots}")
296 resolved["n_shots"] = normalized_n
297 resolved["duration_sec"] = n_shots_to_duration(normalized_n) or DEFAULT_DURATION_SEC
298
299 metadata[SESSION_NSHOT_KEY] = resolved["n_shots"]
300 metadata[SESSION_DURATION_KEY] = resolved["duration_sec"]
301 if width is not None:
302 parsed_width = int(width)
303 if parsed_width <= 0:
304 raise ValueError(f"invalid width: {width}")
305 resolved["width"] = parsed_width
306 if height is not None:
307 parsed_height = int(height)
308 if parsed_height <= 0:
309 raise ValueError(f"invalid height: {height}")
310 resolved["height"] = parsed_height
311 metadata[SESSION_VIDEO_WIDTH_KEY] = resolved["width"]
312 metadata[SESSION_VIDEO_HEIGHT_KEY] = resolved["height"]
313 if language is not None:
314 normalized_language = normalize_language(language)
315 if normalized_language is None:
316 raise ValueError(f"invalid language: {language}")
317 resolved["language"] = normalized_language
318 metadata[SESSION_LANGUAGE_KEY] = resolved["language"]
319 return resolved
320
321
322 def resolve_n_shots_from_wire(data: dict[str, Any] | None) -> int | None:
323 """Resolve explicit shot count from WebSocket envelope or inbound message metadata."""
324 if not isinstance(data, dict):
325 return None
326 for key in ("nShot", "nshot", "n_shots", "nShots"):
327 raw = data.get(key)
328 if raw is None or raw == "":
329 continue
330 parsed = normalize_n_shots(raw)
331 if parsed is not None:
332 return parsed
333 duration = normalize_duration_sec(data.get("durationSec") or data.get("duration_sec"))
334 if duration is not None:
335 return duration_to_n_shots(duration)
336 return None
337
338
338 lines PYTHON