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

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

"""完整合并两份结构兼容的普通 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()