176 lines
5.6 KiB
Python
176 lines
5.6 KiB
Python
# 导入所需的库
|
|
import os
|
|
import soundfile as sf
|
|
import redis
|
|
import hashlib
|
|
import json
|
|
from kafka import KafkaConsumer
|
|
from tools.i18n.i18n import I18nAuto
|
|
from GPT_SoVITS.inference_webui import change_gpt_weights, change_sovits_weights, get_tts_wav
|
|
from dotenv import load_dotenv
|
|
import torch
|
|
|
|
"""
|
|
整体设计说明:
|
|
这个脚本实现了一个文本到语音(TTS)的服务。它使用Kafka作为消息队列接收TTS任务,
|
|
使用Redis存储任务状态和结果,并利用GPT-SoVITS模型进行语音合成。
|
|
主要功能包括:
|
|
1. 初始化配置和模型
|
|
2. 提供语音合成功能
|
|
3. 监听Kafka消息并处理TTS任务
|
|
4. 将合成结果存储到Redis并更新任务状态
|
|
"""
|
|
|
|
# 加载环境变量
|
|
load_dotenv()
|
|
|
|
# 设置GPU设备(如果可用)
|
|
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
|
print(f"使用设备: {device}")
|
|
|
|
# 从环境变量中读取Redis配置
|
|
REDIS_HOST = os.getenv('REDIS_HOST')
|
|
REDIS_PORT = int(os.getenv('REDIS_PORT'))
|
|
REDIS_TTS_DB = int(os.getenv('REDIS_TTS_DB')) # DB 2用于存储TTS结果
|
|
REDIS_TASK_DB = int(os.getenv('REDIS_TASK_DB')) # DB 3用于存储任务状态
|
|
REDIS_PASSWORD = os.getenv('REDIS_PASSWORD')
|
|
|
|
# 从环境变量中读取Kafka配置
|
|
KAFKA_BROKER = os.getenv('KAFKA_BROKER')
|
|
KAFKA_TTS_TOPIC = os.getenv('KAFKA_TTS_TOPIC')
|
|
|
|
# 从环境变量中读取TTS相关配置
|
|
GPT_MODEL_PATH = os.getenv('GPT_MODEL_PATH')
|
|
SOVITS_MODEL_PATH = os.getenv('SOVITS_MODEL_PATH')
|
|
REF_AUDIO_PATH = os.getenv('REF_AUDIO_ZN_PATH')
|
|
REF_TEXT_PATH = os.getenv('REF_TEXT_ZN_PATH')
|
|
REF_LANGUAGE = os.getenv('REF_LANGUAGE')
|
|
TARGET_LANGUAGE = os.getenv('TARGET_LANGUAGE')
|
|
OUTPUT_PATH = os.getenv('OUTPUT_PATH')
|
|
|
|
# 初始化Redis客户端
|
|
redis_client = redis.Redis(
|
|
host=REDIS_HOST,
|
|
port=REDIS_PORT,
|
|
db=REDIS_TTS_DB,
|
|
password=REDIS_PASSWORD
|
|
)
|
|
|
|
redis_task_client = redis.Redis(
|
|
host=REDIS_HOST,
|
|
port=REDIS_PORT,
|
|
db=REDIS_TASK_DB,
|
|
password=REDIS_PASSWORD
|
|
)
|
|
|
|
# 初始化国际化工具
|
|
i18n = I18nAuto()
|
|
|
|
def get_audio_hash(text):
|
|
"""
|
|
生成文本的MD5哈希值,用作音频文件名的一部分
|
|
|
|
参数:
|
|
text (str): 需要生成哈希的文本
|
|
|
|
返回:
|
|
str: 文本的MD5哈希值
|
|
"""
|
|
return hashlib.md5(text.encode()).hexdigest()
|
|
|
|
# 初始化模型
|
|
print("正在初始化模型...")
|
|
change_gpt_weights(gpt_path=GPT_MODEL_PATH)
|
|
change_sovits_weights(sovits_path=SOVITS_MODEL_PATH)
|
|
|
|
# 读取参考文本
|
|
with open(REF_TEXT_PATH, 'r', encoding='utf-8') as file:
|
|
ref_text = file.read()
|
|
|
|
print("模型初始化成功。")
|
|
|
|
def synthesize(target_text, output_wav_path):
|
|
"""
|
|
使用GPT-SoVITS模型合成语音
|
|
|
|
参数:
|
|
target_text (str): 需要合成语音的目标文本
|
|
output_wav_path (str): 输出音频文件的路径
|
|
|
|
返回:
|
|
str: 如果成功,返回输出音频文件的路径;如果失败,返回None
|
|
"""
|
|
with torch.cuda.device(device):
|
|
synthesis_result = get_tts_wav(ref_wav_path=REF_AUDIO_PATH,
|
|
prompt_text=ref_text,
|
|
prompt_language=i18n(REF_LANGUAGE),
|
|
text=target_text,
|
|
text_language=i18n(TARGET_LANGUAGE), top_p=1, temperature=1)
|
|
|
|
result_list = list(synthesis_result)
|
|
|
|
if result_list:
|
|
last_sampling_rate, last_audio_data = result_list[-1]
|
|
sf.write(output_wav_path, last_audio_data, last_sampling_rate)
|
|
return output_wav_path
|
|
else:
|
|
return None
|
|
|
|
def kafka_consumer():
|
|
"""
|
|
Kafka消费者函数,用于接收和处理TTS任务
|
|
|
|
该函数会持续监听Kafka的TTS主题,接收任务并进行处理:
|
|
1. 接收任务信息
|
|
2. 更新任务状态
|
|
3. 调用synthesize函数合成语音
|
|
4. 将结果保存到Redis
|
|
5. 更新任务完成状态
|
|
"""
|
|
consumer = KafkaConsumer(
|
|
KAFKA_TTS_TOPIC,
|
|
bootstrap_servers=KAFKA_BROKER,
|
|
auto_offset_reset='latest',
|
|
value_deserializer=lambda m: json.loads(m.decode('utf-8'))
|
|
)
|
|
print(f"TTS消费者已启动")
|
|
for message in consumer:
|
|
try:
|
|
task_id = message.value['task_id']
|
|
target_text = message.value['text']
|
|
text_hash = message.value['text_hash']
|
|
|
|
# 更新任务状态为 "processing"
|
|
redis_task_client.set(f"task_status:tts:{task_id}", "processing")
|
|
|
|
output_wav_path = os.path.join(OUTPUT_PATH, f"{text_hash}.wav")
|
|
|
|
# 再次检查文件是否存在(以防在此期间被其他进程创建)
|
|
if not os.path.exists(output_wav_path):
|
|
output_path = synthesize(target_text, output_wav_path)
|
|
else:
|
|
output_path = output_wav_path
|
|
|
|
if output_path:
|
|
# 将结果保存在 DB 2
|
|
redis_client.set(f"tts:{task_id}", json.dumps({"path": output_path}))
|
|
print(f"音频合成成功: {output_path}")
|
|
|
|
# 更新任务状态为 "completed"
|
|
redis_task_client.set(f"task_status:tts:{task_id}", "completed")
|
|
else:
|
|
print("音频合成失败")
|
|
|
|
# 更新任务状态为 "failed"
|
|
redis_task_client.set(f"task_status:tts:{task_id}", "failed")
|
|
except Exception as e:
|
|
print(f"处理消息时出错: {str(e)}")
|
|
|
|
# 更新任务状态为 "failed"
|
|
redis_task_client.set(f"task_status:tts:{task_id}", "failed")
|
|
|
|
if __name__ == "__main__":
|
|
# 设置CUDA设备
|
|
torch.cuda.set_device(device)
|
|
# 启动Kafka消费者
|
|
kafka_consumer() |