File size: 2,727 Bytes
40147b8 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 | from __future__ import annotations
from typing import List, Sequence
from .configuration_moss_tts_nano import MossTTSNanoConfig
USER_ROLE_PREFIX = "user\n"
USER_TEMPLATE_REFERENCE_PREFIX = (
"<user_inst>\n"
"- Reference(s):\n"
)
USER_TEMPLATE_AFTER_REFERENCE = (
"\n- Instruction:\nNone\n"
"- Tokens:\nNone\n"
"- Quality:\nNone\n"
"- Sound Event:\nNone\n"
"- Ambient Sound:\nNone\n"
"- Language:\nNone\n"
"- Text:\n"
)
USER_TEMPLATE_PREFIX = USER_TEMPLATE_REFERENCE_PREFIX + "None" + USER_TEMPLATE_AFTER_REFERENCE
USER_TEMPLATE_SUFFIX = "\n</user_inst>"
ASSISTANT_TURN_PREFIX = "\n"
ASSISTANT_ROLE_PREFIX = "assistant\n"
def encode_text(tokenizer, text: str) -> List[int]:
try:
return list(tokenizer.encode(text, add_special_tokens=False))
except TypeError:
return list(tokenizer.encode(text))
def decode_text(tokenizer, token_ids: Sequence[int]) -> str:
try:
return str(
tokenizer.decode(
list(token_ids),
skip_special_tokens=False,
clean_up_tokenization_spaces=False,
)
)
except TypeError:
try:
return str(tokenizer.decode(list(token_ids), skip_special_tokens=False))
except TypeError:
return str(tokenizer.decode(list(token_ids)))
def build_user_prompt_prefix(tokenizer, config: MossTTSNanoConfig) -> List[int]:
return [config.im_start_token_id] + encode_text(tokenizer, USER_ROLE_PREFIX) + encode_text(
tokenizer,
USER_TEMPLATE_REFERENCE_PREFIX,
)
def build_user_prompt_after_reference(tokenizer) -> List[int]:
return encode_text(tokenizer, USER_TEMPLATE_AFTER_REFERENCE)
def build_assistant_prompt_prefix(tokenizer, config: MossTTSNanoConfig) -> List[int]:
return encode_text(tokenizer, USER_TEMPLATE_SUFFIX) + [config.im_end_token_id] + encode_text(
tokenizer,
ASSISTANT_TURN_PREFIX,
) + [config.im_start_token_id] + encode_text(
tokenizer,
ASSISTANT_ROLE_PREFIX,
)
def build_prompt_prefix(tokenizer, config: MossTTSNanoConfig) -> List[int]:
return (
build_user_prompt_prefix(tokenizer, config)
+ encode_text(tokenizer, "None")
+ build_user_prompt_after_reference(tokenizer)
)
def build_prompt_suffix(tokenizer, config: MossTTSNanoConfig) -> List[int]:
return build_assistant_prompt_prefix(tokenizer, config)
def build_prompt_token_ids(
tokenizer,
config: MossTTSNanoConfig,
text_token_ids: Sequence[int],
) -> List[int]:
return build_prompt_prefix(tokenizer, config) + [int(token_id) for token_id in text_token_ids] + build_prompt_suffix(
tokenizer,
config,
)
|