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

- 支持根据相态、固定参数和拟合记录生成局部数据
- 展开锚点与邻域样本并保留分组和来源信息
- 训练时使用有效掩码排除无效曲线区间
- 增加压力和压力导数的邻域差分损失
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)) 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: def parse_args() -> argparse.Namespace:
"""解析自动拟合邻域 HDF5 的输入输出路径以及是否只导出邻域样本。""" """解析自动拟合邻域 HDF5 的输入输出路径以及是否只导出邻域样本。"""
parser = argparse.ArgumentParser( parser = argparse.ArgumentParser(
@ -71,15 +82,21 @@ def main() -> None:
anchor_params = np.asarray(src["anchor_params"][:], dtype=np.float32) anchor_params = np.asarray(src["anchor_params"][:], dtype=np.float32)
anchor_schedule = np.asarray(src["anchor_schedule"][:], 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 = 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_schedule_meta = np.asarray(src["anchor_schedule_meta"][:], dtype=np.float32)
anchor_family_name = np.asarray(src["anchor_family_name"][:]).astype(str) anchor_family_name = np.asarray(src["anchor_family_name"][:]).astype(str)
anchor_section_index = np.asarray(src["anchor_section_index"][:], dtype=np.int32) 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) anchor_q_json = np.asarray(src["anchor_q_json"][:]).astype(str)
neighbor_anchor_id = np.asarray(src["neighbor_anchor_id"][:], dtype=np.int32) neighbor_anchor_id = np.asarray(src["neighbor_anchor_id"][:], dtype=np.int32)
neighbor_params = np.asarray(src["neighbor_params"][:], dtype=np.float32) neighbor_params = np.asarray(src["neighbor_params"][:], dtype=np.float32)
neighbor_curve = np.asarray(src["neighbor_curve"][:], 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 = np.asarray(src["neighbor_objective"][:], dtype=np.float32)
neighbor_objective_p = np.asarray(src["neighbor_objective_p"][:], 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) neighbor_objective_d = np.asarray(src["neighbor_objective_d"][:], dtype=np.float32)
@ -100,6 +117,20 @@ def main() -> None:
.tolist() .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_anchors = int(anchor_params.shape[0])
n_neighbors = int(neighbor_params.shape[0]) n_neighbors = int(neighbor_params.shape[0])
@ -125,11 +156,14 @@ def main() -> None:
param_dim = int(anchor_params.shape[1]) param_dim = int(anchor_params.shape[1])
schedule_dim = int(anchor_schedule.shape[1]) schedule_dim = int(anchor_schedule.shape[1])
curve_dim = int(anchor_curve.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]) schedule_meta_dim = int(anchor_schedule_meta.shape[1])
dst.create_dataset("params", shape=(total_rows, param_dim), dtype=np.float32) 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("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", 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("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("schedule_meta", shape=(total_rows, schedule_meta_dim), dtype=np.float32)
dst.create_dataset( dst.create_dataset(
@ -163,6 +197,8 @@ def main() -> None:
dst["params"][write_pos:anchor_end] = anchor_params dst["params"][write_pos:anchor_end] = anchor_params
dst["schedule"][write_pos:anchor_end] = anchor_schedule dst["schedule"][write_pos:anchor_end] = anchor_schedule
dst["curve"][write_pos:anchor_end] = anchor_curve 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["group_id"][write_pos:anchor_end] = np.arange(n_anchors, dtype=np.int32)
dst["schedule_meta"][write_pos:anchor_end] = anchor_schedule_meta dst["schedule_meta"][write_pos:anchor_end] = anchor_schedule_meta
dst["family_name"][write_pos:anchor_end] = anchor_family_name.tolist() 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["params"][write_pos:neighbor_end] = neighbor_params
dst["schedule"][write_pos:neighbor_end] = anchor_schedule[neighbor_anchor_id] dst["schedule"][write_pos:neighbor_end] = anchor_schedule[neighbor_anchor_id]
dst["curve"][write_pos:neighbor_end] = neighbor_curve 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["group_id"][write_pos:neighbor_end] = neighbor_anchor_id
dst["schedule_meta"][write_pos:neighbor_end] = anchor_schedule_meta[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[ dst["family_name"][write_pos:neighbor_end] = anchor_family_name[

@ -19,7 +19,9 @@
from __future__ import annotations from __future__ import annotations
import argparse import argparse
import csv
import json import json
import math
import sys import sys
from collections import Counter from collections import Counter
from pathlib import Path 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 ( from src.data.curve_processing import (
clean_curve_for_dataset, clean_curve_for_dataset,
is_valid_curve, 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.params import Params, Schedule, generate_params_dataset
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 ( 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]: def _family_id_map(cfg: Config) -> dict[str, int]:
"""把流量制度族名称映射为整数 id便于写入模型特征和元数据。""" """把流量制度族名称映射为整数 id便于写入模型特征和元数据。"""
mode = str(cfg.raw["schedule"]["generation_mode"]).lower() mode = str(cfg.raw["schedule"]["generation_mode"]).lower()
@ -184,59 +157,12 @@ def _sample_schedule_family_random(
rng: np.random.RandomState, rng: np.random.RandomState,
family_override: str | None = None, family_override: str | None = None,
) -> tuple[list[float], list[float], dict]: ) -> tuple[list[float], list[float], dict]:
"""按 family_random 配置随机生成不同类别的生产/关井制度。""" """复用正式数据生成器,使局部数据与普通数据的流量制度分布保持一致。"""
fcfg = cfg.raw["schedule"]["family_random"] return sample_schedule_family_random(
if family_override is None: cfg,
fam = _pick_mixture(rng, fcfg["families"]) rng,
fam_name = str(fam.get("name", "inc_tail_shutin")).lower() family_override=family_override,
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}
def sample_schedule_by_mode( 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 = argparse.ArgumentParser(description="Generate anchor-neighborhood autofit dataset for generalized local ranking")
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( parser.add_argument(
"--stage", "--stage",
choices=[ choices=[
@ -292,6 +219,13 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--neighbors-per-anchor", type=int, default=24) parser.add_argument("--neighbors-per-anchor", type=int, default=24)
parser.add_argument("--max-attempts-factor", type=int, default=4) parser.add_argument("--max-attempts-factor", type=int, default=4)
parser.add_argument("--anchor-max-attempts-factor", type=int, default=5) 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("--seed", type=int, default=42)
parser.add_argument("--span-frac", type=float, default=0.08) parser.add_argument("--span-frac", type=float, default=0.08)
parser.add_argument( parser.add_argument(
@ -311,6 +245,18 @@ def parse_args() -> argparse.Namespace:
) )
parser.add_argument("--solver-timeout", type=int, default=120) parser.add_argument("--solver-timeout", type=int, default=120)
parser.add_argument("--well-index", type=int, default=0) 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( parser.add_argument(
"--use-runner-server", "--use-runner-server",
action="store_true", action="store_true",
@ -342,7 +288,18 @@ def resolve_config(args: argparse.Namespace) -> Config:
config_path = args.config config_path = args.config
if config_path is None: if config_path is None:
config_path = str(config_for_stage(args.stage) or Path("configs/data_gen_family_random.yaml")) 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: def resolve_output_path(cfg: Config, args: argparse.Namespace) -> Path:
@ -372,7 +329,7 @@ def run_solver_and_extract_curve(
params: Params, params: Params,
well_index: int, well_index: int,
timeout: int, timeout: int,
) -> tuple[np.ndarray, np.ndarray, dict]: ) -> tuple[np.ndarray, np.ndarray, np.ndarray, dict]:
"""调用 C++ 求解器运行一次正演,并把双对数输出重采样为模型曲线向量。""" """调用 C++ 求解器运行一次正演,并把双对数输出重采样为模型曲线向量。"""
ok = runner.run_simulation(params, timeout=timeout, override_schedule=params.schedule, include_schedule=True) 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 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: if not valid:
raise RuntimeError(f"curve_invalid_{reason}") 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 = { raw = {
"t": t_clean.tolist(), "t": t_clean.tolist(),
"p": p_clean.tolist(), "p": p_clean.tolist(),
@ -405,24 +367,123 @@ def run_solver_and_extract_curve(
"n_steps": int(result["nSteps"]), "n_steps": int(result["nSteps"]),
"n_wells": int(result["nWells"]), "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: def model_param_names(cfg: Config) -> list[str]:
"""按固定参数顺序把 Params 对象转换成数值数组。""" """返回H5中保存并交给模型的原始参数列顺序。"""
return np.asarray( return list(
[params.k, params.skin, params.wellboreC, params.phi, params.h, params.Ct, params.Cf], cfg.raw["params"].get(
dtype=np.float32, "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( def sample_anchor_params_and_schedule(
cfg: Config, cfg: Config,
rng: np.random.RandomState, rng: np.random.RandomState,
family_override: str | None = None, family_override: str | None = None,
anchor_params: Params | None = None,
) -> tuple[Params, np.ndarray, str]: ) -> 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) timeQ, q, sched_info = sample_schedule_by_mode(cfg, rng, family_override=family_override)
section_indices = _resolve_section_indices(cfg, timeQ, q, rng) section_indices = _resolve_section_indices(cfg, timeQ, q, rng)
sec = int(section_indices[int(rng.randint(0, len(section_indices)))]) 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( 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, k: int,
objective_bins: int, objective_bins: int,
span_fracs: list[float], 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: if len(rows) <= k or len(span_fracs) <= 1:
indices = select_objective_stratified_indices( 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( def create_output_file(
output_path: Path, output_path: Path,
param_dim: int, param_names: list[str],
schedule_dim: int, schedule_dim: int,
curve_dim: int, curve_dim: int,
span_fracs: list[float], span_fracs: list[float],
search_names: list[str], search_names: list[str],
dataset_case: str | None,
solver_type: int | None,
) -> h5py.File: ) -> h5py.File:
"""创建邻域 HDF5 文件,并初始化锚点、候选、曲线和元数据数据集。""" """创建邻域 HDF5 文件,并初始化锚点、候选、曲线和元数据数据集。"""
output_path.parent.mkdir(parents=True, exist_ok=True) output_path.parent.mkdir(parents=True, exist_ok=True)
f = h5py.File(output_path, "w") 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["schedule_meta_names"] = np.asarray(SCHEDULE_META_NAMES, dtype="S")
f.attrs["span_fracs"] = np.asarray(span_fracs, dtype=np.float32) f.attrs["span_fracs"] = np.asarray(span_fracs, dtype=np.float32)
f.attrs["search_param_names"] = np.asarray(search_names, dtype="S") 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_anchors = 0
n_neighbors = 0 n_neighbors = 0
@ -680,6 +746,7 @@ def create_output_file(
n_time_points = int(curve_dim // 2) 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", 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_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_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_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,)) 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_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", 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_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", 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_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,)) 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, schedule_vec: np.ndarray,
curve: np.ndarray, curve: np.ndarray,
curve_time: np.ndarray, curve_time: np.ndarray,
curve_valid_mask: np.ndarray,
schedule_meta: np.ndarray, schedule_meta: np.ndarray,
family_name: str, family_name: str,
section_index: int, section_index: int,
@ -719,6 +788,7 @@ def append_anchor(
("anchor_schedule", schedule_vec.reshape(1, -1)), ("anchor_schedule", schedule_vec.reshape(1, -1)),
("anchor_curve", curve.reshape(1, -1)), ("anchor_curve", curve.reshape(1, -1)),
("anchor_curve_time", curve_time.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)), ("anchor_schedule_meta", schedule_meta.reshape(1, -1)),
]: ]:
ds = f[name] ds = f[name]
@ -743,6 +813,7 @@ def append_neighbors(
params_list: list[np.ndarray], params_list: list[np.ndarray],
curve_list: list[np.ndarray], curve_list: list[np.ndarray],
curve_time_list: list[np.ndarray], curve_time_list: list[np.ndarray],
curve_valid_mask_list: list[np.ndarray],
obj_list: list[dict[str, float]], obj_list: list[dict[str, float]],
span_frac_list: list[float], span_frac_list: list[float],
) -> None: ) -> 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_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")[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_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")[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_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) 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): if use_schedule_meta_features(cfg):
schedule_dim += len(schedule_meta_feature_names(cfg)) schedule_dim += len(schedule_meta_feature_names(cfg))
curve_dim = int(cfg.curve_dim) 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) local_search_names = search_param_names(cfg)
f = create_output_file( f = create_output_file(
output_path, output_path,
param_dim=param_dim, param_names=output_param_names,
schedule_dim=schedule_dim, schedule_dim=schedule_dim,
curve_dim=curve_dim, curve_dim=curve_dim,
span_fracs=span_fracs, span_fracs=span_fracs,
search_names=local_search_names, search_names=local_search_names,
dataset_case=cfg.dataset_case,
solver_type=cfg.get("solver_type", default=None),
) )
summary = { summary = {
"config_path": str(cfg.path), "config_path": str(cfg.path),
@ -809,6 +883,13 @@ def main() -> None:
"use_runner_server": bool(args.use_runner_server), "use_runner_server": bool(args.use_runner_server),
"max_perturbed_dims": int(args.max_perturbed_dims), "max_perturbed_dims": int(args.max_perturbed_dims),
"search_param_names": local_search_names, "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), "objective_bins": int(args.objective_bins),
"anchor_fail_reasons": Counter(), "anchor_fail_reasons": Counter(),
"neighbor_fail_reasons": Counter(), "neighbor_fail_reasons": Counter(),
@ -819,14 +900,38 @@ def main() -> None:
} }
try: try:
anchor_attempts = 0
max_anchor_attempts = max(int(args.n_anchors), int(args.n_anchors) * int(args.anchor_max_attempts_factor)) 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)) 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: while int(f.attrs["n_anchors"]) < int(args.n_anchors) and anchor_attempts < max_anchor_attempts:
anchor_idx = anchor_attempts anchor_idx = anchor_attempts
anchor_attempts += 1 anchor_attempts += 1
target_family = family_plan[int(f.attrs["n_anchors"])] if family_plan else None 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 summary["family_counter_attempted"][family_name] += 1
anchor_runner = CppRunner( anchor_runner = CppRunner(
@ -836,7 +941,7 @@ def main() -> None:
use_server=bool(args.use_runner_server), use_server=bool(args.use_runner_server),
) )
try: 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, runner=anchor_runner,
cfg=cfg, cfg=cfg,
params=anchor_params, params=anchor_params,
@ -856,10 +961,11 @@ def main() -> None:
) )
anchor_id = append_anchor( anchor_id = append_anchor(
f=f, f=f,
anchor_params=params_to_array(anchor_params), anchor_params=params_to_model_array(cfg, anchor_params),
schedule_vec=schedule_vec, schedule_vec=schedule_vec,
curve=anchor_curve, curve=anchor_curve,
curve_time=anchor_curve_time, curve_time=anchor_curve_time,
curve_valid_mask=anchor_curve_valid_mask,
schedule_meta=schedule_meta, schedule_meta=schedule_meta,
family_name=family_name, family_name=family_name,
section_index=int(anchor_params.schedule.sectionIndex), section_index=int(anchor_params.schedule.sectionIndex),
@ -874,7 +980,7 @@ def main() -> None:
) )
valid_neighbor_rows: list[ 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() valid_span_counter: Counter[str] = Counter()
max_attempts = max( max_attempts = max(
@ -901,25 +1007,46 @@ def main() -> None:
max_perturbed_dims=int(args.max_perturbed_dims), max_perturbed_dims=int(args.max_perturbed_dims),
) )
try: 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, runner=anchor_runner,
cfg=cfg, cfg=cfg,
params=cand, params=cand,
well_index=int(args.well_index), well_index=int(args.well_index),
timeout=int(args.solver_timeout), 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 n_time_points = curve_dim // 2
obj = dual_log_objective(anchor_curve, cand_curve, {"parts": [ anchor_common = np.concatenate(
{"name": "log_pressure", "start": 0, "end": n_time_points}, [
{"name": "log_derivative", "start": n_time_points, "end": curve_dim}, 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( valid_neighbor_rows.append(
( (
params_to_array(cand), params_to_model_array(cfg, cand),
cand_curve.astype(np.float32), cand_curve.astype(np.float32),
obj, obj,
span_frac, span_frac,
cand_curve_time.astype(np.float32), cand_curve_time.astype(np.float32),
cand_curve_valid_mask.astype(np.uint8),
) )
) )
summary["neighbor_span_counter_valid"][f"{span_frac:g}"] += 1 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_obj_rows = [x[2] for x in selected_rows]
neighbor_span_rows = [float(x[3]) 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_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: for span_frac in neighbor_span_rows:
summary["neighbor_span_counter_kept"][f"{span_frac:g}"] += 1 summary["neighbor_span_counter_kept"][f"{span_frac:g}"] += 1
@ -973,6 +1101,7 @@ def main() -> None:
params_list=neighbor_params_rows, params_list=neighbor_params_rows,
curve_list=neighbor_curve_rows, curve_list=neighbor_curve_rows,
curve_time_list=neighbor_curve_time_rows, curve_time_list=neighbor_curve_time_rows,
curve_valid_mask_list=neighbor_curve_valid_mask_rows,
obj_list=neighbor_obj_rows, obj_list=neighbor_obj_rows,
span_frac_list=neighbor_span_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-derivative-shape", type=float, default=0.10)
parser.add_argument("--w-autofit-pressure", type=float, default=0.0) 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-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("--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("--use-sample-reweight", action="store_true", default=True)
parser.add_argument("--no-sample-reweight", action="store_false", dest="use_sample_reweight") 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, derivative_shape=args.w_derivative_shape,
autofit_pressure=args.w_autofit_pressure, autofit_pressure=args.w_autofit_pressure,
autofit_derivative=args.w_autofit_derivative, autofit_derivative=args.w_autofit_derivative,
delta_pressure=args.w_delta_pressure,
delta_derivative=args.w_delta_derivative,
), ),
huber_beta=args.huber_beta, huber_beta=args.huber_beta,
delta_huber_beta=args.delta_huber_beta,
), ),
sample_reweight=SampleReweightConfig( sample_reweight=SampleReweightConfig(
enabled=args.use_sample_reweight, enabled=args.use_sample_reweight,

@ -25,7 +25,7 @@ import joblib
import numpy as np import numpy as np
import torch import torch
from torch import nn 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 from src.models.forward_surrogate import ForwardSurrogate, ForwardSurrogateConfig
@ -39,6 +39,8 @@ METRIC_KEYS = (
"loss_derivative_shape", "loss_derivative_shape",
"loss_autofit_pressure", "loss_autofit_pressure",
"loss_autofit_derivative", "loss_autofit_derivative",
"loss_delta_pressure",
"loss_delta_derivative",
"sample_weight_mean", "sample_weight_mean",
"sample_weight_max", "sample_weight_max",
) )
@ -47,11 +49,57 @@ METRIC_KEYS = (
class ForwardDataset(Dataset): class ForwardDataset(Dataset):
"""把预处理后的参数、流量制度和曲线数组封装成 PyTorch 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 可直接按样本读取。""" """把三个 numpy 数组转为 float32 张量,后续 DataLoader 可直接按样本读取。"""
self.params_x = torch.tensor(params_x, dtype=torch.float32) self.params_x = torch.tensor(params_x, dtype=torch.float32)
self.schedule_x = torch.tensor(schedule_x, dtype=torch.float32) self.schedule_x = torch.tensor(schedule_x, dtype=torch.float32)
self.curve_y = torch.tensor(curve_y, 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: def __len__(self) -> int:
"""返回数据集或容器中可迭代样本的数量。""" """返回数据集或容器中可迭代样本的数量。"""
@ -59,7 +107,129 @@ class ForwardDataset(Dataset):
def __getitem__(self, idx: int): 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 @dataclass
@ -92,6 +262,8 @@ class LossWeights:
derivative_shape: float = 0.10 derivative_shape: float = 0.10
autofit_pressure: float = 0.0 autofit_pressure: float = 0.0
autofit_derivative: float = 0.0 autofit_derivative: float = 0.0
delta_pressure: float = 0.0
delta_derivative: float = 0.0
@dataclass @dataclass
@ -101,6 +273,7 @@ class LossConfig:
weights: LossWeights = field(default_factory=LossWeights) weights: LossWeights = field(default_factory=LossWeights)
use_huber: bool = True use_huber: bool = True
huber_beta: float = 0.05 huber_beta: float = 0.05
delta_huber_beta: float = 0.01
@dataclass @dataclass
@ -150,6 +323,8 @@ class LossBatchParts:
pred_d: torch.Tensor pred_d: torch.Tensor
true_p: torch.Tensor true_p: torch.Tensor
true_d: torch.Tensor true_d: torch.Tensor
mask_p: torch.Tensor
mask_d: torch.Tensor
@dataclass @dataclass
@ -194,6 +369,23 @@ def load_processed_dataset(path: Path) -> dict:
for key in required_keys: for key in required_keys:
if key not in data: if key not in data:
raise KeyError(f"processed dataset 缺少字段: {key}") 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 return data
@ -234,32 +426,61 @@ def get_part_slices(curve_layout: dict) -> dict[str, slice]:
return out 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 损失,返回每个样本一个损失值。""" """按样本计算 Smooth L1 损失,返回每个样本一个损失值。"""
diff = torch.abs(pred - target) diff = torch.abs(pred - target)
loss = torch.where(diff < beta, 0.5 * diff * diff / beta, diff - 0.5 * beta) 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( def regression_per_sample(
pred: torch.Tensor, pred: torch.Tensor,
target: torch.Tensor, target: torch.Tensor,
loss_cfg: LossConfig, loss_cfg: LossConfig,
valid_mask: torch.Tensor | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
"""按配置在 Smooth L1 和 MSE 之间切换点值损失。""" """按配置在 Smooth L1 和 MSE 之间切换点值损失。"""
if loss_cfg.use_huber: if loss_cfg.use_huber:
return smooth_l1_per_sample(pred, target, beta=float(loss_cfg.huber_beta)) return smooth_l1_per_sample(
return mse_per_sample(pred, target) 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: 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 return x_scaled * scale + mean
def autofit_curve_objective_per_sample(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor: def autofit_curve_objective_per_sample(
"""用 torch 计算自动拟合风格的曲线误差,作为训练附加目标。""" pred: torch.Tensor,
weight_factor = torch.clamp(torch.abs(target) * 0.01, max=100.0) target: torch.Tensor,
weight = 1.0 / (1.0 + weight_factor) valid_mask: torch.Tensor | None = None,
scale = torch.maximum( ) -> torch.Tensor:
torch.maximum(torch.abs(target), torch.abs(pred)), """用 torch 计算与 C++ 当前真实误差公式等价的曲线误差。"""
torch.full_like(target, 1e-12), log_error = torch.abs(target - pred)
) relative_error = -torch.expm1(-log_error)
relative_error = torch.abs(target - pred) / scale point_error = 0.7 * log_error + 0.3 * relative_error
absolute_error = torch.abs(target - pred) squared_mean = masked_mean_per_sample(point_error**2, valid_mask)
point_error = 0.7 * relative_error + 0.3 * absolute_error return torch.sqrt(torch.clamp(squared_mean, min=1.0e-12))
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 build_sample_weight( def build_sample_weight(
true_p: torch.Tensor, true_p: torch.Tensor,
true_d: torch.Tensor, true_d: torch.Tensor,
reweight_cfg: SampleReweightConfig, reweight_cfg: SampleReweightConfig,
mask_p: torch.Tensor | None = None,
mask_d: torch.Tensor | None = None,
) -> torch.Tensor: ) -> torch.Tensor:
"""根据真实曲线幅值构造样本权重,让高幅值样本训练时权重更高。""" """根据真实曲线幅值构造样本权重,让高幅值样本训练时权重更高。"""
p_level = true_p.abs().mean(dim=1) p_level = masked_mean_per_sample(true_p.abs(), mask_p)
d_level = true_d.abs().mean(dim=1) d_level = masked_mean_per_sample(true_d.abs(), mask_d)
p_norm = p_level / (p_level.mean().detach() + 1e-6) p_norm = p_level / (p_level.mean().detach() + 1e-6)
d_norm = d_level / (d_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( def split_curve_parts(
pred: torch.Tensor, pred: torch.Tensor,
target: torch.Tensor, target: torch.Tensor,
valid_mask: torch.Tensor,
slices: dict[str, slice], slices: dict[str, slice],
) -> LossBatchParts: ) -> LossBatchParts:
"""把拼接曲线拆成压力和导数两段。""" """把拼接曲线拆成压力和导数两段。"""
@ -316,6 +537,8 @@ def split_curve_parts(
pred_d=pred[:, slices["log_derivative"]], pred_d=pred[:, slices["log_derivative"]],
true_p=target[:, slices["log_pressure"]], true_p=target[:, slices["log_pressure"]],
true_d=target[:, slices["log_derivative"]], 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]: ) -> dict[str, torch.Tensor]:
"""计算标准化空间中的基础点值、偏置和导数形状损失。""" """计算标准化空间中的基础点值、偏置和导数形状损失。"""
return { return {
"loss_pressure": regression_per_sample(parts.pred_p, parts.true_p, loss_cfg), "loss_pressure": regression_per_sample(
"loss_derivative": regression_per_sample(parts.pred_d, parts.true_d, loss_cfg), 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( "loss_bias_pressure": l1_per_sample(
parts.pred_p.mean(dim=1, keepdim=True), masked_mean_per_sample(parts.pred_p, parts.mask_p).unsqueeze(1),
parts.true_p.mean(dim=1, keepdim=True), masked_mean_per_sample(parts.true_p, parts.mask_p).unsqueeze(1),
), ),
"loss_bias_derivative": l1_per_sample( "loss_bias_derivative": l1_per_sample(
parts.pred_d.mean(dim=1, keepdim=True), masked_mean_per_sample(parts.pred_d, parts.mask_d).unsqueeze(1),
parts.true_d.mean(dim=1, keepdim=True), masked_mean_per_sample(parts.true_d, parts.mask_d).unsqueeze(1),
), ),
"loss_derivative_shape": regression_per_sample( "loss_derivative_shape": regression_per_sample(
first_diff(parts.pred_d), first_diff(parts.pred_d),
first_diff(parts.true_d), first_diff(parts.true_d),
loss_cfg, 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( "loss_autofit_pressure": autofit_curve_objective_per_sample(
affine_restore(parts.pred_p, mean_p, scale_p), affine_restore(parts.pred_p, mean_p, scale_p),
affine_restore(parts.true_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( "loss_autofit_derivative": autofit_curve_objective_per_sample(
affine_restore(parts.pred_d, mean_d, scale_d), affine_restore(parts.pred_d, mean_d, scale_d),
affine_restore(parts.true_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( def weighted_total_vector(
loss_vectors: dict[str, torch.Tensor], loss_vectors: dict[str, torch.Tensor],
weights: LossWeights, weights: LossWeights,
@ -387,21 +725,59 @@ def weighted_total_vector(
def compute_weighted_loss( def compute_weighted_loss(
pred: torch.Tensor, pred: torch.Tensor,
target: torch.Tensor, target: torch.Tensor,
valid_mask: torch.Tensor,
context: LossContext, context: LossContext,
group_id: torch.Tensor | None = None,
is_anchor: torch.Tensor | None = None,
source_id: torch.Tensor | None = None,
) -> dict[str, torch.Tensor]: ) -> 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 = compute_basic_loss_vectors(parts, context.loss_cfg)
loss_vectors.update(compute_autofit_loss_vectors(parts, context)) loss_vectors.update(compute_autofit_loss_vectors(parts, context))
total_vec = weighted_total_vector(loss_vectors, context.loss_cfg.weights) total_vec = weighted_total_vector(loss_vectors, context.loss_cfg.weights)
if context.reweight_cfg.enabled: 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: else:
sample_weight = torch.ones_like(total_vec) sample_weight = torch.ones_like(total_vec)
metrics = {key: value.mean() for key, value in loss_vectors.items()} 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_mean"] = sample_weight.mean()
metrics["sample_weight_max"] = sample_weight.max() metrics["sample_weight_max"] = sample_weight.max()
return metrics return metrics
@ -454,19 +830,43 @@ def run_loader_epoch(
total = init_metric_accumulator() total = init_metric_accumulator()
total_n = 0 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() grad_context = torch.enable_grad() if is_train else torch.no_grad()
with grad_context: 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) params_x = params_x.to(device)
schedule_x = schedule_x.to(device) schedule_x = schedule_x.to(device)
curve_y = curve_y.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: if is_train:
optimizer.zero_grad() optimizer.zero_grad()
pred = model_forward(model, params_x, schedule_x, use_schedule) 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: if is_train:
losses["loss"].backward() losses["loss"].backward()
@ -474,9 +874,27 @@ def run_loader_epoch(
batch_size = params_x.size(0) batch_size = params_x.size(0)
accumulate_metrics(total, losses, batch_size) 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 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( def evaluate(
@ -510,13 +928,89 @@ def build_curve_stats(data: dict, device: str) -> CurveStats:
def build_dataloaders(data: dict, cfg: TrainConfig) -> DatasetBundle: def build_dataloaders(data: dict, cfg: TrainConfig) -> DatasetBundle:
"""根据预处理数组构造训练、验证、测试 DataLoader。""" """根据预处理数组构造训练、验证、测试 DataLoader。"""
train_ds = ForwardDataset(data["X_params_train"], data["X_schedule_train"], data["Y_curve_train"]) delta_enabled = bool(
val_ds = ForwardDataset(data["X_params_val"], data["X_schedule_val"], data["Y_curve_val"]) cfg.loss.weights.delta_pressure != 0.0
test_ds = ForwardDataset(data["X_params_test"], data["X_schedule_test"], data["Y_curve_test"]) 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 = torch.Generator()
loader_generator.manual_seed(int(cfg.runtime.seed)) loader_generator.manual_seed(int(cfg.runtime.seed))
train_loader = DataLoader( train_loader = DataLoader(
train_ds, train_ds,
batch_size=cfg.optim.batch_size, 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" weights: pressure={weights.pressure}, derivative={weights.derivative}, "
f"bias_p={weights.bias_pressure}, " f"bias_p={weights.bias_pressure}, "
f"bias_d={weights.bias_derivative}, d_shape={weights.derivative_shape}, " 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( print(
f" sample_reweight={reweight.enabled}, alpha={reweight.alpha}, " 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"ds={train_metrics['loss_derivative_shape']:.6f}, "
f"ap={train_metrics['loss_autofit_pressure']:.6f}, " f"ap={train_metrics['loss_autofit_pressure']:.6f}, "
f"ad={train_metrics['loss_autofit_derivative']:.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"wmean={train_metrics['sample_weight_mean']:.4f}, "
f"wmax={train_metrics['sample_weight_max']:.4f}) " f"wmax={train_metrics['sample_weight_max']:.4f}) "
f"val={val_metrics['loss']:.6f} " 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"ds={val_metrics['loss_derivative_shape']:.6f}, "
f"ap={val_metrics['loss_autofit_pressure']:.6f}, " f"ap={val_metrics['loss_autofit_pressure']:.6f}, "
f"ad={val_metrics['loss_autofit_derivative']:.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"wmean={val_metrics['sample_weight_mean']:.4f}, "
f"wmax={val_metrics['sample_weight_max']:.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"ds={test_metrics['loss_derivative_shape']:.6f}, "
f"ap={test_metrics['loss_autofit_pressure']:.6f}, " f"ap={test_metrics['loss_autofit_pressure']:.6f}, "
f"ad={test_metrics['loss_autofit_derivative']:.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"wmean={test_metrics['sample_weight_mean']:.4f}, "
f"wmax={test_metrics['sample_weight_max']:.4f})" f"wmax={test_metrics['sample_weight_max']:.4f})"
) )
@ -652,6 +1154,7 @@ def build_checkpoint_payload(
"seed": int(cfg.runtime.seed), "seed": int(cfg.runtime.seed),
"curve_layout": curve_layout, "curve_layout": curve_layout,
"loss_weights": asdict(cfg.loss.weights), "loss_weights": asdict(cfg.loss.weights),
"delta_huber_beta": float(cfg.loss.delta_huber_beta),
"sample_reweight": asdict(cfg.sample_reweight), "sample_reweight": asdict(cfg.sample_reweight),
} }
@ -748,6 +1251,7 @@ def build_metrics_payload(
"use_schedule": cfg.model.use_schedule, "use_schedule": cfg.model.use_schedule,
"seed": int(cfg.runtime.seed), "seed": int(cfg.runtime.seed),
"loss_weights": asdict(cfg.loss.weights), "loss_weights": asdict(cfg.loss.weights),
"delta_huber_beta": float(cfg.loss.delta_huber_beta),
"sample_reweight": asdict(cfg.sample_reweight), "sample_reweight": asdict(cfg.sample_reweight),
"curve_layout": curve_layout, "curve_layout": curve_layout,
} }

Loading…
Cancel
Save