From 0000785a5320e965fba7446ca4d9ec24f863be7a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stardust=C2=B7=E5=87=8F?= Date: Tue, 5 Sep 2023 10:28:07 +0800 Subject: [PATCH] Update server.py --- server.py | 97 +++++++++++++++++++++++++++++++------------------------ 1 file changed, 55 insertions(+), 42 deletions(-) diff --git a/server.py b/server.py index c736ca4..caa6238 100644 --- a/server.py +++ b/server.py @@ -14,9 +14,9 @@ from scipy.io import wavfile # Flask Init app = Flask(__name__) app.config['JSON_AS_ASCII'] = False + def get_text(text, language_str, hps): norm_text, phone, tone, word2ph = clean_text(text, language_str) - 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: @@ -27,25 +27,36 @@ def get_text(text, language_str, hps): word2ph[i] = word2ph[i] * 2 word2ph[0] += 1 bert = get_bert(norm_text, word2ph, language_str) + del word2ph + assert bert.shape[-1] == len(phone), phone - assert bert.shape[-1] == len(phone) - + if language_str=='ZH': + bert = bert + ja_bert = torch.zeros(768, len(phone)) + elif language_str=="JA": + ja_bert = bert + bert = torch.zeros(1024, len(phone)) + else: + bert = torch.zeros(1024, len(phone)) + ja_bert = torch.zeros(768, len(phone)) + assert bert.shape[-1] == len(phone), ( + bert.shape, len(phone), sum(word2ph), p1, p2, t1, t2, pold, pold2, word2ph, text, w2pho) phone = torch.LongTensor(phone) tone = torch.LongTensor(tone) language = torch.LongTensor(language) + return bert, ja_bert, phone, tone, 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,) +def infer(text, sdp_ratio, noise_scale, noise_scale_w, length_scale, sid, language): + bert, ja_bert, phones, tones, lang_ids = get_text(text, language, 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) + ja_bert = ja_bert.to(device).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 + audio = net_g.infer(x_tst, x_tst_lengths, speakers, tones, lang_ids, bert, ja_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 @@ -84,40 +95,42 @@ _ = net_g.eval() _ = utils.load_checkpoint("logs/G_649000.pth", net_g, None,skip_optimizer=True) -@app.route("/",methods=['GET','POST']) +@app.route("/") def main(): - if request.method == 'GET': - try: - speaker = request.args.get('speaker') - 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)) - noisew = float(request.args.get("noisew", 0.6)) - length = float(request.args.get("length", 1.2)) - if length >= 2: - return "Too big length" - if len(text) >=200: - return "Too long text" - fmt = request.args.get("format", "wav") - if None in (speaker, text): - return "Missing Parameter" - if fmt not in ("mp3", "wav", "ogg"): - return "Invalid Format" - except: - return "Invalid Parameter" + try: + speaker = request.args.get('speaker') + 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)) + noisew = float(request.args.get("noisew", 0.6)) + length = float(request.args.get("length", 1.2)) + language = request.args.get('language') + if length >= 2: + return "Too big length" + if len(text) >=250: + return "Too long text" + fmt = request.args.get("format", "wav") + if None in (speaker, text): + return "Missing Parameter" + if fmt not in ("mp3", "wav", "ogg"): + return "Invalid Format" + if language not in ("JA", "ZH") + return "Invalid language" + except: + return "Invalid Parameter" - with torch.no_grad(): - audio = infer(text, sdp_ratio=sdp_ratio, noise_scale=noise, noise_scale_w=noisew, length_scale=length, sid=speaker) + with torch.no_grad(): + audio = infer(text, sdp_ratio=sdp_ratio, noise_scale=noise, noise_scale_w=noisew, length_scale=length, sid=speaker,language = language) - with BytesIO() as wav: - wavfile.write(wav, hps.data.sampling_rate, audio) - torch.cuda.empty_cache() - if fmt == "wav": - return Response(wav.getvalue(), mimetype="audio/wav") - wav.seek(0, 0) - with BytesIO() as ofp: - wav2(wav, ofp, fmt) - return Response( - ofp.getvalue(), - mimetype="audio/mpeg" if fmt == "mp3" else "audio/ogg" - ) + with BytesIO() as wav: + wavfile.write(wav, hps.data.sampling_rate, audio) + torch.cuda.empty_cache() + if fmt == "wav": + return Response(wav.getvalue(), mimetype="audio/wav") + wav.seek(0, 0) + with BytesIO() as ofp: + wav2(wav, ofp, fmt) + return Response( + ofp.getvalue(), + mimetype="audio/mpeg" if fmt == "mp3" else "audio/ogg" + )