From a5fec6653db804200af0780accacdbe8bd6d3637 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=BA=90=E6=96=87=E9=9B=A8?= <41315874+fumiama@users.noreply.github.com> Date: Sat, 6 Jul 2024 00:15:04 +0900 Subject: [PATCH] fix(gpt): compile failed --- ChatTTS/core.py | 13 ++++--------- ChatTTS/model/gpt.py | 23 ++++++++++------------- examples/cmd/stream.py | 32 ++++---------------------------- tools/audio/__init__.py | 2 +- tools/audio/np.py | 19 +++++++++++++++++++ 5 files changed, 38 insertions(+), 51 deletions(-) diff --git a/ChatTTS/core.py b/ChatTTS/core.py index b1aad9d..5683344 100644 --- a/ChatTTS/core.py +++ b/ChatTTS/core.py @@ -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 diff --git a/ChatTTS/model/gpt.py b/ChatTTS/model/gpt.py index ca1aac1..ba6687e 100644 --- a/ChatTTS/model/gpt.py +++ b/ChatTTS/model/gpt.py @@ -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( diff --git a/examples/cmd/stream.py b/examples/cmd/stream.py index 2612173..d85d090 100644 --- a/examples/cmd/stream.py +++ b/examples/cmd/stream.py @@ -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: diff --git a/tools/audio/__init__.py b/tools/audio/__init__.py index f456fa7..6eeb3b3 100644 --- a/tools/audio/__init__.py +++ b/tools/audio/__init__.py @@ -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 diff --git a/tools/audio/np.py b/tools/audio/np.py index b40c7f7..32ba820 100644 --- a/tools/audio/np.py +++ b/tools/audio/np.py @@ -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