From c4c84e44e1a63be9671fcf1f9096127fe98bcfdb Mon Sep 17 00:00:00 2001 From: hrz <1710360675@qq.com> Date: Wed, 21 May 2025 14:52:24 +0800 Subject: [PATCH] =?UTF-8?q?update:=E5=90=88=E5=B9=B6main=E5=88=86=E6=94=AF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../core/providers/asr/aliyun.py | 2 +- .../core/providers/tts/aliyun.py | 99 +++++++---- .../xiaozhi-server/core/providers/tts/base.py | 57 ++----- .../core/providers/tts/cozecn.py | 26 ++- .../core/providers/tts/custom.py | 18 +- .../xiaozhi-server/core/providers/tts/edge.py | 5 +- .../core/providers/tts/fishspeech.py | 155 ++++++++++++------ .../core/providers/tts/gpt_sovits_v2.py | 87 +++++++--- .../core/providers/tts/gpt_sovits_v3.py | 45 +++-- .../core/providers/tts/minimax.py | 58 ++++--- .../core/providers/tts/openai.py | 34 +++- .../core/providers/tts/siliconflow.py | 27 ++- .../core/providers/tts/ttson.py | 78 +++++---- 13 files changed, 458 insertions(+), 233 deletions(-) diff --git a/main/xiaozhi-server/core/providers/asr/aliyun.py b/main/xiaozhi-server/core/providers/asr/aliyun.py index 6606168c..fee62364 100644 --- a/main/xiaozhi-server/core/providers/asr/aliyun.py +++ b/main/xiaozhi-server/core/providers/asr/aliyun.py @@ -239,7 +239,7 @@ class ASRProvider(ASRProviderBase): ) -> Tuple[Optional[str], Optional[str]]: """将语音数据转换为文本""" if self._is_token_expired(): - logger.warning("Token已过期,正在自动刷新...") + logger.bind(tag=TAG).warning("Token已过期,正在自动刷新...") self._refresh_token() file_path = None diff --git a/main/xiaozhi-server/core/providers/tts/aliyun.py b/main/xiaozhi-server/core/providers/tts/aliyun.py index 611f5a2c..a71151d2 100644 --- a/main/xiaozhi-server/core/providers/tts/aliyun.py +++ b/main/xiaozhi-server/core/providers/tts/aliyun.py @@ -6,18 +6,18 @@ import hashlib import base64 import requests from datetime import datetime +from core.providers.tts.base import TTSProviderBase from pydub import AudioSegment -from core.providers.tts.base import TTSProviderBase - -import http.client -import urllib.parse import time import uuid from urllib import parse - from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType +from config.logger import setup_logging + +TAG = __name__ +logger = setup_logging() class AccessToken: @@ -48,7 +48,7 @@ class AccessToken: } # 构造规范化的请求字符串 query_string = AccessToken._encode_dict(parameters) - print("规范化的请求字符串: %s" % query_string) + # print('规范化的请求字符串: %s' % query_string) # 构造待签名字符串 string_to_sign = ( "GET" @@ -57,7 +57,7 @@ class AccessToken: + "&" + AccessToken._encode_text(query_string) ) - print("待签名的字符串: %s" % string_to_sign) + # print('待签名的字符串: %s' % string_to_sign) # 计算签名 secreted_string = hmac.new( bytes(access_key_secret + "&", encoding="utf-8"), @@ -65,16 +65,16 @@ class AccessToken: hashlib.sha1, ).digest() signature = base64.b64encode(secreted_string) - print("签名: %s" % signature) + # print('签名: %s' % signature) # 进行URL编码 signature = AccessToken._encode_text(signature) - print("URL编码后的签名: %s" % signature) + # print('URL编码后的签名: %s' % signature) # 调用服务 full_url = "http://nls-meta.cn-shanghai.aliyuncs.com/?Signature=%s&%s" % ( signature, query_string, ) - print("url: %s" % full_url) + # print('url: %s' % full_url) # 提交HTTP GET请求 response = requests.get(full_url) if response.ok: @@ -84,7 +84,7 @@ class AccessToken: token = root_obj[key]["Id"] expire_time = root_obj[key]["ExpireTime"] return token, expire_time - print(response.text) + # print(response.text) return None, None @@ -94,33 +94,67 @@ class TTSProvider(TTSProviderBase): super().__init__(config, delete_audio_file) # 新增空值判断逻辑 - access_key_id = config.get("access_key_id") - access_key_secret = config.get("access_key_secret") - if access_key_id and access_key_secret: - # 使用密钥对生成临时token - token, expire_time = AccessToken.create_token( - access_key_id, access_key_secret - ) - else: - # 直接使用预生成的长期token - token = config.get("token") - expire_time = None - - print("token: %s, expire time(s): %s" % (token, expire_time)) + self.access_key_id = config.get("access_key_id") + self.access_key_secret = config.get("access_key_secret") self.appkey = config.get("appkey") - self.token = token self.format = config.get("format", "wav") self.sample_rate = config.get("sample_rate", 16000) self.voice = config.get("voice", "xiaoyun") self.volume = config.get("volume", 50) self.speech_rate = config.get("speech_rate", 0) self.pitch_rate = config.get("pitch_rate", 0) - self.host = config.get("host", "nls-gateway-cn-shanghai.aliyuncs.com") self.api_url = f"https://{self.host}/stream/v1/tts" self.header = {"Content-Type": "application/json"} + if self.access_key_id and self.access_key_secret: + # 使用密钥对生成临时token + self._refresh_token() + else: + # 直接使用预生成的长期token + self.token = config.get("token") + self.expire_time = None + + def _refresh_token(self): + """刷新Token并记录过期时间""" + if self.access_key_id and self.access_key_secret: + self.token, expire_time_str = AccessToken.create_token( + self.access_key_id, self.access_key_secret + ) + if not expire_time_str: + raise ValueError("无法获取有效的Token过期时间") + + try: + # 统一转换为字符串处理 + expire_str = str(expire_time_str).strip() + + if expire_str.isdigit(): + expire_time = datetime.fromtimestamp(int(expire_str)) + else: + expire_time = datetime.strptime(expire_str, "%Y-%m-%dT%H:%M:%SZ") + self.expire_time = expire_time.timestamp() - 60 + except Exception as e: + raise ValueError(f"无效的过期时间格式: {expire_str}") from e + + else: + self.expire_time = None + + if not self.token: + raise ValueError("无法获取有效的访问Token") + + def _is_token_expired(self): + """检查Token是否过期""" + if not self.expire_time: + return False # 长期Token不过期 + # 新增调试日志 + # current_time = time.time() + # remaining = self.expire_time - current_time + # print(f"Token过期检查: 当前时间 {datetime.fromtimestamp(current_time)} | " + # f"过期时间 {datetime.fromtimestamp(self.expire_time)} | " + # f"剩余 {remaining:.2f}秒") + return time.time() > self.expire_time + def generate_filename(self, extension=".wav"): return os.path.join( self.output_file, @@ -128,6 +162,9 @@ class TTSProvider(TTSProviderBase): ) async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False): + if self._is_token_expired(): + logger.bind(tag=TAG).warning("Token已过期,正在自动刷新...") + self._refresh_token() request_json = { "appkey": self.appkey, "token": self.token, @@ -141,12 +178,17 @@ class TTSProvider(TTSProviderBase): } print(self.api_url, json.dumps(request_json, ensure_ascii=False)) - tmp_file = self.generate_filename() try: resp = requests.post( self.api_url, json.dumps(request_json), headers=self.header ) + if resp.status_code == 401: # Token过期特殊处理 + self._refresh_token() + resp = requests.post( + self.api_url, json.dumps(request_json), headers=self.header + ) # 检查返回请求数据的mime类型是否是audio/***,是则保存到指定路径下;返回的是binary格式的 + tmp_file = self.generate_filename() if resp.headers["Content-Type"].startswith("audio/"): with open(tmp_file, "wb") as f: f.write(resp.content) @@ -154,7 +196,7 @@ class TTSProvider(TTSProviderBase): raise Exception( f"{__name__} status_code: {resp.status_code} response: {resp.content}" ) - # 使用 pydub 读取临时文件 + # 使用 pydub 读取临时文件 audio = AudioSegment.from_file(tmp_file, format="wav") audio = audio.set_channels(1).set_frame_rate(16000) opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data) @@ -171,5 +213,6 @@ class TTSProvider(TTSProviderBase): except FileNotFoundError: # 若文件不存在,忽略该异常 pass + except Exception as e: raise Exception(f"{__name__} error: {e}") diff --git a/main/xiaozhi-server/core/providers/tts/base.py b/main/xiaozhi-server/core/providers/tts/base.py index 0b08b862..481f1bc5 100644 --- a/main/xiaozhi-server/core/providers/tts/base.py +++ b/main/xiaozhi-server/core/providers/tts/base.py @@ -10,12 +10,11 @@ import torch import torchaudio from config.logger import setup_logging -import os import numpy as np -import opuslib_next from pydub import AudioSegment from abc import ABC, abstractmethod from core.utils import textUtils +from core.utils.util import audio_to_data from core.opus import opus_encoder_utils import queue @@ -34,7 +33,9 @@ class TTSProviderBase(ABC): self.tts_audio_queue = queue.Queue() self.enable_two_way = False self.stop_event = threading.Event() - self.opus_encoder = opus_encoder_utils.OpusEncoderUtils(sample_rate=16000, channels=1, frame_size_ms=60) + self.opus_encoder = opus_encoder_utils.OpusEncoderUtils( + sample_rate=16000, channels=1, frame_size_ms=60 + ) self.tts_text_buff = [] self.punctuations = ( @@ -79,12 +80,12 @@ class TTSProviderBase(ABC): def _get_segment_text(self): # 合并当前全部文本并处理未分割部分 full_text = "".join(self.tts_text_buff) - current_text = full_text[self.processed_chars:] # 从未处理的位置开始 + current_text = full_text[self.processed_chars :] # 从未处理的位置开始 last_punct_pos = -1 for punct in self.punctuations: pos = current_text.rfind(punct) if (pos != -1 and last_punct_pos == -1) or ( - pos != -1 and pos < last_punct_pos + pos != -1 and pos < last_punct_pos ): last_punct_pos = pos if last_punct_pos != -1: @@ -220,6 +221,7 @@ class TTSProviderBase(ABC): ) self.active_tasks.add(future) if self.active_tasks: + async def wrap_future(future): return await asyncio.wrap_future(future) @@ -292,48 +294,13 @@ class TTSProviderBase(ABC): async def text_to_speak_stream(self, text, queue: queue.Queue, text_index=0): raise Exception("该TTS还没有实现stream模式") + def audio_to_pcm_data(self, audio_file_path): + """音频文件转换为PCM编码""" + return audio_to_data(audio_file_path, is_opus=False) + def audio_to_opus_data(self, audio_file_path): """音频文件转换为Opus编码""" - # 获取文件后缀名 - file_type = os.path.splitext(audio_file_path)[1] - if file_type: - file_type = file_type.lstrip(".") - audio = AudioSegment.from_file(audio_file_path, format=file_type) - - # 转换为单声道/16kHz采样率/16位小端编码(确保与编码器匹配) - audio = audio.set_channels(1).set_frame_rate(16000).set_sample_width(2) - - # 音频时长(秒) - duration = len(audio) / 1000.0 - - # 获取原始PCM数据(16位小端) - raw_data = audio.raw_data - - # 初始化Opus编码器 - encoder = opuslib_next.Encoder(16000, 1, opuslib_next.APPLICATION_AUDIO) - - # 编码参数 - frame_duration = 60 # 60ms per frame - frame_size = int(16000 * frame_duration / 1000) # 960 samples/frame - - opus_datas = [] - # 按帧处理所有音频数据(包括最后一帧可能补零) - for i in range(0, len(raw_data), frame_size * 2): # 16bit=2bytes/sample - # 获取当前帧的二进制数据 - chunk = raw_data[i: i + frame_size * 2] - - # 如果最后一帧不足,补零 - if len(chunk) < frame_size * 2: - chunk += b"\x00" * (frame_size * 2 - len(chunk)) - - # 转换为numpy数组处理 - np_frame = np.frombuffer(chunk, dtype=np.int16) - - # 编码Opus数据 - opus_data = encoder.encode(np_frame.tobytes(), frame_size) - opus_datas.append(opus_data) - - return opus_datas, duration + return audio_to_data(audio_file_path, is_opus=True) def get_audio_from_tts(self, data_bytes, src_rate, to_rate=16000): tts_speech = torch.from_numpy( diff --git a/main/xiaozhi-server/core/providers/tts/cozecn.py b/main/xiaozhi-server/core/providers/tts/cozecn.py index 51ab4df9..61221e85 100644 --- a/main/xiaozhi-server/core/providers/tts/cozecn.py +++ b/main/xiaozhi-server/core/providers/tts/cozecn.py @@ -16,14 +16,21 @@ class TTSProvider(TTSProviderBase): super().__init__(config, delete_audio_file) self.model = config.get("model") self.access_token = config.get("access_token") - self.voice = config.get("voice") + if config.get("private_voice"): + self.voice = config.get("private_voice") + else: + self.voice = config.get("voice") + self.response_format = config.get("response_format") self.host = "api.coze.cn" self.api_url = f"https://{self.host}/v1/audio/speech" def generate_filename(self, extension=".wav"): - return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}") + return os.path.join( + self.output_file, + f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}", + ) async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False): request_json = { @@ -34,9 +41,11 @@ class TTSProvider(TTSProviderBase): } headers = { "Authorization": f"Bearer {self.access_token}", - "Content-Type": "application/json" + "Content-Type": "application/json", } - response = requests.request("POST", self.api_url, json=request_json, headers=headers) + response = requests.request( + "POST", self.api_url, json=request_json, headers=headers + ) data = response.content tmp_file = self.generate_filename() file_to_save = open(tmp_file, "wb") @@ -45,8 +54,13 @@ class TTSProvider(TTSProviderBase): audio = AudioSegment.from_file(tmp_file, format="wav") audio = audio.set_channels(1).set_frame_rate(16000) opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data) - yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text, - sentence_type=SentenceType.SENTENCE_START) + yield TTSMessageDTO( + u_id=u_id, + msg_type=MsgType.TTS_TEXT_RESPONSE, + content=opus_datas, + tts_finish_text=text, + sentence_type=SentenceType.SENTENCE_START, + ) # 用完后删除临时文件 try: os.remove(tmp_file) diff --git a/main/xiaozhi-server/core/providers/tts/custom.py b/main/xiaozhi-server/core/providers/tts/custom.py index 695117ab..695ac346 100644 --- a/main/xiaozhi-server/core/providers/tts/custom.py +++ b/main/xiaozhi-server/core/providers/tts/custom.py @@ -22,7 +22,10 @@ class TTSProvider(TTSProviderBase): self.output_file = config.get("output_dir", "tmp/") def generate_filename(self): - return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}.{self.format}") + return os.path.join( + self.output_file, + f"tts-{datetime.now().date()}@{uuid.uuid4().hex}.{self.format}", + ) async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False): request_params = {} @@ -37,13 +40,20 @@ class TTSProvider(TTSProviderBase): with open(tmp_file, "wb") as file: file.write(resp.content) else: - logger.bind(tag=TAG).error(f"Custom TTS请求失败: {resp.status_code} - {resp.text}") + logger.bind(tag=TAG).error( + f"Custom TTS请求失败: {resp.status_code} - {resp.text}" + ) # 使用 pydub 读取临时文件 audio = AudioSegment.from_file(tmp_file, format=self.format) audio = audio.set_channels(1).set_frame_rate(16000) opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data) - yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text, - sentence_type=SentenceType.SENTENCE_START) + yield TTSMessageDTO( + u_id=u_id, + msg_type=MsgType.TTS_TEXT_RESPONSE, + content=opus_datas, + tts_finish_text=text, + sentence_type=SentenceType.SENTENCE_START, + ) # 用完后删除临时文件 try: os.remove(tmp_file) diff --git a/main/xiaozhi-server/core/providers/tts/edge.py b/main/xiaozhi-server/core/providers/tts/edge.py index 8ebba47a..b623c77c 100644 --- a/main/xiaozhi-server/core/providers/tts/edge.py +++ b/main/xiaozhi-server/core/providers/tts/edge.py @@ -17,7 +17,10 @@ logger = setup_logging() class TTSProvider(TTSProviderBase): def __init__(self, config, delete_audio_file): super().__init__(config, delete_audio_file) - self.voice = config.get("voice") + if config.get("private_voice"): + self.voice = config.get("private_voice") + else: + self.voice = config.get("voice") def generate_filename(self, extension=".mp3"): return os.path.join( diff --git a/main/xiaozhi-server/core/providers/tts/fishspeech.py b/main/xiaozhi-server/core/providers/tts/fishspeech.py index 4acac43a..2049193e 100644 --- a/main/xiaozhi-server/core/providers/tts/fishspeech.py +++ b/main/xiaozhi-server/core/providers/tts/fishspeech.py @@ -35,7 +35,7 @@ class ServeReferenceAudio(BaseModel): def decode_audio(cls, values): audio = values.get("audio") if ( - isinstance(audio, str) and len(audio) > 255 + isinstance(audio, str) and len(audio) > 255 ): # Check if audio is a string (Base64) try: values["audio"] = base64.b64decode(audio) @@ -96,32 +96,64 @@ class TTSProvider(TTSProviderBase): def __init__(self, config, delete_audio_file): super().__init__(config, delete_audio_file) - self.reference_id = config.get("reference_id") - self.reference_audio = config.get("reference_audio", []) - self.reference_text = config.get("reference_text", []) - self.format = config.get("format", "wav") - self.channels = config.get("channels", 1) - self.rate = config.get("rate", 44100) + self.reference_id = ( + None if not config.get("reference_id") else config.get("reference_id") + ) + 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("response_format", "wav") + 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 = config.get("max_new_tokens", 1024) - self.chunk_length = config.get("chunk_length", 200) - self.top_p = config.get("top_p", 0.7) - self.repetition_penalty = config.get("repetition_penalty", 1.2) - self.temperature = config.get("temperature", 0.7) - self.streaming = config.get("streaming", False) + self.normalize = str(config.get("normalize", True)).lower() in ( + "true", + "1", + "yes", + ) + + # 处理空字符串的情况 + 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", + "yes", + ) self.use_memory_cache = config.get("use_memory_cache", "on") - self.seed = config.get("seed") + self.seed = int(config.get("seed")) if config.get("seed") else None self.api_url = config.get("api_url", "http://127.0.0.1:8080/v1/tts") def generate_filename(self, extension=".wav"): - return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}") + return os.path.join( + self.output_file, + f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}", + ) def _get_audio_from_tts(self, data_bytes): - tts_speech = torch.from_numpy(np.array(np.frombuffer(data_bytes, dtype=np.int16))).unsqueeze(dim=0) + tts_speech = torch.from_numpy( + np.array(np.frombuffer(data_bytes, dtype=np.int16)) + ).unsqueeze(dim=0) with io.BytesIO() as bf: torchaudio.save(bf, tts_speech, 44100, format="wav") audio = AudioSegment.from_file(bf, format="wav") @@ -147,29 +179,35 @@ class TTSProvider(TTSProviderBase): # Prepare reference data if self.reference_audio and self.reference_text: - byte_audios = [audio_to_bytes(ref_audio) for ref_audio in self.reference_audio] - ref_texts = [read_ref_text(ref_text) for ref_text in self.reference_text] - data["references"] = [ - ServeReferenceAudio( - audio=audio if audio else b"", text=text - ) - for text, audio in zip(ref_texts, byte_audios) - ], + byte_audios = [ + audio_to_bytes(ref_audio) for ref_audio in self.reference_audio + ] + ref_texts = [ + read_ref_text(ref_text) for ref_text in self.reference_text + ] + data["references"] = ( + [ + ServeReferenceAudio(audio=audio if audio else b"", text=text) + for text, audio in zip(ref_texts, byte_audios) + ], + ) data["reference_id"] = None pydantic_data = ServeTTSRequest(**data) audio_buff = None - chunk_total = b'' - last_raw = b'' - audio_raw = b'' + chunk_total = b"" + last_raw = b"" + audio_raw = b"" print("请求tts") with requests.post( - self.api_url, - data=ormsgpack.packb(pydantic_data, option=ormsgpack.OPT_SERIALIZE_PYDANTIC), - headers={ - "Authorization": f"Bearer {self.api_key}", - "Content-Type": "application/msgpack", - }, + self.api_url, + data=ormsgpack.packb( + pydantic_data, option=ormsgpack.OPT_SERIALIZE_PYDANTIC + ), + headers={ + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/msgpack", + }, ) as response: if response.status_code == 200: index = 0 @@ -177,34 +215,55 @@ class TTSProvider(TTSProviderBase): # 拼接当前块和上一块数据 chunk_total += chunk # 最后一个是静音,说明是一个完整的音频 - if len(chunk_total) % 2 == 0 and chunk_total[-2:] == b'\x00\x00': + if ( + len(chunk_total) % 2 == 0 + and chunk_total[-2:] == b"\x00\x00" + ): audio = self._get_audio_from_tts(chunk_total) audio_raw = audio_raw + audio.raw_data # 长度凑够2贞开始发送,60ms*2=120ms if len(audio_raw) >= 3840: opus_datas = self.wav_to_opus_data_audio_raw(audio_raw) if index == 0: - yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, - content=opus_datas, - tts_finish_text=text, sentence_type=SentenceType.SENTENCE_START) + yield TTSMessageDTO( + u_id=u_id, + msg_type=MsgType.TTS_TEXT_RESPONSE, + content=opus_datas, + tts_finish_text=text, + sentence_type=SentenceType.SENTENCE_START, + ) else: - yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, - content=opus_datas, - tts_finish_text=text, sentence_type=None) - audio_raw = b'' - chunk_total = b'' + yield TTSMessageDTO( + u_id=u_id, + msg_type=MsgType.TTS_TEXT_RESPONSE, + content=opus_datas, + tts_finish_text=text, + sentence_type=None, + ) + audio_raw = b"" + chunk_total = b"" if len(chunk_total) > 0: audio = self._get_audio_from_tts(chunk_total) audio_raw = audio_raw + audio.raw_data opus_datas = self.wav_to_opus_data_audio_raw(audio_raw) - yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, - tts_finish_text=text, sentence_type=SentenceType.SENTENCE_END) + yield TTSMessageDTO( + u_id=u_id, + msg_type=MsgType.TTS_TEXT_RESPONSE, + content=opus_datas, + tts_finish_text=text, + sentence_type=SentenceType.SENTENCE_END, + ) else: - yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=[], - tts_finish_text=text, sentence_type=SentenceType.SENTENCE_END) + yield TTSMessageDTO( + u_id=u_id, + msg_type=MsgType.TTS_TEXT_RESPONSE, + content=[], + tts_finish_text=text, + sentence_type=SentenceType.SENTENCE_END, + ) else: - print('请求失败:', response.status_code, response.text) + print("请求失败:", response.status_code, response.text) except Exception as e: logger.bind(tag=TAG).error("tts发生错误") traceback.print_exc() diff --git a/main/xiaozhi-server/core/providers/tts/gpt_sovits_v2.py b/main/xiaozhi-server/core/providers/tts/gpt_sovits_v2.py index 014c2649..5afad8b6 100644 --- a/main/xiaozhi-server/core/providers/tts/gpt_sovits_v2.py +++ b/main/xiaozhi-server/core/providers/tts/gpt_sovits_v2.py @@ -1,10 +1,8 @@ import os import uuid -import json -import base64 import requests from pydub import AudioSegment - +from core.utils.util import parse_string_to_list from config.logger import setup_logging from datetime import datetime from core.providers.tts.base import TTSProviderBase @@ -13,6 +11,7 @@ from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType TAG = __name__ logger = setup_logging() + class TTSProvider(TTSProviderBase): def __init__(self, config, delete_audio_file): super().__init__(config, delete_audio_file) @@ -21,23 +20,62 @@ 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 = config.get("top_k", 5) - self.top_p = config.get("top_p", 1) - self.temperature = 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 = config.get("batch_size", 1) - self.batch_threshold = config.get("batch_threshold", 0.75) - self.split_bucket = config.get("split_bucket", True) - self.return_fragment = config.get("return_fragment", False) - self.speed_factor = config.get("speed_factor", 1.0) - self.streaming_mode = config.get("streaming_mode", False) - self.seed = config.get("seed", -1) - self.parallel_infer = config.get("parallel_infer", True) - self.repetition_penalty = config.get("repetition_penalty", 1.35) - self.aux_ref_audio_paths = config.get("aux_ref_audio_paths", []) + + self.split_bucket = str(config.get("split_bucket", True)).lower() in ( + "true", + "1", + "yes", + ) + self.return_fragment = str(config.get("return_fragment", False)).lower() in ( + "true", + "1", + "yes", + ) + + self.streaming_mode = str(config.get("streaming_mode", False)).lower() in ( + "true", + "1", + "yes", + ) + + self.parallel_infer = str(config.get("parallel_infer", True)).lower() in ( + "true", + "1", + "yes", + ) + + self.aux_ref_audio_paths = parse_string_to_list( + config.get("aux_ref_audio_paths") + ) def generate_filename(self, extension=".wav"): - return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}") + return os.path.join( + self.output_file, + f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}", + ) async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False): tmp_file = self.generate_filename() @@ -60,7 +98,7 @@ class TTSProvider(TTSProviderBase): "streaming_mode": self.streaming_mode, "seed": self.seed, "parallel_infer": self.parallel_infer, - "repetition_penalty": self.repetition_penalty + "repetition_penalty": self.repetition_penalty, } resp = requests.post(self.url, json=request_json) @@ -68,13 +106,20 @@ class TTSProvider(TTSProviderBase): with open(tmp_file, "wb") as file: file.write(resp.content) else: - logger.bind(tag=TAG).error(f"GPT_SoVITS_V2 TTS请求失败: {resp.status_code} - {resp.text}") + logger.bind(tag=TAG).error( + f"GPT_SoVITS_V2 TTS请求失败: {resp.status_code} - {resp.text}" + ) # 使用 pydub 读取临时文件 audio = AudioSegment.from_file(tmp_file, format="wav") audio = audio.set_channels(1).set_frame_rate(16000) opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data) - yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text, - sentence_type=SentenceType.SENTENCE_START) + yield TTSMessageDTO( + u_id=u_id, + msg_type=MsgType.TTS_TEXT_RESPONSE, + content=opus_datas, + tts_finish_text=text, + sentence_type=SentenceType.SENTENCE_START, + ) # 用完后删除临时文件 try: os.remove(tmp_file) diff --git a/main/xiaozhi-server/core/providers/tts/gpt_sovits_v3.py b/main/xiaozhi-server/core/providers/tts/gpt_sovits_v3.py index 57b63ae1..2391bbb3 100644 --- a/main/xiaozhi-server/core/providers/tts/gpt_sovits_v3.py +++ b/main/xiaozhi-server/core/providers/tts/gpt_sovits_v3.py @@ -2,7 +2,7 @@ import os import uuid import requests from pydub import AudioSegment - +from core.utils.util import parse_string_to_list from config.logger import setup_logging from datetime import datetime from core.providers.tts.base import TTSProviderBase @@ -11,6 +11,7 @@ from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType TAG = __name__ logger = setup_logging() + class TTSProvider(TTSProviderBase): def __init__(self, config, delete_audio_file): super().__init__(config, delete_audio_file) @@ -19,18 +20,29 @@ 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 = config.get("top_k", 15) - self.top_p = config.get("top_p", 1.0) - self.temperature = config.get("temperature", 1.0) - self.cut_punc = config.get("cut_punc","") - self.speed = config.get("speed", 1.0) - self.inp_refs = config.get("inp_refs",[]) - self.sample_steps = config.get("sample_steps",32) - self.if_sr = config.get("if_sr",False) + # 处理空字符串的情况 + 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.inp_refs = parse_string_to_list(config.get("inp_refs")) + self.if_sr = str(config.get("if_sr", False)).lower() in ("true", "1", "yes") def generate_filename(self, extension=".wav"): - return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}") + return os.path.join( + self.output_file, + f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}", + ) async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False): tmp_file = self.generate_filename() @@ -55,13 +67,20 @@ class TTSProvider(TTSProviderBase): with open(tmp_file, "wb") as file: file.write(resp.content) else: - logger.bind(tag=TAG).error(f"GPT_SoVITS_V3 TTS请求失败: {resp.status_code} - {resp.text}") + logger.bind(tag=TAG).error( + f"GPT_SoVITS_V3 TTS请求失败: {resp.status_code} - {resp.text}" + ) # 使用 pydub 读取临时文件 audio = AudioSegment.from_file(tmp_file, format="wav") audio = audio.set_channels(1).set_frame_rate(16000) opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data) - yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text, - sentence_type=SentenceType.SENTENCE_START) + yield TTSMessageDTO( + u_id=u_id, + msg_type=MsgType.TTS_TEXT_RESPONSE, + content=opus_datas, + tts_finish_text=text, + sentence_type=SentenceType.SENTENCE_START, + ) # 用完后删除临时文件 try: os.remove(tmp_file) diff --git a/main/xiaozhi-server/core/providers/tts/minimax.py b/main/xiaozhi-server/core/providers/tts/minimax.py index 33f12075..06d87da2 100644 --- a/main/xiaozhi-server/core/providers/tts/minimax.py +++ b/main/xiaozhi-server/core/providers/tts/minimax.py @@ -2,10 +2,9 @@ import os import uuid import json import requests +from core.utils.util import parse_string_to_list from datetime import datetime - from pydub import AudioSegment - from core.providers.tts.base import TTSProviderBase from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType @@ -16,30 +15,35 @@ class TTSProvider(TTSProviderBase): self.group_id = config.get("group_id") self.api_key = config.get("api_key") self.model = config.get("model") - self.voice_id = config.get("voice_id") + if config.get("private_voice"): + self.voice_id = config.get("private_voice") + else: + self.voice_id = config.get("voice_id") default_voice_setting = { "voice_id": "female-shaonv", "speed": 1, "vol": 1, "pitch": 0, - "emotion": "happy" - } - default_pronunciation_dict = { - "tone": [ - "处理/(chu3)(li3)", "危险/dangerous" - ] + "emotion": "happy", } + default_pronunciation_dict = {"tone": ["处理/(chu3)(li3)", "危险/dangerous"]} defult_audio_setting = { "sample_rate": 32000, "bitrate": 128000, "format": "mp3", - "channel": 1 + "channel": 1, + } + self.voice_setting = { + **default_voice_setting, + **config.get("voice_setting", {}), + } + self.pronunciation_dict = { + **default_pronunciation_dict, + **config.get("pronunciation_dict", {}), } - self.voice_setting = {**default_voice_setting, **config.get("voice_setting", {})} - self.pronunciation_dict = {**default_pronunciation_dict, **config.get("pronunciation_dict", {})} self.audio_setting = {**defult_audio_setting, **config.get("audio_setting", {})} - self.timber_weights = config.get("timber_weights", []) + self.timber_weights = parse_string_to_list(config.get("timber_weights")) if self.voice_id: self.voice_setting["voice_id"] = self.voice_id @@ -48,11 +52,14 @@ class TTSProvider(TTSProviderBase): self.api_url = f"https://{self.host}/v1/t2a_v2?GroupId={self.group_id}" self.header = { "Content-Type": "application/json", - "Authorization": f"Bearer {self.api_key}" + "Authorization": f"Bearer {self.api_key}", } def generate_filename(self, extension=".mp3"): - return os.path.join(self.output_file, f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}") + return os.path.join( + self.output_file, + f"tts-{__name__}{datetime.now().date()}@{uuid.uuid4().hex}{extension}", + ) async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False): tmp_file = self.generate_filename() @@ -70,25 +77,34 @@ class TTSProvider(TTSProviderBase): request_json["voice_setting"]["voice_id"] = "" try: - resp = requests.post(self.api_url, json.dumps(request_json), headers=self.header) + resp = requests.post( + self.api_url, json.dumps(request_json), headers=self.header + ) # 检查返回请求数据的status_code是否为0 if resp.json()["base_resp"]["status_code"] == 0: - data = resp.json()['data']['audio'] + data = resp.json()["data"]["audio"] file_to_save = open(tmp_file, "wb") file_to_save.write(bytes.fromhex(data)) else: - raise Exception(f"{__name__} status_code: {resp.status_code} response: {resp.content}") + raise Exception( + f"{__name__} status_code: {resp.status_code} response: {resp.content}" + ) except Exception as e: raise Exception(f"{__name__} error: {e}") # 使用 pydub 读取临时文件 audio = AudioSegment.from_file(tmp_file, format="mp3") audio = audio.set_channels(1).set_frame_rate(16000) opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data) - yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text, - sentence_type=SentenceType.SENTENCE_START) + yield TTSMessageDTO( + u_id=u_id, + msg_type=MsgType.TTS_TEXT_RESPONSE, + content=opus_datas, + tts_finish_text=text, + sentence_type=SentenceType.SENTENCE_START, + ) # 用完后删除临时文件 try: os.remove(tmp_file) except FileNotFoundError: # 若文件不存在,忽略该异常 - pass \ No newline at end of file + pass diff --git a/main/xiaozhi-server/core/providers/tts/openai.py b/main/xiaozhi-server/core/providers/tts/openai.py index 6ea5ea5e..64203508 100644 --- a/main/xiaozhi-server/core/providers/tts/openai.py +++ b/main/xiaozhi-server/core/providers/tts/openai.py @@ -9,46 +9,64 @@ from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType from core.utils.util import check_model_key from core.providers.tts.base import TTSProviderBase + class TTSProvider(TTSProviderBase): def __init__(self, config, delete_audio_file): super().__init__(config, delete_audio_file) self.api_key = config.get("api_key") self.api_url = config.get("api_url", "https://api.openai.com/v1/audio/speech") self.model = config.get("model", "tts-1") - self.voice = config.get("voice", "alloy") + if config.get("private_voice"): + self.voice = config.get("private_voice") + else: + self.voice = config.get("voice", "alloy") self.response_format = "wav" - self.speed = 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) def generate_filename(self, extension=".wav"): - return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}") + return os.path.join( + self.output_file, + f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}", + ) async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False): tmp_file = self.generate_filename() headers = { "Authorization": f"Bearer {self.api_key}", - "Content-Type": "application/json" + "Content-Type": "application/json", } data = { "model": self.model, "input": text, "voice": self.voice, "response_format": "wav", - "speed": self.speed + "speed": self.speed, } response = requests.post(self.api_url, json=data, headers=headers) if response.status_code == 200: with open(tmp_file, "wb") as audio_file: audio_file.write(response.content) else: - raise Exception(f"OpenAI TTS请求失败: {response.status_code} - {response.text}") + raise Exception( + f"OpenAI TTS请求失败: {response.status_code} - {response.text}" + ) # 使用 pydub 读取临时文件 audio = AudioSegment.from_file(tmp_file, format="wav") audio = audio.set_channels(1).set_frame_rate(16000) opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data) - yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text, - sentence_type=SentenceType.SENTENCE_START) + yield TTSMessageDTO( + u_id=u_id, + msg_type=MsgType.TTS_TEXT_RESPONSE, + content=opus_datas, + tts_finish_text=text, + sentence_type=SentenceType.SENTENCE_START, + ) # 用完后删除临时文件 try: os.remove(tmp_file) diff --git a/main/xiaozhi-server/core/providers/tts/siliconflow.py b/main/xiaozhi-server/core/providers/tts/siliconflow.py index 6361c412..72f59e66 100644 --- a/main/xiaozhi-server/core/providers/tts/siliconflow.py +++ b/main/xiaozhi-server/core/providers/tts/siliconflow.py @@ -14,17 +14,23 @@ class TTSProvider(TTSProviderBase): super().__init__(config, delete_audio_file) self.model = config.get("model") self.access_token = config.get("access_token") - self.voice = config.get("voice") + if config.get("private_voice"): + self.voice = config.get("private_voice") + else: + self.voice = config.get("voice") self.response_format = config.get("response_format") self.sample_rate = config.get("sample_rate") - self.speed = config.get("speed") + self.speed = float(config.get("speed", 1.0)) self.gain = config.get("gain") self.host = "api.siliconflow.cn" self.api_url = f"https://{self.host}/v1/audio/speech" def generate_filename(self, extension=".wav"): - return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}") + return os.path.join( + self.output_file, + f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}", + ) async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False): tmp_file = self.generate_filename() @@ -36,9 +42,11 @@ class TTSProvider(TTSProviderBase): } headers = { "Authorization": f"Bearer {self.access_token}", - "Content-Type": "application/json" + "Content-Type": "application/json", } - response = requests.request("POST", self.api_url, json=request_json, headers=headers) + response = requests.request( + "POST", self.api_url, json=request_json, headers=headers + ) data = response.content file_to_save = open(tmp_file, "wb") file_to_save.write(data) @@ -46,8 +54,13 @@ class TTSProvider(TTSProviderBase): audio = AudioSegment.from_file(tmp_file, format="wav") audio = audio.set_channels(1).set_frame_rate(16000) opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data) - yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text, - sentence_type=SentenceType.SENTENCE_START) + yield TTSMessageDTO( + u_id=u_id, + msg_type=MsgType.TTS_TEXT_RESPONSE, + content=opus_datas, + tts_finish_text=text, + sentence_type=SentenceType.SENTENCE_START, + ) # 用完后删除临时文件 try: os.remove(tmp_file) diff --git a/main/xiaozhi-server/core/providers/tts/ttson.py b/main/xiaozhi-server/core/providers/tts/ttson.py index 9fd7d2a1..cc458574 100644 --- a/main/xiaozhi-server/core/providers/tts/ttson.py +++ b/main/xiaozhi-server/core/providers/tts/ttson.py @@ -14,49 +14,63 @@ from core.providers.tts.dto.dto import TTSMessageDTO, MsgType, SentenceType class TTSProvider(TTSProviderBase): def __init__(self, config, delete_audio_file): super().__init__(config, delete_audio_file) - self.url = config.get("url", "https://u95167-bd74-2aef8085.westx.seetacloud.com:8443/flashsummary/tts?token=") - self.voice_id = config.get("voice_id", 1695) + self.url = config.get( + "url", + "https://u95167-bd74-2aef8085.westx.seetacloud.com:8443/flashsummary/tts?token=", + ) + if config.get("private_voice"): + self.voice_id = int(config.get("private_voice")) + else: + self.voice_id = int(config.get("voice_id", 1695)) self.token = config.get("token") self.to_lang = config.get("to_lang") - self.volume_change_dB = config.get("volume_change_dB", 0) - self.speed_factor = config.get("speed_factor", 1) - self.stream = config.get("stream", False) + self.volume_change_dB = int(config.get("volume_change_dB", 0)) + self.speed_factor = int(config.get("speed_factor", 1)) + self.stream = str(config.get("stream", False)).lower() in ("true", "1", "yes") self.output_file = config.get("output_dir") - self.pitch_factor = config.get("pitch_factor", 0) + self.pitch_factor = int(config.get("pitch_factor", 0)) self.format = config.get("format", "mp3") - self.emotion = config.get("emotion", 1) - self.header = { - "Content-Type": "application/json" - } + self.emotion = int(config.get("emotion", 1)) + self.header = {"Content-Type": "application/json"} def generate_filename(self, extension=".mp3"): - return os.path.join(self.output_file, f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}") + return os.path.join( + self.output_file, + f"tts-{datetime.now().date()}@{uuid.uuid4().hex}{extension}", + ) async def text_to_speak(self, u_id, text, is_last_text=False, is_first_text=False): tmp_file = self.generate_filename() - url = f'{self.url}{self.token}' + url = f"{self.url}{self.token}" result = "firefly" - payload = json.dumps({ - "to_lang": self.to_lang, - "text": text, - "emotion": self.emotion, - "format": self.format, - "volume_change_dB": self.volume_change_dB, - "voice_id": self.voice_id, - "pitch_factor": self.pitch_factor, - "speed_factor": self.speed_factor, - "token": self.token - }) + payload = json.dumps( + { + "to_lang": self.to_lang, + "text": text, + "emotion": self.emotion, + "format": self.format, + "volume_change_dB": self.volume_change_dB, + "voice_id": self.voice_id, + "pitch_factor": self.pitch_factor, + "speed_factor": self.speed_factor, + "token": self.token, + } + ) resp = requests.request("POST", url, data=payload) if resp.status_code != 200: return resp_json = resp.json() try: - result = resp_json['url'] + ':' + str( - resp_json[ - 'port']) + '/flashsummary/retrieveFileData?stream=True&token=' + self.token + '&voice_audio_path=' + \ - resp_json['voice_path'] + result = ( + resp_json["url"] + + ":" + + str(resp_json["port"]) + + "/flashsummary/retrieveFileData?stream=True&token=" + + self.token + + "&voice_audio_path=" + + resp_json["voice_path"] + ) except Exception as e: print("error:", e) @@ -67,8 +81,13 @@ class TTSProvider(TTSProviderBase): audio = AudioSegment.from_file(tmp_file, format="mp3") audio = audio.set_channels(1).set_frame_rate(16000) opus_datas = self.wav_to_opus_data_audio_raw(audio.raw_data) - yield TTSMessageDTO(u_id=u_id, msg_type=MsgType.TTS_TEXT_RESPONSE, content=opus_datas, tts_finish_text=text, - sentence_type=SentenceType.SENTENCE_START) + yield TTSMessageDTO( + u_id=u_id, + msg_type=MsgType.TTS_TEXT_RESPONSE, + content=opus_datas, + tts_finish_text=text, + sentence_type=SentenceType.SENTENCE_START, + ) # 用完后删除临时文件 try: os.remove(tmp_file) @@ -78,4 +97,3 @@ class TTSProvider(TTSProviderBase): voice_path = resp_json.get("voice_path") des_path = tmp_file shutil.move(voice_path, des_path) -