Improve: transc order by path, train.list preserves order

This commit is contained in:
litagin02
2024-05-26 10:47:14 +09:00
parent d8ca829e28
commit e52453040c
2 changed files with 18 additions and 8 deletions

View File

@@ -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:

View File

@@ -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")