From 6d6cd7c0b00f07a057d9b799eeca80866a0d5ff7 Mon Sep 17 00:00:00 2001 From: AIFSH <1509359472@qq.com> Date: Mon, 13 May 2024 06:55:17 +0000 Subject: [PATCH] fix bug --- .../configs/text2semantic_finetune.yaml | 4 +- .../configs/vits_decoder_finetune.yaml | 128 ++++ .../configs/vits_decoder_pretrain.yaml | 127 ++++ fish_speech/configs/vqgan_finetune.yaml | 4 +- fish_speech/configs/vqgan_pretrain.yaml | 4 +- fish_speech/datasets/concat_repeat.py | 52 ++ fish_speech/datasets/vits.py | 195 +++++ fish_speech/i18n/locale/en_US.json | 15 +- fish_speech/i18n/locale/es_ES.json | 15 +- fish_speech/i18n/locale/ja_JP.json | 17 +- fish_speech/i18n/locale/zh_CN.json | 15 +- fish_speech/models/vits_decoder/__init__.py | 3 + fish_speech/models/vits_decoder/lit_module.py | 394 ++++++++++ fish_speech/models/vits_decoder/losses.py | 67 ++ .../models/vits_decoder/modules/attentions.py | 350 +++++++++ .../models/vits_decoder/modules/commons.py | 190 +++++ .../models/vits_decoder/modules/models.py | 686 ++++++++++++++++++ .../models/vits_decoder/modules/modules.py | 619 ++++++++++++++++ .../models/vits_decoder/modules/mrte.py | 58 ++ .../models/vits_decoder/modules/vq_encoder.py | 101 +++ fish_speech/text/clean.py | 7 +- fish_speech/utils/file.py | 1 - fish_speech/utils/rich_utils.py | 10 +- .../{models/vqgan => utils}/spectrogram.py | 0 fish_speech/webui/launch_utils.py | 14 +- fish_speech/webui/manage.py | 423 +++++++++-- nodes.py | 6 +- tools/vqgan/inference.py | 2 +- 28 files changed, 3398 insertions(+), 109 deletions(-) create mode 100644 fish_speech/configs/vits_decoder_finetune.yaml create mode 100644 fish_speech/configs/vits_decoder_pretrain.yaml create mode 100644 fish_speech/datasets/concat_repeat.py create mode 100644 fish_speech/datasets/vits.py create mode 100644 fish_speech/models/vits_decoder/__init__.py create mode 100644 fish_speech/models/vits_decoder/lit_module.py create mode 100644 fish_speech/models/vits_decoder/losses.py create mode 100644 fish_speech/models/vits_decoder/modules/attentions.py create mode 100644 fish_speech/models/vits_decoder/modules/commons.py create mode 100644 fish_speech/models/vits_decoder/modules/models.py create mode 100644 fish_speech/models/vits_decoder/modules/modules.py create mode 100644 fish_speech/models/vits_decoder/modules/mrte.py create mode 100644 fish_speech/models/vits_decoder/modules/vq_encoder.py rename fish_speech/{models/vqgan => utils}/spectrogram.py (100%) diff --git a/fish_speech/configs/text2semantic_finetune.yaml b/fish_speech/configs/text2semantic_finetune.yaml index f347f58..41683e1 100644 --- a/fish_speech/configs/text2semantic_finetune.yaml +++ b/fish_speech/configs/text2semantic_finetune.yaml @@ -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 diff --git a/fish_speech/configs/vits_decoder_finetune.yaml b/fish_speech/configs/vits_decoder_finetune.yaml new file mode 100644 index 0000000..ffdb9d3 --- /dev/null +++ b/fish_speech/configs/vits_decoder_finetune.yaml @@ -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 diff --git a/fish_speech/configs/vits_decoder_pretrain.yaml b/fish_speech/configs/vits_decoder_pretrain.yaml new file mode 100644 index 0000000..4139f1a --- /dev/null +++ b/fish_speech/configs/vits_decoder_pretrain.yaml @@ -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 diff --git a/fish_speech/configs/vqgan_finetune.yaml b/fish_speech/configs/vqgan_finetune.yaml index 991bcbd..e0ee971 100644 --- a/fish_speech/configs/vqgan_finetune.yaml +++ b/fish_speech/configs/vqgan_finetune.yaml @@ -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} diff --git a/fish_speech/configs/vqgan_pretrain.yaml b/fish_speech/configs/vqgan_pretrain.yaml index c6bb9ef..a195de3 100644 --- a/fish_speech/configs/vqgan_pretrain.yaml +++ b/fish_speech/configs/vqgan_pretrain.yaml @@ -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} diff --git a/fish_speech/datasets/concat_repeat.py b/fish_speech/datasets/concat_repeat.py new file mode 100644 index 0000000..ce2671c --- /dev/null +++ b/fish_speech/datasets/concat_repeat.py @@ -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)] diff --git a/fish_speech/datasets/vits.py b/fish_speech/datasets/vits.py new file mode 100644 index 0000000..4b4f5d3 --- /dev/null +++ b/fish_speech/datasets/vits.py @@ -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 diff --git a/fish_speech/i18n/locale/en_US.json b/fish_speech/i18n/locale/en_US.json index 7b51eee..04fe469 100644 --- a/fish_speech/i18n/locale/en_US.json +++ b/fish_speech/i18n/locale/en_US.json @@ -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" } diff --git a/fish_speech/i18n/locale/es_ES.json b/fish_speech/i18n/locale/es_ES.json index 74fa490..718a735 100644 --- a/fish_speech/i18n/locale/es_ES.json +++ b/fish_speech/i18n/locale/es_ES.json @@ -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" } diff --git a/fish_speech/i18n/locale/ja_JP.json b/fish_speech/i18n/locale/ja_JP.json index dc79c92..c5aacdb 100644 --- a/fish_speech/i18n/locale/ja_JP.json +++ b/fish_speech/i18n/locale/ja_JP.json @@ -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": "新規" } diff --git a/fish_speech/i18n/locale/zh_CN.json b/fish_speech/i18n/locale/zh_CN.json index 0b1ab26..fb46176 100644 --- a/fish_speech/i18n/locale/zh_CN.json +++ b/fish_speech/i18n/locale/zh_CN.json @@ -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": "创建新的检查点" } diff --git a/fish_speech/models/vits_decoder/__init__.py b/fish_speech/models/vits_decoder/__init__.py new file mode 100644 index 0000000..fccd6fb --- /dev/null +++ b/fish_speech/models/vits_decoder/__init__.py @@ -0,0 +1,3 @@ +from .lit_module import VITSDecoder + +__all__ = ["VITSDecoder"] diff --git a/fish_speech/models/vits_decoder/lit_module.py b/fish_speech/models/vits_decoder/lit_module.py new file mode 100644 index 0000000..f7d4d06 --- /dev/null +++ b/fish_speech/models/vits_decoder/lit_module.py @@ -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) diff --git a/fish_speech/models/vits_decoder/losses.py b/fish_speech/models/vits_decoder/losses.py new file mode 100644 index 0000000..52e2f37 --- /dev/null +++ b/fish_speech/models/vits_decoder/losses.py @@ -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 diff --git a/fish_speech/models/vits_decoder/modules/attentions.py b/fish_speech/models/vits_decoder/modules/attentions.py new file mode 100644 index 0000000..4a3037f --- /dev/null +++ b/fish_speech/models/vits_decoder/modules/attentions.py @@ -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 diff --git a/fish_speech/models/vits_decoder/modules/commons.py b/fish_speech/models/vits_decoder/modules/commons.py new file mode 100644 index 0000000..0fb7d32 --- /dev/null +++ b/fish_speech/models/vits_decoder/modules/commons.py @@ -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 diff --git a/fish_speech/models/vits_decoder/modules/models.py b/fish_speech/models/vits_decoder/modules/models.py new file mode 100644 index 0000000..0706814 --- /dev/null +++ b/fish_speech/models/vits_decoder/modules/models.py @@ -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) diff --git a/fish_speech/models/vits_decoder/modules/modules.py b/fish_speech/models/vits_decoder/modules/modules.py new file mode 100644 index 0000000..3d54770 --- /dev/null +++ b/fish_speech/models/vits_decoder/modules/modules.py @@ -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) diff --git a/fish_speech/models/vits_decoder/modules/mrte.py b/fish_speech/models/vits_decoder/modules/mrte.py new file mode 100644 index 0000000..8bb84d4 --- /dev/null +++ b/fish_speech/models/vits_decoder/modules/mrte.py @@ -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 diff --git a/fish_speech/models/vits_decoder/modules/vq_encoder.py b/fish_speech/models/vits_decoder/modules/vq_encoder.py new file mode 100644 index 0000000..c05d0ea --- /dev/null +++ b/fish_speech/models/vits_decoder/modules/vq_encoder.py @@ -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 diff --git a/fish_speech/text/clean.py b/fish_speech/text/clean.py index f2c8e72..0a8fdd3 100644 --- a/fish_speech/text/clean.py +++ b/fish_speech/text/clean.py @@ -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 with - text = re.sub(r"", r"", 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 with - text = re.sub(r"", r"", text) return text diff --git a/fish_speech/utils/file.py b/fish_speech/utils/file.py index cb4ba2a..4047aa5 100644 --- a/fish_speech/utils/file.py +++ b/fish_speech/utils/file.py @@ -9,7 +9,6 @@ from natsort import natsorted AUDIO_EXTENSIONS = { ".mp3", ".wav", - ".WAV", ".flac", ".ogg", ".m4a", diff --git a/fish_speech/utils/rich_utils.py b/fish_speech/utils/rich_utils.py index c8f77c7..6a465f5 100644 --- a/fish_speech/utils/rich_utils.py +++ b/fish_speech/utils/rich_utils.py @@ -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`) diff --git a/fish_speech/models/vqgan/spectrogram.py b/fish_speech/utils/spectrogram.py similarity index 100% rename from fish_speech/models/vqgan/spectrogram.py rename to fish_speech/utils/spectrogram.py diff --git a/fish_speech/webui/launch_utils.py b/fish_speech/webui/launch_utils.py index 659e994..2f57b59 100644 --- a/fish_speech/webui/launch_utils.py +++ b/fish_speech/webui/launch_utils.py @@ -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", diff --git a/fish_speech/webui/manage.py b/fish_speech/webui/manage.py index 8e25a7d..baf6f4b 100644 --- a/fish_speech/webui/manage.py +++ b/fish_speech/webui/manage.py @@ -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""" + + {html.escape(msg)} + {desc} + + """ + + 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="", @@ -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, diff --git a/nodes.py b/nodes.py index 348a55b..9b5e225 100644 --- a/nodes.py +++ b/nodes.py @@ -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}" \ diff --git a/tools/vqgan/inference.py b/tools/vqgan/inference.py index 33662cc..2f84cf5 100644 --- a/tools/vqgan/inference.py +++ b/tools/vqgan/inference.py @@ -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