Improve: transc order by path, train.list preserves order
This commit is contained in:
@@ -2,7 +2,7 @@ import argparse
|
|||||||
import json
|
import json
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from random import shuffle
|
from random import shuffle, sample
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
@@ -156,16 +156,26 @@ def preprocess(
|
|||||||
train_list: list[str] = []
|
train_list: list[str] = []
|
||||||
val_list: list[str] = []
|
val_list: list[str] = []
|
||||||
|
|
||||||
# 各話者ごとにシャッフルして、val_per_lang個をval_listに、残りをtrain_listに追加
|
# 各話者ごとに発話リストを処理
|
||||||
for spk, utts in spk_utt_map.items():
|
for spk, utts in spk_utt_map.items():
|
||||||
shuffle(utts)
|
if val_per_lang == 0:
|
||||||
val_list += utts[:val_per_lang]
|
train_list.extend(utts)
|
||||||
train_list += utts[val_per_lang:]
|
continue
|
||||||
|
# ランダムにval_per_lang個のインデックスを選択
|
||||||
|
val_indices = set(sample(range(len(utts)), val_per_lang))
|
||||||
|
# 元の順序を保ちながらリストを分割
|
||||||
|
for index, utt in enumerate(utts):
|
||||||
|
if index in val_indices:
|
||||||
|
val_list.append(utt)
|
||||||
|
else:
|
||||||
|
train_list.append(utt)
|
||||||
|
|
||||||
shuffle(val_list)
|
# バリデーションリストのサイズ調整
|
||||||
if len(val_list) > max_val_total:
|
if len(val_list) > max_val_total:
|
||||||
train_list += val_list[max_val_total:]
|
extra_val = val_list[max_val_total:]
|
||||||
val_list = val_list[:max_val_total]
|
val_list = val_list[:max_val_total]
|
||||||
|
# 余剰のバリデーション発話をトレーニングリストに追加(元の順序を保持)
|
||||||
|
train_list.extend(extra_val)
|
||||||
|
|
||||||
with train_path.open("w", encoding="utf-8") as f:
|
with train_path.open("w", encoding="utf-8") as f:
|
||||||
for line in train_list:
|
for line in train_list:
|
||||||
|
|||||||
@@ -152,7 +152,7 @@ if __name__ == "__main__":
|
|||||||
output_file.parent.mkdir(parents=True, exist_ok=True)
|
output_file.parent.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
wav_files = [f for f in input_dir.rglob("*.wav") if f.is_file()]
|
wav_files = [f for f in input_dir.rglob("*.wav") if f.is_file()]
|
||||||
wav_files = sorted(wav_files, key=lambda x: x.name)
|
wav_files = sorted(wav_files, key=lambda x: str(x))
|
||||||
|
|
||||||
if output_file.exists():
|
if output_file.exists():
|
||||||
logger.warning(f"{output_file} exists, backing up to {output_file}.bak")
|
logger.warning(f"{output_file} exists, backing up to {output_file}.bak")
|
||||||
|
|||||||
Reference in New Issue
Block a user