统一代理评分与自动拟合误差计算

- 按物理时间重叠区对齐目标曲线和预测曲线
- 使用与C++真实求解一致的压力和导数误差公式
- 统一在线评分服务、trace回放和排序验证逻辑
- 增加固定时间轴和有效掩码回归测试
feature/Model-20260625
lvjunjie 4 weeks ago
parent dcc7bb688f
commit 54ea2ef32d

@ -30,6 +30,8 @@ ROOT = Path(__file__).resolve().parents[1]
sys.path.append(str(ROOT)) sys.path.append(str(ROOT))
import joblib import joblib
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt import matplotlib.pyplot as plt
import numpy as np import numpy as np
import torch import torch
@ -107,7 +109,7 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--h", type=float, default=DEFAULT_SINGLE_CASE["params"]["h"]) parser.add_argument("--h", type=float, default=DEFAULT_SINGLE_CASE["params"]["h"])
parser.add_argument("--Ct", type=float, default=DEFAULT_SINGLE_CASE["params"]["Ct"]) parser.add_argument("--Ct", type=float, default=DEFAULT_SINGLE_CASE["params"]["Ct"])
parser.add_argument("--Cf", type=float, default=DEFAULT_SINGLE_CASE["params"]["Cf"]) parser.add_argument("--Cf", type=float, default=DEFAULT_SINGLE_CASE["params"]["Cf"])
parser.add_argument("--solver-type", type=int, choices=[1, 2, 3, 4], default=None) parser.add_argument("--solver-type", type=int, choices=[1, 2, 3, 4, 5], default=None)
parser.add_argument( parser.add_argument(
"--dataset-case", "--dataset-case",
type=str, type=str,

@ -19,6 +19,8 @@ import sys
from pathlib import Path from pathlib import Path
import joblib import joblib
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt import matplotlib.pyplot as plt
import numpy as np import numpy as np
import torch import torch
@ -84,10 +86,19 @@ def parse_args() -> argparse.Namespace:
def calc_metrics( def calc_metrics(
y_true: np.ndarray, y_true: np.ndarray,
y_pred: np.ndarray, y_pred: np.ndarray,
valid_mask: np.ndarray | None = None,
eps_range: float = 1e-3, eps_range: float = 1e-3,
eps_var: float = 1e-6, eps_var: float = 1e-6,
) -> dict: ) -> dict:
"""计算 RMSE、MAE、Bias、NRMSE、R2 等回归指标。""" """计算 RMSE、MAE、Bias、NRMSE、R2 等回归指标。"""
y_true = np.asarray(y_true)
y_pred = np.asarray(y_pred)
if valid_mask is not None:
keep = np.asarray(valid_mask, dtype=bool)
y_true = y_true[keep]
y_pred = y_pred[keep]
if y_true.size == 0:
raise ValueError("metric mask contains no valid points")
err = y_pred - y_true err = y_pred - y_true
mse = np.mean(err**2) mse = np.mean(err**2)
rmse = float(np.sqrt(mse)) rmse = float(np.sqrt(mse))
@ -233,13 +244,19 @@ def plot_sample(
curve_true: np.ndarray, curve_true: np.ndarray,
curve_pred: np.ndarray, curve_pred: np.ndarray,
curve_layout: dict, curve_layout: dict,
valid_time_mask: np.ndarray,
output_dir: Path, output_dir: Path,
title_prefix: str, title_prefix: str,
) -> None: ) -> None:
"""绘制单个样本的真实曲线、预测曲线和误差曲线,并保存为图片。""" """绘制单个样本的真实曲线、预测曲线和误差曲线,并保存为图片。"""
true_parts = split_curve_by_layout(curve_true, curve_layout) true_parts = split_curve_by_layout(curve_true, curve_layout)
pred_parts = split_curve_by_layout(curve_pred, curve_layout) pred_parts = split_curve_by_layout(curve_pred, curve_layout)
overall = calc_metrics(curve_true, curve_pred) valid_time_mask = np.asarray(valid_time_mask, dtype=bool)
overall = calc_metrics(
curve_true,
curve_pred,
valid_mask=np.concatenate([valid_time_mask, valid_time_mask]),
)
nrmse_text = "nan" if np.isnan(overall["nrmse"]) else f"{overall['nrmse']:.4f}" nrmse_text = "nan" if np.isnan(overall["nrmse"]) else f"{overall['nrmse']:.4f}"
r2_text = "nan" if np.isnan(overall["r2"]) else f"{overall['r2']:.4f}" r2_text = "nan" if np.isnan(overall["r2"]) else f"{overall['r2']:.4f}"
@ -263,6 +280,8 @@ def plot_sample(
for row, name in enumerate(plot_order): for row, name in enumerate(plot_order):
y_true = true_parts[name] y_true = true_parts[name]
y_pred = pred_parts[name] y_pred = pred_parts[name]
y_true = y_true[valid_time_mask]
y_pred = y_pred[valid_time_mask]
err = y_pred - y_true err = y_pred - y_true
x = np.arange(len(y_true)) x = np.arange(len(y_true))
m = calc_metrics(y_true, y_pred) m = calc_metrics(y_true, y_pred)
@ -346,18 +365,33 @@ def model_predict(
return model(params_x, None) return model(params_x, None)
def prepare_eval_arrays(eval_data: dict, fit_data: dict | None) -> tuple[np.ndarray, np.ndarray, np.ndarray, object]: def prepare_eval_arrays(
eval_data: dict,
fit_data: dict | None,
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, object]:
"""根据评估划分取出参数、流量制度和真实曲线,并完成必要的数组整理。""" """根据评估划分取出参数、流量制度和真实曲线,并完成必要的数组整理。"""
x_params_test = np.asarray(eval_data["X_params_test"], dtype=np.float32) x_params_test = np.asarray(eval_data["X_params_test"], dtype=np.float32)
x_schedule_test = np.asarray(eval_data["X_schedule_test"], dtype=np.float32) x_schedule_test = np.asarray(eval_data["X_schedule_test"], dtype=np.float32)
y_curve_test = np.asarray(eval_data["Y_curve_test"], dtype=np.float32) y_curve_test = np.asarray(eval_data["Y_curve_test"], dtype=np.float32)
eval_meta = eval_data.get("meta", {}) or {}
if eval_meta.get("curve_time_mode") != "fixed":
raise ValueError("Evaluation requires processed curve_time_mode='fixed'")
n_time_points = int(y_curve_test.shape[1] // 2)
if "curve_valid_mask_test" not in eval_data:
raise ValueError("Evaluation processed dataset is missing curve_valid_mask_test")
curve_valid_mask_test = np.asarray(eval_data["curve_valid_mask_test"], dtype=np.uint8)
expected_mask_shape = (y_curve_test.shape[0], n_time_points)
if curve_valid_mask_test.shape != expected_mask_shape:
raise ValueError(
f"curve_valid_mask_test shape {curve_valid_mask_test.shape} != {expected_mask_shape}"
)
eval_scaler_params = eval_data["scaler_params"] eval_scaler_params = eval_data["scaler_params"]
eval_scaler_schedule = eval_data["scaler_schedule"] eval_scaler_schedule = eval_data["scaler_schedule"]
eval_scaler_curve = eval_data["scaler_curve"] eval_scaler_curve = eval_data["scaler_curve"]
if fit_data is None: if fit_data is None:
return x_params_test, x_schedule_test, y_curve_test, eval_scaler_curve return x_params_test, x_schedule_test, y_curve_test, curve_valid_mask_test, eval_scaler_curve
fit_scaler_params = fit_data["scaler_params"] fit_scaler_params = fit_data["scaler_params"]
fit_scaler_schedule = fit_data["scaler_schedule"] fit_scaler_schedule = fit_data["scaler_schedule"]
@ -370,6 +404,15 @@ def prepare_eval_arrays(eval_data: dict, fit_data: dict | None) -> tuple[np.ndar
"Cross-dataset evaluation requires matching param_feature_transform metadata. " "Cross-dataset evaluation requires matching param_feature_transform metadata. "
"Re-preprocess both datasets with the same transform setting." "Re-preprocess both datasets with the same transform setting."
) )
eval_time = np.asarray((eval_data.get("meta", {}) or {}).get("prediction_curve_time", []))
fit_time = np.asarray((fit_data.get("meta", {}) or {}).get("prediction_curve_time", []))
if eval_time.shape != fit_time.shape or not np.allclose(
eval_time,
fit_time,
rtol=1.0e-6,
atol=1.0e-10,
):
raise ValueError("Cross-dataset evaluation requires the same fixed prediction_curve_time")
# 先还原评估集使用的参数特征尺度,再映射到当前模型训练时保存的 scaler 尺度。 # 先还原评估集使用的参数特征尺度,再映射到当前模型训练时保存的 scaler 尺度。
x_params_raw = eval_scaler_params.inverse_transform(x_params_test) x_params_raw = eval_scaler_params.inverse_transform(x_params_test)
@ -379,7 +422,7 @@ def prepare_eval_arrays(eval_data: dict, fit_data: dict | None) -> tuple[np.ndar
x_schedule_fit = fit_scaler_schedule.transform(x_schedule_raw).astype(np.float32) x_schedule_fit = fit_scaler_schedule.transform(x_schedule_raw).astype(np.float32)
y_true_raw = eval_scaler_curve.inverse_transform(y_curve_test).astype(np.float32) y_true_raw = eval_scaler_curve.inverse_transform(y_curve_test).astype(np.float32)
return x_params_fit, x_schedule_fit, y_true_raw, fit_scaler_curve return x_params_fit, x_schedule_fit, y_true_raw, curve_valid_mask_test, fit_scaler_curve
def main() -> None: def main() -> None:
@ -402,7 +445,13 @@ def main() -> None:
fit_data = joblib.load(fit_processed_path) if fit_processed_path is not None else None fit_data = joblib.load(fit_processed_path) if fit_processed_path is not None else None
# prepare_eval_arrays 会处理“评估集”和“拟合 scaler 的训练集”不一致的情况。 # prepare_eval_arrays 会处理“评估集”和“拟合 scaler 的训练集”不一致的情况。
x_params_test, x_schedule_test, y_curve_test, pred_scaler_curve = prepare_eval_arrays(data, fit_data) (
x_params_test,
x_schedule_test,
y_curve_test,
curve_valid_mask_test,
pred_scaler_curve,
) = prepare_eval_arrays(data, fit_data)
meta = data["meta"] meta = data["meta"]
param_dim = int(meta["param_dim"]) param_dim = int(meta["param_dim"])
@ -476,7 +525,9 @@ def main() -> None:
for idx, (curve_true, curve_pred) in enumerate(zip(all_true, all_pred)): for idx, (curve_true, curve_pred) in enumerate(zip(all_true, all_pred)):
# 同时记录整体指标和各曲线通道指标。 # 同时记录整体指标和各曲线通道指标。
overall_m = calc_metrics(curve_true, curve_pred) time_mask = curve_valid_mask_test[idx].astype(bool)
overall_mask = np.concatenate([time_mask, time_mask])
overall_m = calc_metrics(curve_true, curve_pred, valid_mask=overall_mask)
overall_metric_list.append(overall_m) overall_metric_list.append(overall_m)
true_parts = split_curve_by_layout(curve_true, curve_layout) true_parts = split_curve_by_layout(curve_true, curve_layout)
@ -484,7 +535,7 @@ def main() -> None:
part_ms: dict[str, dict] = {} part_ms: dict[str, dict] = {}
for name in part_names: for name in part_names:
part_m = calc_metrics(true_parts[name], pred_parts[name]) part_m = calc_metrics(true_parts[name], pred_parts[name], valid_mask=time_mask)
part_metric_lists[name].append(part_m) part_metric_lists[name].append(part_m)
part_ms[name] = part_m part_ms[name] = part_m
@ -558,11 +609,11 @@ def main() -> None:
print("Random sample indices:", random_indices) print("Random sample indices:", random_indices)
for idx in random_indices: for idx in random_indices:
plot_sample(idx, all_true[idx], all_pred[idx], curve_layout, output_dir, "random") plot_sample(idx, all_true[idx], all_pred[idx], curve_layout, curve_valid_mask_test[idx], output_dir, "random")
for idx in best_indices: for idx in best_indices:
plot_sample(idx, all_true[idx], all_pred[idx], curve_layout, output_dir, "best") plot_sample(idx, all_true[idx], all_pred[idx], curve_layout, curve_valid_mask_test[idx], output_dir, "best")
for idx in worst_indices: for idx in worst_indices:
plot_sample(idx, all_true[idx], all_pred[idx], curve_layout, output_dir, "worst") plot_sample(idx, all_true[idx], all_pred[idx], curve_layout, curve_valid_mask_test[idx], output_dir, "worst")
print("\nArtifacts written to:", output_dir) print("\nArtifacts written to:", output_dir)
print("1. summary_metrics.json") print("1. summary_metrics.json")

@ -27,9 +27,11 @@ from scripts.validate_autofit_local_ranking import (
_corr_spearman, _corr_spearman,
_rank_positions, _rank_positions,
_sample_local_candidates, _sample_local_candidates,
calculate_candidate_objectives,
infer_curve_layout, infer_curve_layout,
load_model, load_model,
predict_surrogate_curve, predict_surrogate_curve,
resolve_prediction_curve_time,
run_solver_and_extract_curve, run_solver_and_extract_curve,
) )
from src.common.config import Config from src.common.config import Config
@ -41,7 +43,6 @@ from src.common.experiment_paths import (
) )
from src.data.params import Params, Schedule from src.data.params import Params, Schedule
from src.data.runner_client import CppRunner from src.data.runner_client import CppRunner
from src.evaluation.autofit_objective import dual_log_objective
THETA_STAR = { THETA_STAR = {
@ -63,6 +64,8 @@ def parse_args() -> argparse.Namespace:
description="Generate Q schedules around theta* and run first-layer local ranking validation." description="Generate Q schedules around theta* and run first-layer local ranking validation."
) )
parser.add_argument("--config", type=str, default=None) parser.add_argument("--config", type=str, default=None)
parser.add_argument("--dataset-case", type=str, default=None)
parser.add_argument("--solver-type", type=int, choices=[1, 2, 3, 4, 5], default=None)
parser.add_argument( parser.add_argument(
"--stage", "--stage",
choices=[ choices=[
@ -94,6 +97,7 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--h", type=float, default=THETA_STAR["h"]) parser.add_argument("--h", type=float, default=THETA_STAR["h"])
parser.add_argument("--Ct", type=float, default=THETA_STAR["Ct"]) parser.add_argument("--Ct", type=float, default=THETA_STAR["Ct"])
parser.add_argument("--Cf", type=float, default=THETA_STAR["Cf"]) parser.add_argument("--Cf", type=float, default=THETA_STAR["Cf"])
parser.add_argument("--fixed-cf", type=float, default=None)
return parser.parse_args() return parser.parse_args()
@ -111,7 +115,14 @@ def resolve_paths(args: argparse.Namespace) -> tuple[Config, Path, Path, Path]:
if args.output_dir is not None if args.output_dir is not None
else Path("results") / f"q_sweep_local_ranking_{tag}" else Path("results") / f"q_sweep_local_ranking_{tag}"
) )
return Config(config_path), processed_path.resolve(), model_path.resolve(), output_dir.resolve() cfg = Config(config_path, dataset_case=args.dataset_case)
if args.fixed_cf is not None:
cfg.raw["params"].setdefault("fixed_params", {})["Cf"] = {
"enabled": True,
"value": float(args.fixed_cf),
}
args.Cf = float(args.fixed_cf)
return cfg, processed_path.resolve(), model_path.resolve(), output_dir.resolve()
def make_theta_params(args: argparse.Namespace, schedule: Schedule) -> Params: def make_theta_params(args: argparse.Namespace, schedule: Schedule) -> Params:
@ -374,6 +385,15 @@ def main() -> None:
processed = joblib.load(processed_path) processed = joblib.load(processed_path)
curve_layout = infer_curve_layout(processed["meta"], int(processed["meta"]["curve_dim"])) curve_layout = infer_curve_layout(processed["meta"], int(processed["meta"]["curve_dim"]))
model, use_schedule, device = load_model(model_path) model, use_schedule, device = load_model(model_path)
configured_solver_type = cfg.get("solver_type", default=None)
if args.solver_type is not None and configured_solver_type is not None:
if int(args.solver_type) != int(configured_solver_type):
raise ValueError(
f"--solver-type={args.solver_type} conflicts with "
f"dataset case solver_type={configured_solver_type}"
)
solver_type = args.solver_type if args.solver_type is not None else configured_solver_type
prediction_curve_time = resolve_prediction_curve_time(processed, curve_layout, solver_type)
q_cases = generate_q_cases() q_cases = generate_q_cases()
# 每个 q_case 固定一条目标流量制度;后续只扰动物理参数,验证局部排序是否稳健。 # 每个 q_case 固定一条目标流量制度;后续只扰动物理参数,验证局部排序是否稳健。
if args.case_id_contains: if args.case_id_contains:
@ -408,7 +428,7 @@ def main() -> None:
target_runner = make_runner(cfg, output_dir, f"target_{case_idx:03d}") target_runner = make_runner(cfg, output_dir, f"target_{case_idx:03d}")
try: try:
# 先用真实求解器生成该流量制度下的目标曲线,候选曲线都和它比较。 # 先用真实求解器生成该流量制度下的目标曲线,候选曲线都和它比较。
target_curve, _ = run_solver_and_extract_curve( target_curve, target_raw = run_solver_and_extract_curve(
runner=target_runner, runner=target_runner,
cfg=cfg, cfg=cfg,
params=target_params, params=target_params,
@ -443,7 +463,7 @@ def main() -> None:
runner = make_runner(cfg, output_dir, f"case_{case_idx:03d}_cand_{attempt_id:04d}") runner = make_runner(cfg, output_dir, f"case_{case_idx:03d}_cand_{attempt_id:04d}")
try: try:
# 同一候选同时计算 solver 目标和 surrogate 目标,后面只比较排序,不混用数值来源。 # 同一候选同时计算 solver 目标和 surrogate 目标,后面只比较排序,不混用数值来源。
solver_curve, _ = run_solver_and_extract_curve( solver_curve, solver_raw = run_solver_and_extract_curve(
runner=runner, runner=runner,
cfg=cfg, cfg=cfg,
params=cand, params=cand,
@ -459,9 +479,17 @@ def main() -> None:
params=cand, params=cand,
schedule=schedule, schedule=schedule,
cfg=cfg, cfg=cfg,
solver_type=None if solver_type is None else int(solver_type),
)
solver_obj, surrogate_obj = calculate_candidate_objectives(
target_curve=target_curve,
target_raw=target_raw,
solver_curve=solver_curve,
solver_raw=solver_raw,
surrogate_curve=pred_curve,
curve_layout=curve_layout,
prediction_curve_time=prediction_curve_time,
) )
solver_obj = dual_log_objective(target_curve, solver_curve, curve_layout)
surrogate_obj = dual_log_objective(target_curve, pred_curve, curve_layout)
row = { row = {
"case_id": q_case["case_id"], "case_id": q_case["case_id"],
"case_index": case_idx, "case_index": case_idx,
@ -541,6 +569,8 @@ def main() -> None:
"config_path": str(cfg.path), "config_path": str(cfg.path),
"processed_path": str(processed_path), "processed_path": str(processed_path),
"model_path": str(model_path), "model_path": str(model_path),
"dataset_case": cfg.dataset_case,
"solver_type": None if solver_type is None else int(solver_type),
"theta_star": { "theta_star": {
"k": float(args.k), "k": float(args.k),
"skin": float(args.skin), "skin": float(args.skin),

@ -22,6 +22,8 @@ ROOT = Path(__file__).resolve().parents[1]
sys.path.append(str(ROOT)) sys.path.append(str(ROOT))
import joblib import joblib
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt import matplotlib.pyplot as plt
import numpy as np import numpy as np
@ -36,11 +38,15 @@ from src.common.experiment_paths import (
) )
from src.data.curve_processing import clean_curve_for_dataset, is_valid_curve, resample_curve_to_model_features from src.data.curve_processing import clean_curve_for_dataset, is_valid_curve, resample_curve_to_model_features
from src.data.params import Params, Schedule from src.data.params import Params, Schedule
from src.evaluation.autofit_objective import dual_log_objective from src.evaluation.autofit_objective import (
prediction_curve_time_from_meta,
timed_dual_log_objective,
)
PARAM_COLUMNS = ["k", "skin", "wellboreC", "phi", "h", "Ct", "Cf"] PARAM_COLUMNS = ["k", "skin", "wellboreC", "phi", "h", "Ct", "Cf"]
DEFAULT_KEEP_FRACS = [0.5, 0.6, 0.7] DEFAULT_KEEP_FRACS = [0.5, 0.6, 0.7]
TIME_AWARE_COMMON_POINTS = 50
def parse_args() -> argparse.Namespace: def parse_args() -> argparse.Namespace:
@ -55,6 +61,8 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--processed", type=str, default=None) parser.add_argument("--processed", type=str, default=None)
parser.add_argument("--model", type=str, default=None) parser.add_argument("--model", type=str, default=None)
parser.add_argument("--config", type=str, default=None) parser.add_argument("--config", type=str, default=None)
parser.add_argument("--dataset-case", type=str, default=None)
parser.add_argument("--solver-type", type=int, choices=[1, 2, 3, 4, 5], default=None)
parser.add_argument("--output-dir", type=str, default=None) parser.add_argument("--output-dir", type=str, default=None)
parser.add_argument("--keep-fracs", type=str, default="0.5,0.6,0.7") parser.add_argument("--keep-fracs", type=str, default="0.5,0.6,0.7")
return parser.parse_args() return parser.parse_args()
@ -137,7 +145,11 @@ def build_schedule_from_meta(meta: dict) -> Schedule:
return Schedule(sectionIndex=int(target["section_index"]), timeQ=time_q, q=q) return Schedule(sectionIndex=int(target["section_index"]), timeQ=time_q, q=q)
def build_target_curve_from_meta(cfg: Config, processed: dict, meta: dict) -> np.ndarray: def build_target_curve_from_meta(
cfg: Config,
processed: dict,
meta: dict,
) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
"""从 CSV 元数据字段还原目标双对数曲线向量。""" """从 CSV 元数据字段还原目标双对数曲线向量。"""
target = meta["target"]["target_loglog"] target = meta["target"]["target_loglog"]
t = np.asarray(target["time"], dtype=np.float64) t = np.asarray(target["time"], dtype=np.float64)
@ -153,7 +165,12 @@ def build_target_curve_from_meta(cfg: Config, processed: dict, meta: dict) -> np
curve_dim = int(processed["meta"]["curve_dim"]) curve_dim = int(processed["meta"]["curve_dim"])
curve = resample_curve_to_model_features(cfg, t_clean, p_clean, d_clean, curve_dim) curve = resample_curve_to_model_features(cfg, t_clean, p_clean, d_clean, curve_dim)
return curve.astype(np.float32) return (
np.asarray(t_clean, dtype=np.float64),
np.asarray(p_clean, dtype=np.float64),
np.asarray(d_clean, dtype=np.float64),
curve.astype(np.float32),
)
def params_from_row(row: dict) -> Params: def params_from_row(row: dict) -> Params:
@ -317,15 +334,29 @@ def main() -> None:
trace_csv, trace_meta, config_path, processed_path, model_path, output_dir = resolve_paths(args) trace_csv, trace_meta, config_path, processed_path, model_path, output_dir = resolve_paths(args)
keep_fracs = [float(x) for x in args.keep_fracs.split(",") if x.strip()] keep_fracs = [float(x) for x in args.keep_fracs.split(",") if x.strip()]
cfg = Config(config_path) cfg = Config(config_path, dataset_case=args.dataset_case)
configured_solver_type = cfg.get("solver_type", default=None)
if args.solver_type is not None and configured_solver_type is not None:
if int(args.solver_type) != int(configured_solver_type):
raise ValueError(
f"--solver-type={args.solver_type} conflicts with "
f"dataset case solver_type={configured_solver_type}"
)
solver_type = args.solver_type if args.solver_type is not None else configured_solver_type
processed = joblib.load(processed_path) processed = joblib.load(processed_path)
model, use_schedule, device = load_model(model_path)
curve_layout = infer_curve_layout(processed["meta"], int(processed["meta"]["curve_dim"])) curve_layout = infer_curve_layout(processed["meta"], int(processed["meta"]["curve_dim"]))
prediction_curve_time = prediction_curve_time_from_meta(processed["meta"], curve_layout)
objective_mode = "time_aware_common_log_grid"
model, use_schedule, device = load_model(model_path)
meta = json.loads(trace_meta.read_text(encoding="utf-8")) meta = json.loads(trace_meta.read_text(encoding="utf-8"))
rows = read_trace_csv(trace_csv) rows = read_trace_csv(trace_csv)
schedule = build_schedule_from_meta(meta) schedule = build_schedule_from_meta(meta)
target_curve = build_target_curve_from_meta(cfg, processed, meta) target_time, target_pressure, target_derivative, target_curve = build_target_curve_from_meta(
cfg,
processed,
meta,
)
output_dir.mkdir(parents=True, exist_ok=True) output_dir.mkdir(parents=True, exist_ok=True)
@ -344,8 +375,17 @@ def main() -> None:
params=params, params=params,
schedule=schedule, schedule=schedule,
cfg=cfg, cfg=cfg,
solver_type=None if solver_type is None else int(solver_type),
)
sur_obj = timed_dual_log_objective(
target_time=target_time,
target_pressure=target_pressure,
target_derivative=target_derivative,
pred_time=prediction_curve_time,
pred_curve=pred_curve,
curve_layout=curve_layout,
n_common_points=TIME_AWARE_COMMON_POINTS,
) )
sur_obj = dual_log_objective(target_curve, pred_curve, curve_layout)
replay_row = dict(row) replay_row = dict(row)
replay_row["surrogate_objective"] = sur_obj["dual_log_objective"] replay_row["surrogate_objective"] = sur_obj["dual_log_objective"]
replay_row["surrogate_p_obj"] = sur_obj["log_pressure_objective"] replay_row["surrogate_p_obj"] = sur_obj["log_pressure_objective"]
@ -375,11 +415,23 @@ def main() -> None:
"trace_meta": str(trace_meta), "trace_meta": str(trace_meta),
"processed_path": str(processed_path), "processed_path": str(processed_path),
"model_path": str(model_path), "model_path": str(model_path),
"dataset_case": cfg.dataset_case,
"solver_type": None if solver_type is None else int(solver_type),
"run_id": meta.get("run_id"), "run_id": meta.get("run_id"),
"n_particle_rows": len(replay_rows), "n_particle_rows": len(replay_rows),
"n_generations": len(generation_summaries), "n_generations": len(generation_summaries),
"target_points_raw": len(meta["target"]["target_loglog"]["time"]), "target_points_raw": len(meta["target"]["target_loglog"]["time"]),
"target_points_cleaned": int(target_time.size),
"target_curve_dim_resampled": int(target_curve.size), "target_curve_dim_resampled": int(target_curve.size),
"objective_mode": objective_mode,
"time_aware_objective": True,
"objective_common_time_points": TIME_AWARE_COMMON_POINTS,
"prediction_curve_time_points": int(prediction_curve_time.size),
"prediction_curve_time_range": [
float(prediction_curve_time[0]),
float(prediction_curve_time[-1]),
],
"target_time_range_cleaned": [float(target_time[0]), float(target_time[-1])],
"use_schedule": bool(use_schedule), "use_schedule": bool(use_schedule),
"keep_fracs": keep_fracs, "keep_fracs": keep_fracs,
"generation_summary_mean": {}, "generation_summary_mean": {},

@ -29,9 +29,12 @@ from src.common.experiment_paths import (
normalize_tag, normalize_tag,
processed_path_for_tag, processed_path_for_tag,
) )
from src.data.curve_processing import clean_curve_for_dataset, is_valid_curve, resample_curve_to_model_features from src.data.curve_processing import clean_curve_for_dataset, is_valid_curve
from src.data.params import Params, Schedule from src.data.params import Params, Schedule
from src.evaluation.autofit_objective import dual_log_objective from src.evaluation.autofit_objective import (
prediction_curve_time_from_meta,
timed_dual_log_objective,
)
PARAM_COLUMNS = ["k", "skin", "wellboreC", "phi", "h", "Cf", "solverType"] PARAM_COLUMNS = ["k", "skin", "wellboreC", "phi", "h", "Cf", "solverType"]
@ -120,8 +123,11 @@ def build_schedule_from_meta(meta: dict) -> Schedule:
return Schedule(sectionIndex=int(target["section_index"]), timeQ=time_q, q=q) return Schedule(sectionIndex=int(target["section_index"]), timeQ=time_q, q=q)
def build_target_curve_from_meta(cfg: Config, processed: dict, meta: dict) -> np.ndarray: def build_target_raw_curve_from_meta(
"""从 CSV 元数据字段还原目标双对数曲线向量。""" cfg: Config,
meta: dict,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""读取目标曲线并保留其物理时间轴。"""
target = meta["target"]["target_loglog"] target = meta["target"]["target_loglog"]
t = np.asarray(target["time"], dtype=np.float64) t = np.asarray(target["time"], dtype=np.float64)
p = np.asarray(target["pressure"], dtype=np.float64) p = np.asarray(target["pressure"], dtype=np.float64)
@ -133,9 +139,24 @@ def build_target_curve_from_meta(cfg: Config, processed: dict, meta: dict) -> np
valid, reason = is_valid_curve(cfg, t_clean, p_clean, d_clean) valid, reason = is_valid_curve(cfg, t_clean, p_clean, d_clean)
if not valid: if not valid:
raise RuntimeError(f"Target loglog curve from meta is invalid: {reason}") raise RuntimeError(f"Target loglog curve from meta is invalid: {reason}")
curve_dim = int(processed["meta"]["curve_dim"]) return t_clean, p_clean, d_clean
curve = resample_curve_to_model_features(cfg, t_clean, p_clean, d_clean, curve_dim)
return curve.astype(np.float32)
def evaluate_prediction_objective(
target_raw: tuple[np.ndarray, np.ndarray, np.ndarray],
prediction_time: np.ndarray,
pred_curve: np.ndarray,
curve_layout: dict,
) -> dict[str, float]:
"""按固定物理时间网格和 C++ 规则评价代理曲线。"""
return timed_dual_log_objective(
target_time=target_raw[0],
target_pressure=target_raw[1],
target_derivative=target_raw[2],
pred_time=prediction_time,
pred_curve=pred_curve,
curve_layout=curve_layout,
)
def params_from_row(row: dict) -> Params: def params_from_row(row: dict) -> Params:
@ -178,7 +199,8 @@ def main() -> None:
# trace meta 中保存目标曲线和制度;所有候选粒子都与同一个目标样本比较。 # trace meta 中保存目标曲线和制度;所有候选粒子都与同一个目标样本比较。
meta = json.loads(meta_path.read_text(encoding="utf-8")) meta = json.loads(meta_path.read_text(encoding="utf-8"))
schedule = build_schedule_from_meta(meta) schedule = build_schedule_from_meta(meta)
target_curve = build_target_curve_from_meta(cfg, processed, meta) target_raw = build_target_raw_curve_from_meta(cfg, meta)
prediction_time = prediction_curve_time_from_meta(processed["meta"], curve_layout)
rows = read_candidates(candidates_path) rows = read_candidates(candidates_path)
# 这里仅写出代理评分。哪些粒子进入真实求解器由 C++ 决定, # 这里仅写出代理评分。哪些粒子进入真实求解器由 C++ 决定,
@ -220,7 +242,12 @@ def main() -> None:
cfg=cfg, cfg=cfg,
solver_type=solver_type, solver_type=solver_type,
) )
obj = dual_log_objective(target_curve, pred_curve, curve_layout) obj = evaluate_prediction_objective(
target_raw=target_raw,
prediction_time=prediction_time,
pred_curve=pred_curve,
curve_layout=curve_layout,
)
out["surrogate_objective"] = f"{obj['dual_log_objective']:.17g}" out["surrogate_objective"] = f"{obj['dual_log_objective']:.17g}"
out["surrogate_p_obj"] = f"{obj['log_pressure_objective']:.17g}" out["surrogate_p_obj"] = f"{obj['log_pressure_objective']:.17g}"
out["surrogate_d_obj"] = f"{obj['log_derivative_objective']:.17g}" out["surrogate_d_obj"] = f"{obj['log_derivative_objective']:.17g}"

@ -23,7 +23,8 @@ import joblib
from scripts.compare_single_case import load_model from scripts.compare_single_case import load_model
from scripts.score_pso_candidates import ( from scripts.score_pso_candidates import (
build_schedule_from_meta, build_schedule_from_meta,
build_target_curve_from_meta, build_target_raw_curve_from_meta,
evaluate_prediction_objective,
params_from_row, params_from_row,
read_candidates, read_candidates,
resolve_paths, resolve_paths,
@ -31,7 +32,7 @@ from scripts.score_pso_candidates import (
) )
from scripts.validate_autofit_local_ranking import infer_curve_layout, predict_surrogate_curve from scripts.validate_autofit_local_ranking import infer_curve_layout, predict_surrogate_curve
from src.common.config import Config from src.common.config import Config
from src.evaluation.autofit_objective import dual_log_objective from src.evaluation.autofit_objective import prediction_curve_time_from_meta
FIELDNAMES = [ FIELDNAMES = [
@ -87,7 +88,11 @@ class PsoScoringServer:
self.meta = json.loads(self.meta_path.read_text(encoding="utf-8")) self.meta = json.loads(self.meta_path.read_text(encoding="utf-8"))
# trace meta 提供本轮 PSO 的目标曲线和制度;服务生命周期内保持不变。 # trace meta 提供本轮 PSO 的目标曲线和制度;服务生命周期内保持不变。
self.schedule = build_schedule_from_meta(self.meta) self.schedule = build_schedule_from_meta(self.meta)
self.target_curve = build_target_curve_from_meta(self.cfg, self.processed, self.meta) self.target_raw = build_target_raw_curve_from_meta(self.cfg, self.meta)
self.prediction_time = prediction_curve_time_from_meta(
self.processed["meta"],
self.curve_layout,
)
def score_file(self, candidates_path: Path, output_path: Path) -> int: def score_file(self, candidates_path: Path, output_path: Path) -> int:
"""读取候选文件、逐行调用代理模型评分,并写出代理目标函数结果文件。""" """读取候选文件、逐行调用代理模型评分,并写出代理目标函数结果文件。"""
@ -122,7 +127,12 @@ class PsoScoringServer:
cfg=self.cfg, cfg=self.cfg,
solver_type=solver_type, solver_type=solver_type,
) )
obj = dual_log_objective(self.target_curve, pred_curve, self.curve_layout) obj = evaluate_prediction_objective(
target_raw=self.target_raw,
prediction_time=self.prediction_time,
pred_curve=pred_curve,
curve_layout=self.curve_layout,
)
out["surrogate_objective"] = f"{obj['dual_log_objective']:.17g}" out["surrogate_objective"] = f"{obj['dual_log_objective']:.17g}"
out["surrogate_p_obj"] = f"{obj['log_pressure_objective']:.17g}" out["surrogate_p_obj"] = f"{obj['log_pressure_objective']:.17g}"
out["surrogate_d_obj"] = f"{obj['log_derivative_objective']:.17g}" out["surrogate_d_obj"] = f"{obj['log_derivative_objective']:.17g}"

@ -20,6 +20,8 @@ ROOT = Path(__file__).resolve().parents[1]
sys.path.append(str(ROOT)) sys.path.append(str(ROOT))
import joblib import joblib
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt import matplotlib.pyplot as plt
import numpy as np import numpy as np
import torch import torch
@ -35,7 +37,11 @@ from src.data.param_features import (
from src.data.params import Params, Schedule from src.data.params import Params, Schedule
from src.data.runner_client import CppRunner, read_result_bin from src.data.runner_client import CppRunner, read_result_bin
from src.data.schedule_features import build_schedule_model_vector from src.data.schedule_features import build_schedule_model_vector
from src.evaluation.autofit_objective import dual_log_objective from src.evaluation.autofit_objective import (
prediction_curve_time_from_meta,
timed_dual_log_objective,
timed_raw_dual_objective,
)
from src.models.forward_surrogate import ForwardSurrogate from src.models.forward_surrogate import ForwardSurrogate
@ -89,8 +95,14 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--h", type=float, default=DEFAULT_CASE["params"]["h"]) parser.add_argument("--h", type=float, default=DEFAULT_CASE["params"]["h"])
parser.add_argument("--Ct", type=float, default=DEFAULT_CASE["params"]["Ct"]) parser.add_argument("--Ct", type=float, default=DEFAULT_CASE["params"]["Ct"])
parser.add_argument("--Cf", type=float, default=DEFAULT_CASE["params"]["Cf"]) parser.add_argument("--Cf", type=float, default=DEFAULT_CASE["params"]["Cf"])
parser.add_argument("--solver-type", type=int, choices=[1, 2, 3, 4], default=None) parser.add_argument("--solver-type", type=int, choices=[1, 2, 3, 4, 5], default=None)
parser.add_argument("--dataset-case", type=str, default=None) parser.add_argument("--dataset-case", type=str, default=None)
parser.add_argument(
"--fixed-cf",
type=float,
default=None,
help="Fix Cf instead of perturbing it during local validation",
)
return parser.parse_args() return parser.parse_args()
@ -142,6 +154,12 @@ def resolve_paths(args: argparse.Namespace) -> tuple[Config, Path, Path, Path]:
config_path = str(config_for_stage(args.stage) or Path(DEFAULT_CASE["config"])) config_path = str(config_for_stage(args.stage) or Path(DEFAULT_CASE["config"]))
cfg = Config(config_path, dataset_case=args.dataset_case) cfg = Config(config_path, dataset_case=args.dataset_case)
if args.fixed_cf is not None:
cfg.raw["params"].setdefault("fixed_params", {})["Cf"] = {
"enabled": True,
"value": float(args.fixed_cf),
}
args.Cf = float(args.fixed_cf)
if args.processed is not None: if args.processed is not None:
processed_path = Path(args.processed) processed_path = Path(args.processed)
else: else:
@ -298,6 +316,48 @@ def run_solver_and_extract_curve(
return curve_feat, raw return curve_feat, raw
def resolve_prediction_curve_time(
processed: dict,
curve_layout: dict,
solver_type: int | None,
) -> np.ndarray:
"""从预处理元数据中读取代理模型输出对应的物理时间网格。"""
_ = solver_type
return prediction_curve_time_from_meta(processed.get("meta", {}), curve_layout)
def calculate_candidate_objectives(
target_curve: np.ndarray,
target_raw: dict,
solver_curve: np.ndarray,
solver_raw: dict,
surrogate_curve: np.ndarray,
curve_layout: dict,
prediction_curve_time: np.ndarray,
) -> tuple[dict[str, float], dict[str, float]]:
"""在同一物理时间轴上分别评价真实求解曲线和代理预测曲线。"""
target_time = np.asarray(target_raw["t"], dtype=np.float64)
target_pressure = np.asarray(target_raw["p"], dtype=np.float64)
target_derivative = np.asarray(target_raw["d"], dtype=np.float64)
solver_objective = timed_raw_dual_objective(
target_time=target_time,
target_pressure=target_pressure,
target_derivative=target_derivative,
pred_time=np.asarray(solver_raw["t"], dtype=np.float64),
pred_pressure=np.asarray(solver_raw["p"], dtype=np.float64),
pred_derivative=np.asarray(solver_raw["d"], dtype=np.float64),
)
surrogate_objective = timed_dual_log_objective(
target_time=target_time,
target_pressure=target_pressure,
target_derivative=target_derivative,
pred_time=prediction_curve_time,
pred_curve=surrogate_curve,
curve_layout=curve_layout,
)
return solver_objective, surrogate_objective
def _corr_pearson(x: np.ndarray, y: np.ndarray) -> float: def _corr_pearson(x: np.ndarray, y: np.ndarray) -> float:
"""计算 Pearson 相关系数,用于评估代理分数与真实分数的线性相关。""" """计算 Pearson 相关系数,用于评估代理分数与真实分数的线性相关。"""
if x.size < 2: if x.size < 2:
@ -451,6 +511,7 @@ def main() -> None:
target_params = build_params_from_args(args, schedule) target_params = build_params_from_args(args, schedule)
solver_type = resolve_solver_type(args, cfg) solver_type = resolve_solver_type(args, cfg)
search_names = _search_param_names(cfg) search_names = _search_param_names(cfg)
prediction_curve_time = resolve_prediction_curve_time(processed, curve_layout, solver_type)
model, use_schedule, device = load_model(model_path) model, use_schedule, device = load_model(model_path)
shared_runner = None shared_runner = None
@ -501,7 +562,7 @@ def main() -> None:
runner = shared_runner if shared_runner is not None else make_runner(f"cand_{attempt_idx:04d}") runner = shared_runner if shared_runner is not None else make_runner(f"cand_{attempt_idx:04d}")
try: try:
# 对每个候选同时拿到真实求解器曲线和代理预测曲线,比较二者目标函数排序。 # 对每个候选同时拿到真实求解器曲线和代理预测曲线,比较二者目标函数排序。
solver_curve, _ = run_solver_and_extract_curve( solver_curve, solver_raw = run_solver_and_extract_curve(
runner=runner, runner=runner,
cfg=cfg, cfg=cfg,
params=cand, params=cand,
@ -520,8 +581,15 @@ def main() -> None:
solver_type=solver_type, solver_type=solver_type,
) )
solver_obj = dual_log_objective(target_curve, solver_curve, curve_layout) solver_obj, surrogate_obj = calculate_candidate_objectives(
surrogate_obj = dual_log_objective(target_curve, pred_curve, curve_layout) target_curve=target_curve,
target_raw=target_raw,
solver_curve=solver_curve,
solver_raw=solver_raw,
surrogate_curve=pred_curve,
curve_layout=curve_layout,
prediction_curve_time=prediction_curve_time,
)
rows.append( rows.append(
{ {
@ -584,6 +652,8 @@ def main() -> None:
"config_path": str(cfg.path), "config_path": str(cfg.path),
"processed_path": str(processed_path), "processed_path": str(processed_path),
"model_path": str(model_path), "model_path": str(model_path),
"dataset_case": cfg.dataset_case,
"solver_type": solver_type,
"target_params": { "target_params": {
"k": target_params.k, "k": target_params.k,
"skin": target_params.skin, "skin": target_params.skin,

@ -26,16 +26,17 @@ from scripts.validate_autofit_local_ranking import (
_corr_spearman, _corr_spearman,
_rank_positions, _rank_positions,
_sample_local_candidates, _sample_local_candidates,
calculate_candidate_objectives,
infer_curve_layout, infer_curve_layout,
load_model, load_model,
predict_surrogate_curve, predict_surrogate_curve,
resolve_prediction_curve_time,
run_solver_and_extract_curve, run_solver_and_extract_curve,
) )
from src.common.config import Config from src.common.config import Config
from src.common.experiment_paths import config_for_stage, model_checkpoint_for_tag, normalize_tag, processed_path_for_tag from src.common.experiment_paths import config_for_stage, model_checkpoint_for_tag, normalize_tag, processed_path_for_tag
from src.data.params import Params, Schedule, generate_params_dataset from src.data.params import Params, Schedule, generate_params_dataset
from src.data.runner_client import CppRunner from src.data.runner_client import CppRunner
from src.evaluation.autofit_objective import dual_log_objective
def parse_span_fracs(text: str) -> list[float]: def parse_span_fracs(text: str) -> list[float]:
@ -57,9 +58,11 @@ def parse_args() -> argparse.Namespace:
description="Run local ranking validation on multiple synthetic target cases" description="Run local ranking validation on multiple synthetic target cases"
) )
parser.add_argument("--config", type=str, default=None) parser.add_argument("--config", type=str, default=None)
parser.add_argument("--dataset-case", type=str, default=None)
parser.add_argument("--solver-type", type=int, choices=[1, 2, 3, 4, 5], default=None)
parser.add_argument( parser.add_argument(
"--stage", "--stage",
choices=["fixed_case", "case_neighborhood", "family_random", "family_random_hard"], choices=["fixed_case", "case_neighborhood", "family_random", "family_random_hard", "family_random_v2_q"],
default="family_random", default="family_random",
) )
parser.add_argument("--processed", type=str, default=None) parser.add_argument("--processed", type=str, default=None)
@ -73,6 +76,12 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--well-index", type=int, default=0) parser.add_argument("--well-index", type=int, default=0)
parser.add_argument("--solver-timeout", type=int, default=120) parser.add_argument("--solver-timeout", type=int, default=120)
parser.add_argument("--target-max-attempts-factor", type=int, default=5) parser.add_argument("--target-max-attempts-factor", type=int, default=5)
parser.add_argument("--fixed-cf", type=float, default=None)
parser.add_argument(
"--schedule-mode",
choices=["fixed_case", "case_neighborhood", "family_random"],
default=None,
)
return parser.parse_args() return parser.parse_args()
@ -90,7 +99,15 @@ def resolve_paths(args: argparse.Namespace) -> tuple[Config, Path, Path, Path]:
if args.output_dir is not None if args.output_dir is not None
else Path("results") / f"autofit_local_validation_batch_{tag}" else Path("results") / f"autofit_local_validation_batch_{tag}"
) )
return Config(config_path), processed_path.resolve(), model_path.resolve(), output_dir.resolve() cfg = Config(config_path, dataset_case=args.dataset_case)
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, processed_path.resolve(), model_path.resolve(), output_dir.resolve()
def sample_target_case(cfg: Config, rng: np.random.RandomState, seed: int) -> Params: def sample_target_case(cfg: Config, rng: np.random.RandomState, seed: int) -> Params:
@ -219,6 +236,15 @@ def main() -> None:
model, use_schedule, device = load_model(model_path) model, use_schedule, device = load_model(model_path)
span_fracs = parse_span_fracs(args.span_fracs) span_fracs = parse_span_fracs(args.span_fracs)
rng = np.random.RandomState(int(args.seed)) rng = np.random.RandomState(int(args.seed))
configured_solver_type = cfg.get("solver_type", default=None)
if args.solver_type is not None and configured_solver_type is not None:
if int(args.solver_type) != int(configured_solver_type):
raise ValueError(
f"--solver-type={args.solver_type} conflicts with "
f"dataset case solver_type={configured_solver_type}"
)
solver_type = args.solver_type if args.solver_type is not None else configured_solver_type
prediction_curve_time = resolve_prediction_curve_time(processed, curve_layout, solver_type)
all_candidate_rows: list[dict] = [] all_candidate_rows: list[dict] = []
target_rows: list[dict] = [] target_rows: list[dict] = []
@ -234,7 +260,7 @@ def main() -> None:
target_params = sample_target_case(cfg, rng, seed=int(args.seed) + 100000 + attempt) target_params = sample_target_case(cfg, rng, seed=int(args.seed) + 100000 + attempt)
target_runner = make_runner(cfg, output_dir, f"target_{target_id:03d}") target_runner = make_runner(cfg, output_dir, f"target_{target_id:03d}")
try: try:
target_curve, _target_raw = run_solver_and_extract_curve( target_curve, target_raw = run_solver_and_extract_curve(
runner=target_runner, runner=target_runner,
cfg=cfg, cfg=cfg,
params=target_params, params=target_params,
@ -274,7 +300,7 @@ def main() -> None:
runner = make_runner(cfg, output_dir, f"target_{target_id:03d}_span_{span_frac:g}_cand_{cand_id:03d}") runner = make_runner(cfg, output_dir, f"target_{target_id:03d}_span_{span_frac:g}_cand_{cand_id:03d}")
try: try:
# 候选真实曲线和代理曲线都与同一 target_curve 比较,保证目标函数可排序。 # 候选真实曲线和代理曲线都与同一 target_curve 比较,保证目标函数可排序。
solver_curve, _ = run_solver_and_extract_curve( solver_curve, solver_raw = run_solver_and_extract_curve(
runner=runner, runner=runner,
cfg=cfg, cfg=cfg,
params=cand, params=cand,
@ -290,9 +316,17 @@ def main() -> None:
params=cand, params=cand,
schedule=cand.schedule, schedule=cand.schedule,
cfg=cfg, cfg=cfg,
solver_type=None if solver_type is None else int(solver_type),
)
solver_obj, surrogate_obj = calculate_candidate_objectives(
target_curve=target_curve,
target_raw=target_raw,
solver_curve=solver_curve,
solver_raw=solver_raw,
surrogate_curve=pred_curve,
curve_layout=curve_layout,
prediction_curve_time=prediction_curve_time,
) )
solver_obj = dual_log_objective(target_curve, solver_curve, curve_layout)
surrogate_obj = dual_log_objective(target_curve, pred_curve, curve_layout)
rows.append( rows.append(
{ {
"target_id": target_id, "target_id": target_id,
@ -390,6 +424,8 @@ def main() -> None:
"config_path": str(cfg.path), "config_path": str(cfg.path),
"processed_path": str(processed_path), "processed_path": str(processed_path),
"model_path": str(model_path), "model_path": str(model_path),
"dataset_case": cfg.dataset_case,
"solver_type": None if solver_type is None else int(solver_type),
"n_targets_requested": int(args.n_targets), "n_targets_requested": int(args.n_targets),
"n_target_span_rows": int(len(target_rows)), "n_target_span_rows": int(len(target_rows)),
"span_fracs": span_fracs, "span_fracs": span_fracs,

@ -2,8 +2,8 @@
"""自动拟合候选曲线目标函数。 """自动拟合候选曲线目标函数。
本模块提供轻量级的曲线误差计算用于比较目标曲线和代理模型预测曲线的匹配度 本模块提供轻量级的曲线误差计算用于比较目标曲线和代理模型预测曲线的匹配度
目标函数同时考虑相对误差和绝对误差并分别计算 log_pressure log_derivative 目标函数C++ 真实求解器的当前误差公式保持一致并分别计算 log_pressure
两条曲线的贡献适合在参数筛选PSO 候选排序和局部邻域验证中复用 log_derivative 两条曲线的贡献适合在参数筛选PSO 候选排序和局部邻域验证中复用
""" """
from __future__ import annotations from __future__ import annotations
@ -24,7 +24,7 @@ def split_curve_by_layout(curve: np.ndarray, layout: dict) -> dict[str, np.ndarr
def calculate_curve_objective_1d(target: np.ndarray, pred: np.ndarray) -> float: def calculate_curve_objective_1d(target: np.ndarray, pred: np.ndarray) -> float:
"""计算单段曲线的自动拟合目标,兼顾相对误差和绝对误差。""" """在自然对数曲线上计算与 C++ calculatePointError 等价的均方根误差。"""
target = np.asarray(target, dtype=np.float64).reshape(-1) target = np.asarray(target, dtype=np.float64).reshape(-1)
pred = np.asarray(pred, dtype=np.float64).reshape(-1) pred = np.asarray(pred, dtype=np.float64).reshape(-1)
@ -33,18 +33,12 @@ def calculate_curve_objective_1d(target: np.ndarray, pred: np.ndarray) -> float:
if not (np.isfinite(target).all() and np.isfinite(pred).all()): if not (np.isfinite(target).all() and np.isfinite(pred).all()):
return float("inf") return float("inf")
# 对较大的 log 值适当降权,避免其完全主导排序结果。 # 模型曲线已经是 ln(raw)。因此 C++ 中的 logError 就是两列之差;
weight_factor = np.minimum(100.0, np.abs(target) * 0.01) # raw 相对误差可稳定地化简为 1-exp(-|delta_log|),无需先 exp 回原始尺度。
weight = 1.0 / (1.0 + weight_factor) log_error = np.abs(target - pred)
relative_error = -np.expm1(-log_error)
scale = np.maximum(np.maximum(np.abs(target), np.abs(pred)), 1e-12) point_error = 0.7 * log_error + 0.3 * relative_error
relative_error = np.abs(target - pred) / scale return float(np.sqrt(np.mean(point_error**2)))
absolute_error = np.abs(target - pred)
# 相对误差用于跨尺度比较曲线形态,绝对误差用于保留整体幅值差异。
point_error = 0.7 * relative_error + 0.3 * absolute_error
weighted_mse = np.sum(weight * (point_error**2)) / max(np.sum(weight), 1e-12)
return float(np.sqrt(weighted_mse))
def dual_log_objective( def dual_log_objective(
@ -69,3 +63,215 @@ def dual_log_objective(
"log_derivative_objective": float(d_obj), "log_derivative_objective": float(d_obj),
"dual_log_objective": float(combined), "dual_log_objective": float(combined),
} }
def prediction_curve_time_from_meta(meta: dict, curve_layout: dict) -> np.ndarray:
"""读取预处理模型数据中保存的固定物理时间网格。"""
mode = str(meta.get("curve_time_mode", "missing"))
if mode != "fixed":
raise ValueError(
"surrogate scoring requires curve_time_mode='fixed'; "
f"got curve_time_mode={mode!r}"
)
raw_time = meta.get("prediction_curve_time")
if raw_time is None:
raise ValueError("curve_time_mode is fixed but prediction_curve_time is missing")
prediction_time = np.asarray(raw_time, dtype=np.float64).reshape(-1)
pressure_part = next(
(part for part in curve_layout["parts"] if str(part["name"]) == "log_pressure"),
None,
)
if pressure_part is None:
raise ValueError("curve_layout has no log_pressure part")
expected_size = int(pressure_part["end"]) - int(pressure_part["start"])
if prediction_time.size != expected_size:
raise ValueError(
f"prediction_curve_time has {prediction_time.size} points; expected {expected_size}"
)
if not np.isfinite(prediction_time).all() or np.any(prediction_time <= 0.0):
raise ValueError("prediction_curve_time must contain positive finite values")
if np.any(np.diff(prediction_time) <= 0.0):
raise ValueError("prediction_curve_time must be strictly increasing")
return prediction_time
def _prepare_timed_raw_curve(
time: np.ndarray,
pressure: np.ndarray,
derivative: np.ndarray,
) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
"""清洗并排序一条原始压力/导数曲线,不在此处进行重采样。"""
time = np.asarray(time, dtype=np.float64).reshape(-1)
pressure = np.asarray(pressure, dtype=np.float64).reshape(-1)
derivative = np.asarray(derivative, dtype=np.float64).reshape(-1)
if time.size != pressure.size or time.size != derivative.size:
raise ValueError("time, pressure and derivative lengths do not match")
valid = (
np.isfinite(time)
& np.isfinite(pressure)
& np.isfinite(derivative)
& (time > 0.0)
)
time = time[valid]
pressure = pressure[valid]
derivative = derivative[valid]
if time.size < 3:
raise ValueError("timed curve has fewer than three valid points")
order = np.argsort(time, kind="stable")
time = time[order]
pressure = pressure[order]
derivative = derivative[order]
keep = np.ones(time.size, dtype=bool)
keep[1:] = time[1:] > time[:-1]
time = time[keep]
pressure = pressure[keep]
derivative = derivative[keep]
if time.size < 3:
raise ValueError("timed curve has fewer than three unique time points")
return time, pressure, derivative
def _calculate_raw_curve_objective_1d(target: np.ndarray, pred: np.ndarray) -> float:
"""在原始数值上复现 nmCalculationAutoFitPSO::calculatePointError。"""
target = np.asarray(target, dtype=np.float64).reshape(-1)
pred = np.asarray(pred, dtype=np.float64).reshape(-1)
if target.size == 0 or pred.size != target.size:
return float("inf")
if not (np.isfinite(target).all() and np.isfinite(pred).all()):
return float("inf")
target_abs = np.maximum(np.abs(target), 1.0e-10)
pred_abs = np.maximum(np.abs(pred), 1.0e-10)
log_error = np.abs(np.log(target_abs) - np.log(pred_abs))
relative_error = np.abs(target - pred) / np.maximum(
np.maximum(np.abs(target), np.abs(pred)),
1.0e-10,
)
point_error = 0.7 * log_error + 0.3 * relative_error
return float(np.sqrt(np.mean(point_error**2)))
def timed_raw_dual_objective(
target_time: np.ndarray,
target_pressure: np.ndarray,
target_derivative: np.ndarray,
pred_time: np.ndarray,
pred_pressure: np.ndarray,
pred_derivative: np.ndarray,
n_common_points: int = 50,
w_pressure: float = 0.5,
w_derivative: float = 0.5,
) -> dict[str, float]:
"""按照 C++ 目标函数的规则,在同一物理时间网格上比较两条原始曲线。"""
if int(n_common_points) < 2:
raise ValueError("n_common_points must be at least two")
target_time, target_pressure, target_derivative = _prepare_timed_raw_curve(
target_time,
target_pressure,
target_derivative,
)
pred_time, pred_pressure, pred_derivative = _prepare_timed_raw_curve(
pred_time,
pred_pressure,
pred_derivative,
)
overlap_start = max(float(target_time[0]), float(pred_time[0]))
overlap_end = min(float(target_time[-1]), float(pred_time[-1]))
if overlap_start <= 0.0 or overlap_start >= overlap_end:
return {
"log_pressure_objective": float("inf"),
"log_derivative_objective": float("inf"),
"dual_log_objective": float("inf"),
}
common_time = np.geomspace(overlap_start, overlap_end, int(n_common_points))
target_pressure_common = np.interp(common_time, target_time, target_pressure)
target_derivative_common = np.interp(common_time, target_time, target_derivative)
pred_pressure_common = np.interp(common_time, pred_time, pred_pressure)
pred_derivative_common = np.interp(common_time, pred_time, pred_derivative)
p_obj = _calculate_raw_curve_objective_1d(
target_pressure_common,
pred_pressure_common,
)
d_obj = _calculate_raw_curve_objective_1d(
target_derivative_common,
pred_derivative_common,
)
total_w = max(float(w_pressure) + float(w_derivative), 1.0e-12)
combined = (float(w_pressure) * p_obj + float(w_derivative) * d_obj) / total_w
return {
"log_pressure_objective": float(p_obj),
"log_derivative_objective": float(d_obj),
"dual_log_objective": float(combined),
}
def timed_dual_log_objective(
target_time: np.ndarray,
target_pressure: np.ndarray,
target_derivative: np.ndarray,
pred_time: np.ndarray,
pred_curve: np.ndarray,
curve_layout: dict,
n_common_points: int = 50,
w_pressure: float = 0.5,
w_derivative: float = 0.5,
) -> dict[str, float]:
"""按照 C++ 时间对齐规则,将预测对数曲线与原始目标曲线进行比较。"""
pred_parts = split_curve_by_layout(pred_curve, curve_layout)
pred_log_pressure = pred_parts["log_pressure"]
pred_log_derivative = pred_parts["log_derivative"]
pred_time = np.asarray(pred_time, dtype=np.float64).reshape(-1)
if pred_time.size != pred_log_pressure.size or pred_time.size != pred_log_derivative.size:
raise ValueError("pred_time length does not match predicted curve parts")
target_time, target_pressure, target_derivative = _prepare_timed_raw_curve(
target_time,
target_pressure,
target_derivative,
)
# 不使用超出目标曲线实际时间范围的固定网格输出做插值;这些位置在训练时已由掩码排除。
covered = pred_time <= target_time[-1]
if int(np.sum(covered)) < 3:
raise ValueError("target range contains fewer than three prediction time points")
pred_time = pred_time[covered]
pred_log_pressure = pred_log_pressure[covered]
pred_log_derivative = pred_log_derivative[covered]
with np.errstate(over="ignore", invalid="ignore"):
pred_pressure = np.exp(pred_log_pressure)
pred_derivative = np.exp(pred_log_derivative)
target_end = float(target_time[-1])
if pred_time[-1] < target_end:
dt = float(pred_time[-1] - pred_time[-2])
if dt <= 0.0:
raise ValueError("prediction time grid has a non-positive final interval")
fraction = (target_end - float(pred_time[-1])) / dt
pressure_end = pred_pressure[-1] + fraction * (pred_pressure[-1] - pred_pressure[-2])
derivative_end = pred_derivative[-1] + fraction * (
pred_derivative[-1] - pred_derivative[-2]
)
pred_time = np.append(pred_time, target_end)
pred_pressure = np.append(pred_pressure, max(float(pressure_end), 1.0e-300))
pred_derivative = np.append(pred_derivative, max(float(derivative_end), 1.0e-300))
return timed_raw_dual_objective(
target_time=target_time,
target_pressure=target_pressure,
target_derivative=target_derivative,
pred_time=pred_time,
pred_pressure=pred_pressure,
pred_derivative=pred_derivative,
n_common_points=n_common_points,
w_pressure=w_pressure,
w_derivative=w_derivative,
)

@ -0,0 +1,331 @@
from __future__ import annotations
import sys
import tempfile
import unittest
from pathlib import Path
import h5py
import joblib
import numpy as np
import torch
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from src.common.config import Config
from src.data.curve_processing import resample_curve_to_features_with_time_and_mask
from src.data.preprocess import MaskedStandardScaler
from src.data.preprocess import preprocess_dataset
from src.evaluation.autofit_objective import (
prediction_curve_time_from_meta,
timed_dual_log_objective,
)
from src.training.train_forward import (
CurveStats,
LossConfig,
LossContext,
SampleReweightConfig,
compute_weighted_loss,
load_processed_dataset,
smooth_l1_per_sample,
)
def _write_minimal_raw_h5(
path: Path,
curve_time: np.ndarray,
*,
include_mask: bool,
) -> None:
n_samples, n_time_points = curve_time.shape
with h5py.File(path, "w") as f:
f.create_dataset("params", data=np.ones((n_samples, 7), dtype=np.float32))
f.create_dataset("schedule", data=np.zeros((n_samples, 4), dtype=np.float32))
f.create_dataset(
"curve",
data=np.ones((n_samples, 2 * n_time_points), dtype=np.float32),
)
f.create_dataset("curve_time", data=curve_time.astype(np.float32))
if include_mask:
f.create_dataset(
"curve_valid_mask",
data=np.ones((n_samples, n_time_points), dtype=np.uint8),
)
class FixedTimePipelineTests(unittest.TestCase):
def setUp(self) -> None:
self.cfg = Config(ROOT / "configs" / "data_gen_family_random_v2_q.yaml")
def test_all_curves_use_one_physical_time_grid(self) -> None:
grids = []
masks = []
for end_time in (12.0, 30.0, 72.0, 120.0, 140.0):
time = np.geomspace(1.6e-3, end_time, 240)
pressure = 10.0 + np.log1p(time)
derivative = 0.1 + np.sqrt(time)
features, grid, valid_mask = resample_curve_to_features_with_time_and_mask(
self.cfg,
time,
pressure,
derivative,
)
self.assertEqual(features.shape, (320,))
grids.append(grid)
masks.append(valid_mask)
for grid in grids[1:]:
np.testing.assert_array_equal(grid, grids[0])
self.assertAlmostEqual(float(grids[0][0]), 1.6e-3, places=8)
self.assertAlmostEqual(float(grids[0][-1]), 139.0, places=5)
self.assertLess(int(masks[0].sum()), int(masks[-1].sum()))
def test_fixed_grid_point_before_curve_start_is_invalid(self) -> None:
time = np.geomspace(1.61e-3, 139.0, 240)
_, grid, valid_mask = resample_curve_to_features_with_time_and_mask(
self.cfg,
time,
np.ones_like(time),
np.ones_like(time),
)
self.assertAlmostEqual(float(grid[0]), 1.6e-3, places=8)
self.assertEqual(int(valid_mask[0]), 0)
self.assertEqual(int(valid_mask[1]), 1)
def test_masked_scaler_ignores_placeholder_values(self) -> None:
rng = np.random.RandomState(7)
values = rng.normal(size=(12, 8))
mask = np.ones_like(values, dtype=bool)
mask[1:, 5:] = False
changed = values.copy()
changed[~mask] = 1.0e9
scaler_a = MaskedStandardScaler().fit(values, mask)
scaler_b = MaskedStandardScaler().fit(changed, mask)
np.testing.assert_allclose(scaler_a.mean_, scaler_b.mean_, atol=0.0, rtol=0.0)
np.testing.assert_allclose(scaler_a.scale_, scaler_b.scale_, atol=0.0, rtol=0.0)
def test_masked_loss_has_no_gradient_outside_coverage(self) -> None:
pred = torch.tensor([[0.2, 0.4, 3.0, 4.0]], dtype=torch.float32, requires_grad=True)
target = torch.zeros_like(pred)
mask = torch.tensor([[1.0, 1.0, 0.0, 0.0]], dtype=torch.float32)
loss = smooth_l1_per_sample(pred, target, beta=0.05, valid_mask=mask).mean()
loss.backward()
self.assertGreater(float(pred.grad[0, 0]), 0.0)
self.assertGreater(float(pred.grad[0, 1]), 0.0)
self.assertEqual(float(pred.grad[0, 2]), 0.0)
self.assertEqual(float(pred.grad[0, 3]), 0.0)
def test_composite_training_loss_respects_time_mask(self) -> None:
pred = torch.zeros((2, 320), dtype=torch.float32, requires_grad=True)
target = torch.ones_like(pred)
valid_mask = torch.zeros((2, 160), dtype=torch.float32)
valid_mask[:, :80] = 1.0
context = LossContext(
slices={"log_pressure": slice(0, 160), "log_derivative": slice(160, 320)},
curve_stats=CurveStats(
mean_raw=torch.zeros(320, dtype=torch.float32),
scale_raw=torch.ones(320, dtype=torch.float32),
),
loss_cfg=LossConfig(),
reweight_cfg=SampleReweightConfig(enabled=False),
)
losses = compute_weighted_loss(pred, target, valid_mask, context)
losses["loss"].backward()
self.assertTrue(torch.isfinite(losses["loss"]))
self.assertGreater(float(pred.grad[:, :80].abs().sum()), 0.0)
self.assertGreater(float(pred.grad[:, 160:240].abs().sum()), 0.0)
self.assertEqual(float(pred.grad[:, 80:160].abs().sum()), 0.0)
self.assertEqual(float(pred.grad[:, 240:].abs().sum()), 0.0)
def test_perfect_fit_composite_loss_has_finite_gradient(self) -> None:
pred = torch.zeros((1, 4), dtype=torch.float32, requires_grad=True)
target = torch.zeros_like(pred)
valid_mask = torch.ones((1, 2), dtype=torch.float32)
context = LossContext(
slices={"log_pressure": slice(0, 2), "log_derivative": slice(2, 4)},
curve_stats=CurveStats(
mean_raw=torch.zeros(4, dtype=torch.float32),
scale_raw=torch.ones(4, dtype=torch.float32),
),
loss_cfg=LossConfig(),
reweight_cfg=SampleReweightConfig(enabled=False),
)
losses = compute_weighted_loss(pred, target, valid_mask, context)
losses["loss"].backward()
self.assertTrue(torch.isfinite(losses["loss"]))
self.assertTrue(torch.isfinite(pred.grad).all())
self.assertEqual(float(pred.grad.abs().sum()), 0.0)
def test_surrogate_metadata_requires_fixed_time(self) -> None:
layout = {
"parts": [
{"name": "log_pressure", "start": 0, "end": 3},
{"name": "log_derivative", "start": 3, "end": 6},
]
}
with self.assertRaises(ValueError):
prediction_curve_time_from_meta({"curve_time_mode": "missing"}, layout)
with self.assertRaises(ValueError):
prediction_curve_time_from_meta({"curve_time_mode": "per_sample"}, layout)
result = prediction_curve_time_from_meta(
{
"curve_time_mode": "fixed",
"prediction_curve_time": [0.1, 1.0, 10.0],
},
layout,
)
np.testing.assert_array_equal(result, np.asarray([0.1, 1.0, 10.0]))
def test_fixed_grid_objective_aligns_on_physical_time(self) -> None:
prediction_time = np.asarray([0.1, 1.0, 2.0, 4.0], dtype=np.float64)
prediction_pressure = 2.0 + 3.0 * prediction_time
prediction_derivative = 1.0 + 0.5 * prediction_time
prediction_curve = np.concatenate(
[np.log(prediction_pressure), np.log(prediction_derivative)]
)
target_time = np.asarray([0.3, 0.8, 1.5, 2.3, 3.0], dtype=np.float64)
layout = {
"parts": [
{"name": "log_pressure", "start": 0, "end": 4},
{"name": "log_derivative", "start": 4, "end": 8},
]
}
objective = timed_dual_log_objective(
target_time=target_time,
target_pressure=2.0 + 3.0 * target_time,
target_derivative=1.0 + 0.5 * target_time,
pred_time=prediction_time,
pred_curve=prediction_curve,
curve_layout=layout,
)
self.assertLess(objective["dual_log_objective"], 1.0e-12)
def test_preprocess_preserves_fixed_grid_and_masks(self) -> None:
rng = np.random.RandomState(11)
n_groups = 10
n = n_groups * 5
grid = np.geomspace(1.6e-3, 139.0, 160).astype(np.float32)
params = np.column_stack(
[
rng.uniform(1.0e-3, 1.0, n),
rng.uniform(-2.0, 2.0, n),
rng.uniform(1.0e-4, 0.1, n),
rng.uniform(0.01, 0.2, n),
rng.uniform(2.0, 20.0, n),
rng.uniform(1.0e-4, 2.0e-3, n),
np.tile(np.arange(1, 6, dtype=np.float32), n_groups),
]
).astype(np.float32)
schedule = rng.normal(size=(n, 12)).astype(np.float32)
curve = rng.normal(size=(n, 320)).astype(np.float32)
mask = np.zeros((n, 160), dtype=np.uint8)
for row in range(n):
mask[row, : 80 + 20 * (row % 5)] = 1
with tempfile.TemporaryDirectory() as temp_dir:
h5_path = Path(temp_dir) / "fixed.h5"
pkl_path = Path(temp_dir) / "fixed.pkl"
with h5py.File(h5_path, "w") as f:
f.create_dataset("params", data=params)
f.create_dataset("schedule", data=schedule)
f.create_dataset("curve", data=curve)
f.create_dataset("curve_time", data=np.tile(grid, (n, 1)))
f.create_dataset("curve_valid_mask", data=mask)
f.create_dataset("group_id", data=np.repeat(np.arange(n_groups), 5))
f.attrs["param_names"] = np.asarray(
["k", "skin", "wellboreC", "phi", "h", "Cf", "solverType"],
dtype="S",
)
preprocess_dataset(h5_path, pkl_path, test_size=0.2, val_size=0.2, random_seed=3)
processed = joblib.load(pkl_path)
self.assertEqual(processed["meta"]["curve_time_mode"], "fixed")
self.assertEqual(processed["meta"]["curve_valid_mask_mode"], "explicit")
np.testing.assert_array_equal(
np.asarray(processed["meta"]["prediction_curve_time"], dtype=np.float32),
grid,
)
train_mask = processed["curve_valid_mask_train"].astype(bool)
full_train_mask = np.concatenate([train_mask, train_mask], axis=1)
self.assertTrue(np.all(processed["Y_curve_train"][~full_train_mask] == 0.0))
self.assertTrue(np.all(np.asarray(processed["scaler_curve"].n_samples_seen_) > 0))
def test_preprocess_keeps_240_dim_two_channel_curve(self) -> None:
fixed_grid = np.geomspace(1.6e-3, 139.0, 120)
with tempfile.TemporaryDirectory() as temp_dir:
h5_path = Path(temp_dir) / "two_channel_240.h5"
pkl_path = Path(temp_dir) / "two_channel_240.pkl"
_write_minimal_raw_h5(
h5_path,
np.tile(fixed_grid, (20, 1)),
include_mask=True,
)
preprocess_dataset(h5_path, pkl_path, random_seed=5)
processed = joblib.load(pkl_path)
self.assertEqual(processed["meta"]["curve_dim"], 240)
self.assertEqual(processed["meta"]["curve_layout"]["n_time_points"], 120)
self.assertFalse(processed["meta"]["dropped_legacy_slope"])
self.assertEqual(processed["Y_curve_train"].shape[1], 240)
self.assertEqual(processed["curve_valid_mask_train"].shape[1], 120)
np.testing.assert_allclose(
processed["meta"]["prediction_curve_time"],
fixed_grid,
rtol=1.0e-6,
atol=1.0e-10,
)
def test_preprocess_rejects_per_sample_time_grid(self) -> None:
fixed_grid = np.geomspace(1.6e-3, 139.0, 160)
other_grid = np.geomspace(1.6e-3, 120.0, 160)
with tempfile.TemporaryDirectory() as temp_dir:
h5_path = Path(temp_dir) / "per_sample.h5"
_write_minimal_raw_h5(
h5_path,
np.stack([fixed_grid, other_grid], axis=0),
include_mask=True,
)
with self.assertRaisesRegex(ValueError, "one fixed curve_time grid"):
preprocess_dataset(h5_path, Path(temp_dir) / "unused.pkl")
def test_preprocess_rejects_missing_curve_mask(self) -> None:
fixed_grid = np.geomspace(1.6e-3, 139.0, 160)
with tempfile.TemporaryDirectory() as temp_dir:
h5_path = Path(temp_dir) / "missing_mask.h5"
_write_minimal_raw_h5(
h5_path,
np.tile(fixed_grid, (2, 1)),
include_mask=False,
)
with self.assertRaisesRegex(ValueError, "requires curve_valid_mask"):
preprocess_dataset(h5_path, Path(temp_dir) / "unused.pkl")
def test_training_rejects_legacy_processed_dataset(self) -> None:
legacy = {}
for split in ("train", "val", "test"):
legacy[f"X_params_{split}"] = np.zeros((1, 1), dtype=np.float32)
legacy[f"X_schedule_{split}"] = np.zeros((1, 1), dtype=np.float32)
legacy[f"Y_curve_{split}"] = np.zeros((1, 4), dtype=np.float32)
with tempfile.TemporaryDirectory() as temp_dir:
path = Path(temp_dir) / "legacy.pkl"
joblib.dump(legacy, path)
with self.assertRaisesRegex(ValueError, "curve_time_mode='fixed'"):
load_processed_dataset(path)
if __name__ == "__main__":
unittest.main()
Loading…
Cancel
Save