update:优化配置及文档

This commit is contained in:
hrz
2025-04-30 15:05:42 +08:00
parent 9e9af2a031
commit d2f8f05acb
18 changed files with 167 additions and 137 deletions
@@ -53,6 +53,8 @@ async def handleTextMessage(conn, message):
# 如果是唤醒词,且关闭了唤醒词回复,就不用回答
await send_stt_message(conn, text)
await send_tts_message(conn, "stop", None)
elif is_wakeup_words:
await startToChat(conn, "嘿,你好呀")
else:
# 否则需要LLM对文字内容进行答复
await startToChat(conn, text)
@@ -27,9 +27,14 @@ class TTSProvider(TTSProviderBase):
else:
self.voice = config.get("voice")
self.speed_ratio = float(config.get("speed_ratio", 0.1))
self.volume_ratio = float(config.get("volume_ratio", 0.1))
self.pitch_ratio = float(config.get("pitch_ratio", 0.1))
# 处理空字符串的情况
speed_ratio = config.get("speed_ratio", "1.0")
volume_ratio = config.get("volume_ratio", "1.0")
pitch_ratio = config.get("pitch_ratio", "1.0")
self.speed_ratio = float(speed_ratio) if speed_ratio else 1.0
self.volume_ratio = float(volume_ratio) if volume_ratio else 1.0
self.pitch_ratio = float(pitch_ratio) if pitch_ratio else 1.0
self.api_url = config.get("api_url")
self.authorization = config.get("authorization")
@@ -89,18 +89,35 @@ class TTSProvider(TTSProviderBase):
self.reference_audio = parse_string_to_list(config.get("reference_audio"))
self.reference_text = parse_string_to_list(config.get("reference_text"))
self.format = config.get("format", "wav")
self.channels = int(config.get("channels", 1))
self.rate = int(config.get("rate", 44100))
self.api_key = config.get("api_key", "YOUR_API_KEY")
have_key = check_model_key("FishSpeech TTS", self.api_key)
if not have_key:
return
self.normalize = config.get("normalize", True)
self.max_new_tokens = int(config.get("max_new_tokens", 1024))
self.chunk_length = int(config.get("chunk_length", 200))
self.top_p = float(config.get("top_p", 0.7))
self.repetition_penalty = float(config.get("repetition_penalty", 1.2))
self.temperature = float(config.get("temperature", 0.7))
# 处理空字符串的情况
channels = config.get("channels", "1")
rate = config.get("rate", "44100")
max_new_tokens = config.get("max_new_tokens", "1024")
chunk_length = config.get("chunk_length", "200")
self.channels = int(channels) if channels else 1
self.rate = int(rate) if rate else 44100
self.max_new_tokens = int(max_new_tokens) if max_new_tokens else 1024
self.chunk_length = int(chunk_length) if chunk_length else 200
# 处理空字符串的情况
top_p = config.get("top_p", "0.7")
temperature = config.get("temperature", "0.7")
repetition_penalty = config.get("repetition_penalty", "1.2")
self.top_p = float(top_p) if top_p else 0.7
self.temperature = float(temperature) if temperature else 0.7
self.repetition_penalty = (
float(repetition_penalty) if repetition_penalty else 1.2
)
self.streaming = str(config.get("streaming", False)).lower() in (
"true",
"1",
@@ -20,12 +20,29 @@ class TTSProvider(TTSProviderBase):
self.ref_audio_path = config.get("ref_audio_path")
self.prompt_text = config.get("prompt_text")
self.prompt_lang = config.get("prompt_lang", "zh")
self.top_k = int(config.get("top_k", 5))
self.top_p = float(config.get("top_p", 1))
self.temperature = float(config.get("temperature", 1))
# 处理空字符串的情况
top_k = config.get("top_k", "5")
top_p = config.get("top_p", "1")
temperature = config.get("temperature", "1")
batch_threshold = config.get("batch_threshold", "0.75")
batch_size = config.get("batch_size", "1")
speed_factor = config.get("speed_factor", "1.0")
seed = config.get("seed", "-1")
repetition_penalty = config.get("repetition_penalty", "1.35")
self.top_k = int(top_k) if top_k else 5
self.top_p = float(top_p) if top_p else 1
self.temperature = float(temperature) if temperature else 1
self.batch_threshold = float(batch_threshold) if batch_threshold else 0.75
self.batch_size = int(batch_size) if batch_size else 1
self.speed_factor = float(speed_factor) if speed_factor else 1.0
self.seed = int(seed) if seed else -1
self.repetition_penalty = (
float(repetition_penalty) if repetition_penalty else 1.35
)
self.text_split_method = config.get("text_split_method", "cut0")
self.batch_size = int(config.get("batch_size", 1))
self.batch_threshold = float(config.get("batch_threshold", 0.75))
self.split_bucket = str(config.get("split_bucket", True)).lower() in (
"true",
@@ -37,19 +54,19 @@ class TTSProvider(TTSProviderBase):
"1",
"yes",
)
self.speed_factor = float(config.get("speed_factor", 1.0))
self.streaming_mode = str(config.get("streaming_mode", False)).lower() in (
"true",
"1",
"yes",
)
self.seed = int(config.get("seed", -1))
self.parallel_infer = str(config.get("parallel_infer", True)).lower() in (
"true",
"1",
"yes",
)
self.repetition_penalty = float(config.get("repetition_penalty", 1.35))
self.aux_ref_audio_paths = parse_string_to_list(
config.get("aux_ref_audio_paths")
)
@@ -18,13 +18,22 @@ class TTSProvider(TTSProviderBase):
self.prompt_text = config.get("prompt_text")
self.prompt_language = config.get("prompt_language")
self.text_language = config.get("text_language", "audo")
self.top_k = int(config.get("top_k", 15))
self.top_p = float(config.get("top_p", 1.0))
self.temperature = float(config.get("temperature", 1.0))
# 处理空字符串的情况
top_k = config.get("top_k", "15")
top_p = config.get("top_p", "1.0")
temperature = config.get("temperature", "1.0")
sample_steps = config.get("sample_steps", "32")
speed = config.get("speed", "1.0")
self.top_k = int(top_k) if top_k else 15
self.top_p = float(top_p) if top_p else 1.0
self.temperature = float(temperature) if temperature else 1.0
self.sample_steps = int(sample_steps) if sample_steps else 32
self.speed = float(speed) if speed else 1.0
self.cut_punc = config.get("cut_punc", "")
self.speed = float(config.get("speed", 1.0))
self.inp_refs = parse_string_to_list(config.get("inp_refs"))
self.sample_steps = int(config.get("sample_steps", 32))
self.if_sr = str(config.get("if_sr", False)).lower() in ("true", "1", "yes")
def generate_filename(self, extension=".wav"):
@@ -21,7 +21,11 @@ class TTSProvider(TTSProviderBase):
else:
self.voice = config.get("voice", "alloy")
self.response_format = "wav"
self.speed = float(config.get("speed", 1.0))
# 处理空字符串的情况
speed = config.get("speed", "1.0")
self.speed = float(speed) if speed else 1.0
self.output_file = config.get("output_dir", "tmp/")
check_model_key("TTS", self.api_key)
@@ -21,8 +21,15 @@ class VADProvider(VADProviderBase):
(get_speech_timestamps, _, _, _, _) = self.utils
self.decoder = opuslib_next.Decoder(16000, 1)
self.vad_threshold = float(config.get("threshold", 0.5))
self.silence_threshold_ms = int(config.get("min_silence_duration_ms", 1000))
# 处理空字符串的情况
threshold = config.get("threshold", "0.5")
min_silence_duration_ms = config.get("min_silence_duration_ms", "1000")
self.vad_threshold = float(threshold) if threshold else 0.5
self.silence_threshold_ms = (
int(min_silence_duration_ms) if min_silence_duration_ms else 1000
)
def is_vad(self, conn, opus_packet):
try: