Nos_StyleTTS2-Brais-GL / scripts /download_data.py
cmagui's picture
Initial commit: full repository with code, configs and weights
613ce86
Raw
History Blame Contribute Delete
3.83 kB
import torch
import yaml
import torchaudio
import os
import argparse
import resampy
import sys
from datasets import load_dataset, get_dataset_config_names
from sklearn.model_selection import train_test_split
from tqdm import tqdm
sys.path.insert(0, os.path.join(os.path.dirname(__file__), '..'))
from Utils.ASR.AuxiliaryASR.phonemize import run_cotovia_with_phrase, clean_output
def get_speaker_number(speaker, speaker_map={}):
if speaker not in speaker_map:
speaker_map[speaker] = len(speaker_map)
return speaker_map, speaker_map[speaker]
def download_data(dataset_name: str, output_folder: str = "Data", ood=False, target_sr=24000, download_data=False):
speaker_map = {}
speaker = 1
splits = ["train", "val", "test"]
BASE = f"proxectonos/{dataset_name}"
name = get_dataset_config_names(BASE)[0]
ds_train_full = load_dataset(
BASE, data_files=f"{name}_train.csv", sep="\t")["train"]
ds_test = load_dataset(
BASE, data_files=f"{name}_test.csv", sep="\t")["train"]
train_idx, val_idx = train_test_split(
range(len(ds_train_full)), test_size=0.1, random_state=42
)
ds_train = ds_train_full.select(train_idx)
ds_val = ds_train_full.select(val_idx)
datasets = {
"train": ds_train,
"val": ds_val,
"test": ds_test
}
open_mode = "a" if ood else "w"
audios_output_folder = os.path.join(
output_folder, dataset_name, "audios")
os.makedirs(audios_output_folder, exist_ok=True)
for split in splits:
split_ds = datasets[split]
print(f"Processing split: {split}, number of samples: {len(split_ds)}")
if ood:
output_txt = os.path.join(
output_folder, f"OOD_texts.txt")
else:
output_txt = os.path.join(
output_folder, f"{split}.txt")
with open(output_txt, open_mode, encoding="utf-8") as f:
for item in tqdm(split_ds):
try:
file_path = os.path.join(
audios_output_folder, item["file_name"])
if download_data:
waveform, sr = torchaudio.load(item["audio"])
if sr != target_sr:
waveform = resampy.resample(
waveform.numpy(), sr, target_sr)
waveform = torch.from_numpy(waveform)
torchaudio.save(file_path, waveform, target_sr)
# fonemizar el texto normalizado
phonemized_text = clean_output(run_cotovia_with_phrase(
str(item["normalized"])))
f.write(f"{file_path}|{phonemized_text}|{speaker}\n")
except Exception as e:
print(
f"Error processing sample: {e}")
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description="Download and process datasets for TTS training.")
parser.add_argument("--config", type=str, required=True,
default="Configs/download_data.yaml",
help="Path to the model configuration file.")
parser.add_argument("--download_data", action="store_true",
help="Whether to download the data or not.")
args = parser.parse_args()
config = yaml.safe_load(open(args.config, "r"))
datasets_to_download = {'OOD': config['OOD']['dataset'],
'data': config['data']['dataset']}
output_folder = os.path.join("Data")
for key, dataset in datasets_to_download.items():
ood = (key == 'OOD')
download_data(dataset, output_folder, ood=ood,
target_sr=config['target_sr'], download_data=args.download_data)