fix(gpt): compile failed

This commit is contained in:
源文雨
2024-07-06 00:15:04 +09:00
parent 6e80df2bcc
commit a5fec6653d
5 changed files with 38 additions and 51 deletions
+4 -9
View File
@@ -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
View File
@@ -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
View File
@@ -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 -1
View File
@@ -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
+19
View File
@@ -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