支持局部邻域数据和差分损失训练

- 支持根据相态、固定参数和拟合记录生成局部数据
- 展开锚点与邻域样本并保留分组和来源信息
- 训练时使用有效掩码排除无效曲线区间
- 增加压力和压力导数的邻域差分损失
feature/Model-20260625
lvjunjie 2 days ago
parent 0e83f2cd44
commit be498f8ffc

@ -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[

@ -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,
)

@ -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,

@ -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,13 +928,89 @@ 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"])
delta_enabled = bool(
cfg.loss.weights.delta_pressure != 0.0
or cfg.loss.weights.delta_derivative != 0.0
)
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,
@ -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,
}

Loading…
Cancel
Save