CoolFace
Apppublic

honey126/VoxAI

sourceHugging Facemitupdated 9mo agoView on Hugging Face
0likes
prepare_csv_wavs.py139 linesDownload Raw Back to scripts
1import sys
2import os
3
4sys.path.append(os.getcwd())
5
6from pathlib import Path
7import json
8import shutil
9import argparse
10
11import csv
12import torchaudio
13from tqdm import tqdm
14from datasets.arrow_writer import ArrowWriter
15
16from model.utils import (
17    convert_char_to_pinyin,
18)
19
20PRETRAINED_VOCAB_PATH = Path(__file__).parent.parent / "data/Emilia_ZH_EN_pinyin/vocab.txt"
21
22
23def is_csv_wavs_format(input_dataset_dir):
24    fpath = Path(input_dataset_dir)
25    metadata = fpath / "metadata.csv"
26    wavs = fpath / "wavs"
27    return metadata.exists() and metadata.is_file() and wavs.exists() and wavs.is_dir()
28
29
30def prepare_csv_wavs_dir(input_dir):
31    assert is_csv_wavs_format(input_dir), f"not csv_wavs format: {input_dir}"
32    input_dir = Path(input_dir)
33    metadata_path = input_dir / "metadata.csv"
34    audio_path_text_pairs = read_audio_text_pairs(metadata_path.as_posix())
35
36    sub_result, durations = [], []
37    vocab_set = set()
38    polyphone = True
39    for audio_path, text in audio_path_text_pairs:
40        if not Path(audio_path).exists():
41            print(f"audio {audio_path} not found, skipping")
42            continue
43        audio_duration = get_audio_duration(audio_path)
44        # assume tokenizer = "pinyin"  ("pinyin" | "char")
45        text = convert_char_to_pinyin([text], polyphone=polyphone)[0]
46        sub_result.append({"audio_path": audio_path, "text": text, "duration": audio_duration})
47        durations.append(audio_duration)
48        vocab_set.update(list(text))
49
50    return sub_result, durations, vocab_set
51
52
53def get_audio_duration(audio_path):
54    audio, sample_rate = torchaudio.load(audio_path)
55    num_channels = audio.shape[0]
56    return audio.shape[1] / (sample_rate * num_channels)
57
58
59def read_audio_text_pairs(csv_file_path):
60    audio_text_pairs = []
61
62    parent = Path(csv_file_path).parent
63    with open(csv_file_path, mode="r", newline="", encoding="utf-8") as csvfile:
64        reader = csv.reader(csvfile, delimiter="|")
65        next(reader)  # Skip the header row
66        for row in reader:
67            if len(row) >= 2:
68                audio_file = row[0].strip()  # First column: audio file path
69                text = row[1].strip()  # Second column: text
70                audio_file_path = parent / audio_file
71                audio_text_pairs.append((audio_file_path.as_posix(), text))
72
73    return audio_text_pairs
74
75
76def save_prepped_dataset(out_dir, result, duration_list, text_vocab_set, is_finetune):
77    out_dir = Path(out_dir)
78    # save preprocessed dataset to disk
79    out_dir.mkdir(exist_ok=True, parents=True)
80    print(f"\nSaving to {out_dir} ...")
81
82    # dataset = Dataset.from_dict({"audio_path": audio_path_list, "text": text_list, "duration": duration_list})  # oom
83    # dataset.save_to_disk(f"data/{dataset_name}/raw", max_shard_size="2GB")
84    raw_arrow_path = out_dir / "raw.arrow"
85    with ArrowWriter(path=raw_arrow_path.as_posix(), writer_batch_size=1) as writer:
86        for line in tqdm(result, desc="Writing to raw.arrow ..."):
87            writer.write(line)
88
89    # dup a json separately saving duration in case for DynamicBatchSampler ease
90    dur_json_path = out_dir / "duration.json"
91    with open(dur_json_path.as_posix(), "w", encoding="utf-8") as f:
92        json.dump({"duration": duration_list}, f, ensure_ascii=False)
93
94    # vocab map, i.e. tokenizer
95    # add alphabets and symbols (optional, if plan to ft on de/fr etc.)
96    # if tokenizer == "pinyin":
97    #     text_vocab_set.update([chr(i) for i in range(32, 127)] + [chr(i) for i in range(192, 256)])
98    voca_out_path = out_dir / "vocab.txt"
99    with open(voca_out_path.as_posix(), "w") as f:
100        for vocab in sorted(text_vocab_set):
101            f.write(vocab + "\n")
102
103    if is_finetune:
104        file_vocab_finetune = PRETRAINED_VOCAB_PATH.as_posix()
105        shutil.copy2(file_vocab_finetune, voca_out_path)
106    else:
107        with open(voca_out_path, "w") as f:
108            for vocab in sorted(text_vocab_set):
109                f.write(vocab + "\n")
110
111    dataset_name = out_dir.stem
112    print(f"\nFor {dataset_name}, sample count: {len(result)}")
113    print(f"For {dataset_name}, vocab size is: {len(text_vocab_set)}")
114    print(f"For {dataset_name}, total {sum(duration_list)/3600:.2f} hours")
115
116
117def prepare_and_save_set(inp_dir, out_dir, is_finetune: bool = True):
118    if is_finetune:
119        assert PRETRAINED_VOCAB_PATH.exists(), f"pretrained vocab.txt not found: {PRETRAINED_VOCAB_PATH}"
120    sub_result, durations, vocab_set = prepare_csv_wavs_dir(inp_dir)
121    save_prepped_dataset(out_dir, sub_result, durations, vocab_set, is_finetune)
122
123
124def cli():
125    # finetune: python scripts/prepare_csv_wavs.py /path/to/input_dir /path/to/output_dir_pinyin
126    # pretrain: python scripts/prepare_csv_wavs.py /path/to/output_dir_pinyin --pretrain
127    parser = argparse.ArgumentParser(description="Prepare and save dataset.")
128    parser.add_argument("inp_dir", type=str, help="Input directory containing the data.")
129    parser.add_argument("out_dir", type=str, help="Output directory to save the prepared data.")
130    parser.add_argument("--pretrain", action="store_true", help="Enable for new pretrain, otherwise is a fine-tune")
131
132    args = parser.parse_args()
133
134    prepare_and_save_set(args.inp_dir, args.out_dir, is_finetune=not args.pretrain)
135
136
137if __name__ == "__main__":
138    cli()
139