From f140b7dafeab7f49e6352a1e91b25c249645f21a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stardust=C2=B7=E5=87=8F?= <2225664821@qq.com> Date: Tue, 29 Aug 2023 10:08:30 +0800 Subject: [PATCH] =?UTF-8?q?=E6=9B=B4=E6=96=B0=E9=A2=84=E5=A4=84=E7=90=86?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- preprocess_text.py | 37 ++++++++++++++++++------------------- 1 file changed, 18 insertions(+), 19 deletions(-) diff --git a/preprocess_text.py b/preprocess_text.py index 6d965ee..45dc61a 100644 --- a/preprocess_text.py +++ b/preprocess_text.py @@ -38,25 +38,24 @@ if 2 in stage: if spk not in spk_id_map.keys(): spk_id_map[spk] = current_sid current_sid += 1 - # - # train_list = [] - # val_list = [] - # - # for spk, utts in spk_utt_map.items(): - # shuffle(utts) - # val_list+=utts[:val_per_spk] - # train_list+=utts[val_per_spk:] - # if len(val_list) > max_val_total: - # train_list+=val_list[max_val_total:] - # val_list = val_list[:max_val_total] - # - # with open( train_path,"w", encoding='utf-8') as f: - # for line in train_list: - # f.write(line) - # - # with open(val_path, "w", encoding='utf-8') as f: - # for line in val_list: - # f.write(line) + train_list = [] + val_list = [] + + for spk, utts in spk_utt_map.items(): + shuffle(utts) + val_list+=utts[:val_per_spk] + train_list+=utts[val_per_spk:] + if len(val_list) > max_val_total: + train_list+=val_list[max_val_total:] + val_list = val_list[:max_val_total] + + with open( train_path,"w", encoding='utf-8') as f: + for line in train_list: + f.write(line) + + with open(val_path, "w", encoding='utf-8') as f: + for line in val_list: + f.write(line) if 3 in stage: assert 2 in stage