@@ -93,6 +93,17 @@ def _is_music(metadata: dict[str, Any]) -> bool:
9393 return str (metadata .get ("family" ) or "" ).casefold () == "music"
9494
9595
96+ def _choice_values (definition : Any ) -> list [str ]:
97+ if not isinstance (definition , dict ):
98+ return []
99+ values : list [str ] = []
100+ for choice in definition .get ("choices" , []):
101+ value = choice .get ("value" ) if isinstance (choice , dict ) else choice
102+ if value is not None and str (value ):
103+ values .append (str (value ))
104+ return values
105+
106+
96107def _models (session : Any , model_kind : str ) -> list [dict [str , Any ]]:
97108 """Translate WanGP model metadata into NodeTool provider models."""
98109 if model_kind not in {"video" , "image" , "tts" , "music" }:
@@ -127,7 +138,42 @@ def _models(session: Any, model_kind: str) -> list[dict[str, Any]]:
127138 "provider" : "wangp" ,
128139 }
129140 if model_kind == "tts" :
130- model ["capabilities" ] = supported
141+ base_model_type = str (metadata .get ("base_model_type" ) or model_type )
142+ tts_capabilities = list (supported )
143+ media_inputs = metadata .get ("media_inputs" )
144+ audio_inputs = (
145+ media_inputs .get ("audio" , {}) if isinstance (media_inputs , dict ) else {}
146+ )
147+ if audio_inputs .get ("prompt" ):
148+ tts_capabilities .append ("voice_cloning" )
149+ if base_model_type in {"qwen3_tts_base" , "omnivoice" }:
150+ tts_capabilities .append ("reference_transcript" )
151+ if base_model_type in {
152+ "qwen3_tts_customvoice" ,
153+ "index_tts2" ,
154+ "index_tts25" ,
155+ }:
156+ tts_capabilities .append ("instruction_control" )
157+ if base_model_type in {"qwen3_tts_voicedesign" , "omnivoice" }:
158+ tts_capabilities .append ("voice_design" )
159+
160+ setting_values = metadata .get ("setting_values" )
161+ model_mode = (
162+ setting_values .get ("model_mode" )
163+ if isinstance (setting_values , dict )
164+ else None
165+ )
166+ mode_label = str (
167+ model_mode .get ("label" , "" ) if isinstance (model_mode , dict ) else ""
168+ ).casefold ()
169+ mode_values = _choice_values (model_mode )
170+ if mode_label == "speaker" :
171+ tts_capabilities .append ("preset_voice" )
172+ model ["voices" ] = mode_values
173+ elif mode_label == "language" :
174+ tts_capabilities .append ("language_selection" )
175+ model ["languages" ] = mode_values
176+ model ["capabilities" ] = list (dict .fromkeys (tts_capabilities ))
131177 else :
132178 model ["supportedTasks" ] = supported
133179 models .append (model )
@@ -205,6 +251,9 @@ def _settings(
205251 elif params .get ("durationSeconds" ) is not None and operation .endswith ("video" ):
206252 settings ["video_length" ] = f"{ float (params ['durationSeconds' ]):g} s"
207253
254+ if operation in {"text_to_image" , "image_to_image" }:
255+ settings ["image_mode" ] = 1
256+
208257 if operation in {"image_to_image" , "image_to_video" }:
209258 image_path = str (_input_path (request , "image_path" ))
210259 image_inputs = (metadata or {}).get ("media_inputs" , {}).get ("image" , {})
@@ -226,8 +275,12 @@ def _settings(
226275 if operation == "text_to_audio" :
227276 style_prompt = str (params .get ("prompt" ) or "" )
228277 lyrics = str (params .get ("lyrics" ) or "" ).strip ()
229- settings ["prompt" ] = lyrics or "[Instrumental]"
230- settings ["alt_prompt" ] = style_prompt
278+ base_model_type = str ((metadata or {}).get ("base_model_type" ) or model_type )
279+ if base_model_type .startswith ("stable_audio3" ):
280+ settings ["prompt" ] = style_prompt
281+ else :
282+ settings ["prompt" ] = lyrics or "[Instrumental]"
283+ settings ["alt_prompt" ] = style_prompt
231284 if params .get ("durationSeconds" ) is not None :
232285 settings ["duration_seconds" ] = float (params ["durationSeconds" ])
233286
@@ -236,11 +289,16 @@ def _settings(
236289 settings ["alt_prompt" ] = str (params ["referenceText" ])
237290 elif params .get ("instructions" ) is not None :
238291 settings ["alt_prompt" ] = str (params ["instructions" ])
239- model_mode = params .get ("voice" ) or params .get ("language" )
292+ base_model_type = str ((metadata or {}).get ("base_model_type" ) or model_type )
293+ model_mode = (
294+ params .get ("voice" )
295+ if base_model_type == "qwen3_tts_customvoice"
296+ else params .get ("language" )
297+ )
240298 if model_mode :
241299 settings ["model_mode" ] = str (model_mode )
242- if params .get ("speed" ) is not None :
243- settings ["speech_speed " ] = float (params ["speed" ])
300+ if params .get ("speed" ) is not None and base_model_type == "index_tts25" :
301+ settings ["custom_settings " ] = { "speech_speed" : float (params ["speed" ])}
244302 if request .get ("reference_audio_path" ):
245303 settings ["audio_guide" ] = str (
246304 _input_path (request , "reference_audio_path" )
@@ -276,9 +334,18 @@ def _generated_path(result: Any, media_type: str | None = None) -> str:
276334 path = getattr (artifact , "path" , None )
277335 if path and Path (path ).is_file ():
278336 return str (Path (path ).resolve ())
337+ media_suffixes = {
338+ "image" : {".bmp" , ".gif" , ".jpeg" , ".jpg" , ".png" , ".tif" , ".tiff" , ".webp" },
339+ "video" : {".avi" , ".mkv" , ".mov" , ".mp4" , ".webm" },
340+ "audio" : {".aac" , ".flac" , ".m4a" , ".mp3" , ".ogg" , ".opus" , ".wav" },
341+ }
342+ expected_suffixes = media_suffixes .get (media_type or "" )
279343 for path in getattr (result , "generated_files" , ()):
280344 candidate = Path (str (path ))
281- if candidate .is_file ():
345+ if candidate .is_file () and (
346+ expected_suffixes is None
347+ or candidate .suffix .casefold () in expected_suffixes
348+ ):
282349 return str (candidate .resolve ())
283350 raise RuntimeError ("WanGP completed without a generated media file" )
284351
0 commit comments