diff --git a/ML/nmWTAI-ML/scripts/flatten_autofit_neighborhood_dataset.py b/ML/nmWTAI-ML/scripts/flatten_autofit_neighborhood_dataset.py index 5db05e7..658b702 100644 --- a/ML/nmWTAI-ML/scripts/flatten_autofit_neighborhood_dataset.py +++ b/ML/nmWTAI-ML/scripts/flatten_autofit_neighborhood_dataset.py @@ -20,6 +20,17 @@ ROOT = Path(__file__).resolve().parents[1] sys.path.append(str(ROOT)) +EXTRA_SCHEDULE_META_NAMES = [ + "q_end_start_ratio", + "q_range_ratio", + "max_step_ratio", + "mean_step_ratio", + "trend_slope", + "is_nonincreasing", + "is_strong_decline", +] + + def parse_args() -> argparse.Namespace: """解析自动拟合邻域 HDF5 的输入输出路径以及是否只导出邻域样本。""" parser = argparse.ArgumentParser( @@ -71,15 +82,21 @@ def main() -> None: anchor_params = np.asarray(src["anchor_params"][:], dtype=np.float32) anchor_schedule = np.asarray(src["anchor_schedule"][:], dtype=np.float32) anchor_curve = np.asarray(src["anchor_curve"][:], dtype=np.float32) + anchor_curve_time = np.asarray(src["anchor_curve_time"][:], dtype=np.float32) + if "anchor_curve_valid_mask" not in src or "neighbor_curve_valid_mask" not in src: + raise ValueError("Neighborhood H5 is missing fixed-time curve_valid_mask datasets") + anchor_curve_valid_mask = np.asarray(src["anchor_curve_valid_mask"][:], dtype=np.uint8) anchor_schedule_meta = np.asarray(src["anchor_schedule_meta"][:], dtype=np.float32) anchor_family_name = np.asarray(src["anchor_family_name"][:]).astype(str) anchor_section_index = np.asarray(src["anchor_section_index"][:], dtype=np.int32) - anchor_time_q_json = np.asarray(src["anchor_time_q_json"][:]).astype(str) + anchor_time_q_json = np.asarray(src["anchor_timeQ_json"][:]).astype(str) anchor_q_json = np.asarray(src["anchor_q_json"][:]).astype(str) neighbor_anchor_id = np.asarray(src["neighbor_anchor_id"][:], dtype=np.int32) neighbor_params = np.asarray(src["neighbor_params"][:], dtype=np.float32) neighbor_curve = np.asarray(src["neighbor_curve"][:], dtype=np.float32) + neighbor_curve_time = np.asarray(src["neighbor_curve_time"][:], dtype=np.float32) + neighbor_curve_valid_mask = np.asarray(src["neighbor_curve_valid_mask"][:], dtype=np.uint8) neighbor_objective = np.asarray(src["neighbor_objective"][:], dtype=np.float32) neighbor_objective_p = np.asarray(src["neighbor_objective_p"][:], dtype=np.float32) neighbor_objective_d = np.asarray(src["neighbor_objective_d"][:], dtype=np.float32) @@ -100,6 +117,20 @@ def main() -> None: .tolist() ) + missing_meta_names = [ + name for name in EXTRA_SCHEDULE_META_NAMES + if schedule_meta_names is not None and name not in schedule_meta_names + ] + if missing_meta_names: + # 旧局部数据把这些统计量附在制度向量末尾,展开时补回独立元数据列。 + if len(missing_meta_names) != len(EXTRA_SCHEDULE_META_NAMES): + raise ValueError(f"Partially missing schedule metadata fields: {missing_meta_names}") + anchor_schedule_meta = np.concatenate( + [anchor_schedule_meta, anchor_schedule[:, -len(missing_meta_names):]], + axis=1, + ).astype(np.float32) + schedule_meta_names = schedule_meta_names + missing_meta_names + n_anchors = int(anchor_params.shape[0]) n_neighbors = int(neighbor_params.shape[0]) @@ -125,11 +156,14 @@ def main() -> None: param_dim = int(anchor_params.shape[1]) schedule_dim = int(anchor_schedule.shape[1]) curve_dim = int(anchor_curve.shape[1]) + curve_time_dim = int(anchor_curve_time.shape[1]) schedule_meta_dim = int(anchor_schedule_meta.shape[1]) dst.create_dataset("params", shape=(total_rows, param_dim), dtype=np.float32) dst.create_dataset("schedule", shape=(total_rows, schedule_dim), dtype=np.float32) dst.create_dataset("curve", shape=(total_rows, curve_dim), dtype=np.float32) + dst.create_dataset("curve_time", shape=(total_rows, curve_time_dim), dtype=np.float32) + dst.create_dataset("curve_valid_mask", shape=(total_rows, curve_time_dim), dtype=np.uint8) dst.create_dataset("group_id", shape=(total_rows,), dtype=np.int64) dst.create_dataset("schedule_meta", shape=(total_rows, schedule_meta_dim), dtype=np.float32) dst.create_dataset( @@ -163,6 +197,8 @@ def main() -> None: dst["params"][write_pos:anchor_end] = anchor_params dst["schedule"][write_pos:anchor_end] = anchor_schedule dst["curve"][write_pos:anchor_end] = anchor_curve + dst["curve_time"][write_pos:anchor_end] = anchor_curve_time + dst["curve_valid_mask"][write_pos:anchor_end] = anchor_curve_valid_mask dst["group_id"][write_pos:anchor_end] = np.arange(n_anchors, dtype=np.int32) dst["schedule_meta"][write_pos:anchor_end] = anchor_schedule_meta dst["family_name"][write_pos:anchor_end] = anchor_family_name.tolist() @@ -181,6 +217,8 @@ def main() -> None: dst["params"][write_pos:neighbor_end] = neighbor_params dst["schedule"][write_pos:neighbor_end] = anchor_schedule[neighbor_anchor_id] dst["curve"][write_pos:neighbor_end] = neighbor_curve + dst["curve_time"][write_pos:neighbor_end] = neighbor_curve_time + dst["curve_valid_mask"][write_pos:neighbor_end] = neighbor_curve_valid_mask dst["group_id"][write_pos:neighbor_end] = neighbor_anchor_id dst["schedule_meta"][write_pos:neighbor_end] = anchor_schedule_meta[neighbor_anchor_id] dst["family_name"][write_pos:neighbor_end] = anchor_family_name[ diff --git a/ML/nmWTAI-ML/scripts/generate_autofit_neighborhood_dataset.py b/ML/nmWTAI-ML/scripts/generate_autofit_neighborhood_dataset.py index c48e901..9f3095f 100644 --- a/ML/nmWTAI-ML/scripts/generate_autofit_neighborhood_dataset.py +++ b/ML/nmWTAI-ML/scripts/generate_autofit_neighborhood_dataset.py @@ -19,7 +19,9 @@ from __future__ import annotations import argparse +import csv import json +import math import sys from collections import Counter from pathlib import Path @@ -35,8 +37,9 @@ from src.common.experiment_paths import config_for_stage, normalize_tag from src.data.curve_processing import ( clean_curve_for_dataset, is_valid_curve, - resample_curve_to_features_with_time, + resample_curve_to_features_with_time_and_mask, ) +from src.data.dataset_generation import _sample_schedule_family_random as sample_schedule_family_random from src.data.params import Params, Schedule, generate_params_dataset from src.data.runner_client import CppRunner, read_result_bin from src.data.schedule_features import ( @@ -63,36 +66,6 @@ SCHEDULE_META_NAMES = [ ] -def _pick_mixture(rng: np.random.RandomState, items: list[dict]) -> dict: - """按配置中的概率权重从多个采样组件里抽取一个组件。""" - probs = np.asarray([float(it.get("prob", 0.0)) for it in items], dtype=np.float64) - s = float(np.sum(probs)) - if s <= 0: - return items[int(rng.randint(0, len(items)))] - probs = probs / s - u = float(rng.rand()) - c = 0.0 - for it, p in zip(items, probs): - c += float(p) - if u <= c: - return it - return items[-1] - - -def _normalize_durations_to_total(dt: np.ndarray, total: float, min_dt: float) -> np.ndarray: - """在满足最小时长约束的前提下,将各段时长缩放到指定总时长。""" - dt = np.maximum(dt.astype(np.float64), float(min_dt)) - s = float(np.sum(dt)) - total = max(float(total), float(min_dt) * float(len(dt))) - if s <= 0: - return np.full_like(dt, total / len(dt)) - dt = dt * (total / s) - dt = np.maximum(dt, float(min_dt)) - diff = total - float(np.sum(dt)) - dt[-1] += diff - return np.maximum(dt, float(min_dt)) - - def _family_id_map(cfg: Config) -> dict[str, int]: """把流量制度族名称映射为整数 id,便于写入模型特征和元数据。""" mode = str(cfg.raw["schedule"]["generation_mode"]).lower() @@ -184,59 +157,12 @@ def _sample_schedule_family_random( rng: np.random.RandomState, family_override: str | None = None, ) -> tuple[list[float], list[float], dict]: - """按 family_random 配置随机生成不同类别的生产/关井制度。""" - fcfg = cfg.raw["schedule"]["family_random"] - if family_override is None: - fam = _pick_mixture(rng, fcfg["families"]) - fam_name = str(fam.get("name", "inc_tail_shutin")).lower() - else: - fam_name = str(family_override).lower() - - n_lo, n_hi = fcfg["n_prod_sections_range"] - n_prod = int(rng.randint(int(n_lo), int(n_hi) + 1)) - - prod_total_lo, prod_total_hi = fcfg["prod_total_time_range"] - prod_total = float(prod_total_lo + rng.rand() * (prod_total_hi - prod_total_lo)) - - mu = float(fcfg["duration_lognormal_mu"]) - sigma = float(fcfg["duration_lognormal_sigma"]) - dt_prod = rng.lognormal(mean=mu, sigma=sigma, size=n_prod).astype(np.float64) - dt_prod = _normalize_durations_to_total(dt_prod, total=prod_total, min_dt=0.05) - - q_lo, q_hi = fcfg["q_range"] - q0 = float(q_lo + rng.rand() * (q_hi - q_lo)) - max_rel_step = float(fcfg["max_rel_step"]) - mult_noise_sigma = float(fcfg["mult_noise_sigma"]) - step_jump_lo, step_jump_hi = fcfg["step_jump_rel_range"] - - q_prod = np.zeros((n_prod,), dtype=np.float64) - q_prod[0] = q0 - - if fam_name == "inc_tail_shutin": - for i in range(1, n_prod): - q_prod[i] = q_prod[i - 1] * rng.uniform(1.02, max_rel_step) - q_prod = np.maximum.accumulate(q_prod) - elif fam_name == "dec_tail_shutin": - for i in range(1, n_prod): - q_prod[i] = max(q_prod[i - 1] * rng.uniform(1.0 / max_rel_step, 0.98), q_lo) - q_prod = np.minimum.accumulate(q_prod) - elif fam_name == "mild_step_tail_shutin": - q_prod[:] = q0 - jump_idx = int(rng.randint(1, max(2, n_prod))) - q_prod[jump_idx:] *= float(rng.uniform(step_jump_lo, step_jump_hi)) - else: - for i in range(1, n_prod): - rel = rng.uniform(1.0 / max_rel_step, max_rel_step) - rel = 1.0 + 0.15 * (rel - 1.0) - q_prod[i] = q_prod[i - 1] * rel - - if mult_noise_sigma > 0: - q_prod *= np.exp(rng.normal(loc=0.0, scale=mult_noise_sigma, size=n_prod)) - q_prod = np.clip(q_prod, q_lo, q_hi) - - shut_lo, shut_hi = fcfg["shutin_dt_range"] - dt_shut = float(shut_lo + rng.rand() * (shut_hi - shut_lo)) - return dt_prod.tolist() + [dt_shut], q_prod.tolist() + [0.0], {"family_name": fam_name} + """复用正式数据生成器,使局部数据与普通数据的流量制度分布保持一致。""" + return sample_schedule_family_random( + cfg, + rng, + family_override=family_override, + ) def sample_schedule_by_mode( @@ -275,6 +201,7 @@ def parse_args() -> argparse.Namespace: """解析自动拟合邻域数据生成所需的锚点数量、扰动尺度和输出路径。""" parser = argparse.ArgumentParser(description="Generate anchor-neighborhood autofit dataset for generalized local ranking") parser.add_argument("--config", type=str, default=None) + parser.add_argument("--dataset-case", type=str, default=None) parser.add_argument( "--stage", choices=[ @@ -292,6 +219,13 @@ def parse_args() -> argparse.Namespace: parser.add_argument("--neighbors-per-anchor", type=int, default=24) parser.add_argument("--max-attempts-factor", type=int, default=4) parser.add_argument("--anchor-max-attempts-factor", type=int, default=5) + parser.add_argument( + "--trace-csv", + type=str, + default=None, + help="Optional full-solver PSO trace used to seed representative anchor locations", + ) + parser.add_argument("--trace-anchor-count", type=int, default=16) parser.add_argument("--seed", type=int, default=42) parser.add_argument("--span-frac", type=float, default=0.08) parser.add_argument( @@ -311,6 +245,18 @@ def parse_args() -> argparse.Namespace: ) parser.add_argument("--solver-timeout", type=int, default=120) parser.add_argument("--well-index", type=int, default=0) + parser.add_argument( + "--fixed-cf", + type=float, + default=None, + help="Fix Cf for every anchor and neighbor, e.g. 0.0003 for the current T5 fit", + ) + parser.add_argument( + "--schedule-mode", + choices=["fixed_case", "case_neighborhood", "family_random"], + default=None, + help="Override schedule.generation_mode from the YAML config", + ) parser.add_argument( "--use-runner-server", action="store_true", @@ -342,7 +288,18 @@ def resolve_config(args: argparse.Namespace) -> Config: config_path = args.config if config_path is None: config_path = str(config_for_stage(args.stage) or Path("configs/data_gen_family_random.yaml")) - return Config(config_path) + cfg = Config(config_path, dataset_case=args.dataset_case) + if cfg.dataset_cases and args.dataset_case is None: + available = ", ".join(str(case.get("name")) for case in cfg.dataset_cases) + raise ValueError(f"--dataset-case is required for this config; available cases: {available}") + if args.fixed_cf is not None: + cfg.raw["params"].setdefault("fixed_params", {})["Cf"] = { + "enabled": True, + "value": float(args.fixed_cf), + } + if args.schedule_mode is not None: + cfg.raw["schedule"]["generation_mode"] = str(args.schedule_mode) + return cfg def resolve_output_path(cfg: Config, args: argparse.Namespace) -> Path: @@ -372,7 +329,7 @@ def run_solver_and_extract_curve( params: Params, well_index: int, timeout: int, -) -> tuple[np.ndarray, np.ndarray, dict]: +) -> tuple[np.ndarray, np.ndarray, np.ndarray, dict]: """调用 C++ 求解器运行一次正演,并把双对数输出重采样为模型曲线向量。""" ok = runner.run_simulation(params, timeout=timeout, override_schedule=params.schedule, include_schedule=True) result = read_result_bin(runner.result_bin) if runner.result_bin.exists() else None @@ -397,7 +354,12 @@ def run_solver_and_extract_curve( if not valid: raise RuntimeError(f"curve_invalid_{reason}") - curve_feat, curve_time = resample_curve_to_features_with_time(cfg, t_clean, p_clean, d_clean) + curve_feat, curve_time, curve_valid_mask = resample_curve_to_features_with_time_and_mask( + cfg, + t_clean, + p_clean, + d_clean, + ) raw = { "t": t_clean.tolist(), "p": p_clean.tolist(), @@ -405,24 +367,123 @@ def run_solver_and_extract_curve( "n_steps": int(result["nSteps"]), "n_wells": int(result["nWells"]), } - return curve_feat, curve_time, raw + return curve_feat, curve_time, curve_valid_mask, raw -def params_to_array(params: Params) -> np.ndarray: - """按固定参数顺序把 Params 对象转换成数值数组。""" - return np.asarray( - [params.k, params.skin, params.wellboreC, params.phi, params.h, params.Ct, params.Cf], - dtype=np.float32, +def model_param_names(cfg: Config) -> list[str]: + """返回H5中保存并交给模型的原始参数列顺序。""" + return list( + cfg.raw["params"].get( + "model_param_names", + cfg.raw["params"]["all_physical_param_names"], + ) ) +def params_to_model_array(cfg: Config, params: Params) -> np.ndarray: + """按当前模型列顺序写参数;Ct仍保留在Params中供求解器使用。""" + values: list[float] = [] + for name in model_param_names(cfg): + if name == "solverType": + solver_type = cfg.get("solver_type", default=None) + if solver_type is None: + raise ValueError("model_param_names contains solverType but the selected case has no solver_type") + values.append(float(solver_type)) + continue + if not hasattr(params, name): + raise ValueError(f"Unsupported model parameter column: {name}") + values.append(float(getattr(params, name))) + return np.asarray(values, dtype=np.float32) + + +def load_trace_anchor_params(cfg: Config, trace_csv: str | None, count: int) -> list[Params]: + """从全求解trace选择各代最优点和目标值分位点作为重点锚点。""" + if trace_csv is None or int(count) <= 0: + return [] + + rows: list[dict] = [] + with Path(trace_csv).open("r", newline="", encoding="utf-8-sig") as f: + for row in csv.DictReader(f): + try: + objective = float(row.get("solver_objective", "nan")) + generation = int(row.get("generation", -1)) + success = str(row.get("solver_success", "")) == "1" + if row.get("phase") != "particle_solver" or not success or not math.isfinite(objective): + continue + values = {} + for name in cfg.raw["params"]["all_physical_param_names"]: + fixed = _fixed_param_value(cfg, name) + value = fixed if fixed is not None else float(row[name]) + lo, hi = map(float, cfg.raw["params"]["ranges"][name]) + values[name] = float(np.clip(value, lo, hi)) + rows.append( + { + "generation": generation, + "objective": objective, + "params": Params(**values), + } + ) + except (KeyError, TypeError, ValueError): + continue + + if not rows: + raise ValueError(f"No valid full-solver particle rows found in trace: {trace_csv}") + + best_by_generation: dict[int, dict] = {} + for row in rows: + generation = int(row["generation"]) + current = best_by_generation.get(generation) + if current is None or float(row["objective"]) < float(current["objective"]): + best_by_generation[generation] = row + + selected: list[Params] = [] + seen: set[tuple[float, ...]] = set() + + def add(row: dict) -> None: + params = row["params"] + key = tuple( + round(float(getattr(params, name)), 12) + for name in cfg.raw["params"]["all_physical_param_names"] + ) + if key not in seen and len(selected) < int(count): + seen.add(key) + selected.append(params) + + generation_best_rows = [best_by_generation[key] for key in sorted(best_by_generation)] + if len(generation_best_rows) > int(count): + positions = np.linspace(0, len(generation_best_rows) - 1, num=int(count), dtype=np.int64) + generation_best_rows = [generation_best_rows[int(pos)] for pos in positions] + for row in generation_best_rows: + add(row) + + objective_rows = sorted(rows, key=lambda row: float(row["objective"])) + remaining = int(count) - len(selected) + if remaining > 0: + positions = np.linspace(0, len(objective_rows) - 1, num=max(remaining, 1), dtype=np.int64) + for pos in positions: + add(objective_rows[int(pos)]) + for row in objective_rows: + add(row) + if len(selected) >= int(count): + break + return selected + + def sample_anchor_params_and_schedule( cfg: Config, rng: np.random.RandomState, family_override: str | None = None, + anchor_params: Params | None = None, ) -> tuple[Params, np.ndarray, str]: """为一个锚点同时采样物理参数和流量制度,作为后续邻域搜索中心。""" - params = generate_params_dataset(cfg, n_samples=1, method="sobol", random_seed=int(rng.randint(0, 2**31 - 1)))[0] + params = anchor_params + if params is None: + params = generate_params_dataset( + cfg, + n_samples=1, + method=cfg.raw["params"].get("sampling_method", "sobol"), + random_seed=int(rng.randint(0, 2**31 - 1)), + )[0] timeQ, q, sched_info = sample_schedule_by_mode(cfg, rng, family_override=family_override) section_indices = _resolve_section_indices(cfg, timeQ, q, rng) sec = int(section_indices[int(rng.randint(0, len(section_indices)))]) @@ -588,11 +649,11 @@ def select_objective_stratified_indices(objectives: np.ndarray, k: int, objectiv def select_multiscale_rows( - rows: list[tuple[np.ndarray, np.ndarray, dict[str, float], float, np.ndarray]], + rows: list[tuple[np.ndarray, np.ndarray, dict[str, float], float, np.ndarray, np.ndarray]], k: int, objective_bins: int, span_fracs: list[float], -) -> list[tuple[np.ndarray, np.ndarray, dict[str, float], float, np.ndarray]]: +) -> list[tuple[np.ndarray, np.ndarray, dict[str, float], float, np.ndarray, np.ndarray]]: """从不同扰动尺度的候选中挑选代表性行,形成多尺度邻域数据。""" if len(rows) <= k or len(span_fracs) <= 1: indices = select_objective_stratified_indices( @@ -658,19 +719,24 @@ def build_span_targets(total_keep: int, span_fracs: list[float]) -> dict[float, def create_output_file( output_path: Path, - param_dim: int, + param_names: list[str], schedule_dim: int, curve_dim: int, span_fracs: list[float], search_names: list[str], + dataset_case: str | None, + solver_type: int | None, ) -> h5py.File: """创建邻域 HDF5 文件,并初始化锚点、候选、曲线和元数据数据集。""" output_path.parent.mkdir(parents=True, exist_ok=True) f = h5py.File(output_path, "w") - f.attrs["param_names"] = np.asarray(["k", "skin", "wellboreC", "phi", "h", "Ct", "Cf"], dtype="S") + param_dim = len(param_names) + f.attrs["param_names"] = np.asarray(param_names, dtype="S") f.attrs["schedule_meta_names"] = np.asarray(SCHEDULE_META_NAMES, dtype="S") f.attrs["span_fracs"] = np.asarray(span_fracs, dtype=np.float32) f.attrs["search_param_names"] = np.asarray(search_names, dtype="S") + f.attrs["dataset_case"] = "" if dataset_case is None else str(dataset_case) + f.attrs["solver_type"] = -1 if solver_type is None else int(solver_type) n_anchors = 0 n_neighbors = 0 @@ -680,6 +746,7 @@ def create_output_file( n_time_points = int(curve_dim // 2) f.create_dataset("anchor_curve", shape=(0, curve_dim), maxshape=(None, curve_dim), dtype=np.float32, chunks=(64, curve_dim)) f.create_dataset("anchor_curve_time", shape=(0, n_time_points), maxshape=(None, n_time_points), dtype=np.float32, chunks=(64, n_time_points)) + f.create_dataset("anchor_curve_valid_mask", shape=(0, n_time_points), maxshape=(None, n_time_points), dtype=np.uint8, chunks=(64, n_time_points)) f.create_dataset("anchor_schedule_meta", shape=(0, len(SCHEDULE_META_NAMES)), maxshape=(None, len(SCHEDULE_META_NAMES)), dtype=np.float32, chunks=(64, len(SCHEDULE_META_NAMES))) f.create_dataset("anchor_family_name", shape=(0,), maxshape=(None,), dtype=h5py.string_dtype(encoding="utf-8"), chunks=(64,)) f.create_dataset("anchor_section_index", shape=(0,), maxshape=(None,), dtype=np.int32, chunks=(64,)) @@ -690,6 +757,7 @@ def create_output_file( f.create_dataset("neighbor_params", shape=(0, param_dim), maxshape=(None, param_dim), dtype=np.float32, chunks=(256, param_dim)) f.create_dataset("neighbor_curve", shape=(0, curve_dim), maxshape=(None, curve_dim), dtype=np.float32, chunks=(256, curve_dim)) f.create_dataset("neighbor_curve_time", shape=(0, n_time_points), maxshape=(None, n_time_points), dtype=np.float32, chunks=(256, n_time_points)) + f.create_dataset("neighbor_curve_valid_mask", shape=(0, n_time_points), maxshape=(None, n_time_points), dtype=np.uint8, chunks=(256, n_time_points)) f.create_dataset("neighbor_objective", shape=(0,), maxshape=(None,), dtype=np.float32, chunks=(256,)) f.create_dataset("neighbor_objective_p", shape=(0,), maxshape=(None,), dtype=np.float32, chunks=(256,)) f.create_dataset("neighbor_objective_d", shape=(0,), maxshape=(None,), dtype=np.float32, chunks=(256,)) @@ -705,6 +773,7 @@ def append_anchor( schedule_vec: np.ndarray, curve: np.ndarray, curve_time: np.ndarray, + curve_valid_mask: np.ndarray, schedule_meta: np.ndarray, family_name: str, section_index: int, @@ -719,6 +788,7 @@ def append_anchor( ("anchor_schedule", schedule_vec.reshape(1, -1)), ("anchor_curve", curve.reshape(1, -1)), ("anchor_curve_time", curve_time.reshape(1, -1)), + ("anchor_curve_valid_mask", curve_valid_mask.reshape(1, -1)), ("anchor_schedule_meta", schedule_meta.reshape(1, -1)), ]: ds = f[name] @@ -743,6 +813,7 @@ def append_neighbors( params_list: list[np.ndarray], curve_list: list[np.ndarray], curve_time_list: list[np.ndarray], + curve_valid_mask_list: list[np.ndarray], obj_list: list[dict[str, float]], span_frac_list: list[float], ) -> None: @@ -767,6 +838,7 @@ def append_neighbors( resize_2d("neighbor_params")[start:end] = np.stack(params_list, axis=0).astype(np.float32) resize_2d("neighbor_curve")[start:end] = np.stack(curve_list, axis=0).astype(np.float32) resize_2d("neighbor_curve_time")[start:end] = np.stack(curve_time_list, axis=0).astype(np.float32) + resize_2d("neighbor_curve_valid_mask")[start:end] = np.stack(curve_valid_mask_list, axis=0).astype(np.uint8) resize_1d("neighbor_objective")[start:end] = np.asarray([x["dual_log_objective"] for x in obj_list], dtype=np.float32) resize_1d("neighbor_objective_p")[start:end] = np.asarray([x["log_pressure_objective"] for x in obj_list], dtype=np.float32) resize_1d("neighbor_objective_d")[start:end] = np.asarray([x["log_derivative_objective"] for x in obj_list], dtype=np.float32) @@ -786,16 +858,18 @@ def main() -> None: if use_schedule_meta_features(cfg): schedule_dim += len(schedule_meta_feature_names(cfg)) curve_dim = int(cfg.curve_dim) - param_dim = len(cfg.raw["params"]["all_physical_param_names"]) + output_param_names = model_param_names(cfg) local_search_names = search_param_names(cfg) f = create_output_file( output_path, - param_dim=param_dim, + param_names=output_param_names, schedule_dim=schedule_dim, curve_dim=curve_dim, span_fracs=span_fracs, search_names=local_search_names, + dataset_case=cfg.dataset_case, + solver_type=cfg.get("solver_type", default=None), ) summary = { "config_path": str(cfg.path), @@ -809,6 +883,13 @@ def main() -> None: "use_runner_server": bool(args.use_runner_server), "max_perturbed_dims": int(args.max_perturbed_dims), "search_param_names": local_search_names, + "model_param_names": output_param_names, + "dataset_case": cfg.dataset_case, + "solver_type": cfg.get("solver_type", default=None), + "fixed_cf": _fixed_param_value(cfg, "Cf"), + "schedule_mode": str(cfg.raw["schedule"]["generation_mode"]), + "trace_csv": None if args.trace_csv is None else str(Path(args.trace_csv).resolve()), + "trace_anchor_count_requested": int(args.trace_anchor_count), "objective_bins": int(args.objective_bins), "anchor_fail_reasons": Counter(), "neighbor_fail_reasons": Counter(), @@ -819,14 +900,38 @@ def main() -> None: } try: - anchor_attempts = 0 max_anchor_attempts = max(int(args.n_anchors), int(args.n_anchors) * int(args.anchor_max_attempts_factor)) + trace_anchor_params = load_trace_anchor_params( + cfg, + trace_csv=args.trace_csv, + count=min(max(int(args.trace_anchor_count), 0), int(args.n_anchors)), + ) + global_anchor_count = max_anchor_attempts - len(trace_anchor_params) + global_anchor_params = generate_params_dataset( + cfg, + n_samples=global_anchor_count, + method=cfg.raw["params"].get("sampling_method", "sobol"), + random_seed=int(args.seed), + ) + anchor_param_candidates = trace_anchor_params + global_anchor_params + if len(anchor_param_candidates) < max_anchor_attempts: + raise RuntimeError( + f"Only {len(anchor_param_candidates)} unique anchor parameters were generated; " + f"expected {max_anchor_attempts}" + ) + anchor_attempts = 0 + summary["trace_anchor_count_loaded"] = int(len(trace_anchor_params)) family_plan = build_family_plan(cfg, int(args.n_anchors), balance_families=bool(args.balance_families)) while int(f.attrs["n_anchors"]) < int(args.n_anchors) and anchor_attempts < max_anchor_attempts: anchor_idx = anchor_attempts anchor_attempts += 1 target_family = family_plan[int(f.attrs["n_anchors"])] if family_plan else None - anchor_params, schedule_meta, family_name = sample_anchor_params_and_schedule(cfg, rng, family_override=target_family) + anchor_params, schedule_meta, family_name = sample_anchor_params_and_schedule( + cfg, + rng, + family_override=target_family, + anchor_params=anchor_param_candidates[anchor_idx], + ) summary["family_counter_attempted"][family_name] += 1 anchor_runner = CppRunner( @@ -836,7 +941,7 @@ def main() -> None: use_server=bool(args.use_runner_server), ) try: - anchor_curve, anchor_curve_time, _ = run_solver_and_extract_curve( + anchor_curve, anchor_curve_time, anchor_curve_valid_mask, _ = run_solver_and_extract_curve( runner=anchor_runner, cfg=cfg, params=anchor_params, @@ -856,10 +961,11 @@ def main() -> None: ) anchor_id = append_anchor( f=f, - anchor_params=params_to_array(anchor_params), + anchor_params=params_to_model_array(cfg, anchor_params), schedule_vec=schedule_vec, curve=anchor_curve, curve_time=anchor_curve_time, + curve_valid_mask=anchor_curve_valid_mask, schedule_meta=schedule_meta, family_name=family_name, section_index=int(anchor_params.schedule.sectionIndex), @@ -874,7 +980,7 @@ def main() -> None: ) valid_neighbor_rows: list[ - tuple[np.ndarray, np.ndarray, dict[str, float], float, np.ndarray] + tuple[np.ndarray, np.ndarray, dict[str, float], float, np.ndarray, np.ndarray] ] = [] valid_span_counter: Counter[str] = Counter() max_attempts = max( @@ -901,25 +1007,46 @@ def main() -> None: max_perturbed_dims=int(args.max_perturbed_dims), ) try: - cand_curve, cand_curve_time, _ = run_solver_and_extract_curve( + cand_curve, cand_curve_time, cand_curve_valid_mask, _ = run_solver_and_extract_curve( runner=anchor_runner, cfg=cfg, params=cand, well_index=int(args.well_index), timeout=int(args.solver_timeout), ) + common_mask = anchor_curve_valid_mask.astype(bool) & cand_curve_valid_mask.astype(bool) n_time_points = curve_dim // 2 - obj = dual_log_objective(anchor_curve, cand_curve, {"parts": [ - {"name": "log_pressure", "start": 0, "end": n_time_points}, - {"name": "log_derivative", "start": n_time_points, "end": curve_dim}, - ]}) + anchor_common = np.concatenate( + [ + anchor_curve[:n_time_points][common_mask], + anchor_curve[n_time_points:][common_mask], + ] + ) + cand_common = np.concatenate( + [ + cand_curve[:n_time_points][common_mask], + cand_curve[n_time_points:][common_mask], + ] + ) + n_common = int(np.sum(common_mask)) + obj = dual_log_objective( + anchor_common, + cand_common, + { + "parts": [ + {"name": "log_pressure", "start": 0, "end": n_common}, + {"name": "log_derivative", "start": n_common, "end": 2 * n_common}, + ] + }, + ) valid_neighbor_rows.append( ( - params_to_array(cand), + params_to_model_array(cfg, cand), cand_curve.astype(np.float32), obj, span_frac, cand_curve_time.astype(np.float32), + cand_curve_valid_mask.astype(np.uint8), ) ) summary["neighbor_span_counter_valid"][f"{span_frac:g}"] += 1 @@ -963,6 +1090,7 @@ def main() -> None: neighbor_obj_rows = [x[2] for x in selected_rows] neighbor_span_rows = [float(x[3]) for x in selected_rows] neighbor_curve_time_rows = [x[4] for x in selected_rows] + neighbor_curve_valid_mask_rows = [x[5] for x in selected_rows] for span_frac in neighbor_span_rows: summary["neighbor_span_counter_kept"][f"{span_frac:g}"] += 1 @@ -973,6 +1101,7 @@ def main() -> None: params_list=neighbor_params_rows, curve_list=neighbor_curve_rows, curve_time_list=neighbor_curve_time_rows, + curve_valid_mask_list=neighbor_curve_valid_mask_rows, obj_list=neighbor_obj_rows, span_frac_list=neighbor_span_rows, ) diff --git a/ML/nmWTAI-ML/scripts/train_forward.py b/ML/nmWTAI-ML/scripts/train_forward.py index 8b5febe..b34a103 100644 --- a/ML/nmWTAI-ML/scripts/train_forward.py +++ b/ML/nmWTAI-ML/scripts/train_forward.py @@ -60,7 +60,10 @@ def main() -> None: parser.add_argument("--w-derivative-shape", type=float, default=0.10) parser.add_argument("--w-autofit-pressure", type=float, default=0.0) parser.add_argument("--w-autofit-derivative", type=float, default=0.0) + parser.add_argument("--w-delta-pressure", type=float, default=0.0) + parser.add_argument("--w-delta-derivative", type=float, default=0.0) parser.add_argument("--huber-beta", type=float, default=0.05) + parser.add_argument("--delta-huber-beta", type=float, default=0.01) parser.add_argument("--use-sample-reweight", action="store_true", default=True) parser.add_argument("--no-sample-reweight", action="store_false", dest="use_sample_reweight") @@ -106,8 +109,11 @@ def main() -> None: derivative_shape=args.w_derivative_shape, autofit_pressure=args.w_autofit_pressure, autofit_derivative=args.w_autofit_derivative, + delta_pressure=args.w_delta_pressure, + delta_derivative=args.w_delta_derivative, ), huber_beta=args.huber_beta, + delta_huber_beta=args.delta_huber_beta, ), sample_reweight=SampleReweightConfig( enabled=args.use_sample_reweight, diff --git a/ML/nmWTAI-ML/src/training/train_forward.py b/ML/nmWTAI-ML/src/training/train_forward.py index 7f49da8..c091252 100644 --- a/ML/nmWTAI-ML/src/training/train_forward.py +++ b/ML/nmWTAI-ML/src/training/train_forward.py @@ -25,7 +25,7 @@ import joblib import numpy as np import torch from torch import nn -from torch.utils.data import DataLoader, Dataset +from torch.utils.data import DataLoader, Dataset, Sampler from src.models.forward_surrogate import ForwardSurrogate, ForwardSurrogateConfig @@ -39,6 +39,8 @@ METRIC_KEYS = ( "loss_derivative_shape", "loss_autofit_pressure", "loss_autofit_derivative", + "loss_delta_pressure", + "loss_delta_derivative", "sample_weight_mean", "sample_weight_max", ) @@ -47,11 +49,57 @@ METRIC_KEYS = ( class ForwardDataset(Dataset): """把预处理后的参数、流量制度和曲线数组封装成 PyTorch Dataset。""" - def __init__(self, params_x: np.ndarray, schedule_x: np.ndarray, curve_y: np.ndarray): + def __init__( + self, + params_x: np.ndarray, + schedule_x: np.ndarray, + curve_y: np.ndarray, + curve_valid_mask: np.ndarray | None = None, + group_id: np.ndarray | None = None, + is_anchor: np.ndarray | None = None, + source_id: np.ndarray | None = None, + ): """把三个 numpy 数组转为 float32 张量,后续 DataLoader 可直接按样本读取。""" self.params_x = torch.tensor(params_x, dtype=torch.float32) self.schedule_x = torch.tensor(schedule_x, dtype=torch.float32) self.curve_y = torch.tensor(curve_y, dtype=torch.float32) + n_time_points = int(self.curve_y.shape[1] // 2) + if curve_valid_mask is None: + curve_valid_mask = np.ones((len(self.curve_y), n_time_points), dtype=np.float32) + curve_valid_mask = np.asarray(curve_valid_mask, dtype=np.float32) + expected_shape = (len(self.curve_y), n_time_points) + if curve_valid_mask.shape != expected_shape: + raise ValueError( + f"curve_valid_mask shape mismatch: {curve_valid_mask.shape} != {expected_shape}" + ) + self.curve_valid_mask = torch.tensor(curve_valid_mask, dtype=torch.float32) + self.group_id = self._metadata_tensor("group_id", group_id, default_unique=True) + self.is_anchor = self._metadata_tensor("is_anchor", is_anchor, default_value=-1) + self.source_id = self._metadata_tensor("source_id", source_id, default_value=-1) + + def _metadata_tensor( + self, + name: str, + values: np.ndarray | None, + default_value: int = -1, + default_unique: bool = False, + ) -> torch.Tensor: + """校验可选样本元数据,并以 int64 张量形式保存。""" + if values is None: + if default_unique: + array = np.arange(len(self.params_x), dtype=np.int64) + else: + array = np.full(len(self.params_x), default_value, dtype=np.int64) + else: + raw = np.asarray(values).reshape(-1) + if len(raw) != len(self.params_x): + raise ValueError( + f"{name} length mismatch: {len(raw)} != {len(self.params_x)}" + ) + if np.issubdtype(raw.dtype, np.floating) and not np.all(np.isfinite(raw)): + raise ValueError(f"{name} contains non-finite values") + array = raw.astype(np.int64, copy=False) + return torch.tensor(array, dtype=torch.int64) def __len__(self) -> int: """返回数据集或容器中可迭代样本的数量。""" @@ -59,7 +107,129 @@ class ForwardDataset(Dataset): def __getitem__(self, idx: int): """按索引取出一个训练样本或数据项。""" - return self.params_x[idx], self.schedule_x[idx], self.curve_y[idx] + return ( + self.params_x[idx], + self.schedule_x[idx], + self.curve_y[idx], + self.curve_valid_mask[idx], + self.group_id[idx], + self.is_anchor[idx], + self.source_id[idx], + ) + + +class CompleteGroupBatchSampler(Sampler[list[int]]): + """按完整样本组组织批次,避免同一个锚点邻域组被拆到不同批次。""" + + def __init__( + self, + group_id: np.ndarray, + source_id: np.ndarray | None, + batch_size: int, + shuffle: bool, + seed: int, + drop_last: bool = False, + ): + if batch_size <= 0: + raise ValueError("batch_size must be positive") + + group_values = np.asarray(group_id).reshape(-1) + if source_id is None: + source_values = np.full(len(group_values), -1, dtype=np.int64) + else: + source_values = np.asarray(source_id).reshape(-1) + if len(source_values) != len(group_values): + raise ValueError("source_id length does not match group_id") + + groups: dict[tuple[str, int, int], list[int]] = {} + for index, (value, source) in enumerate(zip(group_values.tolist(), source_values.tolist())): + numeric_value = int(value) + key = ( + ("group", int(source), numeric_value) + if numeric_value >= 0 + else ("row", int(source), index) + ) + groups.setdefault(key, []).append(index) + + self.group_units = list(groups.values()) + largest_group = max((len(unit) for unit in self.group_units), default=0) + if largest_group > batch_size: + raise ValueError( + "complete group does not fit in one batch: " + f"largest_group={largest_group}, batch_size={batch_size}" + ) + + self.batch_size = int(batch_size) + self.shuffle = bool(shuffle) + self.seed = int(seed) + self.drop_last = bool(drop_last) + self.epoch = 0 + self._next_batches: list[list[int]] | None = None + + def _build_batches(self) -> list[list[int]]: + order = np.arange(len(self.group_units), dtype=np.int64) + if self.shuffle and len(order) > 1: + rng = np.random.RandomState(self.seed + self.epoch) + rng.shuffle(order) + + batches: list[list[int]] = [] + batch: list[int] = [] + for position in order.tolist(): + unit = self.group_units[position] + if batch and len(batch) + len(unit) > self.batch_size: + batches.append(batch) + batch = [] + batch.extend(unit) + + if batch and (not self.drop_last or len(batch) == self.batch_size): + batches.append(batch) + return batches + + def _get_next_batches(self) -> list[list[int]]: + if self._next_batches is None: + self._next_batches = self._build_batches() + return self._next_batches + + def __iter__(self): + batches = self._get_next_batches() + self._next_batches = None + self.epoch += 1 + yield from batches + + def __len__(self) -> int: + return len(self._get_next_batches()) + + +def validate_delta_groups(dataset: ForwardDataset, split: str) -> None: + """确保每个局部锚点都具有同一来源下的完整邻域样本组。""" + group_id = dataset.group_id.numpy() + is_anchor = dataset.is_anchor.numpy() + source_id = dataset.source_id.numpy() + schedule_x = dataset.schedule_x.numpy() + local_mask = (is_anchor == 0) | (is_anchor == 1) + if np.any(group_id[local_mask] < 0): + raise ValueError(f"delta groups must use non-negative group_id values in {split} split") + local_indices = np.flatnonzero(local_mask) + if len(local_indices) == 0: + raise ValueError(f"delta loss requires at least one anchor group in {split} split") + + keys = { + (int(source_id[index]), int(group_id[index])) + for index in local_indices.tolist() + } + for key in keys: + same_group = (source_id == key[0]) & (group_id == key[1]) + anchor_indices = np.flatnonzero(same_group & (is_anchor == 1)) + if len(anchor_indices) != 1: + raise ValueError( + f"delta group {key} must have exactly one anchor in {split} split" + ) + neighbor_indices = np.flatnonzero(same_group & (is_anchor == 0)) + if len(neighbor_indices) == 0: + raise ValueError(f"delta group {key} has no neighbors in {split} split") + member_indices = np.flatnonzero(same_group & local_mask) + if not np.all(schedule_x[member_indices] == schedule_x[anchor_indices[0]]): + raise ValueError(f"delta group {key} has inconsistent schedules in {split} split") @dataclass @@ -92,6 +262,8 @@ class LossWeights: derivative_shape: float = 0.10 autofit_pressure: float = 0.0 autofit_derivative: float = 0.0 + delta_pressure: float = 0.0 + delta_derivative: float = 0.0 @dataclass @@ -101,6 +273,7 @@ class LossConfig: weights: LossWeights = field(default_factory=LossWeights) use_huber: bool = True huber_beta: float = 0.05 + delta_huber_beta: float = 0.01 @dataclass @@ -150,6 +323,8 @@ class LossBatchParts: pred_d: torch.Tensor true_p: torch.Tensor true_d: torch.Tensor + mask_p: torch.Tensor + mask_d: torch.Tensor @dataclass @@ -194,6 +369,23 @@ def load_processed_dataset(path: Path) -> dict: for key in required_keys: if key not in data: raise KeyError(f"processed dataset 缺少字段: {key}") + meta = data.get("meta", {}) or {} + if meta.get("curve_time_mode") != "fixed": + raise ValueError("training requires processed curve_time_mode='fixed'") + curve_dim = int(data["Y_curve_train"].shape[1]) + prediction_time = np.asarray(meta.get("prediction_curve_time", []), dtype=np.float64) + if prediction_time.shape != (curve_dim // 2,): + raise ValueError( + f"prediction_curve_time shape {prediction_time.shape} != {(curve_dim // 2,)}" + ) + for split in ("train", "val", "test"): + mask_key = f"curve_valid_mask_{split}" + if mask_key not in data: + raise KeyError(f"processed fixed-time dataset is missing {mask_key}") + expected_shape = (data[f"Y_curve_{split}"].shape[0], curve_dim // 2) + actual_shape = np.asarray(data[mask_key]).shape + if actual_shape != expected_shape: + raise ValueError(f"{mask_key} shape {actual_shape} != {expected_shape}") return data @@ -234,32 +426,61 @@ def get_part_slices(curve_layout: dict) -> dict[str, slice]: return out -def smooth_l1_per_sample(pred: torch.Tensor, target: torch.Tensor, beta: float) -> torch.Tensor: +def masked_mean_per_sample(values: torch.Tensor, valid_mask: torch.Tensor | None) -> torch.Tensor: + """逐样本求均值,仅统计掩码标记的有效时间点。""" + if valid_mask is None: + return values.mean(dim=1) + mask = valid_mask.to(dtype=values.dtype) + if mask.shape != values.shape: + raise ValueError(f"loss mask shape mismatch: {mask.shape} != {values.shape}") + return (values * mask).sum(dim=1) / mask.sum(dim=1).clamp_min(1.0) + + +def smooth_l1_per_sample( + pred: torch.Tensor, + target: torch.Tensor, + beta: float, + valid_mask: torch.Tensor | None = None, +) -> torch.Tensor: """按样本计算 Smooth L1 损失,返回每个样本一个损失值。""" diff = torch.abs(pred - target) loss = torch.where(diff < beta, 0.5 * diff * diff / beta, diff - 0.5 * beta) - return loss.mean(dim=1) + return masked_mean_per_sample(loss, valid_mask) -def mse_per_sample(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor: +def mse_per_sample( + pred: torch.Tensor, + target: torch.Tensor, + valid_mask: torch.Tensor | None = None, +) -> torch.Tensor: """按样本计算均方误差。""" - return ((pred - target) ** 2).mean(dim=1) + return masked_mean_per_sample((pred - target) ** 2, valid_mask) -def l1_per_sample(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor: +def l1_per_sample( + pred: torch.Tensor, + target: torch.Tensor, + valid_mask: torch.Tensor | None = None, +) -> torch.Tensor: """按样本计算平均绝对误差。""" - return torch.abs(pred - target).mean(dim=1) + return masked_mean_per_sample(torch.abs(pred - target), valid_mask) def regression_per_sample( pred: torch.Tensor, target: torch.Tensor, loss_cfg: LossConfig, + valid_mask: torch.Tensor | None = None, ) -> torch.Tensor: """按配置在 Smooth L1 和 MSE 之间切换点值损失。""" if loss_cfg.use_huber: - return smooth_l1_per_sample(pred, target, beta=float(loss_cfg.huber_beta)) - return mse_per_sample(pred, target) + return smooth_l1_per_sample( + pred, + target, + beta=float(loss_cfg.huber_beta), + valid_mask=valid_mask, + ) + return mse_per_sample(pred, target, valid_mask=valid_mask) def first_diff(x: torch.Tensor) -> torch.Tensor: @@ -272,30 +493,29 @@ def affine_restore(x_scaled: torch.Tensor, mean: torch.Tensor, scale: torch.Tens return x_scaled * scale + mean -def autofit_curve_objective_per_sample(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor: - """用 torch 计算自动拟合风格的曲线误差,作为训练附加目标。""" - weight_factor = torch.clamp(torch.abs(target) * 0.01, max=100.0) - weight = 1.0 / (1.0 + weight_factor) - scale = torch.maximum( - torch.maximum(torch.abs(target), torch.abs(pred)), - torch.full_like(target, 1e-12), - ) - relative_error = torch.abs(target - pred) / scale - absolute_error = torch.abs(target - pred) - point_error = 0.7 * relative_error + 0.3 * absolute_error - weighted_mse = (weight * (point_error**2)).sum(dim=1) - weighted_mse = weighted_mse / torch.clamp(weight.sum(dim=1), min=1e-12) - return torch.sqrt(weighted_mse) +def autofit_curve_objective_per_sample( + pred: torch.Tensor, + target: torch.Tensor, + valid_mask: torch.Tensor | None = None, +) -> torch.Tensor: + """用 torch 计算与 C++ 当前真实误差公式等价的曲线误差。""" + log_error = torch.abs(target - pred) + relative_error = -torch.expm1(-log_error) + point_error = 0.7 * log_error + 0.3 * relative_error + squared_mean = masked_mean_per_sample(point_error**2, valid_mask) + return torch.sqrt(torch.clamp(squared_mean, min=1.0e-12)) def build_sample_weight( true_p: torch.Tensor, true_d: torch.Tensor, reweight_cfg: SampleReweightConfig, + mask_p: torch.Tensor | None = None, + mask_d: torch.Tensor | None = None, ) -> torch.Tensor: """根据真实曲线幅值构造样本权重,让高幅值样本训练时权重更高。""" - p_level = true_p.abs().mean(dim=1) - d_level = true_d.abs().mean(dim=1) + p_level = masked_mean_per_sample(true_p.abs(), mask_p) + d_level = masked_mean_per_sample(true_d.abs(), mask_d) p_norm = p_level / (p_level.mean().detach() + 1e-6) d_norm = d_level / (d_level.mean().detach() + 1e-6) @@ -308,6 +528,7 @@ def build_sample_weight( def split_curve_parts( pred: torch.Tensor, target: torch.Tensor, + valid_mask: torch.Tensor, slices: dict[str, slice], ) -> LossBatchParts: """把拼接曲线拆成压力和导数两段。""" @@ -316,6 +537,8 @@ def split_curve_parts( pred_d=pred[:, slices["log_derivative"]], true_p=target[:, slices["log_pressure"]], true_d=target[:, slices["log_derivative"]], + mask_p=valid_mask, + mask_d=valid_mask, ) @@ -325,20 +548,31 @@ def compute_basic_loss_vectors( ) -> dict[str, torch.Tensor]: """计算标准化空间中的基础点值、偏置和导数形状损失。""" return { - "loss_pressure": regression_per_sample(parts.pred_p, parts.true_p, loss_cfg), - "loss_derivative": regression_per_sample(parts.pred_d, parts.true_d, loss_cfg), + "loss_pressure": regression_per_sample( + parts.pred_p, + parts.true_p, + loss_cfg, + valid_mask=parts.mask_p, + ), + "loss_derivative": regression_per_sample( + parts.pred_d, + parts.true_d, + loss_cfg, + valid_mask=parts.mask_d, + ), "loss_bias_pressure": l1_per_sample( - parts.pred_p.mean(dim=1, keepdim=True), - parts.true_p.mean(dim=1, keepdim=True), + masked_mean_per_sample(parts.pred_p, parts.mask_p).unsqueeze(1), + masked_mean_per_sample(parts.true_p, parts.mask_p).unsqueeze(1), ), "loss_bias_derivative": l1_per_sample( - parts.pred_d.mean(dim=1, keepdim=True), - parts.true_d.mean(dim=1, keepdim=True), + masked_mean_per_sample(parts.pred_d, parts.mask_d).unsqueeze(1), + masked_mean_per_sample(parts.true_d, parts.mask_d).unsqueeze(1), ), "loss_derivative_shape": regression_per_sample( first_diff(parts.pred_d), first_diff(parts.true_d), loss_cfg, + valid_mask=parts.mask_d[:, 1:] * parts.mask_d[:, :-1], ), } @@ -360,14 +594,118 @@ def compute_autofit_loss_vectors( "loss_autofit_pressure": autofit_curve_objective_per_sample( affine_restore(parts.pred_p, mean_p, scale_p), affine_restore(parts.true_p, mean_p, scale_p), + valid_mask=parts.mask_p, ), "loss_autofit_derivative": autofit_curve_objective_per_sample( affine_restore(parts.pred_d, mean_d, scale_d), affine_restore(parts.true_d, mean_d, scale_d), + valid_mask=parts.mask_d, ), } +def grouped_channel_delta_loss( + pred_raw: torch.Tensor, + target_raw: torch.Tensor, + valid_mask: torch.Tensor, + group_id: torch.Tensor, + is_anchor: torch.Tensor, + source_id: torch.Tensor, + beta: float, +) -> tuple[torch.Tensor, int]: + """比较锚点与邻域样本的曲线变化,并让每个样本组具有相同权重。""" + if beta <= 0.0: + raise ValueError("delta_huber_beta must be positive") + if group_id.ndim != 1 or is_anchor.ndim != 1 or source_id.ndim != 1: + raise ValueError("group_id, is_anchor and source_id must be one-dimensional") + if ( + len(group_id) != len(pred_raw) + or len(is_anchor) != len(pred_raw) + or len(source_id) != len(pred_raw) + ): + raise ValueError("group metadata length does not match the curve batch") + + zero = pred_raw.sum() * 0.0 + group_losses: list[torch.Tensor] = [] + anchor_indices = torch.nonzero(is_anchor == 1, as_tuple=False).flatten() + for anchor_index in anchor_indices: + current_group = group_id[anchor_index] + current_source = source_id[anchor_index] + same_group = (group_id == current_group) & (source_id == current_source) + anchor_mask = same_group & (is_anchor == 1) + if int(anchor_mask.sum().item()) != 1: + raise ValueError( + f"group_id={int(current_group.item())} must contain exactly one anchor" + ) + + neighbor_indices = torch.nonzero( + same_group & (is_anchor == 0), + as_tuple=False, + ).flatten() + if neighbor_indices.numel() == 0: + continue + + common_mask = valid_mask[neighbor_indices] * valid_mask[anchor_index].unsqueeze(0) + usable = common_mask.sum(dim=1) > 0 + if not bool(torch.any(usable).item()): + continue + + neighbor_indices = neighbor_indices[usable] + common_mask = common_mask[usable] + pred_delta = pred_raw[neighbor_indices] - pred_raw[anchor_index].unsqueeze(0) + target_delta = target_raw[neighbor_indices] - target_raw[anchor_index].unsqueeze(0) + neighbor_losses = smooth_l1_per_sample( + pred_delta, + target_delta, + beta=beta, + valid_mask=common_mask, + ) + group_losses.append(neighbor_losses.mean()) + + if not group_losses: + return zero, 0 + return torch.stack(group_losses).mean(), len(group_losses) + + +def compute_group_delta_losses( + parts: LossBatchParts, + group_id: torch.Tensor, + is_anchor: torch.Tensor, + source_id: torch.Tensor, + context: LossContext, +) -> tuple[torch.Tensor, torch.Tensor, int]: + """在反标准化后的原始对数尺度上计算压力和导数差分损失。""" + pressure_slice = context.slices["log_pressure"] + derivative_slice = context.slices["log_derivative"] + mean_p = context.curve_stats.mean_raw[pressure_slice].unsqueeze(0) + scale_p = context.curve_stats.scale_raw[pressure_slice].unsqueeze(0) + mean_d = context.curve_stats.mean_raw[derivative_slice].unsqueeze(0) + scale_d = context.curve_stats.scale_raw[derivative_slice].unsqueeze(0) + beta = float(context.loss_cfg.delta_huber_beta) + + pressure_loss, pressure_group_count = grouped_channel_delta_loss( + pred_raw=affine_restore(parts.pred_p, mean_p, scale_p), + target_raw=affine_restore(parts.true_p, mean_p, scale_p), + valid_mask=parts.mask_p, + group_id=group_id, + is_anchor=is_anchor, + source_id=source_id, + beta=beta, + ) + derivative_loss, derivative_group_count = grouped_channel_delta_loss( + pred_raw=affine_restore(parts.pred_d, mean_d, scale_d), + target_raw=affine_restore(parts.true_d, mean_d, scale_d), + valid_mask=parts.mask_d, + group_id=group_id, + is_anchor=is_anchor, + source_id=source_id, + beta=beta, + ) + if pressure_group_count != derivative_group_count: + raise RuntimeError("pressure and derivative delta group counts do not match") + return pressure_loss, derivative_loss, pressure_group_count + + def weighted_total_vector( loss_vectors: dict[str, torch.Tensor], weights: LossWeights, @@ -387,21 +725,59 @@ def weighted_total_vector( def compute_weighted_loss( pred: torch.Tensor, target: torch.Tensor, + valid_mask: torch.Tensor, context: LossContext, + group_id: torch.Tensor | None = None, + is_anchor: torch.Tensor | None = None, + source_id: torch.Tensor | None = None, ) -> dict[str, torch.Tensor]: """计算正演代理模型的复合训练损失。""" - parts = split_curve_parts(pred, target, context.slices) + parts = split_curve_parts(pred, target, valid_mask, context.slices) loss_vectors = compute_basic_loss_vectors(parts, context.loss_cfg) loss_vectors.update(compute_autofit_loss_vectors(parts, context)) total_vec = weighted_total_vector(loss_vectors, context.loss_cfg.weights) if context.reweight_cfg.enabled: - sample_weight = build_sample_weight(parts.true_p, parts.true_d, context.reweight_cfg) + sample_weight = build_sample_weight( + parts.true_p, + parts.true_d, + context.reweight_cfg, + mask_p=parts.mask_p, + mask_d=parts.mask_d, + ) else: sample_weight = torch.ones_like(total_vec) metrics = {key: value.mean() for key, value in loss_vectors.items()} - metrics["loss"] = (total_vec * sample_weight).mean() + base_loss = (total_vec * sample_weight).mean() + metrics["loss"] = base_loss + weights = context.loss_cfg.weights + delta_enabled = bool(weights.delta_pressure != 0.0 or weights.delta_derivative != 0.0) + if delta_enabled: + if group_id is None or is_anchor is None: + raise ValueError("delta loss requires group_id and is_anchor metadata") + if source_id is None: + source_id = torch.full_like(group_id, -1) + delta_pressure, delta_derivative, delta_group_count = compute_group_delta_losses( + parts=parts, + group_id=group_id, + is_anchor=is_anchor, + source_id=source_id, + context=context, + ) + metrics["loss"] = ( + metrics["loss"] + + weights.delta_pressure * delta_pressure + + weights.delta_derivative * delta_derivative + ) + else: + delta_pressure = pred.sum() * 0.0 + delta_derivative = pred.sum() * 0.0 + delta_group_count = 0 + metrics["loss_base"] = base_loss + metrics["loss_delta_pressure"] = delta_pressure + metrics["loss_delta_derivative"] = delta_derivative + metrics["delta_group_count"] = pred.new_tensor(float(delta_group_count)) metrics["sample_weight_mean"] = sample_weight.mean() metrics["sample_weight_max"] = sample_weight.max() return metrics @@ -454,19 +830,43 @@ def run_loader_epoch( total = init_metric_accumulator() total_n = 0 + base_loss_total = 0.0 + delta_pressure_total = 0.0 + delta_derivative_total = 0.0 + delta_group_total = 0 grad_context = torch.enable_grad() if is_train else torch.no_grad() with grad_context: - for params_x, schedule_x, curve_y in loader: + for ( + params_x, + schedule_x, + curve_y, + curve_valid_mask, + group_id, + is_anchor, + source_id, + ) in loader: params_x = params_x.to(device) schedule_x = schedule_x.to(device) curve_y = curve_y.to(device) + curve_valid_mask = curve_valid_mask.to(device) + group_id = group_id.to(device) + is_anchor = is_anchor.to(device) + source_id = source_id.to(device) if is_train: optimizer.zero_grad() pred = model_forward(model, params_x, schedule_x, use_schedule) - losses = compute_weighted_loss(pred=pred, target=curve_y, context=context) + losses = compute_weighted_loss( + pred=pred, + target=curve_y, + valid_mask=curve_valid_mask, + context=context, + group_id=group_id, + is_anchor=is_anchor, + source_id=source_id, + ) if is_train: losses["loss"].backward() @@ -474,9 +874,27 @@ def run_loader_epoch( batch_size = params_x.size(0) accumulate_metrics(total, losses, batch_size) + base_loss_total += losses["loss_base"].item() * batch_size + delta_group_count = int(losses["delta_group_count"].item()) + delta_pressure_total += losses["loss_delta_pressure"].item() * delta_group_count + delta_derivative_total += losses["loss_delta_derivative"].item() * delta_group_count + delta_group_total += delta_group_count total_n += batch_size - return average_metrics(total, total_n) + metrics = average_metrics(total, total_n) + if delta_group_total > 0: + metrics["loss_delta_pressure"] = delta_pressure_total / delta_group_total + metrics["loss_delta_derivative"] = delta_derivative_total / delta_group_total + else: + metrics["loss_delta_pressure"] = 0.0 + metrics["loss_delta_derivative"] = 0.0 + weights = context.loss_cfg.weights + metrics["loss"] = ( + base_loss_total / max(total_n, 1) + + weights.delta_pressure * metrics["loss_delta_pressure"] + + weights.delta_derivative * metrics["loss_delta_derivative"] + ) + return metrics def evaluate( @@ -510,21 +928,97 @@ def build_curve_stats(data: dict, device: str) -> CurveStats: def build_dataloaders(data: dict, cfg: TrainConfig) -> DatasetBundle: """根据预处理数组构造训练、验证、测试 DataLoader。""" - train_ds = ForwardDataset(data["X_params_train"], data["X_schedule_train"], data["Y_curve_train"]) - val_ds = ForwardDataset(data["X_params_val"], data["X_schedule_val"], data["Y_curve_val"]) - test_ds = ForwardDataset(data["X_params_test"], data["X_schedule_test"], data["Y_curve_test"]) - - loader_generator = torch.Generator() - loader_generator.manual_seed(int(cfg.runtime.seed)) - - train_loader = DataLoader( - train_ds, - batch_size=cfg.optim.batch_size, - shuffle=True, - generator=loader_generator, + delta_enabled = bool( + cfg.loss.weights.delta_pressure != 0.0 + or cfg.loss.weights.delta_derivative != 0.0 ) - val_loader = DataLoader(val_ds, batch_size=cfg.optim.batch_size, shuffle=False) - test_loader = DataLoader(test_ds, batch_size=cfg.optim.batch_size, shuffle=False) + if cfg.loss.weights.delta_pressure < 0.0 or cfg.loss.weights.delta_derivative < 0.0: + raise ValueError("delta loss weights must be non-negative") + if delta_enabled and cfg.loss.delta_huber_beta <= 0.0: + raise ValueError("delta_huber_beta must be positive when delta loss is enabled") + if delta_enabled: + for split in ("train", "val", "test"): + for field_name in ("group_id", "is_anchor"): + key = f"{field_name}_{split}" + if key not in data: + raise KeyError(f"delta loss requires processed dataset field: {key}") + + train_ds = ForwardDataset( + data["X_params_train"], + data["X_schedule_train"], + data["Y_curve_train"], + data.get("curve_valid_mask_train"), + data.get("group_id_train"), + data.get("is_anchor_train"), + data.get("source_id_train"), + ) + val_ds = ForwardDataset( + data["X_params_val"], + data["X_schedule_val"], + data["Y_curve_val"], + data.get("curve_valid_mask_val"), + data.get("group_id_val"), + data.get("is_anchor_val"), + data.get("source_id_val"), + ) + test_ds = ForwardDataset( + data["X_params_test"], + data["X_schedule_test"], + data["Y_curve_test"], + data.get("curve_valid_mask_test"), + data.get("group_id_test"), + data.get("is_anchor_test"), + data.get("source_id_test"), + ) + + if delta_enabled: + validate_delta_groups(train_ds, "train") + validate_delta_groups(val_ds, "val") + validate_delta_groups(test_ds, "test") + train_loader = DataLoader( + train_ds, + batch_sampler=CompleteGroupBatchSampler( + train_ds.group_id.numpy(), + train_ds.source_id.numpy(), + batch_size=cfg.optim.batch_size, + shuffle=True, + seed=cfg.runtime.seed, + drop_last=False, + ), + ) + val_loader = DataLoader( + val_ds, + batch_sampler=CompleteGroupBatchSampler( + val_ds.group_id.numpy(), + val_ds.source_id.numpy(), + batch_size=cfg.optim.batch_size, + shuffle=False, + seed=cfg.runtime.seed + 1, + drop_last=False, + ), + ) + test_loader = DataLoader( + test_ds, + batch_sampler=CompleteGroupBatchSampler( + test_ds.group_id.numpy(), + test_ds.source_id.numpy(), + batch_size=cfg.optim.batch_size, + shuffle=False, + seed=cfg.runtime.seed + 2, + drop_last=False, + ), + ) + else: + loader_generator = torch.Generator() + loader_generator.manual_seed(int(cfg.runtime.seed)) + train_loader = DataLoader( + train_ds, + batch_size=cfg.optim.batch_size, + shuffle=True, + generator=loader_generator, + ) + val_loader = DataLoader(val_ds, batch_size=cfg.optim.batch_size, shuffle=False) + test_loader = DataLoader(test_ds, batch_size=cfg.optim.batch_size, shuffle=False) return DatasetBundle( train_loader=train_loader, @@ -581,7 +1075,9 @@ def print_training_config(cfg: TrainConfig, curve_layout: dict) -> None: f" weights: pressure={weights.pressure}, derivative={weights.derivative}, " f"bias_p={weights.bias_pressure}, " f"bias_d={weights.bias_derivative}, d_shape={weights.derivative_shape}, " - f"autofit_p={weights.autofit_pressure}, autofit_d={weights.autofit_derivative}" + f"autofit_p={weights.autofit_pressure}, autofit_d={weights.autofit_derivative}, " + f"delta_p={weights.delta_pressure}, delta_d={weights.delta_derivative}, " + f"delta_beta={cfg.loss.delta_huber_beta}" ) print( f" sample_reweight={reweight.enabled}, alpha={reweight.alpha}, " @@ -603,6 +1099,8 @@ def format_metric_line(epoch: int, train_metrics: dict[str, float], val_metrics: f"ds={train_metrics['loss_derivative_shape']:.6f}, " f"ap={train_metrics['loss_autofit_pressure']:.6f}, " f"ad={train_metrics['loss_autofit_derivative']:.6f}, " + f"dp={train_metrics['loss_delta_pressure']:.6f}, " + f"dd={train_metrics['loss_delta_derivative']:.6f}, " f"wmean={train_metrics['sample_weight_mean']:.4f}, " f"wmax={train_metrics['sample_weight_max']:.4f}) " f"val={val_metrics['loss']:.6f} " @@ -613,6 +1111,8 @@ def format_metric_line(epoch: int, train_metrics: dict[str, float], val_metrics: f"ds={val_metrics['loss_derivative_shape']:.6f}, " f"ap={val_metrics['loss_autofit_pressure']:.6f}, " f"ad={val_metrics['loss_autofit_derivative']:.6f}, " + f"dp={val_metrics['loss_delta_pressure']:.6f}, " + f"dd={val_metrics['loss_delta_derivative']:.6f}, " f"wmean={val_metrics['sample_weight_mean']:.4f}, " f"wmax={val_metrics['sample_weight_max']:.4f})" ) @@ -629,6 +1129,8 @@ def format_final_line(test_metrics: dict[str, float]) -> str: f"ds={test_metrics['loss_derivative_shape']:.6f}, " f"ap={test_metrics['loss_autofit_pressure']:.6f}, " f"ad={test_metrics['loss_autofit_derivative']:.6f}, " + f"dp={test_metrics['loss_delta_pressure']:.6f}, " + f"dd={test_metrics['loss_delta_derivative']:.6f}, " f"wmean={test_metrics['sample_weight_mean']:.4f}, " f"wmax={test_metrics['sample_weight_max']:.4f})" ) @@ -652,6 +1154,7 @@ def build_checkpoint_payload( "seed": int(cfg.runtime.seed), "curve_layout": curve_layout, "loss_weights": asdict(cfg.loss.weights), + "delta_huber_beta": float(cfg.loss.delta_huber_beta), "sample_reweight": asdict(cfg.sample_reweight), } @@ -748,6 +1251,7 @@ def build_metrics_payload( "use_schedule": cfg.model.use_schedule, "seed": int(cfg.runtime.seed), "loss_weights": asdict(cfg.loss.weights), + "delta_huber_beta": float(cfg.loss.delta_huber_beta), "sample_reweight": asdict(cfg.sample_reweight), "curve_layout": curve_layout, }