From 3fbf2aadf5dd9d842e16863c304688afd21d3f11 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: Wed, 3 Jul 2024 16:06:45 +0900 Subject: [PATCH] fix: add param `ensure_non_empty` (fix #511) --- .github/workflows/unitest.yml | 3 +- ChatTTS/core.py | 1 + ChatTTS/model/gpt.py | 44 +++++++++++++++++++++++--- tests/#511.py | 58 +++++++++++++++++++++++++++++++++++ tests/testall.sh | 15 +++++++++ 5 files changed, 114 insertions(+), 7 deletions(-) create mode 100644 tests/#511.py create mode 100755 tests/testall.sh diff --git a/.github/workflows/unitest.yml b/.github/workflows/unitest.yml index 89565c9..4903a1d 100644 --- a/.github/workflows/unitest.yml +++ b/.github/workflows/unitest.yml @@ -25,5 +25,4 @@ jobs: run: pip install -r requirements.txt - name: Run Test - run: | - echo "TODO" + run: tests/testall.sh diff --git a/ChatTTS/core.py b/ChatTTS/core.py index 1d18497..95ff538 100644 --- a/ChatTTS/core.py +++ b/ChatTTS/core.py @@ -200,6 +200,7 @@ class Chat: max_new_token: int = 384 min_new_token: int = 0 show_tqdm: bool = True + ensure_non_empty: bool = True @dataclass(repr=False, eq=False) class InferCodeParams(RefineTextParams): diff --git a/ChatTTS/model/gpt.py b/ChatTTS/model/gpt.py index cf82d28..d661062 100644 --- a/ChatTTS/model/gpt.py +++ b/ChatTTS/model/gpt.py @@ -363,6 +363,7 @@ class GPT(nn.Module): return_hidden=False, stream=False, show_tqdm=True, + ensure_non_empty=True, context=Context(), ): @@ -376,13 +377,14 @@ class GPT(nn.Module): ) finish = torch.zeros(inputs_ids.shape[0], device=inputs_ids.device).bool() + old_temperature = temperature + temperature = ( - temperature.unsqueeze_(0) + temperature.unsqueeze(0) .expand(inputs_ids.shape[0], -1) .contiguous() .view(-1, 1) ) - # temperature = rearrange(temperature, "b n -> (b n) 1") attention_mask_cache = torch.ones( ( @@ -464,9 +466,9 @@ class GPT(nn.Module): dtype=torch.float, device=self.device, ) - for i in range(self.num_vq): - x: torch.Tensor = self.head_code[i](hidden_states) - logits[..., i] = x + for num_vq_iter in range(self.num_vq): + x: torch.Tensor = self.head_code[num_vq_iter](hidden_states) + logits[..., num_vq_iter ] = x del x # logits = logits[:, -1].float() @@ -522,6 +524,37 @@ class GPT(nn.Module): ], 1, ) + + if i == 0 and finish.any(): + self.logger.warn( + "unexpected end at index %s", + str([unexpected_idx.item() for unexpected_idx in finish.nonzero()]), + ) + if ensure_non_empty: + if show_tqdm: + pbar.close() + self.logger.warn("regenerate in order to ensure non-empty") + new_gen = self.generate( + emb, + inputs_ids, + old_temperature, + eos_token, + attention_mask, + max_new_token, + min_new_token, + logits_warpers, + logits_processors, + infer_text, + return_attn, + return_hidden, + stream, + show_tqdm, + ensure_non_empty, + context, + ) + for result in new_gen: + yield result + return del inputs_ids inputs_ids = inputs_ids_tmp @@ -529,6 +562,7 @@ class GPT(nn.Module): if stream: minus_prev_end_index = end_idx.neg() + end_idx.add_((finish.logical_not().to(end_idx.device)).int()) if stream: if ( diff --git a/tests/#511.py b/tests/#511.py new file mode 100644 index 0000000..99de544 --- /dev/null +++ b/tests/#511.py @@ -0,0 +1,58 @@ +import os, sys + +if sys.platform == "darwin": + os.environ["PYTORCH_ENABLE_MPS_FALLBACK"] = "1" + +now_dir = os.getcwd() +sys.path.append(now_dir) + +import ChatTTS + +from tools.logger import get_logger + +logger = get_logger("Test #511") + +chat = ChatTTS.Chat(logger) +chat.load(compile=False) # Set to True for better performance + +texts = ["语音太短了会造成生成音频错误, 这是占位占位, 老大爷觉得车夫的想法很有道理", + "评分只是衡量音色的稳定性,不代表音色的好坏, 可以根据自己的需求选择合适的音色", + "举个简单的例子,如果一个沙哑且结巴的音色一直很稳定,那么它的评分就会很高。", + "语音太短了会造成生成音频错误, 这是占位占位。我使用 seed id 去生成音频, 但是生成的音频不稳定", + "seed id只是一个参考ID 不同的环境下音色不一定一致。还是推荐使用 .pt 文件载入音色", + "语音太短了会造成生成音频错误, 这是占位占位。音色标的男女准确吗", + "当前第一批测试的音色有两千条, 根据声纹相似性简单打标, 准确度不高, 特别是特征一项", + "语音太短了会造成生成音频错误, 这是占位占位。仅供参考。如果大家有更好的标注方法,欢迎 PR。", + ] + +rand_spk = chat.sample_random_speaker() + +params_infer_code = ChatTTS.Chat.InferCodeParams( + spk_emb = rand_spk, # add sampled speaker + temperature = .3, # using custom temperature + top_P = 0.005, # top P decode + top_K = 1, # top K decode +) + +params_refine_text = ChatTTS.Chat.RefineTextParams( + prompt='[oral_0][laugh_0][break_4]', +) + +fail = False + +for i in range(4): + + wavs = chat.infer( + texts, + params_refine_text=params_refine_text, + params_infer_code=params_infer_code, + ) + + for k, wav in enumerate(wavs): + if wav is None: + logger.warn("iter", i, "index", k, "is None") + fail = True + +if fail: + import sys + sys.exit(1) diff --git a/tests/testall.sh b/tests/testall.sh new file mode 100755 index 0000000..2f7eec6 --- /dev/null +++ b/tests/testall.sh @@ -0,0 +1,15 @@ +#!/bin/sh + +exitcode=0 + +for file in tests/*.py +do + python "$file" + if [ $? -ne 0 ] + then + echo "Error: $file exited with a non-zero status." + exitcode=1 + fi +done + +exit $exitcode