| import json |
| import os |
| import random |
| from collections import defaultdict |
|
|
| |
| |
| |
| INPUT_FILE_PATH = "/workspace/echoloc/datas/Emotion_Speech_Dataset/mixed_emotion_dataset.json" |
|
|
| |
| |
| |
| SPLIT_TASKS = [ |
| ("1000", 0.0001), |
| ("10000", 0.001), |
| ("100000", 0.01), |
| ("1000000", 0.1), |
| ("20pct", 0.2), |
| ("50pct", 0.5) |
| ] |
|
|
| |
| |
| |
| def split_dataset_no_overlap(): |
| if not os.path.exists(INPUT_FILE_PATH): |
| print(f"错误: 找不到文件 {INPUT_FILE_PATH}") |
| return |
|
|
| |
| print(f"正在读取主文件...") |
| with open(INPUT_FILE_PATH, 'r', encoding='utf-8') as f: |
| all_data = json.load(f) |
| print(f"读取完成。总样本数: {len(all_data)}") |
|
|
| |
| print("正在按长度分桶并打乱...") |
| buckets = defaultdict(list) |
| for item in all_data: |
| seq_len = len(item['segments']) |
| buckets[seq_len].append(item) |
| |
| sorted_lengths = sorted(buckets.keys()) |
| |
| |
| |
| for l in sorted_lengths: |
| print(f" Length {l}: {len(buckets[l])} 条 -> 打乱中...") |
| random.shuffle(buckets[l]) |
|
|
| |
| dir_name = os.path.dirname(INPUT_FILE_PATH) |
| base_name = os.path.basename(INPUT_FILE_PATH) |
| file_name_no_ext = os.path.splitext(base_name)[0] |
|
|
| |
| bucket_offsets = {l: 0 for l in sorted_lengths} |
|
|
| print("\n>>> 开始无重叠切分任务") |
| |
| for suffix, ratio in SPLIT_TASKS: |
| subset_data = [] |
| stats_info = [] |
| |
| print(f"\n正在生成 [{suffix}] (占比 {ratio*100}%)...") |
| |
| for l in sorted_lengths: |
| total_in_bucket = len(buckets[l]) |
| count_to_take = int(total_in_bucket * ratio) |
| |
| start_idx = bucket_offsets[l] |
| end_idx = start_idx + count_to_take |
| |
| |
| if end_idx > total_in_bucket: |
| print(f" [警告] 长度 {l} 的数据不够了!需要至 {end_idx},但只有 {total_in_bucket}。") |
| end_idx = total_in_bucket |
| |
| |
| chunk = buckets[l][start_idx:end_idx] |
| subset_data.extend(chunk) |
| |
| stats_info.append(f"Len{l}:{len(chunk)}") |
| |
| |
| bucket_offsets[l] = end_idx |
| |
| |
| random.shuffle(subset_data) |
| |
| |
| output_filename = f"{file_name_no_ext}_{suffix}.json" |
| output_path = os.path.join(dir_name, output_filename) |
| |
| print(f" - 组成分布: {', '.join(stats_info)}") |
| print(f" - 总样本数: {len(subset_data)}") |
| print(f" - 保存路径: {output_path}") |
| |
| with open(output_path, 'w', encoding='utf-8') as f: |
| json.dump(subset_data, f, ensure_ascii=False, indent=2) |
|
|
| |
| print("\n所有任务完成。") |
| print("剩余未使用的数据量 (按长度):") |
| for l in sorted_lengths: |
| remaining = len(buckets[l]) - bucket_offsets[l] |
| print(f" Length {l}: {remaining} 条") |
|
|
| if __name__ == "__main__": |
| split_dataset_no_overlap() |
|
|