mirror of
https://github.com/2noise/ChatTTS.git
synced 2026-08-30 17:05:27 +08:00
+114
-97
@@ -3,8 +3,9 @@ import json
|
||||
import logging
|
||||
import tempfile
|
||||
from functools import partial
|
||||
from typing import Literal, Optional
|
||||
from typing import Literal, Optional, List, Callable
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from omegaconf import OmegaConf
|
||||
from vocos import Vocos
|
||||
@@ -16,8 +17,8 @@ from .utils.gpu import select_device
|
||||
from .utils.infer import count_invalid_characters, detect_language, apply_character_map, apply_half2full_map, HomophonesReplacer
|
||||
from .utils.io import get_latest_modified_file, del_all
|
||||
from .infer.api import refine_text, infer_code
|
||||
from .utils.download import check_all_assets, download_all_assets
|
||||
from .utils.log import set_utils_logger
|
||||
from .utils.dl import check_all_assets, download_all_assets
|
||||
from .utils.log import logger as utils_logger
|
||||
|
||||
|
||||
class Chat:
|
||||
@@ -26,27 +27,68 @@ class Chat:
|
||||
self.normalizer = {}
|
||||
self.homophones_replacer = None
|
||||
self.logger = logger
|
||||
set_utils_logger(logger)
|
||||
utils_logger.set_logger(logger)
|
||||
|
||||
def check_model(self, level = logging.INFO, use_decoder = False):
|
||||
def has_loaded(self, use_decoder = False):
|
||||
not_finish = False
|
||||
check_list = ['vocos', 'gpt', 'tokenizer']
|
||||
check_list = ['gpt', 'tokenizer']
|
||||
|
||||
if use_decoder:
|
||||
check_list.append('decoder')
|
||||
else:
|
||||
check_list.append('dvae')
|
||||
|
||||
|
||||
for module in check_list:
|
||||
if module not in self.pretrain_models:
|
||||
self.logger.log(logging.WARNING, f'{module} not initialized.')
|
||||
self.logger.warn(f'{module} not initialized.')
|
||||
not_finish = True
|
||||
|
||||
|
||||
if not hasattr(self, "_vocos_decode") or not hasattr(self, "vocos"):
|
||||
self.logger.warn('vocos not initialized.')
|
||||
not_finish = True
|
||||
|
||||
if not not_finish:
|
||||
self.logger.log(level, f'All initialized.')
|
||||
self.logger.info('all models has been initialized.')
|
||||
|
||||
return not not_finish
|
||||
|
||||
def download_models(
|
||||
self,
|
||||
source: Literal['huggingface', 'local', 'custom']='local',
|
||||
force_redownload=False,
|
||||
custom_path: Optional[torch.serialization.FILE_LIKE]=None,
|
||||
) -> Optional[str]:
|
||||
if source == 'local':
|
||||
download_path = os.getcwd()
|
||||
if not check_all_assets(update=True) or force_redownload:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
download_all_assets(tmpdir=tmp)
|
||||
if not check_all_assets(update=False):
|
||||
self.logger.error("download to local path %s failed.", download_path)
|
||||
return None
|
||||
elif source == 'huggingface':
|
||||
hf_home = os.getenv('HF_HOME', os.path.expanduser("~/.cache/huggingface"))
|
||||
try:
|
||||
download_path = get_latest_modified_file(os.path.join(hf_home, 'hub/models--2Noise--ChatTTS/snapshots'))
|
||||
except:
|
||||
download_path = None
|
||||
if download_path is None or force_redownload:
|
||||
self.logger.log(logging.INFO, f'Download from HF: https://huggingface.co/2Noise/ChatTTS')
|
||||
try:
|
||||
download_path = snapshot_download(repo_id="2Noise/ChatTTS", allow_patterns=["*.pt", "*.yaml"])
|
||||
except:
|
||||
download_path = None
|
||||
else:
|
||||
self.logger.log(logging.INFO, f'load latest snapshot from cache: {download_path}')
|
||||
if download_path is None:
|
||||
self.logger.error("download from huggingface failed.")
|
||||
return None
|
||||
elif source == 'custom':
|
||||
self.logger.log(logging.INFO, f'try to load from local: {custom_path}')
|
||||
download_path = custom_path
|
||||
|
||||
return download_path
|
||||
|
||||
def load_models(
|
||||
self,
|
||||
source: Literal['huggingface', 'local', 'custom']='local',
|
||||
@@ -55,36 +97,15 @@ class Chat:
|
||||
custom_path: Optional[torch.serialization.FILE_LIKE]=None,
|
||||
device: Optional[torch.device] = None,
|
||||
coef: Optional[torch.Tensor] = None,
|
||||
):
|
||||
if source == 'local':
|
||||
torch.load
|
||||
download_path = os.getcwd()
|
||||
if not check_all_assets(update=True) or force_redownload:
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
download_all_assets(tmpdir=tmp)
|
||||
if not check_all_assets(update=False):
|
||||
self.logger.error("counld not satisfy all assets needed.")
|
||||
return False
|
||||
elif source == 'huggingface':
|
||||
hf_home = os.getenv('HF_HOME', os.path.expanduser("~/.cache/huggingface"))
|
||||
try:
|
||||
download_path = get_latest_modified_file(os.path.join(hf_home, 'hub/models--2Noise--ChatTTS/snapshots'))
|
||||
except:
|
||||
download_path = None
|
||||
if download_path is None or force_redownload:
|
||||
self.logger.log(logging.INFO, f'Download from HF: https://huggingface.co/2Noise/ChatTTS')
|
||||
download_path = snapshot_download(repo_id="2Noise/ChatTTS", allow_patterns=["*.pt", "*.yaml"])
|
||||
else:
|
||||
self.logger.log(logging.INFO, f'Load from cache: {download_path}')
|
||||
elif source == 'custom':
|
||||
self.logger.log(logging.INFO, f'Load from local: {custom_path}')
|
||||
download_path = custom_path
|
||||
|
||||
) -> bool:
|
||||
download_path = self.download_models(source, force_redownload, custom_path)
|
||||
if download_path is None:
|
||||
return False
|
||||
return self._load(
|
||||
device=device, compile=compile, coef=coef,
|
||||
**{k: os.path.join(download_path, v) for k, v in OmegaConf.load(os.path.join(download_path, 'config', 'path.yaml')).items()},
|
||||
)
|
||||
|
||||
|
||||
def _load(
|
||||
self,
|
||||
vocos_config_path: str = None,
|
||||
@@ -112,9 +133,17 @@ class Chat:
|
||||
).eval()
|
||||
assert vocos_ckpt_path, 'vocos_ckpt_path should not be None'
|
||||
vocos.load_state_dict(torch.load(vocos_ckpt_path))
|
||||
self.pretrain_models['vocos'] = vocos
|
||||
self.vocos = vocos
|
||||
if "mps" in str(self.device):
|
||||
self._vocos_decode: Callable[[torch.Tensor], np.ndarray] = lambda spec: self.vocos.decode(
|
||||
spec.cpu()
|
||||
).cpu().numpy()
|
||||
else:
|
||||
self._vocos_decode: Callable[[torch.Tensor], np.ndarray] = lambda spec: self.vocos.decode(
|
||||
spec
|
||||
).cpu().numpy()
|
||||
self.logger.log(logging.INFO, 'vocos loaded.')
|
||||
|
||||
|
||||
if dvae_config_path:
|
||||
cfg = OmegaConf.load(dvae_config_path)
|
||||
dvae = DVAE(**cfg, coef=coef).to(device).eval()
|
||||
@@ -157,8 +186,13 @@ class Chat:
|
||||
|
||||
self.coef = coef
|
||||
|
||||
return self.check_model()
|
||||
return self.has_loaded()
|
||||
|
||||
def unload(self):
|
||||
logger = self.logger
|
||||
del_all(self)
|
||||
self.__init__(logger)
|
||||
|
||||
def _infer(
|
||||
self,
|
||||
text,
|
||||
@@ -173,23 +207,23 @@ class Chat:
|
||||
do_homophone_replacement=True
|
||||
):
|
||||
|
||||
assert self.check_model(use_decoder=use_decoder)
|
||||
assert self.has_loaded(use_decoder=use_decoder)
|
||||
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
if do_text_normalization:
|
||||
for i, t in enumerate(text):
|
||||
_lang = detect_language(t) if lang is None else lang
|
||||
if self.init_normalizer(_lang):
|
||||
if self._init_normalizer(_lang):
|
||||
text[i] = self.normalizer[_lang](t)
|
||||
if _lang == 'zh':
|
||||
text[i] = apply_half2full_map(text[i])
|
||||
for i, t in enumerate(text):
|
||||
invalid_characters = count_invalid_characters(t)
|
||||
if len(invalid_characters):
|
||||
self.logger.log(logging.WARNING, f'Invalid characters found! : {invalid_characters}')
|
||||
self.logger.warn(f'Invalid characters found! : {invalid_characters}')
|
||||
text[i] = apply_character_map(t)
|
||||
if do_homophone_replacement and self.init_homophones_replacer():
|
||||
if do_homophone_replacement and self._init_homophones_replacer():
|
||||
text[i], replaced_words = self.homophones_replacer.replace(text[i])
|
||||
if replaced_words:
|
||||
repl_res = ', '.join([f'{_[0]}->{_[1]}' for _ in replaced_words])
|
||||
@@ -205,64 +239,25 @@ class Chat:
|
||||
text_tokens = refined.ids
|
||||
text_tokens = [i[i < self.pretrain_models['tokenizer'].convert_tokens_to_ids('[break_0]')] for i in text_tokens]
|
||||
text = self.pretrain_models['tokenizer'].batch_decode(text_tokens)
|
||||
del_all(refined)
|
||||
refined.destroy()
|
||||
if refine_text_only:
|
||||
yield text
|
||||
return
|
||||
|
||||
text = [params_infer_code.get('prompt', '') + i for i in text]
|
||||
params_infer_code.pop('prompt', '')
|
||||
result_gen = infer_code(
|
||||
|
||||
length = [0 for _ in range(len(text))]
|
||||
for result in infer_code(
|
||||
self.pretrain_models,
|
||||
text,
|
||||
device=self.device,
|
||||
**params_infer_code,
|
||||
return_hidden=use_decoder,
|
||||
stream=stream,
|
||||
)
|
||||
if use_decoder:
|
||||
docoder_name = 'decoder'
|
||||
else:
|
||||
docoder_name = 'dvae'
|
||||
if "mps" in str(self.device):
|
||||
vocos_decode = lambda spec: [self.pretrain_models['vocos'].decode(
|
||||
i.cpu()
|
||||
).cpu().numpy() for i in spec]
|
||||
else:
|
||||
vocos_decode = lambda spec: [self.pretrain_models['vocos'].decode(
|
||||
i
|
||||
).cpu().numpy() for i in spec]
|
||||
if stream:
|
||||
|
||||
length = 0
|
||||
for result in result_gen:
|
||||
x = result.hiddens if use_decoder else result.ids
|
||||
assert len(x) == 1
|
||||
chunk_data = x[0]
|
||||
start_seek = length
|
||||
length = len(chunk_data)
|
||||
self.logger.debug(f'{start_seek=} total len: {length}, new len: {length - start_seek = }')
|
||||
chunk_data = chunk_data[start_seek:]
|
||||
if not len(chunk_data):
|
||||
continue
|
||||
self.logger.debug(f'new hidden {len(chunk_data)=}')
|
||||
mel_spec = [self.pretrain_models[docoder_name](i[None].permute(0,2,1).to(self.device)) for i in [chunk_data]]
|
||||
del_all(result)
|
||||
del chunk_data
|
||||
del_all(x)
|
||||
wav = vocos_decode(mel_spec)
|
||||
del_all(mel_spec)
|
||||
self.logger.debug(f'yield wav chunk {len(wav[0])=} {len(wav[0][0])=}')
|
||||
yield wav
|
||||
return
|
||||
result = next(result_gen)
|
||||
x = result.hiddens if use_decoder else result.ids
|
||||
mel_spec = [self.pretrain_models[docoder_name](i[None].permute(0,2,1).to(self.device)) for i in x]
|
||||
del_all(result)
|
||||
del_all(x)
|
||||
wav = vocos_decode(mel_spec)
|
||||
del_all(mel_spec)
|
||||
yield wav
|
||||
):
|
||||
wav = self.decode_to_wavs(result, length, use_decoder)
|
||||
yield wav
|
||||
|
||||
def infer(
|
||||
self,
|
||||
@@ -294,13 +289,35 @@ class Chat:
|
||||
else:
|
||||
return next(res_gen)
|
||||
|
||||
def sample_random_speaker(self, ):
|
||||
|
||||
def sample_random_speaker(self):
|
||||
dim = self.pretrain_models['gpt'].gpt.layers[0].mlp.gate_proj.in_features
|
||||
std, mean = self.pretrain_models['spk_stat'].chunk(2)
|
||||
return torch.randn(dim, device=std.device) * std + mean
|
||||
|
||||
def init_normalizer(self, lang) -> bool:
|
||||
|
||||
def decode_to_wavs(self, result: GPT.GenerationOutputs, start_seeks: List[int], use_decoder: bool):
|
||||
x = result.hiddens if use_decoder else result.ids
|
||||
wavs: List[np.ndarray] = []
|
||||
for i, chunk_data in enumerate(x):
|
||||
start_seek = start_seeks[i]
|
||||
length = len(chunk_data)
|
||||
if length <= start_seek:
|
||||
wavs.append(None)
|
||||
continue
|
||||
start_seeks[i] = length
|
||||
chunk_data = chunk_data[start_seek:]
|
||||
if use_decoder:
|
||||
decoder = self.pretrain_models['decoder']
|
||||
else:
|
||||
decoder = self.pretrain_models['dvae']
|
||||
mel_spec = decoder(chunk_data[None].permute(0,2,1).to(self.device))
|
||||
del chunk_data
|
||||
wavs.append(self._vocos_decode(mel_spec))
|
||||
del_all(mel_spec)
|
||||
result.destroy()
|
||||
del_all(x)
|
||||
return wavs
|
||||
|
||||
def _init_normalizer(self, lang) -> bool:
|
||||
|
||||
if lang in self.normalizer:
|
||||
return True
|
||||
@@ -335,16 +352,16 @@ class Chat:
|
||||
)
|
||||
return False
|
||||
|
||||
def init_homophones_replacer(self):
|
||||
def _init_homophones_replacer(self):
|
||||
if self.homophones_replacer:
|
||||
return True
|
||||
else:
|
||||
try:
|
||||
self.homophones_replacer = HomophonesReplacer(os.path.join(os.path.dirname(__file__), 'res', 'homophones_map.json'))
|
||||
self.logger.log(logging.INFO, 'homophones_replacer loaded.')
|
||||
self.logger.log(logging.INFO, 'successfully loaded HomophonesReplacer.')
|
||||
return True
|
||||
except (IOError, json.JSONDecodeError) as e:
|
||||
self.logger.log(logging.WARNING, f'Error loading homophones map: {e}')
|
||||
self.logger.log(logging.WARNING, f'error loading homophones map: {e}')
|
||||
except Exception as e:
|
||||
self.logger.log(logging.WARNING, f'Error loading homophones_replacer: {e}')
|
||||
self.logger.log(logging.WARNING, f'error loading HomophonesReplacer: {e}')
|
||||
return False
|
||||
|
||||
+14
-13
@@ -164,20 +164,21 @@ class DVAE(nn.Module):
|
||||
return b14.encode_to_string(self.coef.cpu().numpy().astype(np.float32).tobytes())
|
||||
|
||||
def forward(self, inp: torch.Tensor) -> torch.Tensor:
|
||||
with torch.no_grad():
|
||||
|
||||
if self.vq_layer is not None:
|
||||
vq_feats = self.vq_layer._embed(inp)
|
||||
else:
|
||||
vq_feats = inp.detach().clone()
|
||||
if self.vq_layer is not None:
|
||||
vq_feats = self.vq_layer._embed(inp)
|
||||
else:
|
||||
vq_feats = inp.detach().clone()
|
||||
|
||||
vq_feats = vq_feats.view(
|
||||
(vq_feats.size(0), 2, vq_feats.size(1)//2, vq_feats.size(2)),
|
||||
).permute(0, 2, 3, 1).flatten(2)
|
||||
vq_feats = vq_feats.view(
|
||||
(vq_feats.size(0), 2, vq_feats.size(1)//2, vq_feats.size(2)),
|
||||
).permute(0, 2, 3, 1).flatten(2)
|
||||
|
||||
dec_out = self.out_conv(
|
||||
self.decoder(
|
||||
input=vq_feats.transpose_(1, 2),
|
||||
).transpose_(1, 2),
|
||||
)
|
||||
dec_out = self.out_conv(
|
||||
self.decoder(
|
||||
input=vq_feats.transpose_(1, 2),
|
||||
).transpose_(1, 2),
|
||||
)
|
||||
|
||||
return torch.mul(dec_out, self.coef, out=dec_out)
|
||||
return torch.mul(dec_out, self.coef, out=dec_out)
|
||||
|
||||
+66
-48
@@ -1,20 +1,19 @@
|
||||
import os
|
||||
os.environ["TOKENIZERS_PARALLELISM"] = "false"
|
||||
|
||||
from dataclasses import dataclass
|
||||
import logging
|
||||
from typing import Union, List, Optional, Tuple
|
||||
from dataclasses import dataclass
|
||||
|
||||
from tqdm import tqdm
|
||||
from transformers.cache_utils import Cache
|
||||
import omegaconf
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.nn.utils.parametrize as P
|
||||
from torch.nn.utils.parametrizations import weight_norm
|
||||
from tqdm import tqdm
|
||||
from transformers import LlamaModel, LlamaConfig, LogitsWarper
|
||||
from transformers.cache_utils import Cache
|
||||
from transformers.modeling_outputs import BaseModelOutputWithPast
|
||||
|
||||
from ..utils.infer import CustomRepetitionPenaltyLogitsProcessorRepeat
|
||||
@@ -98,23 +97,22 @@ class GPT(nn.Module):
|
||||
|
||||
emb_text: torch.Tensor = self.emb_text(input_ids[text_mask].narrow(1, 0, 1).squeeze_(1).to(self.device_gpt))
|
||||
|
||||
text_mask_inv = ~text_mask
|
||||
text_mask_inv = ~(text_mask.to(self.device_gpt))
|
||||
masked_input_ids: torch.Tensor = input_ids[text_mask_inv].to(self.device_gpt)
|
||||
del text_mask_inv
|
||||
|
||||
emb_code = [self.emb_code[i](masked_input_ids[:, i]) for i in range(self.num_vq)]
|
||||
emb_code = torch.stack(emb_code, 2).sum(2)
|
||||
|
||||
emb = torch.zeros((input_ids.shape[:-1])+(emb_text.shape[-1],), device=emb_text.device, dtype=emb_text.dtype)
|
||||
emb[text_mask] = emb_text
|
||||
emb[~text_mask] = emb_code.to(emb.dtype)
|
||||
emb[text_mask_inv] = emb_code.to(emb.dtype)
|
||||
|
||||
del emb_text, emb_code
|
||||
del emb_text, emb_code, text_mask_inv
|
||||
|
||||
return emb
|
||||
|
||||
@dataclass(repr=False, eq=False)
|
||||
class GenerationInputs():
|
||||
class _GenerationInputs():
|
||||
position_ids: torch.Tensor
|
||||
cache_position: torch.Tensor
|
||||
use_cache: bool
|
||||
@@ -129,7 +127,7 @@ class GPT(nn.Module):
|
||||
if self.inputs_embeds is not None: self.inputs_embeds = self.inputs_embeds.to(device)
|
||||
if self.cache_position is not None: self.cache_position = self.cache_position.to(device)
|
||||
|
||||
def _prepare_inputs_for_generation(
|
||||
def _prepare_generation_inputs(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
past_key_values: Optional[Tuple[Tuple[torch.FloatTensor]]]=None,
|
||||
@@ -138,7 +136,7 @@ class GPT(nn.Module):
|
||||
cache_position: Optional[torch.Tensor]=None,
|
||||
position_ids: Optional[torch.Tensor]=None,
|
||||
use_cache = True,
|
||||
) -> GenerationInputs:
|
||||
) -> _GenerationInputs:
|
||||
# With static cache, the `past_key_values` is None
|
||||
# TODO joao: standardize interface for the different Cache classes and remove of this if
|
||||
has_static_cache = False
|
||||
@@ -197,7 +195,7 @@ class GPT(nn.Module):
|
||||
if has_static_cache:
|
||||
past_key_values = None
|
||||
|
||||
model_inputs = self.GenerationInputs(
|
||||
model_inputs = self._GenerationInputs(
|
||||
position_ids=position_ids,
|
||||
cache_position=cache_position,
|
||||
use_cache=use_cache,
|
||||
@@ -223,7 +221,36 @@ class GPT(nn.Module):
|
||||
attentions: List[Optional[Tuple[torch.FloatTensor, ...]]]
|
||||
hiddens: List[torch.Tensor]
|
||||
|
||||
|
||||
def destroy(self):
|
||||
del_all(self.ids)
|
||||
del_all(self.attentions)
|
||||
del_all(self.hiddens)
|
||||
|
||||
|
||||
def _prepare_generation_outputs(
|
||||
self,
|
||||
inputs_ids: torch.Tensor,
|
||||
start_idx: int,
|
||||
end_idx: torch.Tensor,
|
||||
attentions: List[Optional[Tuple[torch.FloatTensor, ...]]],
|
||||
hiddens: List[torch.Tensor],
|
||||
infer_text: bool,
|
||||
) -> GenerationOutputs:
|
||||
inputs_ids = [inputs_ids[idx].narrow(0, start_idx, i) for idx, i in enumerate(end_idx)]
|
||||
if infer_text:
|
||||
inputs_ids = [i.narrow(1, 0, 1).squeeze_(1) for i in inputs_ids]
|
||||
|
||||
if len(hiddens) > 0:
|
||||
hiddens = torch.stack(hiddens, 1)
|
||||
hiddens = [hiddens[idx].narrow(0, 0, i) for idx, i in enumerate(end_idx.int())]
|
||||
|
||||
return self.GenerationOutputs(
|
||||
ids=inputs_ids,
|
||||
attentions=attentions,
|
||||
hiddens=hiddens,
|
||||
)
|
||||
|
||||
|
||||
def generate(
|
||||
self,
|
||||
emb: torch.Tensor,
|
||||
@@ -256,12 +283,16 @@ class GPT(nn.Module):
|
||||
if attention_mask is not None:
|
||||
attention_mask_cache[:, :attention_mask.shape[1]] = attention_mask
|
||||
|
||||
with tqdm(total=max_new_token) as pbar:
|
||||
with tqdm(
|
||||
total=max_new_token,
|
||||
desc="text" if infer_text else "code",
|
||||
bar_format='{l_bar}{bar}| {n_fmt}/{total_fmt}(max) [{elapsed}, {rate_fmt}{postfix}]',
|
||||
) as pbar:
|
||||
|
||||
past_key_values = None
|
||||
|
||||
for i in range(max_new_token):
|
||||
model_input = self._prepare_inputs_for_generation(
|
||||
model_input = self._prepare_generation_inputs(
|
||||
inputs_ids,
|
||||
past_key_values,
|
||||
attention_mask_cache[:, :inputs_ids.shape[1]],
|
||||
@@ -296,7 +327,7 @@ class GPT(nn.Module):
|
||||
past_key_values = outputs.past_key_values
|
||||
del_all(outputs)
|
||||
if return_hidden:
|
||||
hiddens.append(hidden_states[:, -1])
|
||||
hiddens.append(hidden_states.narrow(1, -1, 1).squeeze_(1))
|
||||
|
||||
with P.cached():
|
||||
if infer_text:
|
||||
@@ -353,52 +384,39 @@ class GPT(nn.Module):
|
||||
finish_or = (idx_next == eos_token).any(1)
|
||||
finish |= finish_or
|
||||
del finish_or
|
||||
inputs_ids = torch.cat([inputs_ids, idx_next.unsqueeze(1)], 1)
|
||||
inputs_ids_tmp = torch.cat([inputs_ids, idx_next.unsqueeze_(1)], 1)
|
||||
else:
|
||||
finish_or = (idx_next == eos_token).any(1)
|
||||
finish |= finish_or
|
||||
del finish_or
|
||||
inputs_ids = torch.cat([inputs_ids, idx_next.unsqueeze(-1).expand(-1, -1, self.num_vq)], 1)
|
||||
inputs_ids_tmp = torch.cat([inputs_ids, idx_next.unsqueeze_(-1).expand(-1, -1, self.num_vq)], 1)
|
||||
|
||||
del idx_next
|
||||
del inputs_ids
|
||||
inputs_ids = inputs_ids_tmp
|
||||
del inputs_ids_tmp, idx_next
|
||||
|
||||
end_idx += (~finish).int().to(end_idx.device)
|
||||
if stream:
|
||||
if end_idx % 24 and not finish.all():
|
||||
continue
|
||||
y_inputs_ids = [inputs_ids[idx, start_idx: start_idx+i] for idx, i in enumerate(end_idx.int())]
|
||||
y_inputs_ids = [i[:, 0] for i in y_inputs_ids] if infer_text else y_inputs_ids
|
||||
y_hiddens = []
|
||||
if return_hidden:
|
||||
y_hiddens = torch.stack(hiddens, 1)
|
||||
y_hiddens = [y_hiddens[idx, :i] for idx, i in enumerate(end_idx.int())]
|
||||
minus_prev_end_index = -end_idx
|
||||
end_idx += (~finish.to(end_idx.device)).int()
|
||||
if stream:
|
||||
if end_idx.all() and (end_idx%24 == 0).any() and torch.add(end_idx, minus_prev_end_index, out=minus_prev_end_index).any():
|
||||
self.logger.debug("yield stream result, end: %d", end_idx)
|
||||
yield self._prepare_generation_outputs(
|
||||
inputs_ids, start_idx, end_idx, attentions, hiddens,
|
||||
infer_text,
|
||||
)
|
||||
del minus_prev_end_index
|
||||
|
||||
yield self.GenerationOutputs(
|
||||
ids=y_inputs_ids,
|
||||
attentions=attentions,
|
||||
hiddens=y_hiddens,
|
||||
)
|
||||
|
||||
if finish.all():
|
||||
pbar.update(max_new_token-i-1)
|
||||
break
|
||||
if finish.all(): break
|
||||
|
||||
pbar.update(1)
|
||||
|
||||
inputs_ids = [inputs_ids[idx, start_idx: start_idx+i] for idx, i in enumerate(end_idx.int())]
|
||||
inputs_ids = [i[:, 0] for i in inputs_ids] if infer_text else inputs_ids
|
||||
|
||||
if return_hidden:
|
||||
hiddens = torch.stack(hiddens, 1)
|
||||
hiddens = [hiddens[idx, :i] for idx, i in enumerate(end_idx.int())]
|
||||
|
||||
if not finish.all():
|
||||
self.logger.warn(f'Incomplete result. hit max_new_token: {max_new_token}')
|
||||
|
||||
del finish
|
||||
|
||||
yield self.GenerationOutputs(
|
||||
ids=inputs_ids,
|
||||
attentions=attentions,
|
||||
hiddens=hiddens,
|
||||
yield self._prepare_generation_outputs(
|
||||
inputs_ids, start_idx, end_idx, attentions, hiddens,
|
||||
infer_text,
|
||||
)
|
||||
|
||||
@@ -3,15 +3,15 @@ from pathlib import Path
|
||||
import hashlib
|
||||
import requests
|
||||
from io import BytesIO
|
||||
from mmap import mmap, ACCESS_READ
|
||||
|
||||
from .log import logger
|
||||
|
||||
def sha256(f) -> str:
|
||||
sha256_hash = hashlib.sha256()
|
||||
# Read and update hash in chunks of 4M
|
||||
for byte_block in iter(lambda: f.read(4 * 1024 * 1024), b""):
|
||||
sha256_hash.update(byte_block)
|
||||
return sha256_hash.hexdigest()
|
||||
def sha256(fileno: int) -> str:
|
||||
data = mmap(fileno, 0, access=ACCESS_READ)
|
||||
h = hashlib.sha256(data).hexdigest()
|
||||
del data
|
||||
return h
|
||||
|
||||
|
||||
def check_model(
|
||||
@@ -19,18 +19,18 @@ def check_model(
|
||||
) -> bool:
|
||||
target = dir_name / model_name
|
||||
relname = target.as_posix()
|
||||
logger.debug(f"checking {relname}...")
|
||||
logger.get_logger().debug(f"checking {relname}...")
|
||||
if not os.path.exists(target):
|
||||
logger.info(f"{target} not exist.")
|
||||
logger.get_logger().info(f"{target} not exist.")
|
||||
return False
|
||||
with open(target, "rb") as f:
|
||||
digest = sha256(f)
|
||||
digest = sha256(f.fileno())
|
||||
bakfile = f"{target}.bak"
|
||||
if digest != hash:
|
||||
logger.warn(f"{target} sha256 hash mismatch.")
|
||||
logger.info(f"expected: {hash}")
|
||||
logger.info(f"real val: {digest}")
|
||||
logger.warn("please add parameter --update to download the latest assets.")
|
||||
logger.get_logger().warn(f"{target} sha256 hash mismatch.")
|
||||
logger.get_logger().info(f"expected: {hash}")
|
||||
logger.get_logger().info(f"real val: {digest}")
|
||||
logger.get_logger().warn("please add parameter --update to download the latest assets.")
|
||||
if remove_incorrect:
|
||||
if not os.path.exists(bakfile):
|
||||
os.rename(str(target), bakfile)
|
||||
@@ -45,7 +45,7 @@ def check_model(
|
||||
def check_all_assets(update=False) -> bool:
|
||||
BASE_DIR = Path(os.getcwd())
|
||||
|
||||
logger.info("checking assets...")
|
||||
logger.get_logger().info("checking assets...")
|
||||
current_dir = BASE_DIR / "asset"
|
||||
names = [
|
||||
"Decoder.pt",
|
||||
@@ -62,7 +62,7 @@ def check_all_assets(update=False) -> bool:
|
||||
):
|
||||
return False
|
||||
|
||||
logger.info("checking configs...")
|
||||
logger.get_logger().info("checking configs...")
|
||||
current_dir = BASE_DIR / "config"
|
||||
names = [
|
||||
"decoder.yaml",
|
||||
@@ -78,44 +78,44 @@ def check_all_assets(update=False) -> bool:
|
||||
):
|
||||
return False
|
||||
|
||||
logger.info("all assets are already latest.")
|
||||
logger.get_logger().info("all assets are already latest.")
|
||||
return True
|
||||
|
||||
|
||||
def download_and_extract_tar_gz(url: str, folder: str):
|
||||
import tarfile
|
||||
|
||||
logger.info(f"downloading {url}")
|
||||
logger.get_logger().info(f"downloading {url}")
|
||||
response = requests.get(url, stream=True, timeout=(5, 10))
|
||||
with BytesIO() as out_file:
|
||||
out_file.write(response.content)
|
||||
out_file.seek(0)
|
||||
logger.info(f"downloaded.")
|
||||
logger.get_logger().info(f"downloaded.")
|
||||
with tarfile.open(fileobj=out_file, mode="r:gz") as tar:
|
||||
tar.extractall(folder)
|
||||
logger.info(f"extracted into {folder}")
|
||||
logger.get_logger().info(f"extracted into {folder}")
|
||||
|
||||
|
||||
def download_and_extract_zip(url: str, folder: str):
|
||||
import zipfile
|
||||
|
||||
logger.info(f"downloading {url}")
|
||||
logger.get_logger().info(f"downloading {url}")
|
||||
response = requests.get(url, stream=True, timeout=(5, 10))
|
||||
with BytesIO() as out_file:
|
||||
out_file.write(response.content)
|
||||
out_file.seek(0)
|
||||
logger.info(f"downloaded.")
|
||||
logger.get_logger().info(f"downloaded.")
|
||||
with zipfile.ZipFile(out_file) as zip_ref:
|
||||
zip_ref.extractall(folder)
|
||||
logger.info(f"extracted into {folder}")
|
||||
logger.get_logger().info(f"extracted into {folder}")
|
||||
|
||||
|
||||
def download_dns_yaml(url: str, folder: str):
|
||||
logger.info(f"downloading {url}")
|
||||
logger.get_logger().info(f"downloading {url}")
|
||||
response = requests.get(url, stream=True, timeout=(5, 10))
|
||||
with open(os.path.join(folder, "dns.yaml"), "wb") as out_file:
|
||||
out_file.write(response.content)
|
||||
logger.info(f"downloaded into {folder}")
|
||||
logger.get_logger().info(f"downloaded into {folder}")
|
||||
|
||||
|
||||
def download_all_assets(tmpdir: str, version="0.2.5"):
|
||||
@@ -140,7 +140,7 @@ def download_all_assets(tmpdir: str, version="0.2.5"):
|
||||
|
||||
architecture = archs.get(architecture, None)
|
||||
if not architecture:
|
||||
logger.error(f"architecture {architecture} is not supported")
|
||||
logger.get_logger().error(f"architecture {architecture} is not supported")
|
||||
exit(1)
|
||||
try:
|
||||
BASE_URL = "https://github.com/fumiama/RVC-Models-Downloader/releases/download/"
|
||||
@@ -14,14 +14,14 @@ def select_device(min_memory=2048):
|
||||
device = torch.device(f'cuda:{selected_gpu}')
|
||||
free_memory_mb = max_free_memory / (1024 * 1024)
|
||||
if free_memory_mb < min_memory:
|
||||
logger.warning(f'GPU {selected_gpu} has {round(free_memory_mb, 2)} MB memory left. Switching to CPU.')
|
||||
logger.get_logger().warning(f'GPU {selected_gpu} has {round(free_memory_mb, 2)} MB memory left. Switching to CPU.')
|
||||
device = torch.device('cpu')
|
||||
elif torch.backends.mps.is_available():
|
||||
# For Apple M1/M2 chips with Metal Performance Shaders
|
||||
logger.info('Apple GPU found, using MPS.')
|
||||
logger.get_logger().info('Apple GPU found, using MPS.')
|
||||
device = torch.device('mps')
|
||||
else:
|
||||
logger.warning('No GPU found, use CPU instead')
|
||||
logger.get_logger().warning('No GPU found, use CPU instead')
|
||||
device = torch.device('cpu')
|
||||
|
||||
return device
|
||||
|
||||
+63
-38
@@ -1,8 +1,12 @@
|
||||
|
||||
import json
|
||||
import re
|
||||
from typing import Dict, Tuple, List
|
||||
import sys
|
||||
|
||||
from numba import jit
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import json
|
||||
|
||||
|
||||
class CustomRepetitionPenaltyLogitsProcessorRepeat():
|
||||
@@ -47,6 +51,26 @@ class CustomRepetitionPenaltyLogitsProcessor():
|
||||
|
||||
return scores
|
||||
|
||||
@jit
|
||||
def _find_index(table: np.ndarray, val: np.uint16):
|
||||
for i in range(table.size):
|
||||
if table[i] == val:
|
||||
return i
|
||||
return -1
|
||||
|
||||
@jit
|
||||
def _fast_replace(table: np.ndarray, text: bytes) -> Tuple[np.ndarray, List[Tuple[str, str]]]:
|
||||
result = np.frombuffer(text, dtype=np.uint16).copy()
|
||||
replaced_words = []
|
||||
for i in range(result.size):
|
||||
ch = result[i]
|
||||
p = _find_index(table[0], ch)
|
||||
if p >= 0:
|
||||
repl_char = table[1][p]
|
||||
result[i] = repl_char
|
||||
replaced_words.append((chr(ch), chr(repl_char)))
|
||||
return result, replaced_words
|
||||
|
||||
class HomophonesReplacer:
|
||||
"""
|
||||
Homophones Replacer
|
||||
@@ -65,37 +89,40 @@ class HomophonesReplacer:
|
||||
[python-pinyin](https://github.com/mozillazg/python-pinyin)
|
||||
|
||||
"""
|
||||
def __init__(self, map_file_path):
|
||||
self.homophones_map = self.load_homophones_map(map_file_path)
|
||||
def __init__(self, map_file_path: str):
|
||||
self.homophones_map = self._load_homophones_map(map_file_path)
|
||||
self.coding = "utf-16-le" if sys.byteorder == "little" else "utf-16-be"
|
||||
|
||||
def load_homophones_map(self, map_file_path):
|
||||
def _load_homophones_map(self, map_file_path: str) -> np.ndarray:
|
||||
with open(map_file_path, 'r', encoding='utf-8') as f:
|
||||
homophones_map = json.load(f)
|
||||
return homophones_map
|
||||
homophones_map: Dict[str, str] = json.load(f)
|
||||
map = np.empty((2, len(homophones_map)), dtype=np.uint32)
|
||||
for i, k in enumerate(homophones_map.keys()):
|
||||
map[:, i] = (ord(k), ord(homophones_map[k]))
|
||||
del homophones_map
|
||||
return map
|
||||
|
||||
def replace(self, text):
|
||||
result = []
|
||||
replaced_words = []
|
||||
for char in text:
|
||||
if char in self.homophones_map:
|
||||
repl_char = self.homophones_map[char]
|
||||
result.append(repl_char)
|
||||
replaced_words.append((char, repl_char))
|
||||
else:
|
||||
result.append(char)
|
||||
return ''.join(result), replaced_words
|
||||
def replace(self, text: str):
|
||||
arr, lst = _fast_replace(
|
||||
self.homophones_map,
|
||||
text.encode(self.coding),
|
||||
)
|
||||
return arr.tobytes().decode(self.coding), lst
|
||||
|
||||
def count_invalid_characters(s):
|
||||
|
||||
s = re.sub(r'\[uv_break\]|\[laugh\]|\[lbreak\]', '', s)
|
||||
pattern = re.compile(r'[^\u4e00-\u9fffA-Za-z,。、,\. ]')
|
||||
non_alphabetic_chinese_chars = pattern.findall(s)
|
||||
accept_pattern = re.compile(r'[^\u4e00-\u9fffA-Za-z,。、,\. ]')
|
||||
sub_pattern = re.compile(r'\[uv_break\]|\[laugh\]|\[lbreak\]')
|
||||
|
||||
def count_invalid_characters(s: str):
|
||||
global accept_pattern, sub_pattern
|
||||
s = sub_pattern.sub('', s)
|
||||
non_alphabetic_chinese_chars = accept_pattern.findall(s)
|
||||
return set(non_alphabetic_chinese_chars)
|
||||
|
||||
def detect_language(sentence):
|
||||
chinese_char_pattern = re.compile(r'[\u4e00-\u9fff]')
|
||||
english_word_pattern = re.compile(r'\b[A-Za-z]+\b')
|
||||
|
||||
chinese_char_pattern = re.compile(r'[\u4e00-\u9fff]')
|
||||
english_word_pattern = re.compile(r'\b[A-Za-z]+\b')
|
||||
def detect_language(sentence):
|
||||
global chinese_char_pattern, english_word_pattern
|
||||
|
||||
chinese_chars = chinese_char_pattern.findall(sentence)
|
||||
english_words = english_word_pattern.findall(sentence)
|
||||
@@ -104,9 +131,9 @@ def detect_language(sentence):
|
||||
return "zh"
|
||||
else:
|
||||
return "en"
|
||||
|
||||
|
||||
character_map = {
|
||||
|
||||
|
||||
character_simplifier = str.maketrans({
|
||||
':': ',',
|
||||
';': ',',
|
||||
'!': '。',
|
||||
@@ -135,9 +162,9 @@ character_map = {
|
||||
'>': ',',
|
||||
'<': ',',
|
||||
'-': ',',
|
||||
}
|
||||
})
|
||||
|
||||
halfwidth_2_fullwidth_map = {
|
||||
halfwidth_2_fullwidth = str.maketrans({
|
||||
'!': '!',
|
||||
'"': '“',
|
||||
"'": '‘',
|
||||
@@ -170,12 +197,10 @@ halfwidth_2_fullwidth_map = {
|
||||
'|': '|',
|
||||
'}': '}',
|
||||
'~': '~'
|
||||
}
|
||||
})
|
||||
|
||||
def apply_half2full_map(text):
|
||||
translation_table = str.maketrans(halfwidth_2_fullwidth_map)
|
||||
return text.translate(translation_table)
|
||||
def apply_half2full_map(text: str) -> str:
|
||||
return text.translate(halfwidth_2_fullwidth)
|
||||
|
||||
def apply_character_map(text):
|
||||
translation_table = str.maketrans(character_map)
|
||||
return text.translate(translation_table)
|
||||
def apply_character_map(text: str) -> str:
|
||||
return text.translate(character_simplifier)
|
||||
|
||||
+4
-3
@@ -10,7 +10,7 @@ def get_latest_modified_file(directory):
|
||||
|
||||
files = [os.path.join(directory, f) for f in os.listdir(directory)]
|
||||
if not files:
|
||||
logger.log(logging.WARNING, f'No files found in the directory: {directory}')
|
||||
logger.get_logger().log(logging.WARNING, f'no files found in the directory: {directory}')
|
||||
return None
|
||||
latest_file = max(files, key=os.path.getmtime)
|
||||
|
||||
@@ -31,11 +31,12 @@ def del_all(d: Union[dict, list]):
|
||||
if isinstance(x, dict) or isinstance(x, list) or is_dataclass(x):
|
||||
del_all(x)
|
||||
del x
|
||||
return
|
||||
elif isinstance(d, list):
|
||||
while len(d):
|
||||
x = d.pop()
|
||||
if isinstance(x, dict) or isinstance(x, list) or is_dataclass(x):
|
||||
del_all(x)
|
||||
del x
|
||||
return
|
||||
else:
|
||||
del d
|
||||
|
||||
|
||||
+10
-4
@@ -1,8 +1,14 @@
|
||||
import logging
|
||||
from pathlib import Path
|
||||
|
||||
logger = logging.getLogger(Path(__file__).parent.name)
|
||||
class Logger():
|
||||
def __init__(self, logger=logging.getLogger(Path(__file__).parent.name)):
|
||||
self.logger = logger
|
||||
|
||||
def set_utils_logger(l: logging.Logger):
|
||||
global logger
|
||||
logger = l
|
||||
def set_logger(self, logger: logging.Logger):
|
||||
self.logger = logger
|
||||
|
||||
def get_logger(self) -> logging.Logger:
|
||||
return self.logger
|
||||
|
||||
logger = Logger()
|
||||
|
||||
Vendored
+16
@@ -132,6 +132,22 @@
|
||||
"chat.load_models(source='custom', custom_path='YOUR CUSTOM PATH')"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### You can also unload models to save the memory"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"chat.unload()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {
|
||||
|
||||
Vendored
+16
@@ -113,6 +113,22 @@
|
||||
"chat.load_models(source='custom', custom_path='YOUR CUSTOM PATH')"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### You can also unload models to save the memory"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"chat.unload()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"metadata": {},
|
||||
|
||||
+56
-36
@@ -1,6 +1,6 @@
|
||||
import random
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import gradio as gr
|
||||
import numpy as np
|
||||
|
||||
@@ -8,22 +8,25 @@ from tools.audio import unsafe_float_to_int16
|
||||
from tools.logger import get_logger
|
||||
logger = get_logger(" WebUI ")
|
||||
|
||||
from tools.seeder import TorchSeedContext
|
||||
|
||||
import ChatTTS
|
||||
chat = ChatTTS.Chat(get_logger("ChatTTS"))
|
||||
|
||||
custom_path: Optional[str] = None
|
||||
|
||||
# 音色选项:用于预置合适的音色
|
||||
voices = {
|
||||
"默认": {"seed": 2},
|
||||
"音色1": {"seed": 1111},
|
||||
"音色2": {"seed": 2222},
|
||||
"音色3": {"seed": 3333},
|
||||
"音色4": {"seed": 4444},
|
||||
"音色5": {"seed": 5555},
|
||||
"音色6": {"seed": 6666},
|
||||
"音色7": {"seed": 7777},
|
||||
"音色8": {"seed": 8888},
|
||||
"音色9": {"seed": 9999},
|
||||
"音色10": {"seed": 11111},
|
||||
"Default": {"seed": 2},
|
||||
"Timbre1": {"seed": 1111},
|
||||
"Timbre2": {"seed": 2222},
|
||||
"Timbre3": {"seed": 3333},
|
||||
"Timbre4": {"seed": 4444},
|
||||
"Timbre5": {"seed": 5555},
|
||||
"Timbre6": {"seed": 6666},
|
||||
"Timbre7": {"seed": 7777},
|
||||
"Timbre8": {"seed": 8888},
|
||||
"Timbre9": {"seed": 9999},
|
||||
}
|
||||
|
||||
def generate_seed():
|
||||
@@ -33,45 +36,62 @@ def generate_seed():
|
||||
def on_voice_change(vocie_selection):
|
||||
return voices.get(vocie_selection)['seed']
|
||||
|
||||
def refine_text(text, audio_seed_input, text_seed_input, refine_text_flag):
|
||||
def reload_chat(coef: Optional[str]) -> str:
|
||||
global custom_path
|
||||
chat.unload()
|
||||
gr.Info("Model unloaded.")
|
||||
try:
|
||||
if len(coef) != 230:
|
||||
gr.Warning("Ingore invalid DVAE coefficient.")
|
||||
coef = None
|
||||
if custom_path == None:
|
||||
ret = chat.load_models(coef=coef)
|
||||
else:
|
||||
logger.info('local model path: %s', custom_path)
|
||||
ret = chat.load_models('custom', custom_path=custom_path, coef=coef)
|
||||
if not ret:
|
||||
raise gr.Error("Unable to load model.")
|
||||
gr.Info("Reload succeess.")
|
||||
return chat.coef
|
||||
except Exception as e:
|
||||
raise gr.Error(str(e))
|
||||
|
||||
def refine_text(text, text_seed_input, refine_text_flag):
|
||||
if not refine_text_flag:
|
||||
return text
|
||||
|
||||
global chat
|
||||
|
||||
torch.manual_seed(audio_seed_input)
|
||||
params_refine_text = {'prompt': '[oral_2][laugh_0][break_6]'}
|
||||
|
||||
torch.manual_seed(text_seed_input)
|
||||
|
||||
text = chat.infer(text,
|
||||
skip_refine_text=False,
|
||||
refine_text_only=True,
|
||||
params_refine_text=params_refine_text,
|
||||
)
|
||||
with TorchSeedContext(text_seed_input):
|
||||
text = chat.infer(text,
|
||||
skip_refine_text=False,
|
||||
refine_text_only=True,
|
||||
params_refine_text=params_refine_text,
|
||||
)
|
||||
return text[0] if isinstance(text, list) else text
|
||||
|
||||
def generate_audio(text, temperature, top_P, top_K, audio_seed_input, text_seed_input, stream):
|
||||
def generate_audio(text, temperature, top_P, top_K, audio_seed_input, stream):
|
||||
if not text: return None
|
||||
|
||||
global chat
|
||||
|
||||
torch.manual_seed(audio_seed_input)
|
||||
rand_spk = chat.sample_random_speaker()
|
||||
params_infer_code = {
|
||||
'spk_emb': rand_spk,
|
||||
'temperature': temperature,
|
||||
'top_P': top_P,
|
||||
'top_K': top_K,
|
||||
with TorchSeedContext(audio_seed_input):
|
||||
rand_spk = chat.sample_random_speaker()
|
||||
params_infer_code = {
|
||||
'spk_emb': rand_spk,
|
||||
'temperature': temperature,
|
||||
'top_P': top_P,
|
||||
'top_K': top_K,
|
||||
}
|
||||
torch.manual_seed(text_seed_input)
|
||||
|
||||
wav = chat.infer(
|
||||
text,
|
||||
skip_refine_text=True,
|
||||
params_infer_code=params_infer_code,
|
||||
stream=stream,
|
||||
)
|
||||
wav = chat.infer(
|
||||
text,
|
||||
skip_refine_text=True,
|
||||
params_infer_code=params_infer_code,
|
||||
stream=stream,
|
||||
)
|
||||
|
||||
if stream:
|
||||
for gen in wav:
|
||||
|
||||
+18
-8
@@ -32,16 +32,22 @@ def main():
|
||||
top_k_slider = gr.Slider(minimum=1, maximum=20, step=1, value=20, label="top_K", interactive=True)
|
||||
|
||||
with gr.Row():
|
||||
voice_selection = gr.Dropdown(label="音色", choices=voices.keys(), value='默认')
|
||||
voice_selection = gr.Dropdown(label="Timbre", choices=voices.keys(), value='Default')
|
||||
audio_seed_input = gr.Number(value=2, label="Audio Seed")
|
||||
generate_audio_seed = gr.Button("\U0001F3B2")
|
||||
text_seed_input = gr.Number(value=42, label="Text Seed")
|
||||
generate_text_seed = gr.Button("\U0001F3B2")
|
||||
|
||||
with gr.Row():
|
||||
dvae_coef_text = gr.Textbox(
|
||||
label="DVAE Coefficient", max_lines=3, show_copy_button=True, scale=4,
|
||||
)
|
||||
reload_chat_button = gr.Button("Reload", scale=1)
|
||||
|
||||
with gr.Row():
|
||||
auto_play_checkbox = gr.Checkbox(label="Auto Play", value=False, scale=1)
|
||||
stream_mode_checkbox = gr.Checkbox(label="Stream Mode", value=False, scale=1)
|
||||
generate_button = gr.Button("Generate", scale=2)
|
||||
generate_button = gr.Button("Generate", scale=2, variant="primary")
|
||||
|
||||
text_output = gr.Textbox(label="Output Text", interactive=False)
|
||||
|
||||
@@ -56,9 +62,11 @@ def main():
|
||||
inputs=[],
|
||||
outputs=text_seed_input)
|
||||
|
||||
reload_chat_button.click(reload_chat, inputs=dvae_coef_text, outputs=dvae_coef_text)
|
||||
|
||||
generate_button.click(fn=lambda: "", outputs=text_output)
|
||||
generate_button.click(refine_text,
|
||||
inputs=[text_input, audio_seed_input, text_seed_input, refine_text_checkbox],
|
||||
inputs=[text_input, text_seed_input, refine_text_checkbox],
|
||||
outputs=text_output)
|
||||
|
||||
@gr.render(inputs=[auto_play_checkbox, stream_mode_checkbox])
|
||||
@@ -72,14 +80,14 @@ def main():
|
||||
show_label=True,
|
||||
)
|
||||
text_output.change(generate_audio,
|
||||
inputs=[text_output, temperature_slider, top_p_slider, top_k_slider, audio_seed_input, text_seed_input, stream_mode_checkbox],
|
||||
inputs=[text_output, temperature_slider, top_p_slider, top_k_slider, audio_seed_input, stream_mode_checkbox],
|
||||
outputs=audio_output)
|
||||
|
||||
gr.Examples(
|
||||
examples=[
|
||||
["四川美食确实以辣闻名,但也有不辣的选择。比如甜水面、赖汤圆、蛋烘糕、叶儿粑等,这些小吃口味温和,甜而不腻,也很受欢迎。", 0.3, 0.7, 20, 2, 42, True],
|
||||
["What is [uv_break]your favorite english food?[laugh][lbreak]", 0.5, 0.5, 10, 245, 531, True],
|
||||
["chat T T S is a text to speech model designed for dialogue applications. [uv_break]it supports mixed language input [uv_break]and offers multi speaker capabilities with precise control over prosodic elements [laugh]like like [uv_break]laughter[laugh], [uv_break]pauses, [uv_break]and intonation. [uv_break]it delivers natural and expressive speech,[uv_break]so please[uv_break] use the project responsibly at your own risk.[uv_break]", 0.2, 0.6, 15, 67, 165, True],
|
||||
["What is [uv_break]your favorite english food?[laugh][lbreak]", 0.5, 0.5, 10, 245, 531, False],
|
||||
["chat T T S is a text to speech model designed for dialogue applications. [uv_break]it supports mixed language input [uv_break]and offers multi speaker capabilities with precise control over prosodic elements [laugh]like like [uv_break]laughter[laugh], [uv_break]pauses, [uv_break]and intonation. [uv_break]it delivers natural and expressive speech,[uv_break]so please[uv_break] use the project responsibly at your own risk.[uv_break]", 0.2, 0.6, 15, 67, 165, False],
|
||||
],
|
||||
inputs=[text_input, temperature_slider, top_p_slider, top_k_slider, audio_seed_input, text_seed_input, refine_text_checkbox],
|
||||
)
|
||||
@@ -93,7 +101,7 @@ def main():
|
||||
|
||||
logger.info("loading ChatTTS model...")
|
||||
|
||||
global chat
|
||||
global chat, custom_path
|
||||
|
||||
if args.custom_path == None:
|
||||
ret = chat.load_models()
|
||||
@@ -106,7 +114,9 @@ def main():
|
||||
else:
|
||||
logger.error("Models load failed.")
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
custom_path = args.custom_path
|
||||
dvae_coef_text.value = chat.coef
|
||||
|
||||
demo.launch(server_name=args.server_name, server_port=args.server_port, root_path=args.root_path, inbrowser=True)
|
||||
|
||||
|
||||
+1
-2
@@ -7,8 +7,7 @@ def unsafe_float_to_int16(audio: np.ndarray) -> np.ndarray:
|
||||
This function will destroy audio, use only once.
|
||||
"""
|
||||
am = np.abs(audio).max() * 32768
|
||||
if am > 32768:
|
||||
am = 32768 * 32768 / am
|
||||
am = 32767 * 32768 / am
|
||||
np.multiply(audio, am, audio)
|
||||
audio16 = audio.astype(np.int16)
|
||||
return audio16
|
||||
|
||||
+2
-2
@@ -4,6 +4,8 @@ from datetime import datetime, timezone
|
||||
|
||||
logging.getLogger("numba").setLevel(logging.WARNING)
|
||||
logging.getLogger("httpx").setLevel(logging.WARNING)
|
||||
logging.getLogger("wetext-zh_normalizer").setLevel(logging.WARNING)
|
||||
logging.getLogger("NeMo-text-processing").setLevel(logging.WARNING)
|
||||
|
||||
# from https://github.com/FloatTech/ZeroBot-Plugin/blob/c70766a989698452e60e5e48fb2f802a2444330d/console/console_windows.go#L89-L96
|
||||
colorCodePanic = "\x1b[1;31m"
|
||||
@@ -48,8 +50,6 @@ class Formatter(logging.Formatter):
|
||||
logstr += f"] {str(record.name)} | {fn} | {str(record.msg)%record.args}"
|
||||
return logstr
|
||||
|
||||
for h in logging.root.handlers:
|
||||
h.setFormatter(Formatter())
|
||||
|
||||
def get_logger(name: str, lv = logging.INFO):
|
||||
logger = logging.getLogger(name)
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
from .ctx import TorchSeedContext
|
||||
@@ -0,0 +1,13 @@
|
||||
import torch
|
||||
|
||||
class TorchSeedContext:
|
||||
def __init__(self, seed):
|
||||
self.seed = seed
|
||||
self.state = None
|
||||
|
||||
def __enter__(self):
|
||||
self.state = torch.random.get_rng_state()
|
||||
torch.manual_seed(self.seed)
|
||||
|
||||
def __exit__(self, type, value, traceback):
|
||||
torch.random.set_rng_state(self.state)
|
||||
Reference in New Issue
Block a user