Add files via upload

This commit is contained in:
Stardust·减
2023-08-03 18:53:01 +08:00
committed by GitHub
parent 7775f24070
commit 8658f9aac5

View File

@@ -17,6 +17,8 @@ import utils
from models import SynthesizerTrn from models import SynthesizerTrn
from text.symbols import symbols from text.symbols import symbols
from text import text_to_sequence from text import text_to_sequence
from text import cleaned_text_to_sequence,_symbol_to_id, get_bert
from text.cleaner import clean_text
from scipy.io import wavfile from scipy.io import wavfile
# Get ffmpeg path # Get ffmpeg path
@@ -25,28 +27,60 @@ ffmpeg_path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "ffmpeg")
# Flask Init # Flask Init
app = Flask(__name__) app = Flask(__name__)
app.config['JSON_AS_ASCII'] = False app.config['JSON_AS_ASCII'] = False
# Text Preprocess def get_text(text, language_str, hps):
def get_text(text, hps): norm_text, phone, tone, word2ph = clean_text(text, language_str)
text_norm = text_to_sequence(text, hps.data.text_cleaners) print([f"{p}{t}" for p, t in zip(phone, tone)])
phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str)
if hps.data.add_blank: if hps.data.add_blank:
text_norm = commons.intersperse(text_norm, 0) phone = commons.intersperse(phone, 0)
text_norm = torch.LongTensor(text_norm) tone = commons.intersperse(tone, 0)
return text_norm language = commons.intersperse(language, 0)
for i in range(len(word2ph)):
word2ph[i] = word2ph[i] * 2
word2ph[0] += 1
bert = get_bert(norm_text, word2ph, language_str)
assert bert.shape[-1] == len(phone)
phone = torch.LongTensor(phone)
tone = torch.LongTensor(tone)
language = torch.LongTensor(language)
return bert, phone, tone, language
def infer(text, sdp_ratio, noise_scale, noise_scale_w,length_scale,sid):
bert, phones, tones, lang_ids = get_text(text,"ZH", hps,)
with torch.no_grad():
x_tst=phones.to(dev).unsqueeze(0)
tones=tones.to(dev).unsqueeze(0)
lang_ids=lang_ids.to(dev).unsqueeze(0)
bert = bert.to(dev).unsqueeze(0)
x_tst_lengths = torch.LongTensor([phones.size(0)]).to(dev)
speakers = torch.LongTensor([hps.data.spk2id[sid]]).to(dev)
audio = net_g.infer(x_tst, x_tst_lengths, speakers, tones, lang_ids,bert, sdp_ratio=sdp_ratio
, noise_scale=noise_scale, noise_scale_w=noise_scale_w, length_scale=length_scale)[0][0,0].data.cpu().float().numpy()
return audio
def replace_punctuation(text, i=2):
punctuation = ",。?!"
for char in punctuation:
text = text.replace(char, char * i)
return text
# Load Generator # Load Generator
hps_mt = utils.get_hparams_from_file("/GPUFS/sysu_hpcedu_123/vits/configs/genshin_xm37.json") hps = utils.get_hparams_from_file("./configs/config.json")
net_g_mt = SynthesizerTrn( dev='cuda'
net_g = SynthesizerTrn(
len(symbols), len(symbols),
hps_mt.data.filter_length // 2 + 1, hps.data.filter_length // 2 + 1,
hps_mt.train.segment_size // hps_mt.data.hop_length, hps.train.segment_size // hps.data.hop_length,
n_speakers=hps_mt.data.n_speakers, n_speakers=hps.data.n_speakers,
**hps_mt.model).cuda() **hps.model).to(dev)
_ = net_g_mt.eval() _ = net_g.eval()
_ = utils.load_checkpoint("/GPUFS/sysu_hpcedu_123/vits/logs/xm37/G_xm37_361200.pth", net_g_mt, None) _ = utils.load_checkpoint("logs/all_in_one/G_521000.pth", net_g, None)
npcList = ['', '', '派蒙', '纳西妲', '阿贝多', '温迪', '枫原万叶', '钟离', '荒泷一斗', '八重神子', '艾尔海森', '提纳里', '迪希雅', '卡维', '宵宫', '莱依拉', '赛诺', '诺艾尔', '托马', '凝光', '莫娜', '北斗', '神里绫华', '雷电将军', '芭芭拉', '鹿野院平藏', '五郎', '迪奥娜', '凯亚', '安柏', '班尼特', '', '柯莱', '夜兰', '妮露', '辛焱', '珐露珊', '', '香菱', '达达利亚', '砂糖', '早柚', '云堇', '刻晴', '丽莎', '迪卢克', '烟绯', '重云', '珊瑚宫心海', '胡桃', '可莉', '流浪者', '久岐忍', '神里绫人', '甘雨', '戴因斯雷布', '优菈', '菲谢尔', '行秋', '白术', '九条裟罗', '雷泽', '申鹤', '迪娜泽黛', '凯瑟琳', '多莉', '坎蒂丝', '萍姥姥', '罗莎莉亚', '留云借风真君', '绮良良', '瑶瑶', '七七', '奥兹', '米卡', '夏洛蒂', '埃洛伊', '博士', '女士', '大慈树王', '三月七', '娜塔莎', '希露瓦', '虎克', '克拉拉', '丹恒', '希儿', '布洛妮娅', '瓦尔特', '杰帕德', '佩拉', '姬子', '艾丝妲', '白露', '', '', '桑博', '伦纳德', '停云', '罗刹', '卡芙卡', '彦卿', '史瓦罗', '螺丝咕姆', '阿兰', '银狼', '素裳', '丹枢', '黑塔', '景元', '帕姆', '可可利亚', '半夏', '符玄', '公输师傅', '奥列格', '青雀', '大毫', '青镞', '费斯曼', '绿芙蓉', '镜流', '信使', '丽塔', '失落迷迭', '缭乱星棘', '伊甸', '伏特加女孩', '狂热蓝调', '莉莉娅', '萝莎莉娅', '八重樱', '八重霞', '卡莲', '第六夜想曲', '卡萝尔', '姬子', '极地战刃', '布洛妮娅', '次生银翼', '理之律者', '真理之律者', '迷城骇兔', '希儿', '魇夜星渊', '黑希儿', '帕朵菲莉丝', '天元骑英', '幽兰黛尔', '德丽莎', '月下初拥', '朔夜观星', '暮光骑士', '明日香', '李素裳', '格蕾修', '梅比乌斯', '渡鸦', '人之律者', '爱莉希雅', '爱衣', '天穹游侠', '琪亚娜', '空之律者', '终焉之律者', '薪炎之律者', '云墨丹心', '符华', '识之律者', '维尔薇', '始源之律者', '芽衣', '雷之律者', '苏莎娜', '阿波尼亚', '陆景和', '莫弈', '夏彦', '左然', '标贝']
@app.route("/",methods=['GET','POST']) @app.route("/",methods=['GET','POST'])
def main(): def main():
@@ -54,9 +88,10 @@ def main():
try: try:
speaker = request.args.get('speaker') speaker = request.args.get('speaker')
text = request.args.get('text').replace("/n","") text = request.args.get('text').replace("/n","")
sdp_ratio = float(request.args.get("sdp_ratio", 0.2))
noise = float(request.args.get("noise", 0.5)) noise = float(request.args.get("noise", 0.5))
noisew = float(request.args.get("noisew", 0.6)) noisew = float(request.args.get("noisew", 0.6))
length = float(request.args.get("length", 1.3)) length = float(request.args.get("length", 1.2))
if length >= 2: if length >= 2:
return "Too big length" return "Too big length"
if len(text) >=200: if len(text) >=200:
@@ -69,16 +104,11 @@ def main():
except: except:
return "Invalid Parameter" return "Invalid Parameter"
stn_tst_mt = get_text(text, hps_mt)
with torch.no_grad(): with torch.no_grad():
x_tst_mt = stn_tst_mt.cuda().unsqueeze(0) audio = infer(text, sdp_ratio=sdp_ratio, noise_scale=noise, noise_scale_w=noisew, length_scale=length, sid=speaker)
x_tst_mt_lengths = torch.LongTensor([stn_tst_mt.size(0)]).cuda()
sid_mt = torch.LongTensor([npcList.index(speaker)]).cuda()
audio_mt = net_g_mt.infer(x_tst_mt, x_tst_mt_lengths, sid=sid_mt, noise_scale=noise, noise_scale_w=noisew, length_scale=length)[0][0,0].data.cpu().float().numpy()
wav = BytesIO() wav = BytesIO()
wavfile.write(wav, hps_mt.data.sampling_rate, audio_mt) wavfile.write(wav, hps.data.sampling_rate, audio)
torch.cuda.empty_cache() torch.cuda.empty_cache()
if fmt == "mp3": if fmt == "mp3":
process = ( process = (
@@ -90,21 +120,3 @@ def main():
out, _ = process.communicate(input=wav.read()) out, _ = process.communicate(input=wav.read())
return Response(out, mimetype="audio/mpeg") return Response(out, mimetype="audio/mpeg")
return Response(wav.read(), mimetype="audio/wav") return Response(wav.read(), mimetype="audio/wav")
elif request.method == 'POST':
receive = request.get_data(as_text=True)
data = json.loads(receive)
speaker = data["speaker"]
text = data["text"].replace("/n","")
stn_tst_mt = get_text(text, hps_mt)
with torch.no_grad():
x_tst_mt = stn_tst_mt.cuda().unsqueeze(0)
x_tst_mt_lengths = torch.LongTensor([stn_tst_mt.size(0)]).cuda()
sid_mt = torch.LongTensor([npcList.index(speaker)]).cuda()
audio_mt = net_g_mt.infer(x_tst_mt, x_tst_mt_lengths, sid=sid_mt, noise_scale=0.667, noise_scale_w=0.8, length_scale=1.15)[0][0,0].data.cpu().float().numpy()
wav = BytesIO()
wavfile.write(wav, hps_mt.data.sampling_rate, audio_mt)
torch.cuda.empty_cache()
return Response(base64.b64encode(wav.read()))
#return Response(wav.read(), mimetype="audio/wav")