mirror of
https://github.com/2noise/ChatTTS.git
synced 2026-08-29 02:10:59 +08:00
chore(format): run black on dev (#589)
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
This commit is contained in:
committed by
GitHub
parent
e69ffacf83
commit
90bb175260
+11
-4
@@ -173,7 +173,8 @@ class Chat:
|
||||
shp = arr.shape
|
||||
assert len(shp) == 2, "prompt must be a 2D tensor"
|
||||
s = b14.encode_to_string(
|
||||
np.array(shp, dtype="<u2").tobytes() + lzma.compress(
|
||||
np.array(shp, dtype="<u2").tobytes()
|
||||
+ lzma.compress(
|
||||
arr.astype("<u2").tobytes(),
|
||||
format=lzma.FORMAT_RAW,
|
||||
filters=[{"id": lzma.FILTER_LZMA2, "preset": 9 | lzma.PRESET_EXTREME}],
|
||||
@@ -222,7 +223,7 @@ class Chat:
|
||||
class InferCodeParams(RefineTextParams):
|
||||
prompt: str = "[speed_5]"
|
||||
spk_emb: Optional[str] = None
|
||||
sample: Optional[str]=None
|
||||
sample: Optional[str] = None
|
||||
temperature: float = 0.3
|
||||
repetition_penalty: float = 1.05
|
||||
max_new_token: int = 2048
|
||||
@@ -580,7 +581,11 @@ class Chat:
|
||||
input_ids, attention_mask, text_mask = self.tokenizer.encode(
|
||||
text,
|
||||
self.gpt.num_vq,
|
||||
prompt=self._decode_prompt(params.sample) if params.sample is not None else None,
|
||||
prompt=(
|
||||
self._decode_prompt(params.sample)
|
||||
if params.sample is not None
|
||||
else None
|
||||
),
|
||||
device=gpt.device_gpt,
|
||||
)
|
||||
|
||||
@@ -641,7 +646,9 @@ class Chat:
|
||||
text = [f"[Sbreak]{i}[Pbreak]{params.prompt}" for i in text]
|
||||
|
||||
input_ids, attention_mask, text_mask = self.tokenizer.encode(
|
||||
text, self.gpt.num_vq, device=gpt.device_gpt,
|
||||
text,
|
||||
self.gpt.num_vq,
|
||||
device=gpt.device_gpt,
|
||||
)
|
||||
|
||||
logits_warpers, logits_processors = gen_logits(
|
||||
|
||||
@@ -31,7 +31,11 @@ class Tokenizer:
|
||||
|
||||
@torch.inference_mode()
|
||||
def encode(
|
||||
self, text: List[str], num_vq: int, prompt: Optional[torch.Tensor]=None, device="cpu",
|
||||
self,
|
||||
text: List[str],
|
||||
num_vq: int,
|
||||
prompt: Optional[torch.Tensor] = None,
|
||||
device="cpu",
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
|
||||
input_ids_lst = []
|
||||
@@ -72,9 +76,11 @@ class Tokenizer:
|
||||
for i in range(len(input_ids_lst)):
|
||||
input_ids.narrow(0, i, 1).narrow(
|
||||
1,
|
||||
max_input_ids_len-prompt_size-input_ids_lst[i].size(0),
|
||||
max_input_ids_len - prompt_size - input_ids_lst[i].size(0),
|
||||
input_ids_lst[i].size(0),
|
||||
).copy_(input_ids_lst[i]) # left padding
|
||||
).copy_(
|
||||
input_ids_lst[i]
|
||||
) # left padding
|
||||
del_all(input_ids_lst)
|
||||
|
||||
attention_mask = torch.zeros(
|
||||
@@ -87,12 +93,16 @@ class Tokenizer:
|
||||
attn = attention_mask.narrow(0, i, 1)
|
||||
attn.narrow(
|
||||
1,
|
||||
max_attention_mask_len-prompt_size-attention_mask_lst[i].size(0),
|
||||
max_attention_mask_len - prompt_size - attention_mask_lst[i].size(0),
|
||||
attention_mask_lst[i].size(0),
|
||||
).copy_(attention_mask_lst[i]) # left padding
|
||||
).copy_(
|
||||
attention_mask_lst[i]
|
||||
) # left padding
|
||||
if prompt_size > 0:
|
||||
attn.narrow(
|
||||
1, max_attention_mask_len-prompt_size, prompt_size,
|
||||
1,
|
||||
max_attention_mask_len - prompt_size,
|
||||
prompt_size,
|
||||
).fill_(1)
|
||||
del_all(attention_mask_lst)
|
||||
|
||||
@@ -100,11 +110,11 @@ class Tokenizer:
|
||||
new_input_ids = input_ids.unsqueeze_(-1).expand(-1, -1, num_vq).clone()
|
||||
|
||||
if prompt_size > 0:
|
||||
text_mask.narrow(1, max_input_ids_len-prompt_size, prompt_size).fill_(0)
|
||||
text_mask.narrow(1, max_input_ids_len - prompt_size, prompt_size).fill_(0)
|
||||
prompt_t = prompt.t().unsqueeze_(0).expand(input_ids.size(0), -1, -1)
|
||||
new_input_ids.narrow(
|
||||
1,
|
||||
max_input_ids_len-prompt_size,
|
||||
max_input_ids_len - prompt_size,
|
||||
prompt_size,
|
||||
).copy_(prompt_t)
|
||||
del prompt_t
|
||||
|
||||
Reference in New Issue
Block a user