mirror of
https://github.com/2noise/ChatTTS.git
synced 2026-08-29 02:10:59 +08:00
fix(gpt): compile failed
This commit is contained in:
+4
-9
@@ -331,7 +331,6 @@ class Chat:
|
||||
|
||||
return self.has_loaded()
|
||||
|
||||
@torch.inference_mode()
|
||||
def _infer(
|
||||
self,
|
||||
text,
|
||||
@@ -506,7 +505,7 @@ class Chat:
|
||||
torch.where(cond, n, emb, out=emb)
|
||||
del cond, n
|
||||
|
||||
@torch.inference_mode()
|
||||
@torch.no_grad()
|
||||
def _infer_code(
|
||||
self,
|
||||
text: Tuple[List[str], str],
|
||||
@@ -549,9 +548,7 @@ class Chat:
|
||||
|
||||
input_ids, attention_mask, text_mask = self._text_to_token(text, gpt.device_gpt)
|
||||
|
||||
with torch.inference_mode(not self.compile):
|
||||
with torch.no_grad():
|
||||
emb = gpt(input_ids, text_mask)
|
||||
emb = gpt(input_ids, text_mask)
|
||||
|
||||
del text_mask
|
||||
|
||||
@@ -592,7 +589,7 @@ class Chat:
|
||||
|
||||
return result
|
||||
|
||||
@torch.inference_mode()
|
||||
@torch.no_grad()
|
||||
def _refine_text(
|
||||
self,
|
||||
text: str,
|
||||
@@ -616,9 +613,7 @@ class Chat:
|
||||
repetition_penalty=params.repetition_penalty,
|
||||
)
|
||||
|
||||
with torch.inference_mode(not self.compile):
|
||||
with torch.no_grad():
|
||||
emb = gpt(input_ids, text_mask)
|
||||
emb = gpt(input_ids, text_mask)
|
||||
|
||||
del text_mask
|
||||
|
||||
|
||||
+10
-13
@@ -347,7 +347,7 @@ class GPT(nn.Module):
|
||||
hiddens=hiddens,
|
||||
)
|
||||
|
||||
@torch.inference_mode()
|
||||
@torch.no_grad()
|
||||
def generate(
|
||||
self,
|
||||
emb: torch.Tensor,
|
||||
@@ -366,7 +366,6 @@ class GPT(nn.Module):
|
||||
show_tqdm=True,
|
||||
ensure_non_empty=True,
|
||||
stream_batch=24,
|
||||
compile=False,
|
||||
context=Context(),
|
||||
):
|
||||
|
||||
@@ -436,17 +435,15 @@ class GPT(nn.Module):
|
||||
|
||||
model_input.to(self.device_gpt, self.gpt.dtype)
|
||||
|
||||
with torch.inference_mode(not compile):
|
||||
with torch.no_grad():
|
||||
outputs: BaseModelOutputWithPast = self.gpt(
|
||||
attention_mask=model_input.attention_mask,
|
||||
position_ids=model_input.position_ids,
|
||||
past_key_values=model_input.past_key_values,
|
||||
inputs_embeds=model_input.inputs_embeds,
|
||||
use_cache=model_input.use_cache,
|
||||
output_attentions=return_attn,
|
||||
cache_position=model_input.cache_position,
|
||||
)
|
||||
outputs: BaseModelOutputWithPast = self.gpt(
|
||||
attention_mask=model_input.attention_mask,
|
||||
position_ids=model_input.position_ids,
|
||||
past_key_values=model_input.past_key_values,
|
||||
inputs_embeds=model_input.inputs_embeds,
|
||||
use_cache=model_input.use_cache,
|
||||
output_attentions=return_attn,
|
||||
cache_position=model_input.cache_position,
|
||||
)
|
||||
del_all(model_input)
|
||||
attentions.append(outputs.attentions)
|
||||
hidden_states = outputs.last_hidden_state.to(
|
||||
|
||||
+4
-28
@@ -1,37 +1,13 @@
|
||||
import sys
|
||||
import torch
|
||||
import numpy as np
|
||||
import ChatTTS
|
||||
from IPython.display import Audio
|
||||
|
||||
|
||||
import io
|
||||
import threading
|
||||
import time
|
||||
import pyaudio
|
||||
|
||||
import random
|
||||
|
||||
import pyaudio
|
||||
import numpy as np
|
||||
import ChatTTS
|
||||
|
||||
# 如果不对batch进行统一量纲,同batch的多段话之间量纲会有较大差异,以致于无法在后处理中获取正常的历史语音结果?
|
||||
def batch_unsafe_float_to_int16(audios: list[np.ndarray], am=None) -> list[np.ndarray]:
|
||||
"""
|
||||
This function will destroy audio, use only once.
|
||||
"""
|
||||
|
||||
valid_audios = [i for i in audios if i is not None]
|
||||
if len(valid_audios) > 1:
|
||||
am = np.abs(np.concatenate(valid_audios, axis=1)).max() * 32768
|
||||
else:
|
||||
am = np.abs(valid_audios[0]).max() * 32768
|
||||
am = 32767 * 32768 / am
|
||||
|
||||
for i in range(len(audios)):
|
||||
if audios[i] is not None:
|
||||
np.multiply(audios[i], am, audios[i])
|
||||
audios[i] = audios[i].astype(np.int16)
|
||||
return audios
|
||||
|
||||
from tools.audio import batch_unsafe_float_to_int16
|
||||
|
||||
# 流式声音处理器
|
||||
class AudioStreamer:
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
from .mp3 import wav_arr_to_mp3_view
|
||||
from .ffmpeg import has_ffmpeg_installed
|
||||
from .np import unsafe_float_to_int16
|
||||
from .np import unsafe_float_to_int16, batch_unsafe_float_to_int16
|
||||
|
||||
@@ -12,3 +12,22 @@ def unsafe_float_to_int16(audio: np.ndarray) -> np.ndarray:
|
||||
np.multiply(audio, am, audio)
|
||||
audio16 = audio.astype(np.int16)
|
||||
return audio16
|
||||
|
||||
@jit
|
||||
def batch_unsafe_float_to_int16(audios: list[np.ndarray]) -> list[np.ndarray]:
|
||||
"""
|
||||
This function will destroy audio, use only once.
|
||||
"""
|
||||
|
||||
valid_audios = [i for i in audios if i is not None]
|
||||
if len(valid_audios) > 1:
|
||||
am = np.abs(np.concatenate(valid_audios, axis=1)).max() * 32768
|
||||
else:
|
||||
am = np.abs(valid_audios[0]).max() * 32768
|
||||
am = 32767 * 32768 / am
|
||||
|
||||
for i in range(len(audios)):
|
||||
if audios[i] is not None:
|
||||
np.multiply(audios[i], am, audios[i])
|
||||
audios[i] = audios[i].astype(np.int16)
|
||||
return audios
|
||||
|
||||
Reference in New Issue
Block a user