Alignment-Lab-AI/ttestv0.1
1115
1import torch2import torchaudio3import time4import os5import numpy as np6import json7from datasets import load_dataset, Audio8from snac import SNAC9from torch.nn import functional as F10from tqdm import tqdm11import wandb12 13# Constants14SNAC_SAMPLE_RATE = 2400015OUTPUT_DIR = "processed_common_voice"16BATCH_SIZE = 100017 18# Ensure CUDA is available19device = torch.device("cuda" if torch.cuda.is_available() else "cpu")20 21def load_snac_model(sample_rate):22 if sample_rate == 24000:23 model = SNAC.from_pretrained("hubertsiuzdak/snac_24khz").eval().to(device)24 else:25 raise ValueError("Unsupported sample rate. Please use 24000.")26 return model27 28snac_model = load_snac_model(SNAC_SAMPLE_RATE)29 30def chunk_and_pad_audio(audio, chunk_size):31 length = audio.shape[-1]32 padded_length = ((length + chunk_size - 1) // chunk_size) * chunk_size33 padded_audio = F.pad(audio, (0, padded_length - length), mode="constant", value=0)34 batched_audio = padded_audio.unfold(-1, size=chunk_size, step=chunk_size)35 return batched_audio36 37def generate_snac_encoding(audio):38 waveform = torch.tensor(audio["array"]).float().to(device)39 if audio["sampling_rate"] != SNAC_SAMPLE_RATE:40 resampler = torchaudio.transforms.Resample(41 orig_freq=audio["sampling_rate"], new_freq=SNAC_SAMPLE_RATE42 )43 waveform = resampler(waveform)44 45 if waveform.dim() == 2:46 waveform = waveform.mean(dim=0, keepdim=True)47 elif waveform.dim() == 1:48 waveform = waveform.unsqueeze(0)49 50 num_second = 151 chunk_size_initial = num_second * SNAC_SAMPLE_RATE52 lcm = np.lcm.reduce([snac_model.vq_strides[0], snac_model.attn_window_size or 1])53 pad_to = snac_model.hop_length * lcm54 chunk_size = int(np.ceil(chunk_size_initial / pad_to) * pad_to)55 56 audio = chunk_and_pad_audio(waveform, chunk_size)57 audio = audio.permute(1, 0, 2)58 59 codes_list = []60 with torch.no_grad():61 for chunk in audio:62 codes = snac_model.encode(chunk.unsqueeze(0))63 codes = [c.cpu() for c in codes]64 codes_list.append(codes)65 66 codes_list = [torch.cat(codes_list, dim=0) for codes_list in zip(*codes_list)]67 codes_list = [code.reshape(-1).cpu().tolist() for code in codes_list]68 69 string_codes = " ".join(map(str, codes_list[0]))70 return string_codes71 72def process_audio(item):73 start_time = time.time()74 try:75 snac_tokens = generate_snac_encoding(item["audio"])76 77 if not snac_tokens:78 raise ValueError("Generated SNAC tokens are empty")79 80 except Exception as e:81 return None82 83 processing_time = time.time() - start_time84 85 return {86 "path": item["path"],87 "sentence": item["sentence"],88 "age": item["age"],89 "gender": item["gender"],90 "accent": item["accent"],91 "locale": item["locale"],92 "snac": snac_tokens,93 "processing_time": processing_time,94 "audio_duration": len(item["audio"]["array"]) / item["audio"]["sampling_rate"],95 }96 97def save_to_jsonl(data, file_path):98 # Open the file in append mode to add new data to the existing language-specific JSONL file99 with open(file_path, "a") as f:100 for item in data:101 json.dump(item, f)102 f.write("\n")103 104def process_language(language):105 # Ensure output directory exists106 language_dir = os.path.join(OUTPUT_DIR, language)107 os.makedirs(language_dir, exist_ok=True)108 jsonl_path = os.path.join(language_dir, f"{language}_processed.jsonl")109 110 # Read existing data111 existing_data = set()112 if os.path.exists(jsonl_path):113 with open(jsonl_path, "r") as f:114 existing_data = set(f.readlines())115 116 # Load the Common Voice dataset for this language117 dataset = load_dataset(118 "mozilla-foundation/common_voice_16_1", language, split="train", streaming=True119 )120 121 # Cast the dataset to include audio122 dataset = dataset.cast_column("audio", Audio(sampling_rate=SNAC_SAMPLE_RATE))123 124 processed_data = []125 total_processed = 0126 report_counter = 0127 128 for item in tqdm(dataset, desc=f"Processing {language}"):129 result = process_audio(item)130 if result:131 json_line = json.dumps(result) + "\n"132 if json_line not in existing_data:133 processed_data.append(result)134 existing_data.add(json_line)135 total_processed += 1136 report_counter += 1137 138 if report_counter % 1000 == 0: # Report to wandb every 1000 rows139 wandb.log(140 {141 "language": language,142 "average_processing_time": np.mean(143 [item["processing_time"] for item in processed_data]144 ),145 "average_audio_duration": np.mean(146 [item["audio_duration"] for item in processed_data]147 ),148 "average_snac_token_count": np.mean(149 [len(item["snac"].split()) for item in processed_data]150 ),151 }152 )153 report_counter = 0 # Reset the counter154 155 # Save every BATCH_SIZE items156 if len(processed_data) >= BATCH_SIZE:157 save_to_jsonl(processed_data, jsonl_path)158 processed_data = [] # Clear the list after saving159 160 # Save any remaining processed data161 if processed_data:162 save_to_jsonl(processed_data, jsonl_path)163 164 return total_processed165 166def main():167 # Initialize wandb168 wandb.init(project="common-voice-processing", job_type="data-processing")169 170 # List of languages to process, starting with English171 languages = ['ckb', 'cnh', 'cs', 'cv', 'cy', 'da', 'de']172 # languages = ['dv', 'dyu', 'el', 'en', 'eo', 'es', 'et']173 # languages = ['eu', 'fa', 'fi', 'fr', 'fy-NL', 'ga-IE', 'gl']174 # languages = ['gn', 'ha', 'he', 'hi', 'hsb', 'hu', 'hy-AM']175 # languages = ['ia', 'id', 'ig', 'is', 'it', 'ja', 'ka']176 # languages = ['kab', 'kk', 'kmr', 'ko', 'ky', 'lg', 'lij']177 # languages = ['lo', 'lt', 'ltg', 'lv', 'mdf', 'mhr', 'mk']178 # languages = ['ml', 'mn', 'mr', 'mrj', 'mt', 'myv', 'nan-tw']179 # languages = ['ne-NP', 'nhi', 'nl', 'nn-NO', 'oc', 'or', 'os']180 # languages = ['pa-IN', 'pl', 'ps', 'pt', 'quy', 'rm-sursilv', 'rm-vallader']181 # languages = ['ro', 'ru', 'rw', 'sah', 'sat', 'sc', 'sk']182 # languages = ['skr', 'sl', 'sq', 'sr', 'sv-SE', 'sw', 'ta']183 # languages = ['te', 'th', 'ti', 'tig', 'tk', 'tok', 'tr']184 # languages = ['tt', 'tw', 'ug', 'uk', 'ur', 'uz', 'vi', 'vot', 'yi', 'yo', 'yue', 'zgh', 'zh-CN', 'zh-HK', 'zh-TW']185 186 total_processed_all_languages = 0187 188 # Process each language189 for language in languages:190 total_processed = process_language(language)191 total_processed_all_languages += total_processed192 193 print(194 f"\nCompleted processing all languages. Total files processed across all languages: {total_processed_all_languages}"195 )196 197 wandb.finish()198 199if __name__ == "__main__":200 main()201 