From edca6a67f1dfdc43a1f5d79bcf38d68ade58b680 Mon Sep 17 00:00:00 2001 From: lvjunjie Date: Fri, 17 Jul 2026 17:08:30 +0800 Subject: [PATCH] =?UTF-8?q?=E4=B8=A4=E4=B8=AA=E6=99=AE=E9=80=9A=E6=95=B0?= =?UTF-8?q?=E6=8D=AE=E9=9B=86=E5=90=88=E5=B9=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ML/nmWTAI-ML/scripts/merge_datasets.py | 1 + ML/nmWTAI-ML/scripts/merge_two_h5.py | 83 ++++++++++++++++++++++++++ 2 files changed, 84 insertions(+) create mode 100644 ML/nmWTAI-ML/scripts/merge_two_h5.py diff --git a/ML/nmWTAI-ML/scripts/merge_datasets.py b/ML/nmWTAI-ML/scripts/merge_datasets.py index 67512b2..020f2c7 100644 --- a/ML/nmWTAI-ML/scripts/merge_datasets.py +++ b/ML/nmWTAI-ML/scripts/merge_datasets.py @@ -27,6 +27,7 @@ sys.path.append(str(ROOT)) REQUIRED_DATASETS = ("params", "schedule", "curve") OPTIONAL_DATASETS = ( "group_id", + "curve_time", "schedule_meta", "family_name", "is_anchor", diff --git a/ML/nmWTAI-ML/scripts/merge_two_h5.py b/ML/nmWTAI-ML/scripts/merge_two_h5.py new file mode 100644 index 0000000..3e17ca7 --- /dev/null +++ b/ML/nmWTAI-ML/scripts/merge_two_h5.py @@ -0,0 +1,83 @@ +"""完整合并两份结构兼容的普通 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()