Update server.py
This commit is contained in:
97
server.py
97
server.py
@@ -14,9 +14,9 @@ from scipy.io import wavfile
|
|||||||
# Flask Init
|
# Flask Init
|
||||||
app = Flask(__name__)
|
app = Flask(__name__)
|
||||||
app.config['JSON_AS_ASCII'] = False
|
app.config['JSON_AS_ASCII'] = False
|
||||||
|
|
||||||
def get_text(text, language_str, hps):
|
def get_text(text, language_str, hps):
|
||||||
norm_text, phone, tone, word2ph = clean_text(text, language_str)
|
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)
|
phone, tone, language = cleaned_text_to_sequence(phone, tone, language_str)
|
||||||
|
|
||||||
if hps.data.add_blank:
|
if hps.data.add_blank:
|
||||||
@@ -27,25 +27,36 @@ def get_text(text, language_str, hps):
|
|||||||
word2ph[i] = word2ph[i] * 2
|
word2ph[i] = word2ph[i] * 2
|
||||||
word2ph[0] += 1
|
word2ph[0] += 1
|
||||||
bert = get_bert(norm_text, word2ph, language_str)
|
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)
|
phone = torch.LongTensor(phone)
|
||||||
tone = torch.LongTensor(tone)
|
tone = torch.LongTensor(tone)
|
||||||
language = torch.LongTensor(language)
|
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, language):
|
||||||
|
bert, ja_bert, phones, tones, lang_ids = get_text(text, language, hps)
|
||||||
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():
|
with torch.no_grad():
|
||||||
x_tst=phones.to(dev).unsqueeze(0)
|
x_tst=phones.to(dev).unsqueeze(0)
|
||||||
tones=tones.to(dev).unsqueeze(0)
|
tones=tones.to(dev).unsqueeze(0)
|
||||||
lang_ids=lang_ids.to(dev).unsqueeze(0)
|
lang_ids=lang_ids.to(dev).unsqueeze(0)
|
||||||
bert = bert.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)
|
x_tst_lengths = torch.LongTensor([phones.size(0)]).to(dev)
|
||||||
speakers = torch.LongTensor([hps.data.spk2id[sid]]).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()
|
, noise_scale=noise_scale, noise_scale_w=noise_scale_w, length_scale=length_scale)[0][0,0].data.cpu().float().numpy()
|
||||||
return audio
|
return audio
|
||||||
|
|
||||||
@@ -84,40 +95,42 @@ _ = net_g.eval()
|
|||||||
|
|
||||||
_ = utils.load_checkpoint("logs/G_649000.pth", net_g, None,skip_optimizer=True)
|
_ = utils.load_checkpoint("logs/G_649000.pth", net_g, None,skip_optimizer=True)
|
||||||
|
|
||||||
@app.route("/",methods=['GET','POST'])
|
@app.route("/")
|
||||||
def main():
|
def main():
|
||||||
if request.method == 'GET':
|
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))
|
||||||
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.2))
|
||||||
length = float(request.args.get("length", 1.2))
|
language = request.args.get('language')
|
||||||
if length >= 2:
|
if length >= 2:
|
||||||
return "Too big length"
|
return "Too big length"
|
||||||
if len(text) >=200:
|
if len(text) >=250:
|
||||||
return "Too long text"
|
return "Too long text"
|
||||||
fmt = request.args.get("format", "wav")
|
fmt = request.args.get("format", "wav")
|
||||||
if None in (speaker, text):
|
if None in (speaker, text):
|
||||||
return "Missing Parameter"
|
return "Missing Parameter"
|
||||||
if fmt not in ("mp3", "wav", "ogg"):
|
if fmt not in ("mp3", "wav", "ogg"):
|
||||||
return "Invalid Format"
|
return "Invalid Format"
|
||||||
except:
|
if language not in ("JA", "ZH")
|
||||||
return "Invalid Parameter"
|
return "Invalid language"
|
||||||
|
except:
|
||||||
|
return "Invalid Parameter"
|
||||||
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
audio = infer(text, sdp_ratio=sdp_ratio, noise_scale=noise, noise_scale_w=noisew, length_scale=length, sid=speaker)
|
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:
|
with BytesIO() as wav:
|
||||||
wavfile.write(wav, hps.data.sampling_rate, audio)
|
wavfile.write(wav, hps.data.sampling_rate, audio)
|
||||||
torch.cuda.empty_cache()
|
torch.cuda.empty_cache()
|
||||||
if fmt == "wav":
|
if fmt == "wav":
|
||||||
return Response(wav.getvalue(), mimetype="audio/wav")
|
return Response(wav.getvalue(), mimetype="audio/wav")
|
||||||
wav.seek(0, 0)
|
wav.seek(0, 0)
|
||||||
with BytesIO() as ofp:
|
with BytesIO() as ofp:
|
||||||
wav2(wav, ofp, fmt)
|
wav2(wav, ofp, fmt)
|
||||||
return Response(
|
return Response(
|
||||||
ofp.getvalue(),
|
ofp.getvalue(),
|
||||||
mimetype="audio/mpeg" if fmt == "mp3" else "audio/ogg"
|
mimetype="audio/mpeg" if fmt == "mp3" else "audio/ogg"
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user