"""批量生成正演代理模型的原始数值试井数据集。 脚本读取数据生成配置,调用并行数据集生成器批量采样地层/井筒参数和流量制度, 运行底层数值求解器并把有效曲线写入 HDF5。它是训练前最上游的数据生产入口。 """ # pylint: disable=import-error,wrong-import-position,broad-exception-caught from __future__ import annotations import argparse import multiprocessing as mp import sys from pathlib import Path ROOT = Path(__file__).resolve().parents[1] sys.path.append(str(ROOT)) from src.common.config import Config from src.common.experiment_paths import config_for_stage from src.data.dataset_generation import ParallelDatasetGenerator def main() -> None: """按配置阶段启动并行数值试井样本生成,输出原始 HDF5 数据集路径。""" parser = argparse.ArgumentParser() parser.add_argument("--config", default=None) parser.add_argument( "--stage", choices=[ "fixed_case", "case_neighborhood", "family_random", "family_random_hard", "family_random_v2_q", ], default=None, ) parser.add_argument("--n-samples", type=int, default=None) parser.add_argument("--n-workers", type=int, default=None) parser.add_argument("--seed", type=int, default=None) parser.add_argument("--method", type=str, default=None) parser.add_argument( "--dataset-case", type=str, default=None, help="Generate only the named dataset_cases entry; omit to generate all cases", ) parser.add_argument( "--dataset-tag", type=str, default=None, help="Optional tag injected into output dataset filename", ) args = parser.parse_args() config_path = args.config if config_path is None: config_path = str(config_for_stage(args.stage) or Path("configs/data_gen.yaml")) # 没有dataset_cases时保持原来的单case行为;有case时按声明顺序依次生成。 base_cfg = Config(config_path) case_names = [str(case.get("name", "")).strip() for case in base_cfg.dataset_cases] if any(not name for name in case_names) or len(case_names) != len(set(case_names)): raise ValueError("dataset_cases中的name必须非空且不能重复") if args.dataset_case is None: selected_cases = case_names or [None] elif args.dataset_case not in case_names: available = ", ".join(case_names) if case_names else "(none)" parser.error(f"unknown --dataset-case {args.dataset_case!r}; available: {available}") else: selected_cases = [args.dataset_case] for case_name in selected_cases: cfg = Config(config_path, dataset_case=case_name) cfg.ensure_dirs() if case_name is None: output_tag = args.dataset_tag else: output_tag = f"{args.dataset_tag}_{case_name}" if args.dataset_tag else case_name print(f"\n=== Generating dataset case: {case_name} ===") path = ParallelDatasetGenerator( cfg=cfg, n_workers=args.n_workers, ).generate( n_samples=args.n_samples, method=args.method, random_seed=args.seed, dataset_tag=output_tag, ) print(path) if __name__ == "__main__": mp.freeze_support() try: mp.set_start_method("spawn", force=True) except RuntimeError: pass main()