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