mirror of
https://github.com/2noise/ChatTTS.git
synced 2026-08-30 17:05:27 +08:00
add sample random spk method
This commit is contained in:
@@ -95,6 +95,9 @@ class Chat:
|
||||
assert gpt_ckpt_path, 'gpt_ckpt_path should not be None'
|
||||
gpt.load_state_dict(torch.load(gpt_ckpt_path, map_location='cpu'))
|
||||
self.pretrain_models['gpt'] = gpt
|
||||
spk_stat_path = os.path.join(os.path.dirname(gpt_ckpt_path), 'spk_stat.pt')
|
||||
assert os.path.exists(spk_stat_path), f'Missing spk_stat.pt: {spk_stat_path}'
|
||||
self.pretrain_models['spk_stat'] = torch.load(spk_stat_path).to(device)
|
||||
self.logger.log(logging.INFO, 'gpt loaded.')
|
||||
|
||||
if decoder_config_path:
|
||||
@@ -144,6 +147,11 @@ class Chat:
|
||||
wav = [self.pretrain_models['vocos'].decode(i).cpu().numpy() for i in mel_spec]
|
||||
|
||||
return wav
|
||||
|
||||
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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user