Inflect-Micro-v2-zh / python /vocoder_utils.py
inoryQwQ's picture
Release 1.0.0: Chinese TTS (Inflect-Micro-v2 arch + BigVGAN), ONNX + AX650 axmodel + checkpoints
0c723b3 verified
Raw
History Blame Contribute Delete
1.23 kB
"""Pretrained mel-conditioned vocoder helpers (NVIDIA BigVGAN)."""
import json
import os
import torch
from bigvgan.env import AttrDict
from bigvgan.bigvgan import BigVGAN
VOC_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "vocoder")
BUILTIN = {
"bigvgan_base": ("bigvgan_base_24k_100band.json",
"bigvgan_base_24k_100band.pt"),
"bigvgan_v2": ("bigvgan_v2_24k_100band.json",
"bigvgan_v2_24k_100band.pt"),
}
_cache = {}
def load_bigvgan(name="bigvgan_base", device="cpu"):
"""Load a pretrained BigVGAN generator (frozen, eval mode)."""
if name in _cache:
return _cache[name]
if name not in BUILTIN:
raise ValueError(f"Unknown vocoder '{name}', choose from {list(BUILTIN)}")
cfg_file, ckpt_file = BUILTIN[name]
h = AttrDict(json.load(open(os.path.join(VOC_DIR, cfg_file))))
model = BigVGAN(h)
ckpt = torch.load(os.path.join(VOC_DIR, ckpt_file), map_location="cpu",
weights_only=True)
model.load_state_dict(ckpt["generator"])
model.remove_weight_norm()
model.eval().to(device)
for p in model.parameters():
p.requires_grad = False
_cache[name] = model
return model