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/generate_dataset.py

101 lines
3.4 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。它是训练前最上游的数据生产入口。
"""
# 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()