You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
nmWTAI-Platform/ML/nmWTAI-ML/scripts/merge_two_h5.py

84 lines
3.0 KiB
Python

"""完整合并两份结构兼容的普通 HDF5 数据集。
本脚本只是面向普通数据集的命令行入口schema 校验分批复制随机混排和来源
记录仍复用 merge_datasets.py 的实现 normal/hard 混合模式不同这里会保留
两个输入文件中的全部样本不按比例抽样也不执行标准化
"""
from __future__ import annotations
import argparse
import sys
from pathlib import Path
import h5py
ROOT = Path(__file__).resolve().parents[1]
sys.path.append(str(ROOT))
from scripts.merge_datasets import merge_datasets
def parse_args() -> argparse.Namespace:
"""Parse two ordinary HDF5 inputs and the merged output path."""
parser = argparse.ArgumentParser(
description="Merge all rows from two compatible generated HDF5 datasets"
)
parser.add_argument("--input-a", required=True, type=str)
parser.add_argument("--input-b", required=True, type=str)
parser.add_argument("--output", required=True, type=str)
parser.add_argument("--label-a", default=None, type=str)
parser.add_argument("--label-b", default=None, type=str)
parser.add_argument("--seed", default=42, type=int)
parser.add_argument("--batch-size", default=4096, type=int)
return parser.parse_args()
def row_count(path: Path) -> int:
"""读取 params 的行数params 是各字段共同使用的样本数基准。"""
if not path.exists():
raise FileNotFoundError(f"Input file not found: {path}")
with h5py.File(path, "r") as source:
if "params" not in source:
raise KeyError(f"HDF5 dataset is missing required key 'params': {path}")
return int(source["params"].shape[0])
def main() -> None:
"""校验两个输入,保留所有样本,并写出一份合并后的 HDF5。"""
args = parse_args()
input_a = Path(args.input_a).resolve()
input_b = Path(args.input_b).resolve()
output = Path(args.output).resolve()
if output in {input_a, input_b}:
raise ValueError("--output must differ from both input files")
count_a = row_count(input_a)
count_b = row_count(input_b)
if count_a <= 0 or count_b <= 0:
raise ValueError(f"Both inputs must contain samples: input_a={count_a}, input_b={count_b}")
# merge_datasets 的 normal/hard 是历史参数名,在这里仅代表输入 A/B。
# 显式传入两个文件的完整行数,可确保双方全部保留而不是按比例抽样。
result = merge_datasets(
normal_input=input_a,
hard_input=input_b,
output=output,
tag=output.stem,
normal_count=count_a,
hard_count=count_b,
seed=args.seed,
normal_label=args.label_a or input_a.stem,
hard_label=args.label_b or input_b.stem,
batch_size=args.batch_size,
)
print(f"Merged dataset written to: {result['output_path']}")
print(f"input_a={count_a}, input_b={count_b}, total={result['total_samples']}")
print(f"Merge summary written to: {result['summary_path']}")
if __name__ == "__main__":
main()