mirror of
https://github.com/2noise/ChatTTS.git
synced 2026-08-28 17:45:02 +08:00
feat(core): dvae&vocos switch to safetensors
This commit is contained in:
@@ -3,10 +3,10 @@ from dataclasses import dataclass
|
||||
|
||||
@dataclass(repr=False, eq=False)
|
||||
class Path:
|
||||
vocos_ckpt_path: str = "asset/Vocos.pt"
|
||||
dvae_ckpt_path: str = "asset/DVAE_full.pt"
|
||||
vocos_ckpt_path: str = "asset/Vocos.safetensors"
|
||||
dvae_ckpt_path: str = "asset/DVAE.safetensors"
|
||||
gpt_ckpt_path: str = "asset/gpt"
|
||||
decoder_ckpt_path: str = "asset/Decoder.pt"
|
||||
decoder_ckpt_path: str = "asset/Decoder.safetensors"
|
||||
tokenizer_path: str = "asset/tokenizer"
|
||||
embed_path: str = "asset/Embed.safetensors"
|
||||
|
||||
|
||||
+21
-30
@@ -15,6 +15,7 @@ from huggingface_hub import snapshot_download
|
||||
from .config import Config
|
||||
from .model import DVAE, Embed, GPT, gen_logits, Tokenizer, Speaker
|
||||
from .utils import (
|
||||
load_safetensors,
|
||||
check_all_assets,
|
||||
download_all_assets,
|
||||
select_device,
|
||||
@@ -97,7 +98,7 @@ class Chat:
|
||||
try:
|
||||
download_path = snapshot_download(
|
||||
repo_id="2Noise/ChatTTS",
|
||||
allow_patterns=["*.pt", "*.yaml", "*.json", "*.safetensors"],
|
||||
allow_patterns=["*.yaml", "*.json", "*.safetensors"],
|
||||
)
|
||||
except:
|
||||
download_path = None
|
||||
@@ -263,26 +264,22 @@ class Chat:
|
||||
.eval()
|
||||
)
|
||||
assert vocos_ckpt_path, "vocos_ckpt_path should not be None"
|
||||
vocos.load_state_dict(torch.load(vocos_ckpt_path, weights_only=True, mmap=True))
|
||||
vocos.load_state_dict(load_safetensors(vocos_ckpt_path))
|
||||
self.vocos = vocos
|
||||
self.logger.log(logging.INFO, "vocos loaded.")
|
||||
|
||||
dvae = (
|
||||
DVAE(
|
||||
decoder_config=asdict(self.config.dvae.decoder),
|
||||
encoder_config=asdict(self.config.dvae.encoder),
|
||||
vq_config=asdict(self.config.dvae.vq),
|
||||
dim=self.config.dvae.decoder.idim,
|
||||
coef=coef,
|
||||
device=device,
|
||||
)
|
||||
.to(device)
|
||||
.eval()
|
||||
dvae = DVAE(
|
||||
decoder_config=asdict(self.config.dvae.decoder),
|
||||
encoder_config=asdict(self.config.dvae.encoder),
|
||||
vq_config=asdict(self.config.dvae.vq),
|
||||
dim=self.config.dvae.decoder.idim,
|
||||
coef=coef,
|
||||
device=device,
|
||||
)
|
||||
coef = str(dvae)
|
||||
assert dvae_ckpt_path, "dvae_ckpt_path should not be None"
|
||||
dvae.load_state_dict(torch.load(dvae_ckpt_path, weights_only=True, mmap=True))
|
||||
self.dvae = dvae
|
||||
dvae.load_pretrained(dvae_ckpt_path, device)
|
||||
self.dvae = dvae.eval()
|
||||
self.logger.log(logging.INFO, "dvae loaded.")
|
||||
|
||||
embed = Embed(
|
||||
@@ -291,7 +288,7 @@ class Chat:
|
||||
self.config.embed.num_text_tokens,
|
||||
self.config.embed.num_vq,
|
||||
)
|
||||
embed.from_pretrained(embed_path, device=device)
|
||||
embed.load_pretrained(embed_path, device=device)
|
||||
self.embed = embed.to(device)
|
||||
self.logger.log(logging.INFO, "embed loaded.")
|
||||
|
||||
@@ -305,7 +302,7 @@ class Chat:
|
||||
logger=self.logger,
|
||||
).eval()
|
||||
assert gpt_ckpt_path, "gpt_ckpt_path should not be None"
|
||||
gpt.from_pretrained(gpt_ckpt_path, embed_path, experimental=experimental)
|
||||
gpt.load_pretrained(gpt_ckpt_path, embed_path, experimental=experimental)
|
||||
gpt.prepare(compile=compile and "cuda" in str(device))
|
||||
self.gpt = gpt
|
||||
self.logger.log(logging.INFO, "gpt loaded.")
|
||||
@@ -315,22 +312,16 @@ class Chat:
|
||||
)
|
||||
self.logger.log(logging.INFO, "speaker loaded.")
|
||||
|
||||
decoder = (
|
||||
DVAE(
|
||||
decoder_config=asdict(self.config.decoder),
|
||||
dim=self.config.decoder.idim,
|
||||
coef=coef,
|
||||
device=device,
|
||||
)
|
||||
.to(device)
|
||||
.eval()
|
||||
decoder = DVAE(
|
||||
decoder_config=asdict(self.config.decoder),
|
||||
dim=self.config.decoder.idim,
|
||||
coef=coef,
|
||||
device=device,
|
||||
)
|
||||
coef = str(decoder)
|
||||
assert decoder_ckpt_path, "decoder_ckpt_path should not be None"
|
||||
decoder.load_state_dict(
|
||||
torch.load(decoder_ckpt_path, weights_only=True, mmap=True)
|
||||
)
|
||||
self.decoder = decoder
|
||||
decoder.load_pretrained(decoder_ckpt_path, device)
|
||||
self.decoder = decoder.eval()
|
||||
self.logger.log(logging.INFO, "decoder loaded.")
|
||||
|
||||
if tokenizer_path:
|
||||
|
||||
@@ -5,10 +5,11 @@ import numpy as np
|
||||
import pybase16384 as b14
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torchaudio
|
||||
from vector_quantize_pytorch import GroupedResidualFSQ
|
||||
|
||||
from ..utils import load_safetensors
|
||||
|
||||
|
||||
class ConvNeXtBlock(nn.Module):
|
||||
def __init__(
|
||||
@@ -250,6 +251,12 @@ class DVAE(nn.Module):
|
||||
self, inp: torch.Tensor, mode: Literal["encode", "decode"] = "decode"
|
||||
) -> torch.Tensor:
|
||||
return super().__call__(inp, mode)
|
||||
|
||||
@torch.inference_mode()
|
||||
def load_pretrained(self, filename: str, device: torch.device):
|
||||
state_dict_tensors = load_safetensors(filename)
|
||||
self.load_state_dict(state_dict_tensors)
|
||||
self.to(device)
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from safetensors.torch import safe_open
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn.utils.parametrizations import weight_norm
|
||||
|
||||
from ..utils import load_safetensors
|
||||
|
||||
|
||||
class Embed(nn.Module):
|
||||
def __init__(
|
||||
@@ -34,11 +35,8 @@ class Embed(nn.Module):
|
||||
)
|
||||
|
||||
@torch.inference_mode()
|
||||
def from_pretrained(self, filename: str, device: torch.device):
|
||||
state_dict_tensors = {}
|
||||
with safe_open(filename, framework="pt") as f:
|
||||
for k in f.keys():
|
||||
state_dict_tensors[k] = f.get_tensor(k)
|
||||
def load_pretrained(self, filename: str, device: torch.device):
|
||||
state_dict_tensors = load_safetensors(filename)
|
||||
self.load_state_dict(state_dict_tensors)
|
||||
self.to(device)
|
||||
|
||||
|
||||
@@ -56,7 +56,7 @@ class GPT(nn.Module):
|
||||
self.head_text = embed.head_text.__call__
|
||||
self.head_code = [hc.__call__ for hc in embed.head_code]
|
||||
|
||||
def from_pretrained(
|
||||
def load_pretrained(
|
||||
self, gpt_folder: str, embed_file_path: str, experimental=False
|
||||
):
|
||||
if self.is_vllm and platform.system().lower() == "linux":
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
{
|
||||
"sha256_asset_Decoder_pt" : "9964e36e840f0e3a748c5f716fe6de6490d2135a5f5155f4a642d51860e2ec38",
|
||||
"sha256_asset_DVAE_full_pt" : "553eb75763511e23f3e5f86303e2163c5ca775489d637fb635d979c8ae58bbe5",
|
||||
"sha256_asset_Embed_safetensors" : "2ff0be7134934155741b643b74e32fb6bf3eec41257984459b2ed60cdb4c48b0",
|
||||
"sha256_asset_Vocos_pt" : "09a670eda1c08b740013679c7a90ebb7f1a97646ea7673069a6838e6b51d6c58",
|
||||
"sha256_asset_Decoder_safetensors": "77aa55e0a977949c4733df3c6f876fa85860d3298cba63295a7bc6901729d4e0",
|
||||
"sha256_asset_DVAE_safetensors" : "1d0b044a8368c0513100a2eca98456b289e6be6a18b7a63be1bcaa315ea874d9",
|
||||
"sha256_asset_Embed_safetensors" : "2ff0be7134934155741b643b74e32fb6bf3eec41257984459b2ed60cdb4c48b0",
|
||||
"sha256_asset_Vocos_safetensors" : "07e5561491cce41f7f90cfdb94b2ff263ff5742c3d89339db99b17ad82cc3f44",
|
||||
|
||||
"sha256_asset_gpt_config_json" : "0aaa1ecd96c49ad4f473459eb1982fa7ad79fa5de08cde2781bf6ad1f9a0c236",
|
||||
"sha256_asset_gpt_model_safetensors" : "cd0806fd971f52f6a22c923ec64982b305e817bcc41ca83417fcf9141b984a0f",
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from .dl import check_all_assets, download_all_assets
|
||||
from .gpu import select_device
|
||||
from .io import get_latest_modified_file, del_all
|
||||
from .io import load_safetensors, get_latest_modified_file, del_all
|
||||
from .log import logger
|
||||
|
||||
+3
-3
@@ -70,10 +70,10 @@ def check_all_assets(base_dir: Path, sha256_map: Dict[str, str], update=False) -
|
||||
base_dir,
|
||||
"asset",
|
||||
names=(
|
||||
"Decoder.pt",
|
||||
"DVAE_full.pt",
|
||||
"Decoder.safetensors",
|
||||
"DVAE.safetensors",
|
||||
"Embed.safetensors",
|
||||
"Vocos.pt",
|
||||
"Vocos.safetensors",
|
||||
),
|
||||
sha256_map=sha256_map,
|
||||
update=update,
|
||||
|
||||
@@ -3,9 +3,20 @@ import logging
|
||||
from typing import Union
|
||||
from dataclasses import is_dataclass
|
||||
|
||||
from safetensors import safe_open
|
||||
import torch
|
||||
|
||||
from .log import logger
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def load_safetensors(filename: str):
|
||||
state_dict_tensors = {}
|
||||
with safe_open(filename, framework="pt") as f:
|
||||
for k in f.keys():
|
||||
state_dict_tensors[k] = f.get_tensor(k)
|
||||
return state_dict_tensors
|
||||
|
||||
def get_latest_modified_file(directory):
|
||||
|
||||
files = [os.path.join(directory, f) for f in os.listdir(directory)]
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
package main
|
||||
|
||||
var files = [...]string{
|
||||
"asset/Decoder.pt",
|
||||
"asset/DVAE_full.pt",
|
||||
"asset/Decoder.safetensors",
|
||||
"asset/DVAE.safetensors",
|
||||
"asset/Embed.safetensors",
|
||||
"asset/Vocos.pt",
|
||||
"asset/Vocos.safetensors",
|
||||
|
||||
"asset/gpt/config.json",
|
||||
"asset/gpt/model.safetensors",
|
||||
@@ -15,10 +15,10 @@ var files = [...]string{
|
||||
}
|
||||
|
||||
const jsontmpl = `{
|
||||
"sha256_asset_Decoder_pt" : "%s",
|
||||
"sha256_asset_DVAE_full_pt" : "%s",
|
||||
"sha256_asset_Embed_safetensors" : "%s",
|
||||
"sha256_asset_Vocos_pt" : "%s",
|
||||
"sha256_asset_Decoder_safetensors": "%s",
|
||||
"sha256_asset_DVAE_safetensors" : "%s",
|
||||
"sha256_asset_Embed_safetensors" : "%s",
|
||||
"sha256_asset_Vocos_safetensors" : "%s",
|
||||
|
||||
"sha256_asset_gpt_config_json" : "%s",
|
||||
"sha256_asset_gpt_model_safetensors" : "%s",
|
||||
|
||||
Reference in New Issue
Block a user