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:
github-actions[bot]
2024-07-19 00:35:25 +09:00
committed by GitHub
parent e69ffacf83
commit 90bb175260
2 changed files with 29 additions and 12 deletions
+11 -4
View File
@@ -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(
+18 -8
View File
@@ -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