This commit is contained in:
AIFSH
2024-05-13 06:55:17 +00:00
parent 0855ceaab0
commit 6d6cd7c0b0
28 changed files with 3398 additions and 109 deletions
@@ -5,14 +5,14 @@ defaults:
project: text2semantic_finetune_dual_ar
max_length: 2048
ckpt_path: checkpoints/text2semantic-sft-medium-v1-4k.pth
ckpt_path: checkpoints/text2semantic-sft-medium-v1.1-4k.pth
resume_weights_only: true
# Lightning Trainer
trainer:
accumulate_grad_batches: 1
gradient_clip_val: 1.0
gradient_clip_algorithm: 'norm'
gradient_clip_algorithm: "norm"
max_steps: 1000
precision: bf16-true
limit_val_batches: 10
@@ -0,0 +1,128 @@
defaults:
- base
- _self_
project: vits_decoder
ckpt_path: checkpoints/vits_decoder_v1.1.ckpt
resume_weights_only: true
# Lightning Trainer
trainer:
accelerator: gpu
devices: auto
strategy:
find_unused_parameters: true
precision: 32
max_steps: 100_000
val_check_interval: 100
benchmark: false
sample_rate: 44100
hop_length: 512
num_mels: 128
n_fft: 2048
win_length: 2048
# Dataset Configuration
tokenizer:
_target_: transformers.AutoTokenizer.from_pretrained
pretrained_model_name_or_path: fishaudio/fish-speech-1
# Dataset Configuration
train_dataset:
_target_: fish_speech.datasets.vits.VITSDataset
filelist: data/source/Genshin/filelist.train.txt
sample_rate: ${sample_rate}
hop_length: ${hop_length}
suffix: ".lab"
tokenizer: ${tokenizer}
sentence_mask_ratio: 0.2
val_dataset:
_target_: fish_speech.datasets.vits.VITSDataset
filelist: data/source/Genshin/filelist.test.txt
sample_rate: ${sample_rate}
hop_length: ${hop_length}
suffix: ".lab"
tokenizer: ${tokenizer}
data:
_target_: fish_speech.datasets.vits.VITSDataModule
train_dataset: ${train_dataset}
val_dataset: ${val_dataset}
num_workers: 4
batch_size: 8
val_batch_size: 4
tokenizer: ${tokenizer}
# Model Configuration
model:
_target_: fish_speech.models.vits_decoder.VITSDecoder
sample_rate: ${sample_rate}
hop_length: ${hop_length}
freeze_discriminator: false
weight_mel: 45.0
weight_kl: 1.0
generator:
_target_: fish_speech.models.vits_decoder.modules.models.SynthesizerTrn
spec_channels: 1025
segment_size: 32
inter_channels: 192
hidden_channels: 192
filter_channels: 768
n_heads: 2
n_layers: 6
kernel_size: 3
p_dropout: 0.1
resblock: "1"
resblock_kernel_sizes: [3, 7, 11]
resblock_dilation_sizes: [[1, 3, 5], [1, 3, 5], [1, 3, 5]]
upsample_rates: [8, 8, 2, 2, 2]
upsample_initial_channel: 512
upsample_kernel_sizes: [16, 16, 8, 2, 2]
gin_channels: 512
vq_mask_ratio: 0.2
ref_mask_ratio: 0.2
discriminator:
_target_: fish_speech.models.vits_decoder.modules.models.EnsembledDiscriminator
periods: [2, 3, 5, 7, 11]
mel_transform:
_target_: fish_speech.utils.spectrogram.LogMelSpectrogram
sample_rate: ${sample_rate}
n_fft: ${n_fft}
hop_length: ${hop_length}
win_length: ${win_length}
n_mels: ${num_mels}
spec_transform:
_target_: fish_speech.utils.spectrogram.LinearSpectrogram
n_fft: ${n_fft}
hop_length: ${hop_length}
win_length: ${win_length}
mode: pow2_sqrt
optimizer:
_target_: torch.optim.AdamW
_partial_: true
lr: 1e-4
betas: [0.8, 0.99]
eps: 1e-6
lr_scheduler:
_target_: torch.optim.lr_scheduler.ExponentialLR
_partial_: true
gamma: 0.999999
callbacks:
grad_norm_monitor:
sub_module:
- generator
- discriminator
model_checkpoint:
every_n_train_steps: ${trainer.val_check_interval}
save_top_k: 10
@@ -0,0 +1,127 @@
defaults:
- base
- _self_
project: vits_decoder
ckpt_path: checkpoints/Bert-VITS2/ensemble.pth
resume_weights_only: true
# Lightning Trainer
trainer:
accelerator: gpu
devices: auto
strategy: ddp_find_unused_parameters_true
precision: 32
max_steps: 1_000_000
val_check_interval: 1000
benchmark: false
sample_rate: 44100
hop_length: 512
num_mels: 128
n_fft: 2048
win_length: 2048
# Dataset Configuration
tokenizer:
_target_: transformers.AutoTokenizer.from_pretrained
pretrained_model_name_or_path: fishaudio/fish-speech-1
# Dataset Configuration
train_dataset:
_target_: fish_speech.datasets.vits.VITSDataset
filelist: data/source/Genshin/filelist.train.txt
sample_rate: ${sample_rate}
hop_length: ${hop_length}
suffix: ".lab"
tokenizer: ${tokenizer}
sentence_mask_ratio: 0.2
val_dataset:
_target_: fish_speech.datasets.vits.VITSDataset
filelist: data/source/Genshin/filelist.test.txt
sample_rate: ${sample_rate}
hop_length: ${hop_length}
suffix: ".lab"
tokenizer: ${tokenizer}
data:
_target_: fish_speech.datasets.vits.VITSDataModule
train_dataset: ${train_dataset}
val_dataset: ${val_dataset}
num_workers: 4
batch_size: 8
val_batch_size: 4
tokenizer: ${tokenizer}
# Model Configuration
model:
_target_: fish_speech.models.vits_decoder.VITSDecoder
sample_rate: ${sample_rate}
hop_length: ${hop_length}
freeze_discriminator: false
weight_mel: 45.0
weight_kl: 1.0
generator:
_target_: fish_speech.models.vits_decoder.modules.models.SynthesizerTrn
spec_channels: 1025
segment_size: 32
inter_channels: 192
hidden_channels: 192
filter_channels: 768
n_heads: 2
n_layers: 6
kernel_size: 3
p_dropout: 0.1
resblock: "1"
resblock_kernel_sizes: [3, 7, 11]
resblock_dilation_sizes: [[1, 3, 5], [1, 3, 5], [1, 3, 5]]
upsample_rates: [8, 8, 2, 2, 2]
upsample_initial_channel: 512
upsample_kernel_sizes: [16, 16, 8, 2, 2]
gin_channels: 512
vq_mask_ratio: 0.2
ref_mask_ratio: 0.2
discriminator:
_target_: fish_speech.models.vits_decoder.modules.models.EnsembledDiscriminator
periods: [2, 3, 5, 7, 11]
mel_transform:
_target_: fish_speech.utils.spectrogram.LogMelSpectrogram
sample_rate: ${sample_rate}
n_fft: ${n_fft}
hop_length: ${hop_length}
win_length: ${win_length}
n_mels: ${num_mels}
spec_transform:
_target_: fish_speech.utils.spectrogram.LinearSpectrogram
n_fft: ${n_fft}
hop_length: ${hop_length}
win_length: ${win_length}
mode: pow2_sqrt
optimizer:
_target_: torch.optim.AdamW
_partial_: true
lr: 1e-4
betas: [0.8, 0.99]
eps: 1e-6
lr_scheduler:
_target_: torch.optim.lr_scheduler.ExponentialLR
_partial_: true
gamma: 0.999999
callbacks:
grad_norm_monitor:
sub_module:
- generator
- discriminator
model_checkpoint:
every_n_train_steps: 1000
save_top_k: 10
+2 -2
View File
@@ -86,7 +86,7 @@ model:
ckpt_path: null # You may download the pretrained vocoder and set the path here
encode_mel_transform:
_target_: fish_speech.models.vqgan.spectrogram.LogMelSpectrogram
_target_: fish_speech.utils.spectrogram.LogMelSpectrogram
sample_rate: ${sample_rate}
n_fft: ${n_fft}
hop_length: ${hop_length}
@@ -96,7 +96,7 @@ model:
f_max: 8000.0
gt_mel_transform:
_target_: fish_speech.models.vqgan.spectrogram.LogMelSpectrogram
_target_: fish_speech.utils.spectrogram.LogMelSpectrogram
sample_rate: ${sample_rate}
n_fft: ${n_fft}
hop_length: ${hop_length}
+2 -2
View File
@@ -89,7 +89,7 @@ model:
ckpt_path: null # You may download the pretrained vocoder and set the path here
encode_mel_transform:
_target_: fish_speech.models.vqgan.spectrogram.LogMelSpectrogram
_target_: fish_speech.utils.spectrogram.LogMelSpectrogram
sample_rate: ${sample_rate}
n_fft: ${n_fft}
hop_length: ${hop_length}
@@ -99,7 +99,7 @@ model:
f_max: 8000.0
gt_mel_transform:
_target_: fish_speech.models.vqgan.spectrogram.LogMelSpectrogram
_target_: fish_speech.utils.spectrogram.LogMelSpectrogram
sample_rate: ${sample_rate}
n_fft: ${n_fft}
hop_length: ${hop_length}
+52
View File
@@ -0,0 +1,52 @@
import bisect
from typing import Iterable
from torch.utils.data import Dataset, IterableDataset
class ConcatRepeatDataset(Dataset):
datasets: list[Dataset]
cumulative_sizes: list[int]
repeats: list[int]
@staticmethod
def cumsum(sequence, repeats):
r, s = [], 0
for dataset, repeat in zip(sequence, repeats):
l = len(dataset) * repeat
r.append(l + s)
s += l
return r
def __init__(self, datasets: Iterable[Dataset], repeats: list[int]):
super().__init__()
self.datasets = list(datasets)
self.repeats = repeats
assert len(self.datasets) > 0, "datasets should not be an empty iterable"
assert len(self.datasets) == len(
repeats
), "datasets and repeats should have the same length"
for d in self.datasets:
assert not isinstance(
d, IterableDataset
), "ConcatDataset does not support IterableDataset"
self.cumulative_sizes = self.cumsum(self.datasets, self.repeats)
def __len__(self):
return self.cumulative_sizes[-1]
def __getitem__(self, idx):
dataset_idx = bisect.bisect_right(self.cumulative_sizes, idx)
if dataset_idx == 0:
sample_idx = idx
else:
sample_idx = idx - self.cumulative_sizes[dataset_idx - 1]
dataset = self.datasets[dataset_idx]
return dataset[sample_idx % len(dataset)]
+195
View File
@@ -0,0 +1,195 @@
import random
from dataclasses import dataclass
from pathlib import Path
from typing import Optional
import librosa
import numpy as np
import torch
import torch.distributed as dist
from lightning import LightningDataModule
from torch.utils.data import DataLoader, Dataset
from torch.utils.data.distributed import DistributedSampler
from transformers import AutoTokenizer
from fish_speech.utils import RankedLogger
logger = RankedLogger(__name__, rank_zero_only=False)
class VITSDataset(Dataset):
def __init__(
self,
filelist: str,
tokenizer: AutoTokenizer,
sample_rate: int = 44100,
hop_length: int = 512,
min_duration: float = 1.5,
max_duration: float = 30.0,
suffix: str = ".lab",
sentence_mask_ratio: float = 0.0,
):
super().__init__()
filelist = Path(filelist)
root = filelist.parent
self.files = []
for line in filelist.read_text(encoding="utf-8").splitlines():
path = root / line
self.files.append(path)
self.sample_rate = sample_rate
self.hop_length = hop_length
self.min_duration = min_duration
self.max_duration = max_duration
self.tokenizer = tokenizer
self.suffix = suffix
self.sentence_mask_ratio = sentence_mask_ratio
def __len__(self):
return len(self.files)
def get_item(self, idx):
audio_file = self.files[idx]
text_file = audio_file.with_suffix(self.suffix)
if text_file.exists() is False or audio_file.exists() is False:
return None
audio, _ = librosa.load(audio_file, sr=self.sample_rate, mono=True)
duration = len(audio) / self.sample_rate
# Pad to minimum duration
if duration < self.min_duration:
pad_duration = self.min_duration - duration
pad_samples = int(pad_duration * self.sample_rate)
audio = np.pad(audio, (0, pad_samples))
# Truncate to maximum duration
if duration > self.max_duration:
random_start = random.randint(
0, len(audio) - int(self.max_duration * self.sample_rate) - 1
)
audio = audio[
random_start : random_start + int(self.max_duration * self.sample_rate)
]
max_value = np.abs(audio).max()
if max_value > 1.0:
audio = audio / max_value
if random.random() < self.sentence_mask_ratio:
text = "-"
else:
text = text_file.read_text(encoding="utf-8")
input_ids = self.tokenizer(text, return_tensors="pt").input_ids.squeeze(0)
return {
"audio": torch.from_numpy(audio),
"text": input_ids,
}
def __getitem__(self, idx):
try:
return self.get_item(idx)
except Exception as e:
import traceback
traceback.print_exc()
logger.error(f"Error loading {self.files[idx]}: {e}")
return None
@dataclass
class VITSCollator:
tokenizer: AutoTokenizer
def __call__(self, batch):
batch = [x for x in batch if x is not None]
audio_lengths = torch.tensor([len(x["audio"]) for x in batch])
audio_maxlen = audio_lengths.max()
text_lengths = torch.tensor([len(x["text"]) for x in batch])
text_maxlen = text_lengths.max()
# Rounds up to nearest multiple of 2 (audio_lengths)
audios = []
texts = []
for x in batch:
audios.append(
torch.nn.functional.pad(x["audio"], (0, audio_maxlen - len(x["audio"])))
)
texts.append(
torch.nn.functional.pad(
x["text"],
(0, text_maxlen - len(x["text"])),
value=self.tokenizer.eos_token_id,
)
)
return {
"audios": torch.stack(audios),
"audio_lengths": audio_lengths,
"texts": torch.stack(texts),
"text_lengths": text_lengths,
}
class VITSDataModule(LightningDataModule):
def __init__(
self,
train_dataset: VITSDataset,
val_dataset: VITSDataset,
tokenizer: AutoTokenizer,
batch_size: int = 32,
num_workers: int = 4,
val_batch_size: Optional[int] = None,
):
super().__init__()
self.train_dataset = train_dataset
self.val_dataset = val_dataset
self.batch_size = batch_size
self.val_batch_size = val_batch_size or batch_size
self.num_workers = num_workers
self.tokenizer = tokenizer
def train_dataloader(self):
return DataLoader(
self.train_dataset,
batch_size=self.batch_size,
collate_fn=VITSCollator(self.tokenizer),
num_workers=self.num_workers,
shuffle=False,
persistent_workers=True,
)
def val_dataloader(self):
return DataLoader(
self.val_dataset,
batch_size=self.val_batch_size,
collate_fn=VITSCollator(self.tokenizer),
num_workers=self.num_workers,
persistent_workers=True,
)
if __name__ == "__main__":
tokenizer = AutoTokenizer.from_pretrained("fishaudio/fish-speech-1")
dataset = VITSDataset(
"data/source/Genshin/filelist.train.txt", tokenizer=tokenizer, suffix=".lab"
)
dataloader = DataLoader(
dataset, batch_size=4, shuffle=False, collate_fn=VITSCollator(tokenizer)
)
for batch in dataloader:
print(batch["audios"].shape)
print(batch["audio_lengths"])
print(batch["texts"].shape)
print(batch["text_lengths"])
break
+13 -2
View File
@@ -14,6 +14,8 @@
"Data Preprocessing": "Data Preprocessing",
"Data Preprocessing Path": "Data Preprocessing Path",
"Data Source": "Data Source",
"Decoder Model Config": "Decoder Model Config",
"Decoder Model Path": "Decoder Model Path",
"Disabled": "Disabled",
"Enable Reference Audio": "Enable Reference Audio",
"English": "English",
@@ -39,16 +41,19 @@
"LLAMA Model Path": "LLAMA Model Path",
"Labeling Device": "Labeling Device",
"LoRA Model to be merged": "LoRA Model to be merged",
"Maximum Audio Duration": "Maximum Audio Duration",
"Maximum Length per Sample": "Maximum Length per Sample",
"Maximum Training Steps": "Maximum Training Steps",
"Maximum tokens per batch, 0 means no limit": "Maximum tokens per batch, 0 means no limit",
"Merge": "Merge",
"Merge LoRA": "Merge LoRA",
"Merge successfully": "Merge successfully",
"Minimum Audio Duration": "Minimum Audio Duration",
"Model Output Path": "Model Output Path",
"Model Size": "Model Size",
"Move": "Move",
"Move files successfully": "Move files successfully",
"No audio generated, please check the input text.": "No audio generated, please check the input text.",
"No selected options": "No selected options",
"Number of Workers": "Number of Workers",
"Open Inference Server": "Open Inference Server",
@@ -56,6 +61,7 @@
"Open Tensorboard": "Open Tensorboard",
"Opened labeler in browser": "Opened labeler in browser",
"Optional Label Language": "Optional Label Language",
"Optional online ver": "Optional online ver",
"Output Path": "Output Path",
"Path error, please check the model file exists in the corresponding path": "Path error, please check the model file exists in the corresponding path",
"Precision": "Precision",
@@ -68,6 +74,9 @@
"Removed path successfully!": "Removed path successfully!",
"Repetition Penalty": "Repetition Penalty",
"Save model every n steps": "Save model every n steps",
"Select LLAMA ckpt": "Select LLAMA ckpt",
"Select VITS ckpt": "Select VITS ckpt",
"Select VQGAN ckpt": "Select VQGAN ckpt",
"Select source file processing method": "Select source file processing method",
"Select the model to be trained": "Select the model to be trained",
"Selected: {}": "Selected: {}",
@@ -92,8 +101,8 @@
"Use LoRA can save GPU memory, but may reduce the quality of the model": "Use LoRA can save GPU memory, but may reduce the quality of the model",
"Use filelist": "Use filelist",
"Use large for 10G+ GPU, medium for 5G, small for 2G": "Use large for 10G+ GPU, medium for 5G, small for 2G",
"VITS Configuration": "VITS Configuration",
"VQGAN Configuration": "VQGAN Configuration",
"VQGAN Model Path": "VQGAN Model Path",
"Validation Batch Size": "Validation Batch Size",
"View the status of the preprocessing folder (use the slider to control the depth of the tree)": "View the status of the preprocessing folder (use the slider to control the depth of the tree)",
"We are not responsible for any misuse of the model, please consider your local laws and regulations before using it.": "We are not responsible for any misuse of the model, please consider your local laws and regulations before using it.",
@@ -101,5 +110,7 @@
"WebUI Port": "WebUI Port",
"Whisper Model": "Whisper Model",
"You can find the source code [here](https://github.com/fishaudio/fish-speech) and models [here](https://huggingface.co/fishaudio/fish-speech-1).": "You can find the source code [here](https://github.com/fishaudio/fish-speech) and models [here](https://huggingface.co/fishaudio/fish-speech-1).",
"bf16-true is recommended for 30+ series GPU, 16-mixed is recommended for 10+ series GPU": "bf16-true is recommended for 30+ series GPU, 16-mixed is recommended for 10+ series GPU"
"bf16-true is recommended for 30+ series GPU, 16-mixed is recommended for 10+ series GPU": "bf16-true is recommended for 30+ series GPU, 16-mixed is recommended for 10+ series GPU",
"latest": "latest",
"new": "new"
}
+13 -2
View File
@@ -14,6 +14,8 @@
"Data Preprocessing": "Preprocesamiento de Datos",
"Data Preprocessing Path": "Ruta de Preprocesamiento de Datos",
"Data Source": "Fuente de Datos",
"Decoder Model Config": "Configuración del modelo decodificador",
"Decoder Model Path": "Ruta del modelo decodificador",
"Disabled": "Desactivado",
"Enable Reference Audio": "Habilitar Audio de Referencia",
"English": "Inglés",
@@ -39,16 +41,19 @@
"LLAMA Model Path": "Ruta del Modelo LLAMA",
"Labeling Device": "Dispositivo de Etiquetado",
"LoRA Model to be merged": "Modelo LoRA a fusionar",
"Maximum Audio Duration": "Duración máxima de audio",
"Maximum Length per Sample": "Longitud Máxima por Muestra",
"Maximum Training Steps": "Pasos Máximos de Entrenamiento",
"Maximum tokens per batch, 0 means no limit": "Máximo de tokens por lote, 0 significa sin límite",
"Merge": "Fusionar",
"Merge LoRA": "Fusionar LoRA",
"Merge successfully": "Fusionado exitosamente",
"Minimum Audio Duration": "Duración mínima de audio",
"Model Output Path": "Ruta de Salida del Modelo",
"Model Size": "Tamaño del Modelo",
"Move": "Mover",
"Move files successfully": "Archivos movidos exitosamente",
"No audio generated, please check the input text.": "No se generó audio, por favor verifique el texto de entrada.",
"No selected options": "No hay opciones seleccionadas",
"Number of Workers": "Número de Trabajadores",
"Open Inference Server": "Abrir Servidor de Inferencia",
@@ -56,6 +61,7 @@
"Open Tensorboard": "Abrir Tensorboard",
"Opened labeler in browser": "Se abrió el etiquetador en el navegador",
"Optional Label Language": "Idioma de Etiquetado Opcional",
"Optional online ver": "Ver en línea opcional",
"Output Path": "Ruta de Salida",
"Path error, please check the model file exists in the corresponding path": "Error de ruta, por favor verifique que el archivo del modelo exista en la ruta correspondiente",
"Precision": "Precisión",
@@ -68,6 +74,9 @@
"Removed path successfully!": "¡Ruta eliminada exitosamente!",
"Repetition Penalty": "Penalización por Repetición",
"Save model every n steps": "Guardar modelo cada n pasos",
"Select LLAMA ckpt": "Seleccionar punto de control LLAMA",
"Select VITS ckpt": "Seleccionar punto de control VITS",
"Select VQGAN ckpt": "Seleccionar punto de control VQGAN",
"Select source file processing method": "Seleccione el método de procesamiento de archivos fuente",
"Select the model to be trained": "Seleccione el modelo a ser entrenado",
"Selected: {}": "Seleccionado: {}",
@@ -92,8 +101,8 @@
"Use LoRA can save GPU memory, but may reduce the quality of the model": "Usar LoRA puede ahorrar memoria GPU, pero puede reducir la calidad del modelo",
"Use filelist": "Usar lista de archivos",
"Use large for 10G+ GPU, medium for 5G, small for 2G": "Use grande para GPU de 10G+, mediano para 5G, pequeño para 2G",
"VITS Configuration": "Configuración de VITS",
"VQGAN Configuration": "Configuración de VQGAN",
"VQGAN Model Path": "Ruta del Modelo VQGAN",
"Validation Batch Size": "Tamaño del Lote de Validación",
"View the status of the preprocessing folder (use the slider to control the depth of the tree)": "Vea el estado de la carpeta de preprocesamiento (use el control deslizante para controlar la profundidad del árbol)",
"We are not responsible for any misuse of the model, please consider your local laws and regulations before using it.": "No somos responsables de ningún mal uso del modelo, por favor considere sus leyes y regulaciones locales antes de usarlo.",
@@ -101,5 +110,7 @@
"WebUI Port": "Puerto de WebUI",
"Whisper Model": "Modelo Whisper",
"You can find the source code [here](https://github.com/fishaudio/fish-speech) and models [here](https://huggingface.co/fishaudio/fish-speech-1).": "Puede encontrar el código fuente [aquí](https://github.com/fishaudio/fish-speech) y los modelos [aquí](https://huggingface.co/fishaudio/fish-speech-1).",
"bf16-true is recommended for 30+ series GPU, 16-mixed is recommended for 10+ series GPU": "Se recomienda bf16-true para GPU de la serie 30+, se recomienda 16-mixed para GPU de la serie 10+"
"bf16-true is recommended for 30+ series GPU, 16-mixed is recommended for 10+ series GPU": "Se recomienda bf16-true para GPU de la serie 30+, se recomienda 16-mixed para GPU de la serie 10+",
"latest": "más reciente",
"new": "nuevo"
}
+14 -3
View File
@@ -14,6 +14,8 @@
"Data Preprocessing": "データ前処理",
"Data Preprocessing Path": "データ前処理パス",
"Data Source": "データソース",
"Decoder Model Config": "デコーダーモデルの構成",
"Decoder Model Path": "デコーダーモデルのパス",
"Disabled": "無効",
"Enable Reference Audio": "リファレンスオーディオを有効にする",
"English": "英語",
@@ -39,16 +41,19 @@
"LLAMA Model Path": "LLAMAモデルパス",
"Labeling Device": "ラベリングデバイス",
"LoRA Model to be merged": "マージするLoRAモデル",
"Maximum Audio Duration": "最大オーディオの長さ",
"Maximum Length per Sample": "サンプルあたりの最大長",
"Maximum Training Steps": "最大トレーニングステップ数",
"Maximum tokens per batch, 0 means no limit": "バッチあたりの最大トークン数。0は制限なしを意味します",
"Merge": "マージ",
"Merge LoRA": "LoRAのマージ",
"Merge successfully": "マージに成功しました",
"Minimum Audio Duration": "最小オーディオの長さ",
"Model Output Path": "モデル出力パス",
"Model Size": "モデルサイズ",
"Move": "移動",
"Move files successfully": "ファイルの移動に成功しました",
"No audio generated, please check the input text.": "オーディオが生成されていません。入力テキストを確認してください。",
"No selected options": "選択されたオプションはありません",
"Number of Workers": "ワーカー数",
"Open Inference Server": "推論サーバーを開く",
@@ -56,6 +61,7 @@
"Open Tensorboard": "Tensorboardを開く",
"Opened labeler in browser": "ブラウザでラベラーを開きました",
"Optional Label Language": "オプションのラベル言語",
"Optional online ver": "オプションのオンラインバージョン",
"Output Path": "出力パス",
"Path error, please check the model file exists in the corresponding path": "パスエラー。対応するパスにモデルファイルが存在するか確認してください",
"Precision": "精度",
@@ -68,6 +74,9 @@
"Removed path successfully!": "パスの削除に成功しました!",
"Repetition Penalty": "反復ペナルティ",
"Save model every n steps": "nステップごとにモデルを保存",
"Select LLAMA ckpt": " LLAMA チェックポイントを選択",
"Select VITS ckpt": "VITS チェックポイントを選択",
"Select VQGAN ckpt": "VQGAN チェックポイントを選択",
"Select source file processing method": "ソースファイルの処理方法を選択",
"Select the model to be trained": "トレーニングするモデルを選択",
"Selected: {}": "選択済み: {}",
@@ -92,8 +101,8 @@
"Use LoRA can save GPU memory, but may reduce the quality of the model": "LoRAを使用するとGPUメモリを節約できますが、モデルの品質が低下する可能性があります",
"Use filelist": "ファイルリストを使用",
"Use large for 10G+ GPU, medium for 5G, small for 2G": "10G以上のGPUには大、5Gには中、2Gには小を使用してください",
"VQGAN Configuration": "VQGAN設定",
"VQGAN Model Path": "VQGANモデルパス",
"VITS Configuration": "VITS の構成",
"VQGAN Configuration": "VQGAN の構成",
"Validation Batch Size": "検証バッチサイズ",
"View the status of the preprocessing folder (use the slider to control the depth of the tree)": "前処理フォルダの状態を表示(スライダーを使用してツリーの深さを制御)",
"We are not responsible for any misuse of the model, please consider your local laws and regulations before using it.": "モデルの誤用については一切責任を負いません。使用する前に、現地の法律と規制を考慮してください。",
@@ -101,5 +110,7 @@
"WebUI Port": "WebUIポート",
"Whisper Model": "Whisperモデル",
"You can find the source code [here](https://github.com/fishaudio/fish-speech) and models [here](https://huggingface.co/fishaudio/fish-speech-1).": "ソースコードは[こちら](https://github.com/fishaudio/fish-speech)、モデルは[こちら](https://huggingface.co/fishaudio/fish-speech-1)にあります。",
"bf16-true is recommended for 30+ series GPU, 16-mixed is recommended for 10+ series GPU": "30シリーズ以降のGPUにはbf16-trueを、10シリーズ以降のGPUには16-mixedをお勧めします"
"bf16-true is recommended for 30+ series GPU, 16-mixed is recommended for 10+ series GPU": "30シリーズ以降のGPUにはbf16-trueを、10シリーズ以降のGPUには16-mixedをお勧めします",
"latest": "最新",
"new": "新規"
}
+13 -2
View File
@@ -14,6 +14,8 @@
"Data Preprocessing": "数据预处理",
"Data Preprocessing Path": "数据预处理路径",
"Data Source": "数据源",
"Decoder Model Config": "解码器模型配置",
"Decoder Model Path": "解码器模型路径",
"Disabled": "禁用",
"Enable Reference Audio": "启用参考音频",
"English": "英文",
@@ -39,16 +41,19 @@
"LLAMA Model Path": "LLAMA 模型路径",
"Labeling Device": "标注加速设备",
"LoRA Model to be merged": "要合并的 LoRA 模型",
"Maximum Audio Duration": "最大音频时长",
"Maximum Length per Sample": "每个样本的最大长度",
"Maximum Training Steps": "最大训练步数",
"Maximum tokens per batch, 0 means no limit": "每批最大令牌数,0 表示无限制",
"Merge": "合并",
"Merge LoRA": "合并 LoRA",
"Merge successfully": "合并成功",
"Minimum Audio Duration": "最小音频时长",
"Model Output Path": "模型输出路径",
"Model Size": "模型规模",
"Move": "移动",
"Move files successfully": "移动文件成功",
"No audio generated, please check the input text.": "没有生成音频,请检查输入文本.",
"No selected options": "没有选择的选项",
"Number of Workers": "数据加载进程数",
"Open Inference Server": "打开推理服务器",
@@ -56,6 +61,7 @@
"Open Tensorboard": "打开 Tensorboard",
"Opened labeler in browser": "在浏览器中打开标注工具",
"Optional Label Language": "[可选] 标注语言",
"Optional online ver": "[可选] 使用在线版",
"Output Path": "输出路径",
"Path error, please check the model file exists in the corresponding path": "路径错误,请检查模型文件是否存在于相应路径",
"Precision": "精度",
@@ -68,6 +74,9 @@
"Removed path successfully!": "移除路径成功!",
"Repetition Penalty": "重复惩罚",
"Save model every n steps": "每 n 步保存模型",
"Select LLAMA ckpt": "选择 LLAMA 检查点",
"Select VITS ckpt": "选择 VITS 检查点",
"Select VQGAN ckpt": "选择 VQGAN 检查点",
"Select source file processing method": "选择源文件处理方法",
"Select the model to be trained": "选择要训练的模型",
"Selected: {}": "已选择: {}",
@@ -92,8 +101,8 @@
"Use LoRA can save GPU memory, but may reduce the quality of the model": "使用 LoRA 可以节省 GPU 内存,但可能会降低模型质量",
"Use filelist": "使用文件列表",
"Use large for 10G+ GPU, medium for 5G, small for 2G": "10G+ GPU 使用 large, 5G 使用 medium, 2G 使用 small",
"VITS Configuration": "VITS 配置",
"VQGAN Configuration": "VQGAN 配置",
"VQGAN Model Path": "VQGAN 模型路径",
"Validation Batch Size": "验证批次大小",
"View the status of the preprocessing folder (use the slider to control the depth of the tree)": "查看预处理文件夹的状态 (使用滑块控制树的深度)",
"We are not responsible for any misuse of the model, please consider your local laws and regulations before using it.": "我们不对模型的任何滥用负责,请在使用之前考虑您当地的法律法规.",
@@ -101,5 +110,7 @@
"WebUI Port": "WebUI 端口",
"Whisper Model": "Whisper 模型",
"You can find the source code [here](https://github.com/fishaudio/fish-speech) and models [here](https://huggingface.co/fishaudio/fish-speech-1).": "你可以在 [这里](https://github.com/fishaudio/fish-speech) 找到源代码和 [这里](https://huggingface.co/fishaudio/fish-speech-1) 找到模型.",
"bf16-true is recommended for 30+ series GPU, 16-mixed is recommended for 10+ series GPU": "30+ 系列 GPU 建议使用 bf16-true, 10+ 系列 GPU 建议使用 16-mixed"
"bf16-true is recommended for 30+ series GPU, 16-mixed is recommended for 10+ series GPU": "30+ 系列 GPU 建议使用 bf16-true, 10+ 系列 GPU 建议使用 16-mixed",
"latest": "最近的检查点",
"new": "创建新的检查点"
}
@@ -0,0 +1,3 @@
from .lit_module import VITSDecoder
__all__ = ["VITSDecoder"]
@@ -0,0 +1,394 @@
from typing import Any, Callable
import lightning as L
import torch
import torch.nn.functional as F
import wandb
from lightning.pytorch.loggers import TensorBoardLogger, WandbLogger
from matplotlib import pyplot as plt
from torch import nn
from fish_speech.models.vits_decoder.losses import (
discriminator_loss,
feature_loss,
generator_loss,
kl_loss,
)
from fish_speech.models.vqgan.utils import (
avg_with_mask,
plot_mel,
sequence_mask,
slice_segments,
)
class VITSDecoder(L.LightningModule):
def __init__(
self,
optimizer: Callable,
lr_scheduler: Callable,
generator: nn.Module,
discriminator: nn.Module,
mel_transform: nn.Module,
spec_transform: nn.Module,
hop_length: int = 512,
sample_rate: int = 44100,
freeze_discriminator: bool = False,
weight_mel: float = 45,
weight_kl: float = 0.1,
):
super().__init__()
# Model parameters
self.optimizer_builder = optimizer
self.lr_scheduler_builder = lr_scheduler
# Generator and discriminator
self.generator = generator
self.discriminator = discriminator
self.mel_transform = mel_transform
self.spec_transform = spec_transform
self.freeze_discriminator = freeze_discriminator
# Loss weights
self.weight_mel = weight_mel
self.weight_kl = weight_kl
# Other parameters
self.hop_length = hop_length
self.sampling_rate = sample_rate
# Disable automatic optimization
self.automatic_optimization = False
if self.freeze_discriminator:
for p in self.discriminator.parameters():
p.requires_grad = False
def configure_optimizers(self):
# Need two optimizers and two schedulers
optimizer_generator = self.optimizer_builder(self.generator.parameters())
optimizer_discriminator = self.optimizer_builder(
self.discriminator.parameters()
)
lr_scheduler_generator = self.lr_scheduler_builder(optimizer_generator)
lr_scheduler_discriminator = self.lr_scheduler_builder(optimizer_discriminator)
return (
{
"optimizer": optimizer_generator,
"lr_scheduler": {
"scheduler": lr_scheduler_generator,
"interval": "step",
"name": "optimizer/generator",
},
},
{
"optimizer": optimizer_discriminator,
"lr_scheduler": {
"scheduler": lr_scheduler_discriminator,
"interval": "step",
"name": "optimizer/discriminator",
},
},
)
def training_step(self, batch, batch_idx):
optim_g, optim_d = self.optimizers()
audios, audio_lengths = batch["audios"], batch["audio_lengths"]
texts, text_lengths = batch["texts"], batch["text_lengths"]
audios = audios.float()
audios = audios[:, None, :]
with torch.no_grad():
gt_mels = self.mel_transform(audios)
gt_specs = self.spec_transform(audios)
spec_lengths = audio_lengths // self.hop_length
spec_masks = torch.unsqueeze(
sequence_mask(spec_lengths, gt_mels.shape[2]), 1
).to(gt_mels.dtype)
(
fake_audios,
ids_slice,
y_mask,
(z, z_p, m_p, logs_p, m_q, logs_q),
) = self.generator(
audios,
audio_lengths,
gt_specs,
spec_lengths,
texts,
text_lengths,
)
gt_mels = slice_segments(gt_mels, ids_slice, self.generator.segment_size)
spec_masks = slice_segments(spec_masks, ids_slice, self.generator.segment_size)
audios = slice_segments(
audios,
ids_slice * self.hop_length,
self.generator.segment_size * self.hop_length,
)
fake_mels = self.mel_transform(fake_audios.squeeze(1))
assert (
audios.shape == fake_audios.shape
), f"{audios.shape} != {fake_audios.shape}"
# Discriminator
if self.freeze_discriminator is False:
y_d_hat_r, y_d_hat_g, _, _ = self.discriminator(
audios, fake_audios.detach()
)
with torch.autocast(device_type=audios.device.type, enabled=False):
loss_disc, _, _ = discriminator_loss(y_d_hat_r, y_d_hat_g)
self.log(
f"train/discriminator/loss",
loss_disc,
on_step=True,
on_epoch=False,
prog_bar=False,
logger=True,
sync_dist=True,
)
optim_d.zero_grad()
self.manual_backward(loss_disc)
self.clip_gradients(
optim_d, gradient_clip_val=1000.0, gradient_clip_algorithm="norm"
)
optim_d.step()
# Adv Loss
y_d_hat_r, y_d_hat_g, _, _ = self.discriminator(audios, fake_audios)
# Adversarial Loss
with torch.autocast(device_type=audios.device.type, enabled=False):
loss_adv, _ = generator_loss(y_d_hat_g)
self.log(
f"train/generator/adv",
loss_adv,
on_step=True,
on_epoch=False,
prog_bar=False,
logger=True,
sync_dist=True,
)
with torch.autocast(device_type=audios.device.type, enabled=False):
loss_fm = feature_loss(y_d_hat_r, y_d_hat_g)
self.log(
f"train/generator/adv_fm",
loss_fm,
on_step=True,
on_epoch=False,
prog_bar=False,
logger=True,
sync_dist=True,
)
with torch.autocast(device_type=audios.device.type, enabled=False):
loss_mel = avg_with_mask(
F.l1_loss(gt_mels, fake_mels, reduction="none"), spec_masks
)
self.log(
"train/generator/loss_mel",
loss_mel,
on_step=True,
on_epoch=False,
prog_bar=False,
logger=True,
sync_dist=True,
)
loss_kl = kl_loss(z_p, logs_q, m_p, logs_p, y_mask)
self.log(
"train/generator/loss_kl",
loss_kl,
on_step=True,
on_epoch=False,
prog_bar=False,
logger=True,
sync_dist=True,
)
loss = (
loss_mel * self.weight_mel + loss_kl * self.weight_kl + loss_adv + loss_fm
)
self.log(
"train/generator/loss",
loss,
on_step=True,
on_epoch=False,
prog_bar=True,
logger=True,
sync_dist=True,
)
# Backward
optim_g.zero_grad()
self.manual_backward(loss)
self.clip_gradients(
optim_g, gradient_clip_val=1000.0, gradient_clip_algorithm="norm"
)
optim_g.step()
# Manual LR Scheduler
scheduler_g, scheduler_d = self.lr_schedulers()
scheduler_g.step()
scheduler_d.step()
def validation_step(self, batch: Any, batch_idx: int):
audios, audio_lengths = batch["audios"], batch["audio_lengths"]
texts, text_lengths = batch["texts"], batch["text_lengths"]
audios = audios.float()
audios = audios[:, None, :]
gt_mels = self.mel_transform(audios)
gt_specs = self.spec_transform(audios)
spec_lengths = audio_lengths // self.hop_length
spec_masks = torch.unsqueeze(
sequence_mask(spec_lengths, gt_mels.shape[2]), 1
).to(gt_mels.dtype)
prior_audios = self.generator.infer(
audios, audio_lengths, gt_specs, spec_lengths, texts, text_lengths
)
posterior_audios = self.generator.infer_posterior(gt_specs, spec_lengths)
prior_mels = self.mel_transform(prior_audios.squeeze(1))
posterior_mels = self.mel_transform(posterior_audios.squeeze(1))
min_mel_length = min(
gt_mels.shape[-1], prior_mels.shape[-1], posterior_mels.shape[-1]
)
gt_mels = gt_mels[:, :, :min_mel_length]
prior_mels = prior_mels[:, :, :min_mel_length]
posterior_mels = posterior_mels[:, :, :min_mel_length]
prior_mel_loss = avg_with_mask(
F.l1_loss(gt_mels, prior_mels, reduction="none"), spec_masks
)
posterior_mel_loss = avg_with_mask(
F.l1_loss(gt_mels, posterior_mels, reduction="none"), spec_masks
)
self.log(
"val/prior_mel_loss",
prior_mel_loss,
on_step=False,
on_epoch=True,
prog_bar=False,
logger=True,
sync_dist=True,
)
self.log(
"val/posterior_mel_loss",
posterior_mel_loss,
on_step=False,
on_epoch=True,
prog_bar=False,
logger=True,
sync_dist=True,
)
# only log the first batch
if batch_idx != 0:
return
for idx, (
mel,
prior_mel,
posterior_mel,
audio,
prior_audio,
posterior_audio,
audio_len,
) in enumerate(
zip(
gt_mels,
prior_mels,
posterior_mels,
audios.detach().float(),
prior_audios.detach().float(),
posterior_audios.detach().float(),
audio_lengths,
)
):
mel_len = audio_len // self.hop_length
image_mels = plot_mel(
[
prior_mel[:, :mel_len],
posterior_mel[:, :mel_len],
mel[:, :mel_len],
],
[
"Prior (VQ)",
"Posterior (Reconstruction)",
"Ground-Truth",
],
)
if isinstance(self.logger, WandbLogger):
self.logger.experiment.log(
{
"reconstruction_mel": wandb.Image(image_mels, caption="mels"),
"wavs": [
wandb.Audio(
audio[0, :audio_len],
sample_rate=self.sampling_rate,
caption="gt",
),
wandb.Audio(
prior_audio[0, :audio_len],
sample_rate=self.sampling_rate,
caption="prior",
),
wandb.Audio(
posterior_audio[0, :audio_len],
sample_rate=self.sampling_rate,
caption="posterior",
),
],
},
)
if isinstance(self.logger, TensorBoardLogger):
self.logger.experiment.add_figure(
f"sample-{idx}/mels",
image_mels,
global_step=self.global_step,
)
self.logger.experiment.add_audio(
f"sample-{idx}/wavs/gt",
audio[0, :audio_len],
self.global_step,
sample_rate=self.sampling_rate,
)
self.logger.experiment.add_audio(
f"sample-{idx}/wavs/prior",
prior_audio[0, :audio_len],
self.global_step,
sample_rate=self.sampling_rate,
)
self.logger.experiment.add_audio(
f"sample-{idx}/wavs/posterior",
posterior_audio[0, :audio_len],
self.global_step,
sample_rate=self.sampling_rate,
)
plt.close(image_mels)
+67
View File
@@ -0,0 +1,67 @@
import torch
import torch.nn.functional as F
from torch import nn
def feature_loss(fmap_r: list[torch.Tensor], fmap_g: list[torch.Tensor]):
loss = 0
for dr, dg in zip(fmap_r, fmap_g):
dr = dr.float().detach()
dg = dg.float()
loss += torch.mean(torch.abs(dr - dg))
return loss * 2
def discriminator_loss(
disc_real_outputs: list[torch.Tensor], disc_generated_outputs: list[torch.Tensor]
):
loss = 0
r_losses = []
g_losses = []
for dr, dg in zip(disc_real_outputs, disc_generated_outputs):
dr = dr.float()
dg = dg.float()
r_loss = torch.mean((1 - dr) ** 2)
g_loss = torch.mean(dg**2)
loss += r_loss + g_loss
r_losses.append(r_loss.item())
g_losses.append(g_loss.item())
return loss, r_losses, g_losses
def generator_loss(disc_outputs: list[torch.Tensor]):
loss = 0
gen_losses = []
for dg in disc_outputs:
dg = dg.float()
l = torch.mean((1 - dg) ** 2)
gen_losses.append(l)
loss += l
return loss, gen_losses
def kl_loss(
z_p: torch.Tensor,
logs_q: torch.Tensor,
m_p: torch.Tensor,
logs_p: torch.Tensor,
z_mask: torch.Tensor,
):
"""
z_p, logs_q: [b, h, t_t]
m_p, logs_p: [b, h, t_t]
"""
z_p = z_p.float()
logs_q = logs_q.float()
m_p = m_p.float()
logs_p = logs_p.float()
z_mask = z_mask.float()
kl = logs_p - logs_q - 0.5
kl += 0.5 * ((z_p - m_p) ** 2) * torch.exp(-2.0 * logs_p)
kl = torch.sum(kl * z_mask)
l = kl / torch.sum(z_mask)
return l
@@ -0,0 +1,350 @@
import math
import torch
from torch import nn
from torch.nn import functional as F
from torch.nn.utils import remove_weight_norm, weight_norm
from fish_speech.models.vits_decoder.modules import commons
from .modules import LayerNorm
class Encoder(nn.Module):
def __init__(
self,
hidden_channels,
filter_channels,
n_heads,
n_layers,
kernel_size=1,
p_dropout=0.0,
window_size=4,
isflow=False,
gin_channels=0,
):
super().__init__()
self.hidden_channels = hidden_channels
self.filter_channels = filter_channels
self.n_heads = n_heads
self.n_layers = n_layers
self.kernel_size = kernel_size
self.p_dropout = p_dropout
self.window_size = window_size
self.drop = nn.Dropout(p_dropout)
self.attn_layers = nn.ModuleList()
self.norm_layers_1 = nn.ModuleList()
self.ffn_layers = nn.ModuleList()
self.norm_layers_2 = nn.ModuleList()
for i in range(self.n_layers):
self.attn_layers.append(
MultiHeadAttention(
hidden_channels,
hidden_channels,
n_heads,
p_dropout=p_dropout,
window_size=window_size,
)
)
self.norm_layers_1.append(LayerNorm(hidden_channels))
self.ffn_layers.append(
FFN(
hidden_channels,
hidden_channels,
filter_channels,
kernel_size,
p_dropout=p_dropout,
)
)
self.norm_layers_2.append(LayerNorm(hidden_channels))
if isflow:
cond_layer = torch.nn.Conv1d(
gin_channels, 2 * hidden_channels * n_layers, 1
)
self.cond_pre = torch.nn.Conv1d(hidden_channels, 2 * hidden_channels, 1)
self.cond_layer = weight_norm(cond_layer, "weight")
self.gin_channels = gin_channels
def forward(self, x, x_mask, g=None):
attn_mask = x_mask.unsqueeze(2) * x_mask.unsqueeze(-1)
x = x * x_mask
if g is not None:
g = self.cond_layer(g)
for i in range(self.n_layers):
if g is not None:
x = self.cond_pre(x)
cond_offset = i * 2 * self.hidden_channels
g_l = g[:, cond_offset : cond_offset + 2 * self.hidden_channels, :]
x = commons.fused_add_tanh_sigmoid_multiply(
x, g_l, torch.IntTensor([self.hidden_channels])
)
y = self.attn_layers[i](x, x, attn_mask)
y = self.drop(y)
x = self.norm_layers_1[i](x + y)
y = self.ffn_layers[i](x, x_mask)
y = self.drop(y)
x = self.norm_layers_2[i](x + y)
x = x * x_mask
return x
class MultiHeadAttention(nn.Module):
def __init__(
self,
channels,
out_channels,
n_heads,
p_dropout=0.0,
window_size=None,
heads_share=True,
block_length=None,
proximal_bias=False,
proximal_init=False,
):
super().__init__()
assert channels % n_heads == 0
self.channels = channels
self.out_channels = out_channels
self.n_heads = n_heads
self.p_dropout = p_dropout
self.window_size = window_size
self.heads_share = heads_share
self.block_length = block_length
self.proximal_bias = proximal_bias
self.proximal_init = proximal_init
self.attn = None
self.k_channels = channels // n_heads
self.conv_q = nn.Conv1d(channels, channels, 1)
self.conv_k = nn.Conv1d(channels, channels, 1)
self.conv_v = nn.Conv1d(channels, channels, 1)
self.conv_o = nn.Conv1d(channels, out_channels, 1)
self.drop = nn.Dropout(p_dropout)
if window_size is not None:
n_heads_rel = 1 if heads_share else n_heads
rel_stddev = self.k_channels**-0.5
self.emb_rel_k = nn.Parameter(
torch.randn(n_heads_rel, window_size * 2 + 1, self.k_channels)
* rel_stddev
)
self.emb_rel_v = nn.Parameter(
torch.randn(n_heads_rel, window_size * 2 + 1, self.k_channels)
* rel_stddev
)
nn.init.xavier_uniform_(self.conv_q.weight)
nn.init.xavier_uniform_(self.conv_k.weight)
nn.init.xavier_uniform_(self.conv_v.weight)
if proximal_init:
with torch.no_grad():
self.conv_k.weight.copy_(self.conv_q.weight)
self.conv_k.bias.copy_(self.conv_q.bias)
def forward(self, x, c, attn_mask=None):
q = self.conv_q(x)
k = self.conv_k(c)
v = self.conv_v(c)
x, self.attn = self.attention(q, k, v, mask=attn_mask)
x = self.conv_o(x)
return x
def attention(self, query, key, value, mask=None):
# reshape [b, d, t] -> [b, n_h, t, d_k]
b, d, t_s, t_t = (*key.size(), query.size(2))
query = query.view(b, self.n_heads, self.k_channels, t_t).transpose(2, 3)
key = key.view(b, self.n_heads, self.k_channels, t_s).transpose(2, 3)
value = value.view(b, self.n_heads, self.k_channels, t_s).transpose(2, 3)
scores = torch.matmul(query / math.sqrt(self.k_channels), key.transpose(-2, -1))
if self.window_size is not None:
assert (
t_s == t_t
), "Relative attention is only available for self-attention."
key_relative_embeddings = self._get_relative_embeddings(self.emb_rel_k, t_s)
rel_logits = self._matmul_with_relative_keys(
query / math.sqrt(self.k_channels), key_relative_embeddings
)
scores_local = self._relative_position_to_absolute_position(rel_logits)
scores = scores + scores_local
if self.proximal_bias:
assert t_s == t_t, "Proximal bias is only available for self-attention."
scores = scores + self._attention_bias_proximal(t_s).to(
device=scores.device, dtype=scores.dtype
)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e4)
if self.block_length is not None:
assert (
t_s == t_t
), "Local attention is only available for self-attention."
block_mask = (
torch.ones_like(scores)
.triu(-self.block_length)
.tril(self.block_length)
)
scores = scores.masked_fill(block_mask == 0, -1e4)
p_attn = F.softmax(scores, dim=-1) # [b, n_h, t_t, t_s]
p_attn = self.drop(p_attn)
output = torch.matmul(p_attn, value)
if self.window_size is not None:
relative_weights = self._absolute_position_to_relative_position(p_attn)
value_relative_embeddings = self._get_relative_embeddings(
self.emb_rel_v, t_s
)
output = output + self._matmul_with_relative_values(
relative_weights, value_relative_embeddings
)
output = (
output.transpose(2, 3).contiguous().view(b, d, t_t)
) # [b, n_h, t_t, d_k] -> [b, d, t_t]
return output, p_attn
def _matmul_with_relative_values(self, x, y):
"""
x: [b, h, l, m]
y: [h or 1, m, d]
ret: [b, h, l, d]
"""
ret = torch.matmul(x, y.unsqueeze(0))
return ret
def _matmul_with_relative_keys(self, x, y):
"""
x: [b, h, l, d]
y: [h or 1, m, d]
ret: [b, h, l, m]
"""
ret = torch.matmul(x, y.unsqueeze(0).transpose(-2, -1))
return ret
def _get_relative_embeddings(self, relative_embeddings, length):
max_relative_position = 2 * self.window_size + 1
# Pad first before slice to avoid using cond ops.
pad_length = max(length - (self.window_size + 1), 0)
slice_start_position = max((self.window_size + 1) - length, 0)
slice_end_position = slice_start_position + 2 * length - 1
if pad_length > 0:
padded_relative_embeddings = F.pad(
relative_embeddings,
commons.convert_pad_shape([[0, 0], [pad_length, pad_length], [0, 0]]),
)
else:
padded_relative_embeddings = relative_embeddings
used_relative_embeddings = padded_relative_embeddings[
:, slice_start_position:slice_end_position
]
return used_relative_embeddings
def _relative_position_to_absolute_position(self, x):
"""
x: [b, h, l, 2*l-1]
ret: [b, h, l, l]
"""
batch, heads, length, _ = x.size()
# Concat columns of pad to shift from relative to absolute indexing.
x = F.pad(x, commons.convert_pad_shape([[0, 0], [0, 0], [0, 0], [0, 1]]))
# Concat extra elements so to add up to shape (len+1, 2*len-1).
x_flat = x.view([batch, heads, length * 2 * length])
x_flat = F.pad(
x_flat, commons.convert_pad_shape([[0, 0], [0, 0], [0, length - 1]])
)
# Reshape and slice out the padded elements.
x_final = x_flat.view([batch, heads, length + 1, 2 * length - 1])[
:, :, :length, length - 1 :
]
return x_final
def _absolute_position_to_relative_position(self, x):
"""
x: [b, h, l, l]
ret: [b, h, l, 2*l-1]
"""
batch, heads, length, _ = x.size()
# pad along column
x = F.pad(
x, commons.convert_pad_shape([[0, 0], [0, 0], [0, 0], [0, length - 1]])
)
x_flat = x.view([batch, heads, length**2 + length * (length - 1)])
# add 0's in the beginning that will skew the elements after reshape
x_flat = F.pad(x_flat, commons.convert_pad_shape([[0, 0], [0, 0], [length, 0]]))
x_final = x_flat.view([batch, heads, length, 2 * length])[:, :, :, 1:]
return x_final
def _attention_bias_proximal(self, length):
"""Bias for self-attention to encourage attention to close positions.
Args:
length: an integer scalar.
Returns:
a Tensor with shape [1, 1, length, length]
"""
r = torch.arange(length, dtype=torch.float32)
diff = torch.unsqueeze(r, 0) - torch.unsqueeze(r, 1)
return torch.unsqueeze(torch.unsqueeze(-torch.log1p(torch.abs(diff)), 0), 0)
class FFN(nn.Module):
def __init__(
self,
in_channels,
out_channels,
filter_channels,
kernel_size,
p_dropout=0.0,
activation=None,
causal=False,
):
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.filter_channels = filter_channels
self.kernel_size = kernel_size
self.p_dropout = p_dropout
self.activation = activation
self.causal = causal
if causal:
self.padding = self._causal_padding
else:
self.padding = self._same_padding
self.conv_1 = nn.Conv1d(in_channels, filter_channels, kernel_size)
self.conv_2 = nn.Conv1d(filter_channels, out_channels, kernel_size)
self.drop = nn.Dropout(p_dropout)
def forward(self, x, x_mask):
x = self.conv_1(self.padding(x * x_mask))
if self.activation == "gelu":
x = x * torch.sigmoid(1.702 * x)
else:
x = torch.relu(x)
x = self.drop(x)
x = self.conv_2(self.padding(x * x_mask))
return x * x_mask
def _causal_padding(self, x):
if self.kernel_size == 1:
return x
pad_l = self.kernel_size - 1
pad_r = 0
padding = [[0, 0], [0, 0], [pad_l, pad_r]]
x = F.pad(x, commons.convert_pad_shape(padding))
return x
def _same_padding(self, x):
if self.kernel_size == 1:
return x
pad_l = (self.kernel_size - 1) // 2
pad_r = self.kernel_size // 2
padding = [[0, 0], [0, 0], [pad_l, pad_r]]
x = F.pad(x, commons.convert_pad_shape(padding))
return x
@@ -0,0 +1,190 @@
import math
import torch
from torch.nn import functional as F
def init_weights(m, mean=0.0, std=0.01):
classname = m.__class__.__name__
if classname.find("Conv") != -1:
m.weight.data.normal_(mean, std)
def get_padding(kernel_size, dilation=1):
return int((kernel_size * dilation - dilation) / 2)
def convert_pad_shape(pad_shape):
l = pad_shape[::-1]
pad_shape = [item for sublist in l for item in sublist]
return pad_shape
def intersperse(lst, item):
result = [item] * (len(lst) * 2 + 1)
result[1::2] = lst
return result
def kl_divergence(m_p, logs_p, m_q, logs_q):
"""KL(P||Q)"""
kl = (logs_q - logs_p) - 0.5
kl += (
0.5 * (torch.exp(2.0 * logs_p) + ((m_p - m_q) ** 2)) * torch.exp(-2.0 * logs_q)
)
return kl
def rand_gumbel(shape):
"""Sample from the Gumbel distribution, protect from overflows."""
uniform_samples = torch.rand(shape) * 0.99998 + 0.00001
return -torch.log(-torch.log(uniform_samples))
def rand_gumbel_like(x):
g = rand_gumbel(x.size()).to(dtype=x.dtype, device=x.device)
return g
def slice_segments(x, ids_str, segment_size=4):
ret = torch.zeros_like(x[:, :, :segment_size])
for i in range(x.size(0)):
idx_str = ids_str[i]
idx_end = idx_str + segment_size
ret[i] = x[i, :, idx_str:idx_end]
return ret
def rand_slice_segments(x, x_lengths=None, segment_size=4):
b, d, t = x.size()
if x_lengths is None:
x_lengths = t
ids_str_max = x_lengths - segment_size + 1
ids_str = (torch.rand([b]).to(device=x.device) * ids_str_max).to(dtype=torch.long)
ret = slice_segments(x, ids_str, segment_size)
return ret, ids_str
def get_timing_signal_1d(length, channels, min_timescale=1.0, max_timescale=1.0e4):
position = torch.arange(length, dtype=torch.float)
num_timescales = channels // 2
log_timescale_increment = math.log(float(max_timescale) / float(min_timescale)) / (
num_timescales - 1
)
inv_timescales = min_timescale * torch.exp(
torch.arange(num_timescales, dtype=torch.float) * -log_timescale_increment
)
scaled_time = position.unsqueeze(0) * inv_timescales.unsqueeze(1)
signal = torch.cat([torch.sin(scaled_time), torch.cos(scaled_time)], 0)
signal = F.pad(signal, [0, 0, 0, channels % 2])
signal = signal.view(1, channels, length)
return signal
def add_timing_signal_1d(x, min_timescale=1.0, max_timescale=1.0e4):
b, channels, length = x.size()
signal = get_timing_signal_1d(length, channels, min_timescale, max_timescale)
return x + signal.to(dtype=x.dtype, device=x.device)
def cat_timing_signal_1d(x, min_timescale=1.0, max_timescale=1.0e4, axis=1):
b, channels, length = x.size()
signal = get_timing_signal_1d(length, channels, min_timescale, max_timescale)
return torch.cat([x, signal.to(dtype=x.dtype, device=x.device)], axis)
def subsequent_mask(length):
mask = torch.tril(torch.ones(length, length)).unsqueeze(0).unsqueeze(0)
return mask
@torch.jit.script
def fused_add_tanh_sigmoid_multiply(input_a, input_b, n_channels):
n_channels_int = n_channels[0]
in_act = input_a + input_b
t_act = torch.tanh(in_act[:, :n_channels_int, :])
s_act = torch.sigmoid(in_act[:, n_channels_int:, :])
acts = t_act * s_act
return acts
def convert_pad_shape(pad_shape):
l = pad_shape[::-1]
pad_shape = [item for sublist in l for item in sublist]
return pad_shape
def shift_1d(x):
x = F.pad(x, convert_pad_shape([[0, 0], [0, 0], [1, 0]]))[:, :, :-1]
return x
def sequence_mask(length, max_length=None):
if max_length is None:
max_length = length.max()
x = torch.arange(max_length, dtype=length.dtype, device=length.device)
return x.unsqueeze(0) < length.unsqueeze(1)
def generate_path(duration, mask):
"""
duration: [b, 1, t_x]
mask: [b, 1, t_y, t_x]
"""
device = duration.device
b, _, t_y, t_x = mask.shape
cum_duration = torch.cumsum(duration, -1)
cum_duration_flat = cum_duration.view(b * t_x)
path = sequence_mask(cum_duration_flat, t_y).to(mask.dtype)
path = path.view(b, t_x, t_y)
path = path - F.pad(path, convert_pad_shape([[0, 0], [1, 0], [0, 0]]))[:, :-1]
path = path.unsqueeze(1).transpose(2, 3) * mask
return path
def clip_grad_value_(parameters, clip_value, norm_type=2):
if isinstance(parameters, torch.Tensor):
parameters = [parameters]
parameters = list(filter(lambda p: p.grad is not None, parameters))
norm_type = float(norm_type)
if clip_value is not None:
clip_value = float(clip_value)
total_norm = 0
for p in parameters:
param_norm = p.grad.data.norm(norm_type)
total_norm += param_norm.item() ** norm_type
if clip_value is not None:
p.grad.data.clamp_(min=-clip_value, max=clip_value)
total_norm = total_norm ** (1.0 / norm_type)
return total_norm
def squeeze(x, x_mask=None, n_sqz=2):
b, c, t = x.size()
t = (t // n_sqz) * n_sqz
x = x[:, :, :t]
x_sqz = x.view(b, c, t // n_sqz, n_sqz)
x_sqz = x_sqz.permute(0, 3, 1, 2).contiguous().view(b, c * n_sqz, t // n_sqz)
if x_mask is not None:
x_mask = x_mask[:, :, n_sqz - 1 :: n_sqz]
else:
x_mask = torch.ones(b, 1, t // n_sqz).to(device=x.device, dtype=x.dtype)
return x_sqz * x_mask, x_mask
def unsqueeze(x, x_mask=None, n_sqz=2):
b, c, t = x.size()
x_unsqz = x.view(b, n_sqz, c // n_sqz, t)
x_unsqz = x_unsqz.permute(0, 2, 3, 1).contiguous().view(b, c // n_sqz, t * n_sqz)
if x_mask is not None:
x_mask = x_mask.unsqueeze(-1).repeat(1, 1, 1, n_sqz).view(b, 1, t * n_sqz)
else:
x_mask = torch.ones(b, 1, t * n_sqz).to(device=x.device, dtype=x.dtype)
return x_unsqz * x_mask, x_mask
@@ -0,0 +1,686 @@
import torch
from torch import nn
from torch.nn import Conv1d, Conv2d, ConvTranspose1d
from torch.nn import functional as F
from torch.nn.utils import remove_weight_norm, spectral_norm, weight_norm
from fish_speech.models.vits_decoder.modules import attentions, commons, modules
from .commons import get_padding, init_weights
from .mrte import MRTE
from .vq_encoder import VQEncoder
class TextEncoder(nn.Module):
def __init__(
self,
out_channels,
hidden_channels,
filter_channels,
n_heads,
n_layers,
kernel_size,
p_dropout,
latent_channels=192,
codebook_size=264,
):
super().__init__()
self.out_channels = out_channels
self.hidden_channels = hidden_channels
self.filter_channels = filter_channels
self.n_heads = n_heads
self.n_layers = n_layers
self.kernel_size = kernel_size
self.p_dropout = p_dropout
self.latent_channels = latent_channels
self.ssl_proj = nn.Conv1d(768, hidden_channels, 1)
self.encoder_ssl = attentions.Encoder(
hidden_channels,
filter_channels,
n_heads,
n_layers // 2,
kernel_size,
p_dropout,
)
self.encoder_text = attentions.Encoder(
hidden_channels, filter_channels, n_heads, n_layers, kernel_size, p_dropout
)
self.text_embedding = nn.Embedding(codebook_size, hidden_channels)
self.mrte = MRTE()
self.encoder2 = attentions.Encoder(
hidden_channels,
filter_channels,
n_heads,
n_layers // 2,
kernel_size,
p_dropout,
)
self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
def forward(self, y, y_lengths, text, text_lengths, ge):
y_mask = torch.unsqueeze(commons.sequence_mask(y_lengths, y.size(2)), 1).to(
y.dtype
)
y = self.ssl_proj(y * y_mask) * y_mask
y = self.encoder_ssl(y * y_mask, y_mask)
text_mask = torch.unsqueeze(
commons.sequence_mask(text_lengths, text.size(1)), 1
).to(y.dtype)
text = self.text_embedding(text).transpose(1, 2)
text = self.encoder_text(text * text_mask, text_mask)
y = self.mrte(y, y_mask, text, text_mask, ge)
y = self.encoder2(y * y_mask, y_mask)
stats = self.proj(y) * y_mask
m, logs = torch.split(stats, self.out_channels, dim=1)
return y, m, logs, y_mask
class ResidualCouplingBlock(nn.Module):
def __init__(
self,
channels,
hidden_channels,
kernel_size,
dilation_rate,
n_layers,
n_flows=4,
gin_channels=0,
):
super().__init__()
self.channels = channels
self.hidden_channels = hidden_channels
self.kernel_size = kernel_size
self.dilation_rate = dilation_rate
self.n_layers = n_layers
self.n_flows = n_flows
self.gin_channels = gin_channels
self.flows = nn.ModuleList()
for i in range(n_flows):
self.flows.append(
modules.ResidualCouplingLayer(
channels,
hidden_channels,
kernel_size,
dilation_rate,
n_layers,
gin_channels=gin_channels,
mean_only=True,
)
)
self.flows.append(modules.Flip())
def forward(self, x, x_mask, g=None, reverse=False):
if not reverse:
for flow in self.flows:
x, _ = flow(x, x_mask, g=g, reverse=reverse)
else:
for flow in reversed(self.flows):
x = flow(x, x_mask, g=g, reverse=reverse)
return x
class PosteriorEncoder(nn.Module):
def __init__(
self,
in_channels,
out_channels,
hidden_channels,
kernel_size,
dilation_rate,
n_layers,
gin_channels=0,
):
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.hidden_channels = hidden_channels
self.kernel_size = kernel_size
self.dilation_rate = dilation_rate
self.n_layers = n_layers
self.gin_channels = gin_channels
self.pre = nn.Conv1d(in_channels, hidden_channels, 1)
self.enc = modules.WN(
hidden_channels,
kernel_size,
dilation_rate,
n_layers,
gin_channels=gin_channels,
)
self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
def forward(self, x, x_lengths, g=None):
if g != None:
g = g.detach()
x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(
x.dtype
)
x = self.pre(x) * x_mask
x = self.enc(x, x_mask, g=g)
stats = self.proj(x) * x_mask
m, logs = torch.split(stats, self.out_channels, dim=1)
z = (m + torch.randn_like(m) * torch.exp(logs)) * x_mask
return z, m, logs, x_mask
class Generator(torch.nn.Module):
def __init__(
self,
initial_channel,
resblock,
resblock_kernel_sizes,
resblock_dilation_sizes,
upsample_rates,
upsample_initial_channel,
upsample_kernel_sizes,
gin_channels=0,
):
super(Generator, self).__init__()
self.num_kernels = len(resblock_kernel_sizes)
self.num_upsamples = len(upsample_rates)
self.conv_pre = Conv1d(
initial_channel, upsample_initial_channel, 7, 1, padding=3
)
resblock = modules.ResBlock1 if resblock == "1" else modules.ResBlock2
self.ups = nn.ModuleList()
for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):
self.ups.append(
weight_norm(
ConvTranspose1d(
upsample_initial_channel // (2**i),
upsample_initial_channel // (2 ** (i + 1)),
k,
u,
padding=(k - u) // 2,
)
)
)
self.resblocks = nn.ModuleList()
for i in range(len(self.ups)):
ch = upsample_initial_channel // (2 ** (i + 1))
for j, (k, d) in enumerate(
zip(resblock_kernel_sizes, resblock_dilation_sizes)
):
self.resblocks.append(resblock(ch, k, d))
self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)
self.ups.apply(init_weights)
if gin_channels != 0:
self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)
def forward(self, x, g=None):
x = self.conv_pre(x)
if g is not None:
x = x + self.cond(g)
for i in range(self.num_upsamples):
x = F.leaky_relu(x, modules.LRELU_SLOPE)
x = self.ups[i](x)
xs = None
for j in range(self.num_kernels):
if xs is None:
xs = self.resblocks[i * self.num_kernels + j](x)
else:
xs += self.resblocks[i * self.num_kernels + j](x)
x = xs / self.num_kernels
x = F.leaky_relu(x)
x = self.conv_post(x)
x = torch.tanh(x)
return x
def remove_weight_norm(self):
print("Removing weight norm...")
for l in self.ups:
remove_weight_norm(l)
for l in self.resblocks:
l.remove_weight_norm()
class DiscriminatorP(torch.nn.Module):
def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False):
super(DiscriminatorP, self).__init__()
self.period = period
self.use_spectral_norm = use_spectral_norm
norm_f = weight_norm if use_spectral_norm == False else spectral_norm
self.convs = nn.ModuleList(
[
norm_f(
Conv2d(
1,
32,
(kernel_size, 1),
(stride, 1),
padding=(get_padding(kernel_size, 1), 0),
)
),
norm_f(
Conv2d(
32,
128,
(kernel_size, 1),
(stride, 1),
padding=(get_padding(kernel_size, 1), 0),
)
),
norm_f(
Conv2d(
128,
512,
(kernel_size, 1),
(stride, 1),
padding=(get_padding(kernel_size, 1), 0),
)
),
norm_f(
Conv2d(
512,
1024,
(kernel_size, 1),
(stride, 1),
padding=(get_padding(kernel_size, 1), 0),
)
),
norm_f(
Conv2d(
1024,
1024,
(kernel_size, 1),
1,
padding=(get_padding(kernel_size, 1), 0),
)
),
]
)
self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))
def forward(self, x):
fmap = []
# 1d to 2d
b, c, t = x.shape
if t % self.period != 0: # pad first
n_pad = self.period - (t % self.period)
x = F.pad(x, (0, n_pad), "reflect")
t = t + n_pad
x = x.view(b, c, t // self.period, self.period)
for l in self.convs:
x = l(x)
x = F.leaky_relu(x, modules.LRELU_SLOPE)
fmap.append(x)
x = self.conv_post(x)
fmap.append(x)
x = torch.flatten(x, 1, -1)
return x, fmap
class DiscriminatorS(torch.nn.Module):
def __init__(self, use_spectral_norm=False):
super(DiscriminatorS, self).__init__()
norm_f = weight_norm if use_spectral_norm == False else spectral_norm
self.convs = nn.ModuleList(
[
norm_f(Conv1d(1, 16, 15, 1, padding=7)),
norm_f(Conv1d(16, 64, 41, 4, groups=4, padding=20)),
norm_f(Conv1d(64, 256, 41, 4, groups=16, padding=20)),
norm_f(Conv1d(256, 1024, 41, 4, groups=64, padding=20)),
norm_f(Conv1d(1024, 1024, 41, 4, groups=256, padding=20)),
norm_f(Conv1d(1024, 1024, 5, 1, padding=2)),
]
)
self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))
def forward(self, x):
fmap = []
for l in self.convs:
x = l(x)
x = F.leaky_relu(x, modules.LRELU_SLOPE)
fmap.append(x)
x = self.conv_post(x)
fmap.append(x)
x = torch.flatten(x, 1, -1)
return x, fmap
class EnsembledDiscriminator(torch.nn.Module):
def __init__(self, periods=(2, 3, 5, 7, 11), use_spectral_norm=False):
super().__init__()
discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]
discs = discs + [
DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods
]
self.discriminators = nn.ModuleList(discs)
def forward(self, y, y_hat):
y_d_rs = []
y_d_gs = []
fmap_rs = []
fmap_gs = []
for i, d in enumerate(self.discriminators):
y_d_r, fmap_r = d(y)
y_d_g, fmap_g = d(y_hat)
y_d_rs.append(y_d_r)
y_d_gs.append(y_d_g)
fmap_rs.append(fmap_r)
fmap_gs.append(fmap_g)
return y_d_rs, y_d_gs, fmap_rs, fmap_gs
class SynthesizerTrn(nn.Module):
"""
Synthesizer for Training
"""
def __init__(
self,
*,
spec_channels,
segment_size,
inter_channels,
hidden_channels,
filter_channels,
n_heads,
n_layers,
kernel_size,
p_dropout,
resblock,
resblock_kernel_sizes,
resblock_dilation_sizes,
upsample_rates,
upsample_initial_channel,
upsample_kernel_sizes,
gin_channels=0,
codebook_size=264,
vq_mask_ratio=0.0,
ref_mask_ratio=0.0,
):
super().__init__()
self.spec_channels = spec_channels
self.inter_channels = inter_channels
self.hidden_channels = hidden_channels
self.filter_channels = filter_channels
self.n_heads = n_heads
self.n_layers = n_layers
self.kernel_size = kernel_size
self.p_dropout = p_dropout
self.resblock = resblock
self.resblock_kernel_sizes = resblock_kernel_sizes
self.resblock_dilation_sizes = resblock_dilation_sizes
self.upsample_rates = upsample_rates
self.upsample_initial_channel = upsample_initial_channel
self.upsample_kernel_sizes = upsample_kernel_sizes
self.segment_size = segment_size
self.gin_channels = gin_channels
self.vq_mask_ratio = vq_mask_ratio
self.ref_mask_ratio = ref_mask_ratio
self.enc_p = TextEncoder(
inter_channels,
hidden_channels,
filter_channels,
n_heads,
n_layers,
kernel_size,
p_dropout,
codebook_size=codebook_size,
)
self.dec = Generator(
inter_channels,
resblock,
resblock_kernel_sizes,
resblock_dilation_sizes,
upsample_rates,
upsample_initial_channel,
upsample_kernel_sizes,
gin_channels=gin_channels,
)
self.enc_q = PosteriorEncoder(
spec_channels,
inter_channels,
hidden_channels,
5,
1,
16,
gin_channels=gin_channels,
)
self.flow = ResidualCouplingBlock(
inter_channels, hidden_channels, 5, 1, 4, gin_channels=gin_channels
)
self.ref_enc = modules.MelStyleEncoder(
spec_channels, style_vector_dim=gin_channels
)
self.vq = VQEncoder()
for param in self.vq.parameters():
param.requires_grad = False
def forward(
self, audio, audio_lengths, gt_specs, gt_spec_lengths, text, text_lengths
):
y_mask = torch.unsqueeze(
commons.sequence_mask(gt_spec_lengths, gt_specs.size(2)), 1
).to(gt_specs.dtype)
ge = self.ref_enc(gt_specs * y_mask, y_mask)
if self.training and self.ref_mask_ratio > 0:
bs = audio.size(0)
mask_speaker_len = int(bs * self.ref_mask_ratio)
mask_indices = torch.randperm(bs)[:mask_speaker_len]
audio[mask_indices] = 0
quantized = self.vq(audio, audio_lengths)
# Block masking, block_size = 4
block_size = 4
if self.training and self.vq_mask_ratio > 0:
reduced_length = quantized.size(-1) // block_size
mask_length = int(reduced_length * self.vq_mask_ratio)
mask_indices = torch.randperm(reduced_length)[:mask_length]
short_mask = torch.zeros(
quantized.size(0),
quantized.size(1),
reduced_length,
device=quantized.device,
dtype=torch.float,
)
short_mask[:, :, mask_indices] = 1.0
long_mask = short_mask.repeat_interleave(block_size, dim=-1)
long_mask = F.interpolate(
long_mask, size=quantized.size(-1), mode="nearest"
)
quantized = quantized.masked_fill(long_mask > 0.5, 0)
x, m_p, logs_p, y_mask = self.enc_p(
quantized, gt_spec_lengths, text, text_lengths, ge
)
z, m_q, logs_q, y_mask = self.enc_q(gt_specs, gt_spec_lengths, g=ge)
z_p = self.flow(z, y_mask, g=ge)
z_slice, ids_slice = commons.rand_slice_segments(
z, gt_spec_lengths, self.segment_size
)
o = self.dec(z_slice, g=ge)
return (
o,
ids_slice,
y_mask,
(z, z_p, m_p, logs_p, m_q, logs_q),
)
@torch.no_grad()
def infer(
self,
audio,
audio_lengths,
gt_specs,
gt_spec_lengths,
text,
text_lengths,
noise_scale=0.5,
):
quantized = self.vq(audio, audio_lengths)
quantized_lengths = audio_lengths // 512
ge = self.encode_ref(gt_specs, gt_spec_lengths)
return self.decode(
quantized,
quantized_lengths,
text,
text_lengths,
noise_scale=noise_scale,
ge=ge,
)
@torch.no_grad()
def infer_posterior(
self,
gt_specs,
gt_spec_lengths,
):
y_mask = torch.unsqueeze(
commons.sequence_mask(gt_spec_lengths, gt_specs.size(2)), 1
).to(gt_specs.dtype)
ge = self.ref_enc(gt_specs * y_mask, y_mask)
z, m_q, logs_q, y_mask = self.enc_q(gt_specs, gt_spec_lengths, g=ge)
o = self.dec(z * y_mask, g=ge)
return o
@torch.no_grad()
def decode(
self,
quantized,
quantized_lengths,
text,
text_lengths,
noise_scale=0.5,
ge=None,
):
x, m_p, logs_p, y_mask = self.enc_p(
quantized, quantized_lengths, text, text_lengths, ge
)
z_p = m_p + torch.randn_like(m_p) * torch.exp(logs_p) * noise_scale
z = self.flow(z_p, y_mask, g=ge, reverse=True)
o = self.dec(z * y_mask, g=ge)
return o
@torch.no_grad()
def encode_ref(self, gt_specs, gt_spec_lengths):
y_mask = torch.unsqueeze(
commons.sequence_mask(gt_spec_lengths, gt_specs.size(2)), 1
).to(gt_specs.dtype)
ge = self.ref_enc(gt_specs * y_mask, y_mask)
return ge
if __name__ == "__main__":
import librosa
from transformers import AutoTokenizer
from fish_speech.utils.spectrogram import LinearSpectrogram
model = SynthesizerTrn(
spec_channels=1025,
segment_size=20480 // 640,
inter_channels=192,
hidden_channels=192,
filter_channels=768,
n_heads=2,
n_layers=6,
kernel_size=3,
p_dropout=0.1,
resblock="1",
resblock_kernel_sizes=[3, 7, 11],
resblock_dilation_sizes=[[1, 3, 5], [1, 3, 5], [1, 3, 5]],
upsample_rates=[8, 8, 2, 2, 2],
upsample_initial_channel=512,
upsample_kernel_sizes=[16, 16, 8, 2, 2],
gin_channels=512,
)
ckpt = "checkpoints/Bert-VITS2/G_0.pth"
# Try to load the model
print(f"Loading model from {ckpt}")
checkpoint = torch.load(ckpt, map_location="cpu", weights_only=True)["model"]
# d_checkpoint = torch.load(
# "checkpoints/Bert-VITS2/D_0.pth", map_location="cpu", weights_only=True
# )["model"]
# print(checkpoint.keys())
checkpoint.pop("dec.cond.weight")
checkpoint.pop("enc_q.enc.cond_layer.weight_v")
# new_checkpoint = {}
# for k, v in checkpoint.items():
# new_checkpoint["generator." + k] = v
# for k, v in d_checkpoint.items():
# new_checkpoint["discriminator." + k] = v
# torch.save(new_checkpoint, "checkpoints/Bert-VITS2/ensemble.pth")
# exit()
print(model.load_state_dict(checkpoint, strict=False))
# Test
ref_audio = librosa.load(
"data/source/云天河/云天河-旁白/《薄太太》第0025集-yth_24.wav", sr=32000
)[0]
input_audio = librosa.load(
"data/source/云天河/云天河-旁白/《薄太太》第0025集-yth_24.wav", sr=32000
)[0]
ref_audio = input_audio
text = "博兴只知道身边的小女人没睡着,他又凑到她耳边压低了声线。阮苏眉睁眼,不觉得你老公像英雄吗?阮苏还是没反应,这男人是不是有病?刚才那冰冷又强势的样子,和现在这幼稚无赖的样子,根本就判若二人。"
encoded_text = AutoTokenizer.from_pretrained("fishaudio/fish-speech-1")
spec = LinearSpectrogram(n_fft=2048, hop_length=640, win_length=2048)
ref_audio = torch.tensor(ref_audio).unsqueeze(0).unsqueeze(0)
ref_spec = spec(ref_audio)
input_audio = torch.tensor(input_audio).unsqueeze(0).unsqueeze(0)
text = encoded_text(text, return_tensors="pt")["input_ids"]
print(ref_audio.size(), ref_spec.size(), input_audio.size(), text.size())
o, y_mask, (z, z_p, m_p, logs_p) = model.infer(
input_audio,
torch.LongTensor([input_audio.size(2)]),
ref_spec,
torch.LongTensor([ref_spec.size(2)]),
text,
torch.LongTensor([text.size(1)]),
)
print(o.size(), y_mask.size(), z.size(), z_p.size(), m_p.size(), logs_p.size())
# Save output
# import soundfile as sf
# sf.write("output.wav", o.squeeze().detach().numpy(), 32000)
@@ -0,0 +1,619 @@
import numpy as np
import torch
from torch import nn
from torch.nn import Conv1d
from torch.nn import functional as F
from torch.nn.utils import remove_weight_norm, weight_norm
from .commons import fused_add_tanh_sigmoid_multiply, get_padding, init_weights
LRELU_SLOPE = 0.1
class LayerNorm(nn.Module):
def __init__(self, channels, eps=1e-5):
super().__init__()
self.channels = channels
self.eps = eps
self.gamma = nn.Parameter(torch.ones(channels))
self.beta = nn.Parameter(torch.zeros(channels))
def forward(self, x):
x = x.transpose(1, -1)
x = F.layer_norm(x, (self.channels,), self.gamma, self.beta, self.eps)
return x.transpose(1, -1)
class ConvReluNorm(nn.Module):
def __init__(
self,
in_channels,
hidden_channels,
out_channels,
kernel_size,
n_layers,
p_dropout,
):
super().__init__()
self.in_channels = in_channels
self.hidden_channels = hidden_channels
self.out_channels = out_channels
self.kernel_size = kernel_size
self.n_layers = n_layers
self.p_dropout = p_dropout
assert n_layers > 1, "Number of layers should be larger than 0."
self.conv_layers = nn.ModuleList()
self.norm_layers = nn.ModuleList()
self.conv_layers.append(
nn.Conv1d(
in_channels, hidden_channels, kernel_size, padding=kernel_size // 2
)
)
self.norm_layers.append(LayerNorm(hidden_channels))
self.relu_drop = nn.Sequential(nn.ReLU(), nn.Dropout(p_dropout))
for _ in range(n_layers - 1):
self.conv_layers.append(
nn.Conv1d(
hidden_channels,
hidden_channels,
kernel_size,
padding=kernel_size // 2,
)
)
self.norm_layers.append(LayerNorm(hidden_channels))
self.proj = nn.Conv1d(hidden_channels, out_channels, 1)
self.proj.weight.data.zero_()
self.proj.bias.data.zero_()
def forward(self, x, x_mask):
x_org = x
for i in range(self.n_layers):
x = self.conv_layers[i](x * x_mask)
x = self.norm_layers[i](x)
x = self.relu_drop(x)
x = x_org + self.proj(x)
return x * x_mask
class WN(torch.nn.Module):
def __init__(
self,
hidden_channels,
kernel_size,
dilation_rate,
n_layers,
gin_channels=0,
p_dropout=0,
):
super(WN, self).__init__()
assert kernel_size % 2 == 1
self.hidden_channels = hidden_channels
self.kernel_size = (kernel_size,)
self.dilation_rate = dilation_rate
self.n_layers = n_layers
self.gin_channels = gin_channels
self.p_dropout = p_dropout
self.in_layers = torch.nn.ModuleList()
self.res_skip_layers = torch.nn.ModuleList()
self.drop = nn.Dropout(p_dropout)
if gin_channels != 0:
cond_layer = torch.nn.Conv1d(
gin_channels, 2 * hidden_channels * n_layers, 1
)
self.cond_layer = torch.nn.utils.weight_norm(cond_layer, name="weight")
for i in range(n_layers):
dilation = dilation_rate**i
padding = int((kernel_size * dilation - dilation) / 2)
in_layer = torch.nn.Conv1d(
hidden_channels,
2 * hidden_channels,
kernel_size,
dilation=dilation,
padding=padding,
)
in_layer = torch.nn.utils.weight_norm(in_layer, name="weight")
self.in_layers.append(in_layer)
# last one is not necessary
if i < n_layers - 1:
res_skip_channels = 2 * hidden_channels
else:
res_skip_channels = hidden_channels
res_skip_layer = torch.nn.Conv1d(hidden_channels, res_skip_channels, 1)
res_skip_layer = torch.nn.utils.weight_norm(res_skip_layer, name="weight")
self.res_skip_layers.append(res_skip_layer)
def forward(self, x, x_mask, g=None, **kwargs):
output = torch.zeros_like(x)
n_channels_tensor = torch.IntTensor([self.hidden_channels])
if g is not None:
g = self.cond_layer(g)
for i in range(self.n_layers):
x_in = self.in_layers[i](x)
if g is not None:
cond_offset = i * 2 * self.hidden_channels
g_l = g[:, cond_offset : cond_offset + 2 * self.hidden_channels, :]
else:
g_l = torch.zeros_like(x_in)
acts = fused_add_tanh_sigmoid_multiply(x_in, g_l, n_channels_tensor)
acts = self.drop(acts)
res_skip_acts = self.res_skip_layers[i](acts)
if i < self.n_layers - 1:
res_acts = res_skip_acts[:, : self.hidden_channels, :]
x = (x + res_acts) * x_mask
output = output + res_skip_acts[:, self.hidden_channels :, :]
else:
output = output + res_skip_acts
return output * x_mask
def remove_weight_norm(self):
if self.gin_channels != 0:
torch.nn.utils.remove_weight_norm(self.cond_layer)
for l in self.in_layers:
torch.nn.utils.remove_weight_norm(l)
for l in self.res_skip_layers:
torch.nn.utils.remove_weight_norm(l)
class ResBlock1(torch.nn.Module):
def __init__(self, channels, kernel_size=3, dilation=(1, 3, 5)):
super(ResBlock1, self).__init__()
self.convs1 = nn.ModuleList(
[
weight_norm(
Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[0],
padding=get_padding(kernel_size, dilation[0]),
)
),
weight_norm(
Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[1],
padding=get_padding(kernel_size, dilation[1]),
)
),
weight_norm(
Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[2],
padding=get_padding(kernel_size, dilation[2]),
)
),
]
)
self.convs1.apply(init_weights)
self.convs2 = nn.ModuleList(
[
weight_norm(
Conv1d(
channels,
channels,
kernel_size,
1,
dilation=1,
padding=get_padding(kernel_size, 1),
)
),
weight_norm(
Conv1d(
channels,
channels,
kernel_size,
1,
dilation=1,
padding=get_padding(kernel_size, 1),
)
),
weight_norm(
Conv1d(
channels,
channels,
kernel_size,
1,
dilation=1,
padding=get_padding(kernel_size, 1),
)
),
]
)
self.convs2.apply(init_weights)
def forward(self, x, x_mask=None):
for c1, c2 in zip(self.convs1, self.convs2):
xt = F.leaky_relu(x, LRELU_SLOPE)
if x_mask is not None:
xt = xt * x_mask
xt = c1(xt)
xt = F.leaky_relu(xt, LRELU_SLOPE)
if x_mask is not None:
xt = xt * x_mask
xt = c2(xt)
x = xt + x
if x_mask is not None:
x = x * x_mask
return x
def remove_weight_norm(self):
for l in self.convs1:
remove_weight_norm(l)
for l in self.convs2:
remove_weight_norm(l)
class ResBlock2(torch.nn.Module):
def __init__(self, channels, kernel_size=3, dilation=(1, 3)):
super(ResBlock2, self).__init__()
self.convs = nn.ModuleList(
[
weight_norm(
Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[0],
padding=get_padding(kernel_size, dilation[0]),
)
),
weight_norm(
Conv1d(
channels,
channels,
kernel_size,
1,
dilation=dilation[1],
padding=get_padding(kernel_size, dilation[1]),
)
),
]
)
self.convs.apply(init_weights)
def forward(self, x, x_mask=None):
for c in self.convs:
xt = F.leaky_relu(x, LRELU_SLOPE)
if x_mask is not None:
xt = xt * x_mask
xt = c(xt)
x = xt + x
if x_mask is not None:
x = x * x_mask
return x
def remove_weight_norm(self):
for l in self.convs:
remove_weight_norm(l)
class Flip(nn.Module):
def forward(self, x, *args, reverse=False, **kwargs):
x = torch.flip(x, [1])
if not reverse:
logdet = torch.zeros(x.size(0)).to(dtype=x.dtype, device=x.device)
return x, logdet
else:
return x
class ResidualCouplingLayer(nn.Module):
def __init__(
self,
channels,
hidden_channels,
kernel_size,
dilation_rate,
n_layers,
p_dropout=0,
gin_channels=0,
mean_only=False,
):
assert channels % 2 == 0, "channels should be divisible by 2"
super().__init__()
self.channels = channels
self.hidden_channels = hidden_channels
self.kernel_size = kernel_size
self.dilation_rate = dilation_rate
self.n_layers = n_layers
self.half_channels = channels // 2
self.mean_only = mean_only
self.pre = nn.Conv1d(self.half_channels, hidden_channels, 1)
self.enc = WN(
hidden_channels,
kernel_size,
dilation_rate,
n_layers,
p_dropout=p_dropout,
gin_channels=gin_channels,
)
self.post = nn.Conv1d(hidden_channels, self.half_channels * (2 - mean_only), 1)
self.post.weight.data.zero_()
self.post.bias.data.zero_()
def forward(self, x, x_mask, g=None, reverse=False):
x0, x1 = torch.split(x, [self.half_channels] * 2, 1)
h = self.pre(x0) * x_mask
h = self.enc(h, x_mask, g=g)
stats = self.post(h) * x_mask
if not self.mean_only:
m, logs = torch.split(stats, [self.half_channels] * 2, 1)
else:
m = stats
logs = torch.zeros_like(m)
if not reverse:
x1 = m + x1 * torch.exp(logs) * x_mask
x = torch.cat([x0, x1], 1)
logdet = torch.sum(logs, [1, 2])
return x, logdet
else:
x1 = (x1 - m) * torch.exp(-logs) * x_mask
x = torch.cat([x0, x1], 1)
return x
class LinearNorm(nn.Module):
def __init__(
self,
in_channels,
out_channels,
bias=True,
spectral_norm=False,
):
super(LinearNorm, self).__init__()
self.fc = nn.Linear(in_channels, out_channels, bias)
if spectral_norm:
self.fc = nn.utils.spectral_norm(self.fc)
def forward(self, input):
out = self.fc(input)
return out
class Mish(nn.Module):
def __init__(self):
super(Mish, self).__init__()
def forward(self, x):
return x * torch.tanh(F.softplus(x))
class Conv1dGLU(nn.Module):
"""
Conv1d + GLU(Gated Linear Unit) with residual connection.
For GLU refer to https://arxiv.org/abs/1612.08083 paper.
"""
def __init__(self, in_channels, out_channels, kernel_size, dropout):
super(Conv1dGLU, self).__init__()
self.out_channels = out_channels
self.conv1 = ConvNorm(in_channels, 2 * out_channels, kernel_size=kernel_size)
self.dropout = nn.Dropout(dropout)
def forward(self, x):
residual = x
x = self.conv1(x)
x1, x2 = torch.split(x, split_size_or_sections=self.out_channels, dim=1)
x = x1 * torch.sigmoid(x2)
x = residual + self.dropout(x)
return x
class ConvNorm(nn.Module):
def __init__(
self,
in_channels,
out_channels,
kernel_size=1,
stride=1,
padding=None,
dilation=1,
bias=True,
spectral_norm=False,
):
super(ConvNorm, self).__init__()
if padding is None:
assert kernel_size % 2 == 1
padding = int(dilation * (kernel_size - 1) / 2)
self.conv = torch.nn.Conv1d(
in_channels,
out_channels,
kernel_size=kernel_size,
stride=stride,
padding=padding,
dilation=dilation,
bias=bias,
)
if spectral_norm:
self.conv = nn.utils.spectral_norm(self.conv)
def forward(self, input):
out = self.conv(input)
return out
class MultiHeadAttention(nn.Module):
"""Multi-Head Attention module"""
def __init__(self, n_head, d_model, d_k, d_v, dropout=0.0, spectral_norm=False):
super().__init__()
self.n_head = n_head
self.d_k = d_k
self.d_v = d_v
self.w_qs = nn.Linear(d_model, n_head * d_k)
self.w_ks = nn.Linear(d_model, n_head * d_k)
self.w_vs = nn.Linear(d_model, n_head * d_v)
self.attention = ScaledDotProductAttention(
temperature=np.power(d_model, 0.5), dropout=dropout
)
self.fc = nn.Linear(n_head * d_v, d_model)
self.dropout = nn.Dropout(dropout)
if spectral_norm:
self.w_qs = nn.utils.spectral_norm(self.w_qs)
self.w_ks = nn.utils.spectral_norm(self.w_ks)
self.w_vs = nn.utils.spectral_norm(self.w_vs)
self.fc = nn.utils.spectral_norm(self.fc)
def forward(self, x, mask=None):
d_k, d_v, n_head = self.d_k, self.d_v, self.n_head
sz_b, len_x, _ = x.size()
residual = x
q = self.w_qs(x).view(sz_b, len_x, n_head, d_k)
k = self.w_ks(x).view(sz_b, len_x, n_head, d_k)
v = self.w_vs(x).view(sz_b, len_x, n_head, d_v)
q = q.permute(2, 0, 1, 3).contiguous().view(-1, len_x, d_k) # (n*b) x lq x dk
k = k.permute(2, 0, 1, 3).contiguous().view(-1, len_x, d_k) # (n*b) x lk x dk
v = v.permute(2, 0, 1, 3).contiguous().view(-1, len_x, d_v) # (n*b) x lv x dv
if mask is not None:
slf_mask = mask.repeat(n_head, 1, 1) # (n*b) x .. x ..
else:
slf_mask = None
output, attn = self.attention(q, k, v, mask=slf_mask)
output = output.view(n_head, sz_b, len_x, d_v)
output = (
output.permute(1, 2, 0, 3).contiguous().view(sz_b, len_x, -1)
) # b x lq x (n*dv)
output = self.fc(output)
output = self.dropout(output) + residual
return output, attn
class ScaledDotProductAttention(nn.Module):
"""Scaled Dot-Product Attention"""
def __init__(self, temperature, dropout):
super().__init__()
self.temperature = temperature
self.softmax = nn.Softmax(dim=2)
self.dropout = nn.Dropout(dropout)
def forward(self, q, k, v, mask=None):
attn = torch.bmm(q, k.transpose(1, 2))
attn = attn / self.temperature
if mask is not None:
attn = attn.masked_fill(mask, -np.inf)
attn = self.softmax(attn)
p_attn = self.dropout(attn)
output = torch.bmm(p_attn, v)
return output, attn
class MelStyleEncoder(nn.Module):
"""MelStyleEncoder"""
def __init__(
self,
n_mel_channels=80,
style_hidden=128,
style_vector_dim=256,
style_kernel_size=5,
style_head=2,
dropout=0.1,
):
super(MelStyleEncoder, self).__init__()
self.in_dim = n_mel_channels
self.hidden_dim = style_hidden
self.out_dim = style_vector_dim
self.kernel_size = style_kernel_size
self.n_head = style_head
self.dropout = dropout
self.spectral = nn.Sequential(
LinearNorm(self.in_dim, self.hidden_dim),
Mish(),
nn.Dropout(self.dropout),
LinearNorm(self.hidden_dim, self.hidden_dim),
Mish(),
nn.Dropout(self.dropout),
)
self.temporal = nn.Sequential(
Conv1dGLU(self.hidden_dim, self.hidden_dim, self.kernel_size, self.dropout),
Conv1dGLU(self.hidden_dim, self.hidden_dim, self.kernel_size, self.dropout),
)
self.slf_attn = MultiHeadAttention(
self.n_head,
self.hidden_dim,
self.hidden_dim // self.n_head,
self.hidden_dim // self.n_head,
self.dropout,
)
self.fc = LinearNorm(self.hidden_dim, self.out_dim)
def temporal_avg_pool(self, x, mask=None):
if mask is None:
out = torch.mean(x, dim=1)
else:
len_ = (~mask).sum(dim=1).unsqueeze(1)
x = x.masked_fill(mask.unsqueeze(-1), 0)
x = x.sum(dim=1)
out = torch.div(x, len_)
return out
def forward(self, x, mask=None):
x = x.transpose(1, 2)
if mask is not None:
mask = (mask.int() == 0).squeeze(1)
max_len = x.shape[1]
slf_attn_mask = (
mask.unsqueeze(1).expand(-1, max_len, -1) if mask is not None else None
)
# spectral
x = self.spectral(x)
# temporal
x = x.transpose(1, 2)
x = self.temporal(x)
x = x.transpose(1, 2)
# self-attention
if mask is not None:
x = x.masked_fill(mask.unsqueeze(-1), 0)
x, _ = self.slf_attn(x, mask=slf_attn_mask)
# fc
x = self.fc(x)
# temoral average pooling
w = self.temporal_avg_pool(x, mask=mask)
return w.unsqueeze(-1)
@@ -0,0 +1,58 @@
import torch
from torch import nn
from torch.nn.utils import remove_weight_norm, weight_norm
from fish_speech.models.vits_decoder.modules.attentions import MultiHeadAttention
class MRTE(nn.Module):
def __init__(
self,
content_enc_channels=192,
hidden_size=512,
out_channels=192,
n_heads=4,
):
super(MRTE, self).__init__()
self.cross_attention = MultiHeadAttention(hidden_size, hidden_size, n_heads)
self.c_pre = nn.Conv1d(content_enc_channels, hidden_size, 1)
self.text_pre = nn.Conv1d(content_enc_channels, hidden_size, 1)
self.c_post = nn.Conv1d(hidden_size, out_channels, 1)
def forward(self, ssl_enc, ssl_mask, text, text_mask, ge, test=None):
if ge == None:
ge = 0
attn_mask = text_mask.unsqueeze(2) * ssl_mask.unsqueeze(-1)
ssl_enc = self.c_pre(ssl_enc * ssl_mask)
text_enc = self.text_pre(text * text_mask)
if test != None:
if test == 0:
x = (
self.cross_attention(
ssl_enc * ssl_mask, text_enc * text_mask, attn_mask
)
+ ssl_enc
+ ge
)
elif test == 1:
x = ssl_enc + ge
elif test == 2:
x = (
self.cross_attention(
ssl_enc * 0 * ssl_mask, text_enc * text_mask, attn_mask
)
+ ge
)
else:
raise ValueError("test should be 0,1,2")
else:
x = (
self.cross_attention(
ssl_enc * ssl_mask, text_enc * text_mask, attn_mask
)
+ ssl_enc
+ ge
)
x = self.c_post(x * ssl_mask)
return x
@@ -0,0 +1,101 @@
import math
import torch
from torch import nn
from fish_speech.models.vqgan.modules.fsq import DownsampleFiniteScalarQuantize
from fish_speech.models.vqgan.modules.wavenet import WaveNet
from fish_speech.models.vqgan.utils import sequence_mask
from fish_speech.utils.spectrogram import LogMelSpectrogram
class VQEncoder(nn.Module):
def __init__(
self,
):
super().__init__()
self.encoder = WaveNet(
input_channels=128,
residual_channels=768,
residual_layers=20,
dilation_cycle=4,
)
self.quantizer = DownsampleFiniteScalarQuantize(
input_dim=768, n_codebooks=1, n_groups=2, levels=[8, 5, 5, 5]
)
self.spec = LogMelSpectrogram(
sample_rate=44100,
n_fft=2048,
win_length=2048,
hop_length=512,
n_mels=128,
f_min=0.0,
f_max=8000.0,
)
self.eval()
e = self.load_state_dict(
torch.load("checkpoints/vq-gan-group-fsq-2x1024.pth", map_location="cpu"),
strict=False,
)
assert len(e.missing_keys) == 0, e.missing_keys
assert all(
k.startswith("decoder.")
or k.startswith("quality_projection.")
or k.startswith("discriminator.")
for k in e.unexpected_keys
), e.unexpected_keys
@torch.no_grad()
def forward(self, audios, audio_lengths, sr=None):
mel_spec = self.spec(audios, sample_rate=sr)
if sr is not None:
audio_lengths = audio_lengths * 44100 // sr
mel_lengths = audio_lengths // self.spec.hop_length
mel_masks = (
torch.arange(mel_spec.shape[2], device=mel_spec.device)
< mel_lengths[:, None]
)
mel_masks_float_conv = mel_masks[:, None, :].float()
mels = mel_spec * mel_masks_float_conv
# Encode
encoded_features = self.encoder(mels) * mel_masks_float_conv
encoded_features = self.quantizer(encoded_features).z * mel_masks_float_conv
return encoded_features
@torch.no_grad()
def indicies_to_vq_features(
self,
indices,
feature_lengths,
):
factor = math.prod(self.quantizer.downsample_factor)
mel_masks = sequence_mask(feature_lengths * factor, indices.shape[2] * factor)
mel_masks_float_conv = mel_masks[:, None, :].float()
z = self.quantizer.decode(indices) * mel_masks_float_conv
return z
@torch.no_grad()
def encode(self, audios, audio_lengths, sr=None):
audios = audios.float()
mels = self.spec(audios, sample_rate=sr)
mel_lengths = audio_lengths // self.spec.hop_length
mel_masks = sequence_mask(mel_lengths, mels.shape[2])
mel_masks_float_conv = mel_masks[:, None, :].float()
mels = mels * mel_masks_float_conv
# Encode
encoded_features = self.encoder(mels) * mel_masks_float_conv
feature_lengths = mel_lengths // math.prod(self.quantizer.downsample_factor)
return self.quantizer.encode(encoded_features), feature_lengths
+2 -5
View File
@@ -1,5 +1,6 @@
import itertools
import re
import string
LANGUAGE_UNICODE_RANGE_MAP = {
"ZH": [(0x4E00, 0x9FFF)],
@@ -18,7 +19,6 @@ SYMBOLS_MAPPING = {
"·": ",",
"、": ",",
"...": "…",
"$": ".",
"“": "'",
"”": "'",
"‘": "'",
@@ -62,12 +62,9 @@ REMOVE_UNKNOWN_SYMBOL_REGEX = re.compile(
def clean_text(text):
# Clean the text
text = text.strip()
# Replace <p:(.*?)> with <PPP(.*?)PPP>
text = re.sub(r"<p:(.*?)>", r"<PPP\1PPP>", text)
# Replace all chinese symbols with their english counterparts
text = REPLACE_SYMBOL_REGEX.sub(lambda x: SYMBOLS_MAPPING[x.group()], text)
text = REMOVE_UNKNOWN_SYMBOL_REGEX.sub("", text)
# Replace <PPP(.*?)PPP> with <p:(.*?)>
text = re.sub(r"<PPP(.*?)PPP>", r"<p:\1>", text)
return text
-1
View File
@@ -9,7 +9,6 @@ from natsort import natsorted
AUDIO_EXTENSIONS = {
".mp3",
".wav",
".WAV",
".flac",
".ogg",
".m4a",
+7 -3
View File
@@ -43,9 +43,13 @@ def print_config_tree(
# add fields from `print_order` to queue
for field in print_order:
queue.append(field) if field in cfg else log.warning(
f"Field '{field}' not found in config. "
+ f"Skipping '{field}' config printing..."
(
queue.append(field)
if field in cfg
else log.warning(
f"Field '{field}' not found in config. "
+ f"Skipping '{field}' config printing..."
)
)
# add all the other fields to queue (not specified in `print_order`)
+8 -6
View File
@@ -1,3 +1,4 @@
import importlib.util
import os
import subprocess
import sys
@@ -17,6 +18,11 @@ GIT = (
GIT = str(GIT)
def is_module_installed(module_name: str) -> bool:
spec = importlib.util.find_spec(module_name)
return spec is not None
@lru_cache()
def commit_hash():
try:
@@ -77,16 +83,12 @@ class Seafoam(Base):
spacing_size: sizes.Size | str = sizes.spacing_md,
radius_size: sizes.Size | str = sizes.radius_md,
text_size: sizes.Size | str = sizes.text_lg,
font: fonts.Font
| str
| Iterable[fonts.Font | str] = (
font: fonts.Font | str | Iterable[fonts.Font | str] = (
fonts.GoogleFont("Quicksand"),
"ui-sans-serif",
"sans-serif",
),
font_mono: fonts.Font
| str
| Iterable[fonts.Font | str] = (
font_mono: fonts.Font | str | Iterable[fonts.Font | str] = (
fonts.GoogleFont("IBM Plex Mono"),
"ui-monospace",
"monospace",
+347 -76
View File
@@ -8,7 +8,6 @@ import shutil
import signal
import subprocess
import sys
import webbrowser
from pathlib import Path
import gradio as gr
@@ -18,7 +17,7 @@ from loguru import logger
from tqdm import tqdm
from fish_speech.i18n import i18n
from fish_speech.webui.launch_utils import Seafoam, versions_html
from fish_speech.webui.launch_utils import Seafoam, is_module_installed, versions_html
PYTHON = os.path.join(os.environ.get("PYTHON_FOLDERPATH", ""), "python")
sys.path.insert(0, "")
@@ -28,6 +27,7 @@ print("You are in ", str(cur_work_dir))
config_path = cur_work_dir / "fish_speech" / "configs"
vqgan_yml_path = config_path / "vqgan_finetune.yaml"
llama_yml_path = config_path / "text2semantic_finetune.yaml"
vits_yml_path = config_path / "vits_decoder_finetune.yaml"
env = os.environ.copy()
env["no_proxy"] = "127.0.0.1, localhost, 0.0.0.0"
@@ -51,6 +51,15 @@ def build_html_ok_message(msg):
"""
def build_html_href(link, desc, msg):
return f"""
<span style="color: green; font-weight: bold; display: inline-block">
{html.escape(msg)}
<a href="{link}">{desc}</a>
</span>
"""
def load_data_in_raw(path):
with open(path, "r", encoding="utf-8") as file:
data = file.read()
@@ -94,21 +103,56 @@ def kill_process(pid):
def change_label(if_label):
global p_label
if if_label == True:
# 设置要访问的URL
url = "https://text-labeler.pages.dev/"
webbrowser.open(url)
yield i18n("Opened labeler in browser")
elif if_label == False:
if if_label == True and p_label is None:
url = "http://localhost:3000"
remote_url = "https://text-labeler.pages.dev/"
try:
p_label = subprocess.Popen(
[
(
"asr-label-linux-x64"
if sys.platform == "linux"
else "asr-label-win-x64.exe"
)
]
)
except FileNotFoundError:
logger.warning("asr-label execution not found!")
yield build_html_href(
link=remote_url,
desc=i18n("Optional online ver"),
msg=i18n("Opened labeler in browser"),
)
elif if_label == False and p_label is not None:
kill_process(p_label.pid)
p_label = None
yield "Nothing"
yield build_html_ok_message("Nothing")
def clean_infer_cache():
import tempfile
temp_dir = Path(tempfile.gettempdir())
gradio_dir = str(temp_dir / "gradio")
try:
shutil.rmtree(gradio_dir)
logger.info(f"Deleted cached audios: {gradio_dir}")
except PermissionError:
logger.info(f"Permission denied: Unable to delete {gradio_dir}")
except FileNotFoundError:
logger.info(f"{gradio_dir} was not found")
except Exception as e:
logger.info(f"An error occurred: {e}")
def change_infer(
if_infer,
host,
port,
infer_vqgan_model,
infer_decoder_model,
infer_decoder_config,
infer_llama_model,
infer_llama_config,
infer_compile,
@@ -124,12 +168,17 @@ def change_infer(
yield build_html_ok_message(
i18n("Inferring interface is launched at {}").format(url)
)
clean_infer_cache()
p_infer = subprocess.Popen(
[
PYTHON,
"tools/webui.py",
"--vqgan-checkpoint-path",
infer_vqgan_model,
"--decoder-checkpoint-path",
infer_decoder_model,
"--decoder-config-name",
infer_decoder_config,
"--llama-checkpoint-path",
infer_llama_model,
"--llama-config-name",
@@ -141,7 +190,7 @@ def change_infer(
env=env,
)
elif if_infer == False and p_infer != None:
elif if_infer == False and p_infer is not None:
kill_process(p_infer.pid)
p_infer = None
yield build_html_error_message(i18n("Infer interface is closed"))
@@ -357,6 +406,8 @@ def check_files(data_path: str, max_depth: int, label_model: str, label_device:
def train_process(
data_path: str,
option: str,
min_duration: float,
max_duration: float,
# vq-gan config
vqgan_ckpt,
vqgan_lr,
@@ -366,6 +417,15 @@ def train_process(
vqgan_data_val_batch_size,
vqgan_precision,
vqgan_check_interval,
# vits config
vits_ckpt,
vits_lr,
vits_maxsteps,
vits_data_num_workers,
vits_data_batch_size,
vits_data_val_batch_size,
vits_precision,
vits_check_interval,
# llama config
llama_ckpt,
llama_base_config,
@@ -393,29 +453,39 @@ def train_process(
print("New Project Name: ", new_project)
if option == "VQGAN" or option == "all":
if min_duration > max_duration:
min_duration, max_duration = max_duration, min_duration
if option == "VQGAN" or option == "VITS":
subprocess.run(
[
PYTHON,
"tools/vqgan/create_train_split.py",
str(data_pre_output.relative_to(cur_work_dir)),
"--min-duration",
str(min_duration),
"--max-duration",
str(max_duration),
]
)
latest = list(
sorted(
[
str(p.relative_to("results"))
for p in Path("results").glob("vqgan_*/")
],
reverse=True,
)
)[0]
if option == "VQGAN":
latest = next(
iter(
sorted(
[
str(p.relative_to("results"))
for p in Path("results").glob("vqgan_*/")
],
reverse=True,
)
),
("vqgan_" + new_project),
)
project = (
("vqgan_" + new_project)
if vqgan_ckpt == "new"
else latest
if vqgan_ckpt == "latest"
else vqgan_ckpt
if vqgan_ckpt == i18n("new")
else latest if vqgan_ckpt == i18n("latest") else vqgan_ckpt
)
logger.info(project)
train_cmd = [
@@ -438,7 +508,49 @@ def train_process(
logger.info(train_cmd)
subprocess.run(train_cmd)
if option == "LLAMA" or option == "all":
if option == "VITS":
latest = next(
iter(
sorted(
[
str(p.relative_to("results"))
for p in Path("results").glob("vits_*/")
],
reverse=True,
)
),
("vits_" + new_project),
)
project = (
("vits_" + new_project)
if vits_ckpt == i18n("new")
else latest if vits_ckpt == i18n("latest") else vits_ckpt
)
ckpt_path = str(Path("checkpoints/vits_decoder_v1.1.ckpt"))
logger.info(project)
train_cmd = [
PYTHON,
"fish_speech/train.py",
"--config-name",
"vits_decoder_finetune",
f"project={project}",
f"ckpt_path={ckpt_path}",
f"trainer.strategy.process_group_backend={backend}",
"tokenizer.pretrained_model_name_or_path=checkpoints",
f"model.optimizer.lr={vits_lr}",
f"trainer.max_steps={vits_maxsteps}",
f"data.num_workers={vits_data_num_workers}",
f"data.batch_size={vits_data_batch_size}",
f"data.val_batch_size={vits_data_val_batch_size}",
f"trainer.precision={vits_precision}",
f"trainer.val_check_interval={vits_check_interval}",
f"train_dataset.filelist={str(data_pre_output / 'vq_train_filelist.txt')}",
f"val_dataset.filelist={str(data_pre_output / 'vq_val_filelist.txt')}",
]
logger.info(train_cmd)
subprocess.run(train_cmd)
if option == "LLAMA":
subprocess.run(
[
PYTHON,
@@ -468,26 +580,27 @@ def train_process(
]
)
ckpt_path = (
"text2semantic-pretrain-medium-2k-v1.pth"
"text2semantic-sft-medium-v1.1-4k.pth"
if llama_base_config == "dual_ar_2_codebook_medium"
else "text2semantic-sft-medium-v1-4k.pth"
else "text2semantic-sft-large-v1.1-4k.pth"
)
latest = list(
sorted(
[
str(p.relative_to("results"))
for p in Path("results").glob("text2sem*/")
],
reverse=True,
)
)[0]
latest = next(
iter(
sorted(
[
str(p.relative_to("results"))
for p in Path("results").glob("text2sem*/")
],
reverse=True,
)
),
("text2semantic_" + new_project),
)
project = (
("text2semantic_" + new_project)
if llama_ckpt == "new"
else latest
if llama_ckpt == "latest"
else llama_ckpt
if llama_ckpt == i18n("new")
else latest if llama_ckpt == i18n("latest") else llama_ckpt
)
logger.info(project)
train_cmd = [
@@ -559,22 +672,33 @@ def fresh_tb_dir():
)
def fresh_vqgan_model():
def fresh_decoder_model():
return gr.Dropdown(
choices=[init_vqgan_yml["ckpt_path"]]
+ [str(Path("checkpoints/vits_decoder_v1.1.ckpt"))]
+ [str(p) for p in Path("results").glob("vqgan*/**/*.ckpt")]
+ [str(p) for p in Path("results").glob("vits*/**/*.ckpt")]
)
def fresh_vqgan_ckpt():
return gr.Dropdown(
choices=["latest", "new"] + [str(p) for p in Path("results").glob("vqgan_*/")]
choices=[i18n("latest"), i18n("new")]
+ [str(p) for p in Path("results").glob("vqgan_*/")]
)
def fresh_vits_ckpt():
return gr.Dropdown(
choices=[i18n("latest"), i18n("new")]
+ [str(p) for p in Path("results").glob("vits_*/")]
)
def fresh_llama_ckpt():
return gr.Dropdown(
choices=["latest", "new"] + [str(p) for p in Path("results").glob("text2sem*/")]
choices=[i18n("latest"), i18n("new")]
+ [str(p) for p in Path("results").glob("text2sem*/")]
)
@@ -585,7 +709,7 @@ def fresh_llama_model():
)
def llama_lora_merge(llama_weight, lora_weight, llama_lora_output):
def llama_lora_merge(llama_weight, lora_llama_config, lora_weight, llama_lora_output):
if (
lora_weight is None
or not Path(lora_weight).exists()
@@ -601,7 +725,7 @@ def llama_lora_merge(llama_weight, lora_weight, llama_lora_output):
PYTHON,
"tools/llama/merge_lora.py",
"--llama-config",
"dual_ar_2_codebook_large",
lora_llama_config,
"--lora-config",
"r_8_alpha_16",
"--llama-weight",
@@ -618,6 +742,7 @@ def llama_lora_merge(llama_weight, lora_weight, llama_lora_output):
init_vqgan_yml = load_yaml_data_in_fact(vqgan_yml_path)
init_llama_yml = load_yaml_data_in_fact(llama_yml_path)
init_vits_yml = load_yaml_data_in_fact(vits_yml_path)
with gr.Blocks(
head="<style>\n" + css + "\n</style>",
@@ -650,6 +775,22 @@ with gr.Blocks(
if_label = gr.Checkbox(
label=i18n("Open Labeler WebUI"), scale=0, show_label=True
)
with gr.Row():
min_duration = gr.Slider(
label=i18n("Minimum Audio Duration"),
value=1.5,
step=0.1,
minimum=0.4,
maximum=30,
)
max_duration = gr.Slider(
label=i18n("Maximum Audio Duration"),
value=30,
step=0.1,
minimum=0.4,
maximum=30,
)
with gr.Row():
add_button = gr.Button(
"\U000027A1 " + i18n("Add to Processing Area"),
@@ -698,17 +839,17 @@ with gr.Blocks(
model_type_radio = gr.Radio(
label=i18n("Select the model to be trained"),
interactive=True,
choices=["VQGAN", "LLAMA", "all"],
value="all",
choices=["VQGAN", "VITS", "LLAMA"],
value="VITS",
)
with gr.Row():
with gr.Tab(label=i18n("VQGAN Configuration")):
with gr.Row(equal_height=False):
vqgan_ckpt = gr.Dropdown(
label="Select VQGAN ckpt",
choices=["latest", "new"]
label=i18n("Select VQGAN ckpt"),
choices=[i18n("latest"), i18n("new")]
+ [str(p) for p in Path("results").glob("vqgan_*/")],
value="latest",
value=i18n("latest"),
interactive=True,
)
with gr.Row(equal_height=False):
@@ -775,6 +916,79 @@ with gr.Blocks(
value=init_vqgan_yml["trainer"]["val_check_interval"],
)
with gr.Tab(label=i18n("VITS Configuration")):
with gr.Row(equal_height=False):
vits_ckpt = gr.Dropdown(
label=i18n("Select VITS ckpt"),
choices=[i18n("latest"), i18n("new")]
+ [str(p) for p in Path("results").glob("vits_*/")],
value=i18n("latest"),
interactive=True,
)
with gr.Row(equal_height=False):
vits_lr_slider = gr.Slider(
label=i18n("Initial Learning Rate"),
interactive=True,
minimum=1e-5,
maximum=1e-4,
step=1e-5,
value=init_vits_yml["model"]["optimizer"]["lr"],
)
vits_maxsteps_slider = gr.Slider(
label=i18n("Maximum Training Steps"),
interactive=True,
minimum=1000,
maximum=100000,
step=1000,
value=init_vits_yml["trainer"]["max_steps"],
)
with gr.Row(equal_height=False):
vits_data_num_workers_slider = gr.Slider(
label=i18n("Number of Workers"),
interactive=True,
minimum=1,
maximum=16,
step=1,
value=init_vits_yml["data"]["num_workers"],
)
vits_data_batch_size_slider = gr.Slider(
label=i18n("Batch Size"),
interactive=True,
minimum=1,
maximum=32,
step=1,
value=init_vits_yml["data"]["batch_size"],
)
with gr.Row(equal_height=False):
vits_data_val_batch_size_slider = gr.Slider(
label=i18n("Validation Batch Size"),
interactive=True,
minimum=1,
maximum=32,
step=1,
value=init_vits_yml["data"]["val_batch_size"],
)
vits_precision_dropdown = gr.Dropdown(
label=i18n("Precision"),
interactive=True,
choices=["32", "bf16-mixed"],
info=i18n(
"bf16-true is recommended for 30+ series GPU, 16-mixed is recommended for 10+ series GPU"
),
value=str(init_vits_yml["trainer"]["precision"]),
)
with gr.Row(equal_height=False):
vits_check_interval_slider = gr.Slider(
label=i18n("Save model every n steps"),
interactive=True,
minimum=500,
maximum=10000,
step=500,
value=init_vits_yml["trainer"]["val_check_interval"],
)
with gr.Tab(label=i18n("LLAMA Configuration")):
with gr.Row(equal_height=False):
llama_use_lora = gr.Checkbox(
@@ -785,10 +999,10 @@ with gr.Blocks(
value=True,
)
llama_ckpt = gr.Dropdown(
label="Select LLAMA ckpt",
choices=["latest", "new"]
label=i18n("Select LLAMA ckpt"),
choices=[i18n("latest"), i18n("new")]
+ [str(p) for p in Path("results").glob("text2sem*/")],
value="latest",
value=i18n("latest"),
interactive=True,
)
with gr.Row(equal_height=False):
@@ -822,9 +1036,11 @@ with gr.Blocks(
minimum=0,
maximum=16,
step=1,
value=init_llama_yml["data"]["num_workers"]
if sys.platform == "linux"
else 0,
value=(
init_llama_yml["data"]["num_workers"]
if sys.platform == "linux"
else 0
),
)
with gr.Row(equal_height=False):
llama_data_batch_size_slider = gr.Slider(
@@ -902,6 +1118,16 @@ with gr.Blocks(
allow_custom_value=True,
interactive=True,
)
lora_llama_config = gr.Dropdown(
label=i18n("LLAMA Model Config"),
info=i18n("Type the path or select from the dropdown"),
choices=[
"dual_ar_2_codebook_large",
"dual_ar_2_codebook_medium",
],
value="dual_ar_2_codebook_large",
allow_custom_value=True,
)
with gr.Row(equal_height=False):
llama_lora_output = gr.Dropdown(
label=i18n("Output Path"),
@@ -956,18 +1182,39 @@ with gr.Blocks(
label=i18n("WebUI Port"), value="7862"
)
with gr.Row():
infer_vqgan_model = gr.Dropdown(
label=i18n("VQGAN Model Path"),
infer_decoder_model = gr.Dropdown(
label=i18n("Decoder Model Path"),
info=i18n(
"Type the path or select from the dropdown"
),
value=init_vqgan_yml["ckpt_path"],
value=str(
Path("checkpoints/vits_decoder_v1.1.ckpt")
),
choices=[init_vqgan_yml["ckpt_path"]]
+ [str(Path("checkpoints/vits_decoder_v1.1.ckpt"))]
+ [
str(p)
for p in Path("results").glob(
"vqgan*/**/*.ckpt"
)
]
+ [
str(p)
for p in Path("results").glob("vits*/**/*.ckpt")
],
allow_custom_value=True,
)
infer_decoder_config = gr.Dropdown(
label=i18n("Decoder Model Config"),
info=i18n(
"Type the path or select from the dropdown"
),
value="vits_decoder_finetune",
choices=[
"vits_decoder_finetune",
"vits_decoder_pretrain",
"vqgan_finetune",
"vqgan_pretrain",
],
allow_custom_value=True,
)
@@ -987,6 +1234,18 @@ with gr.Blocks(
],
allow_custom_value=True,
)
infer_llama_config = gr.Dropdown(
label=i18n("LLAMA Model Config"),
info=i18n(
"Type the path or select from the dropdown"
),
choices=[
"dual_ar_2_codebook_large",
"dual_ar_2_codebook_medium",
],
value="dual_ar_2_codebook_large",
allow_custom_value=True,
)
with gr.Row():
infer_compile = gr.Radio(
label=i18n("Compile Model"),
@@ -994,16 +1253,15 @@ with gr.Blocks(
"Compile the model can significantly reduce the inference time, but will increase cold start time"
),
choices=["Yes", "No"],
value="Yes",
)
infer_llama_config = gr.Dropdown(
label=i18n("LLAMA Model Config"),
choices=[
"dual_ar_2_codebook_large",
"dual_ar_2_codebook_medium",
],
value="dual_ar_2_codebook_large",
allow_custom_value=True,
value=(
"Yes"
if (
sys.platform == "linux"
or is_module_installed("triton")
)
else "No"
),
interactive=is_module_installed("triton"),
)
with gr.Row():
@@ -1084,6 +1342,8 @@ with gr.Blocks(
inputs=[
train_box,
model_type_radio,
min_duration,
max_duration,
# vq-gan config
vqgan_ckpt,
vqgan_lr_slider,
@@ -1093,6 +1353,15 @@ with gr.Blocks(
vqgan_data_val_batch_size_slider,
vqgan_precision_dropdown,
vqgan_check_interval_slider,
# vits config
vits_ckpt,
vits_lr_slider,
vits_maxsteps_slider,
vits_data_num_workers_slider,
vits_data_batch_size_slider,
vits_data_val_batch_size_slider,
vits_precision_dropdown,
vits_check_interval_slider,
# llama config
llama_ckpt,
llama_base_config,
@@ -1115,8 +1384,8 @@ with gr.Blocks(
outputs=[train_error],
)
tb_dir.change(fn=fresh_tb_dir, inputs=[], outputs=[tb_dir])
infer_vqgan_model.change(
fn=fresh_vqgan_model, inputs=[], outputs=[infer_vqgan_model]
infer_decoder_model.change(
fn=fresh_decoder_model, inputs=[], outputs=[infer_decoder_model]
)
infer_llama_model.change(
fn=fresh_llama_model, inputs=[], outputs=[infer_llama_model]
@@ -1131,10 +1400,11 @@ with gr.Blocks(
fn=new_explorer, inputs=[train_box, tree_slider], outputs=[file_markdown]
)
vqgan_ckpt.change(fn=fresh_vqgan_ckpt, inputs=[], outputs=[vqgan_ckpt])
vits_ckpt.change(fn=fresh_vits_ckpt, inputs=[], outputs=[vits_ckpt])
llama_ckpt.change(fn=fresh_llama_ckpt, inputs=[], outputs=[llama_ckpt])
llama_lora_merge_btn.click(
fn=llama_lora_merge,
inputs=[llama_weight, lora_weight, llama_lora_output],
inputs=[llama_weight, lora_llama_config, lora_weight, llama_lora_output],
outputs=[train_error],
)
infer_checkbox.change(
@@ -1143,7 +1413,8 @@ with gr.Blocks(
infer_checkbox,
infer_host_textbox,
infer_port_textbox,
infer_vqgan_model,
infer_decoder_model,
infer_decoder_config,
infer_llama_model,
infer_llama_config,
infer_compile,
+4 -2
View File
@@ -219,7 +219,7 @@ class FishSpeech_INFER:
"prompt_audio": ("AUDIO",),
"text":("STRING",{
"multiline": True,
"default": "你好,世界!"
"default": "你好啊,世界!"
}),
"prompt_text_by_srt":("SRT",{
"multiline": True,
@@ -296,13 +296,15 @@ class FishSpeech_INFER:
hf_hub_download(repo_id="fishaudio/fish-speech-1",filename="tokenizer_config.json",local_dir=checkpoint_path,token=hf_token)
hf_hub_download(repo_id="fishaudio/fish-speech-1",filename="special_tokens_map.json",local_dir=checkpoint_path,token=hf_token)
npy_path = os.path.join(fish_tmp_out, os.path.basename(prompt_audio))
python_exec = sys.executable or "python"
npy_path = os.path.join(fish_tmp_out, os.path.basename(prompt_audio))
step_1 = f"{python_exec} {parent_directory}/tools/vqgan/inference.py -i {prompt_audio} -o {npy_path} -ckpt {vq_model_path} -d {device}"
print("step 1 ",step_1)
p = Popen(step_1,shell=True)
p.wait()
config_name = f"dual_ar_2_codebook_{text2semantic_type}"
npy_path = os.path.join(fish_tmp_out, os.path.basename(prompt_audio)[:-4]+".npy")
step_2 = f'{python_exec} {parent_directory}/tools/llama/generate.py --text "{text}" --prompt-text "{prompt_text}" \
+1 -1
View File
@@ -65,7 +65,7 @@ def load_model(config_name, checkpoint_path, device="cuda"):
def main(input_path, output_path, config_name, checkpoint_path, device):
model = load_model(config_name, checkpoint_path, device=device)
if input_path.suffix in AUDIO_EXTENSIONS:
if input_path.suffix.lower() in AUDIO_EXTENSIONS:
logger.info(f"Processing in-place reconstruction of {input_path}")
# Load audio