fix: resolve KeyError 'pure_1' in CosyVoice 3.0 by adding speaker validation and better defaults

This commit is contained in:
Hawk Lee
2025-12-31 00:34:08 +08:00
parent b217dcee89
commit 12e589394e
+22 -9
View File
@@ -483,7 +483,7 @@ class AIIA_CosyVoice_TTS:
"model": ("COSYVOICE_MODEL",),
"tts_text": ("STRING", {"multiline": True, "default": "你好,这是 CosyVoice 3.0 的全能模式测试。"}),
"instruct_text": ("STRING", {"multiline": True, "default": "一个沉稳、磁性的成熟男性声音,语法标准,情感饱满。"}),
"spk_id": ("STRING", {"default": "pure_1", "tooltip": "V1/V2 时代的固定 ID,如 pure_1, joy_1 等。如果提供参考音频且非 SFT 模型,此项可能被忽略。"}),
"spk_id": ("STRING", {"default": "", "tooltip": "固定音色 ID (如 pure_1)。对于 0.5B/V3 等 Zero-Shot 模型,此项通常为空,需配合参考音频使用。"}),
"speed": ("FLOAT", {"default": 1.0, "min": 0.5, "max": 2.0, "step": 0.1}),
"seed": ("INT", {"default": 42, "min": -1, "max": 2147483647}),
},
@@ -512,7 +512,24 @@ class AIIA_CosyVoice_TTS:
is_v3 = "CosyVoice3" in type(cosyvoice_model).__name__
is_v2 = "CosyVoice2" in type(cosyvoice_model).__name__
# 1. Hybrid / Cross-Lingual / Zero-Shot (Reference Audio provided)
# 1. Identity Validation
available_spks = list(cosyvoice_model.frontend.spk2info.keys())
# Case A: User specified a Speaker ID
if spk_id:
if spk_id not in available_spks:
raise ValueError(f"Speaker ID '{spk_id}' not found. Available: {available_spks if available_spks else 'None (Zero-Shot model)'}")
# Case B: No ID and no Audio
elif reference_audio is None:
if available_spks:
spk_id = available_spks[0]
print(f"[AIIA] Auto-selecting first available speaker: {spk_id}")
else:
# 0.5B models typically have no spk2info
raise ValueError("此模型 (CosyVoice 0.5B/V3) 没有内置音色库。请连接 'reference_audio' 进行零样本克隆或生成。")
# 2. Hybrid / Cross-Lingual / Zero-Shot (Reference Audio provided)
if reference_audio is not None:
ref_wav = reference_audio["waveform"]
if reference_audio["sample_rate"] != sample_rate:
@@ -529,20 +546,17 @@ class AIIA_CosyVoice_TTS:
try:
if is_v3 or is_v2:
# V3/V2: Support Instruct + Audio (Hybrid Mode)
print(f"[AIIA] CosyVoice V3/V2: Hybrid/Zero-Shot Mode. Instruct: {instruct_text[:20]}...")
output = cosyvoice_model.inference_instruct2(
tts_text=tts_text,
instruct_text=instruct_text,
prompt_wav=ref_path,
zero_shot_spk_id=spk_id, # Optional placeholder
zero_shot_spk_id=spk_id,
stream=False,
speed=speed
)
else:
# V1: Source-Prompt based inference
print("[AIIA] CosyVoice V1: Zero-Shot / Cross-Lingual Mode")
# Some versions use inference_vc or inference_zero_shot
if hasattr(cosyvoice_model, 'inference_zero_shot'):
output = cosyvoice_model.inference_zero_shot(tts_text, ref_path, stream=False, speed=speed)
else:
@@ -553,11 +567,10 @@ class AIIA_CosyVoice_TTS:
finally:
if os.path.exists(ref_path): os.unlink(ref_path)
# 2. Instruct / SFT / Random (No Reference Audio)
# 3. Instruct / SFT / Random (No Reference Audio, verified we have spk_id)
else:
if is_v3 or is_v2:
print(f"[AIIA] CosyVoice V3/V2: Instruct Generation. Description: {instruct_text[:20]}...")
# For pure generation in V3, prompt_wav can be None
output = cosyvoice_model.inference_instruct2(
tts_text=tts_text,
instruct_text=instruct_text,
@@ -567,7 +580,6 @@ class AIIA_CosyVoice_TTS:
speed=speed
)
else:
# V1 SFT / Instruct
print(f"[AIIA] CosyVoice V1: SFT/Instruct. Speaker ID: {spk_id}")
if "SFT" in type(cosyvoice_model).__name__ or spk_id:
output = cosyvoice_model.inference_sft(tts_text, spk_id, stream=False, speed=speed)
@@ -578,6 +590,7 @@ class AIIA_CosyVoice_TTS:
final_waveform = torch.cat(all_speech, dim=-1)
except Exception as e:
if isinstance(e, ValueError): raise e
raise RuntimeError(f"CosyVoice generation failed: {e}")
return ({"waveform": final_waveform.unsqueeze(0).cpu(), "sample_rate": sample_rate},)