update:格式化代码

This commit is contained in:
hrz
2025-02-14 00:54:59 +08:00
parent eae590eb26
commit c5180df9ac
11 changed files with 34 additions and 34 deletions
+2 -1
View File
@@ -7,4 +7,5 @@ docs/
tmp/
LICENSE
README.md
README_en.md
README_en.md
.config.yaml
+1
View File
@@ -140,3 +140,4 @@ cython_debug/
model.pt
tmp
.DS_Store
.config.yaml
+7 -2
View File
@@ -1,10 +1,15 @@
import os
import argparse
from core.utils.util import read_config
from core.utils.util import read_config, get_project_dir
def load_config():
"""加载配置文件"""
parser = argparse.ArgumentParser(description="Server configuration")
parser.add_argument("--config_path", type=str, default="config.yaml")
default_config_file = "config.yaml"
# 判断是否存在私有的配置文件
if os.path.exists(get_project_dir() + "." + default_config_file):
default_config_file = "." + default_config_file
parser.add_argument("--config_path", type=str, default=default_config_file)
args = parser.parse_args()
return read_config(args.config_path)
+2 -1
View File
@@ -1,7 +1,8 @@
from abc import ABC, abstractmethod
class LLMProviderBase(ABC):
@abstractmethod
def response(self, session_id, dialogue):
"""LLM response generator"""
pass
pass
+2 -1
View File
@@ -5,6 +5,7 @@ from core.providers.llm.base import LLMProviderBase
logger = logging.getLogger(__name__)
class LLMProvider(LLMProviderBase):
def __init__(self, config):
self.api_key = config["api_key"]
@@ -35,4 +36,4 @@ class LLMProvider(LLMProviderBase):
except Exception as e:
logger.error(f"Error in response generation: {e}")
yield "【服务响应异常】"
yield "【服务响应异常】"
+4 -1
View File
@@ -4,6 +4,7 @@ from core.providers.llm.base import LLMProviderBase
logger = logging.getLogger(__name__)
class LLMProvider(LLMProviderBase):
def __init__(self, config):
self.model_name = config.get("model_name")
@@ -12,6 +13,8 @@ class LLMProvider(LLMProviderBase):
self.base_url = config.get("base_url")
else:
self.base_url = config.get("url")
if "" in self.api_key:
logger.error("你还没配置LLM的密钥,请在配置文件中配置密钥,否则无法正常工作")
self.client = openai.OpenAI(api_key=self.api_key, base_url=self.base_url)
def response(self, session_id, dialogue):
@@ -29,4 +32,4 @@ class LLMProvider(LLMProviderBase):
if content: # 仅在content非空时生成
yield content
except Exception as e:
logger.error(f"Error in response generation: {e}")
logger.error(f"Error in response generation: {e}")
+4 -10
View File
@@ -1,20 +1,14 @@
import asyncio
import logging
import os
import json
import uuid
import base64
# from datetime import datetime
# import edge_tts
import numpy as np
import opuslib
# import requests
# from core.utils.util import read_config, get_project_dir
from pydub import AudioSegment
from abc import ABC, abstractmethod
logger = logging.getLogger(__name__)
class TTSProviderBase(ABC):
def __init__(self, config, delete_audio_file):
self.delete_audio_file = delete_audio_file
@@ -34,8 +28,8 @@ class TTSProviderBase(ABC):
max_repeat_time = max_repeat_time - 1
logger.error(f"语音生成失败: {text}:{tmp_file},再试{max_repeat_time}")
if max_repeat_time>0:
logger.info(f"语音生成成功: {text}:{tmp_file},重试{5-max_repeat_time}")
if max_repeat_time > 0:
logger.info(f"语音生成成功: {text}:{tmp_file},重试{5 - max_repeat_time}")
return tmp_file
except Exception as e:
@@ -86,4 +80,4 @@ class TTSProviderBase(ABC):
opus_data = encoder.encode(np_frame.tobytes(), frame_size)
opus_datas.append(opus_data)
return opus_datas, duration
return opus_datas, duration
+2 -1
View File
@@ -6,6 +6,7 @@ import requests
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
@@ -52,4 +53,4 @@ class TTSProvider(TTSProviderBase):
if "data" in resp.json():
data = resp.json()["data"]
file_to_save = open(output_file, "wb")
file_to_save.write(base64.b64decode(data))
file_to_save.write(base64.b64decode(data))
+2 -1
View File
@@ -4,6 +4,7 @@ import edge_tts
from datetime import datetime
from core.providers.tts.base import TTSProviderBase
class TTSProvider(TTSProviderBase):
def __init__(self, config, delete_audio_file):
super().__init__(config, delete_audio_file)
@@ -14,4 +15,4 @@ class TTSProvider(TTSProviderBase):
async def text_to_speak(self, text, output_file):
communicate = edge_tts.Communicate(text, voice=self.voice) # Use your preferred voice
await communicate.save(output_file)
await communicate.save(output_file)
+1 -11
View File
@@ -1,19 +1,9 @@
import asyncio
import logging
import os
import sys
import json
import uuid
import base64
import logging
import importlib
from datetime import datetime
import edge_tts
import numpy as np
import opuslib
import requests
from core.utils.util import read_config, get_project_dir
from pydub import AudioSegment
from abc import ABC, abstractmethod
logger = logging.getLogger(__name__)
+7 -5
View File
@@ -1,8 +1,7 @@
import yaml
import unicodedata
import socket
import os
import json
import yaml
import socket
def get_project_dir():
@@ -41,6 +40,7 @@ def is_segment(tokens):
else:
return False
def is_punctuation_or_emoji(char):
"""检查字符是否为空格、指定标点或表情符号"""
# 定义需要去除的中英文标点(包括全角/半角)
@@ -49,7 +49,7 @@ def is_punctuation_or_emoji(char):
'', '.', # 中文句号 + 英文句号
'', '!', # 中文感叹号 + 英文感叹号
'-', '', # 英文连字符 + 中文全角横线
'' # 中文顿号
'' # 中文顿号
}
if char.isspace() or char in punctuation_set:
return True
@@ -63,6 +63,7 @@ def is_punctuation_or_emoji(char):
]
return any(start <= code_point <= end for start, end in emoji_ranges)
def get_string_no_punctuation_or_emoji(s):
"""去除字符串首尾的空格、标点符号和表情符号"""
chars = list(s)
@@ -74,7 +75,8 @@ def get_string_no_punctuation_or_emoji(s):
end = len(chars) - 1
while end >= start and is_punctuation_or_emoji(chars[end]):
end -= 1
return ''.join(chars[start:end+1])
return ''.join(chars[start:end + 1])
def remove_punctuation_and_length(text):
# 全角符号和半角符号的Unicode范围