zsy814's picture
Initial EchoLoc code release
d8bfe4a verified
Raw
History Blame Contribute Delete
3.8 kB
import json
import os
import random
from collections import defaultdict
# ==========================================
# 配置部分
# ==========================================
INPUT_FILE_PATH = "/workspace/echoloc/datas/Emotion_Speech_Dataset/mixed_emotion_dataset.json"
# 定义需要的切分任务
# 格式: ("文件名后缀", 比例)
# 注意:所有比例之和不能超过 1.0
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
# 1. 读取数据
print(f"正在读取主文件...")
with open(INPUT_FILE_PATH, 'r', encoding='utf-8') as f:
all_data = json.load(f)
print(f"读取完成。总样本数: {len(all_data)}")
# 2. 分桶:按长度将数据归类
print("正在按长度分桶并打乱...")
buckets = defaultdict(list)
for item in all_data:
seq_len = len(item['segments'])
buckets[seq_len].append(item)
sorted_lengths = sorted(buckets.keys())
# 3. 桶内打乱 (Shuffle)
# 这是关键步骤:先彻底打乱,后面顺序切片也就等于随机切片了
for l in sorted_lengths:
print(f" Length {l}: {len(buckets[l])} 条 -> 打乱中...")
random.shuffle(buckets[l])
# 4. 依次切分 (Slicing)
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]
# 初始化每个桶的起始游标 (Offset),初始都是 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
# 切片:取 [start : end] 这一段
chunk = buckets[l][start_idx:end_idx]
subset_data.extend(chunk)
stats_info.append(f"Len{l}:{len(chunk)}")
# 更新游标,供下一个任务使用
bucket_offsets[l] = end_idx
# 再次打乱 subset_data,避免数据是按长度排序的 (222...333...444...)
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()