You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
nmWTAI-Platform/ML/nmWTAI-ML/src/data/preprocess.py

790 lines
33 KiB
Python

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

# -*- coding: utf-8 -*-
"""原始 HDF5 样本预处理。
数据生成阶段写出的 HDF5 文件仍处在“原始样本”状态:参数尚未做特征变换,输入和
监督目标尚未标准化,也没有训练/验证/测试划分。本模块完成这些准备工作,并把
结果保存为 joblib 文件供训练脚本直接读取。
预处理只在训练集上拟合 scaler随后复用到验证集和测试集避免数据泄漏。输出
payload 中保留 curve_layout、参数变换配置、流量制度元数据和附加字段使后续
训练、评估和误差分析可以追溯每个样本的来源。
"""
from __future__ import annotations
import copy
from pathlib import Path
import h5py
import joblib
import numpy as np
from sklearn.model_selection import GroupShuffleSplit, train_test_split
from sklearn.preprocessing import StandardScaler
from src.data.param_features import build_param_feature_transform, transform_param_features
class SelectiveStandardScaler:
"""只标准化指定连续列,同时保持 one-hot 等类别列原值不变。"""
def __init__(self, scaled_indices: np.ndarray | list[int]):
self.scaled_indices = np.asarray(scaled_indices, dtype=np.int64)
def fit(self, x: np.ndarray) -> "SelectiveStandardScaler":
"""在训练集的连续特征列上拟合均值和标准差。"""
values = _ensure_2d("params", x).astype(np.float64)
indices = np.unique(self.scaled_indices)
if np.any(indices < 0) or np.any(indices >= values.shape[1]):
raise ValueError(f"scaled parameter indices out of range: {indices.tolist()}")
self.scaled_indices_ = indices
self.n_features_in_ = int(values.shape[1])
self.n_samples_seen_ = int(values.shape[0])
self.mean_ = np.zeros(self.n_features_in_, dtype=np.float64)
self.var_ = np.ones(self.n_features_in_, dtype=np.float64)
self.scale_ = np.ones(self.n_features_in_, dtype=np.float64)
if indices.size:
inner = StandardScaler().fit(values[:, indices])
self.mean_[indices] = inner.mean_
self.var_[indices] = inner.var_
self.scale_[indices] = inner.scale_
return self
def _validate_transform_input(self, x: np.ndarray) -> np.ndarray:
"""检查待变换数组维度是否与拟合时一致。"""
if not hasattr(self, "n_features_in_"):
raise RuntimeError("SelectiveStandardScaler is not fitted")
values = _ensure_2d("params", x).astype(np.float64)
if values.shape[1] != self.n_features_in_:
raise ValueError(
f"parameter feature dimension mismatch: {values.shape[1]} != {self.n_features_in_}"
)
return values
def transform(self, x: np.ndarray) -> np.ndarray:
"""标准化连续列未选择的类别列保持0/1。"""
values = self._validate_transform_input(x)
return (values - self.mean_) / self.scale_
def inverse_transform(self, x: np.ndarray) -> np.ndarray:
"""恢复标准化前的连续特征,同时保持类别列不变。"""
values = self._validate_transform_input(x)
return values * self.scale_ + self.mean_
def fit_transform(self, x: np.ndarray) -> np.ndarray:
"""拟合后立即变换训练集。"""
return self.fit(x).transform(x)
class MaskedStandardScaler:
"""每个曲线列只使用实际物理时间覆盖范围内的有效值进行标准化。"""
def fit(self, x: np.ndarray, valid_mask: np.ndarray) -> "MaskedStandardScaler":
values = _ensure_2d("curve", x).astype(np.float64)
mask = np.asarray(valid_mask, dtype=bool)
if mask.shape != values.shape:
raise ValueError(f"curve mask shape mismatch: {mask.shape} != {values.shape}")
if not np.isfinite(values[mask]).all():
raise ValueError("valid curve values contain non-finite values")
counts = mask.sum(axis=0).astype(np.int64)
sums = np.where(mask, values, 0.0).sum(axis=0)
means = np.divide(sums, counts, out=np.zeros_like(sums), where=counts > 0)
centered = np.where(mask, values - means, 0.0)
variances = np.divide(
(centered * centered).sum(axis=0),
counts,
out=np.ones_like(sums),
where=counts > 0,
)
scales = np.sqrt(np.maximum(variances, 0.0))
scales[(counts == 0) | (scales < 1.0e-12)] = 1.0
self.n_features_in_ = int(values.shape[1])
self.n_samples_seen_ = counts
self.mean_ = means
self.var_ = variances
self.scale_ = scales
return self
def _validate(self, x: np.ndarray) -> np.ndarray:
if not hasattr(self, "n_features_in_"):
raise RuntimeError("MaskedStandardScaler is not fitted")
values = _ensure_2d("curve", x).astype(np.float64)
if values.shape[1] != self.n_features_in_:
raise ValueError(
f"curve feature dimension mismatch: {values.shape[1]} != {self.n_features_in_}"
)
return values
def transform(self, x: np.ndarray) -> np.ndarray:
values = self._validate(x)
return (values - self.mean_) / self.scale_
def inverse_transform(self, x: np.ndarray) -> np.ndarray:
values = self._validate(x)
return values * self.scale_ + self.mean_
def _value_counts(values: np.ndarray | None, indices: np.ndarray) -> dict[str, int]:
"""按字符串键记录某个划分中的类别样本数。"""
if values is None:
return {}
unique, counts = np.unique(np.asarray(values)[indices], return_counts=True)
result = {}
for value, count in zip(unique.tolist(), counts.tolist()):
key = str(int(value)) if float(value).is_integer() else str(value)
result[key] = int(count)
return result
def _split_sample_indices(
n_samples: int,
test_size: float,
val_size: float,
random_seed: int,
group_id: np.ndarray | None,
solver_type: np.ndarray | None,
) -> tuple[np.ndarray, np.ndarray, np.ndarray, str]:
"""优先按group整体划分没有可靠group时按solverType分层随机划分。"""
if test_size <= 0.0 or val_size <= 0.0 or test_size + val_size >= 1.0:
raise ValueError("test_size and val_size must be positive and sum to less than 1")
idx = np.arange(n_samples)
val_ratio_in_train_val = val_size / (1.0 - test_size)
groups = None if group_id is None else np.asarray(group_id).reshape(-1)
use_groups = bool(
groups is not None
and len(groups) == n_samples
and np.all(np.isfinite(groups))
and np.all(groups >= 0)
and np.unique(groups).size >= 3
)
if use_groups:
first = GroupShuffleSplit(n_splits=1, test_size=test_size, random_state=random_seed)
train_val_pos, test_pos = next(first.split(idx, groups=groups))
idx_train_val = idx[train_val_pos]
idx_test = idx[test_pos]
second = GroupShuffleSplit(
n_splits=1,
test_size=val_ratio_in_train_val,
random_state=random_seed + 1,
)
train_pos, val_pos = next(
second.split(idx_train_val, groups=groups[idx_train_val])
)
idx_train = idx_train_val[train_pos]
idx_val = idx_train_val[val_pos]
return idx_train, idx_val, idx_test, "group_id"
stratify = None if solver_type is None else np.asarray(solver_type).reshape(-1)
idx_train_val, idx_test = train_test_split(
idx,
test_size=test_size,
random_state=random_seed,
shuffle=True,
stratify=stratify,
)
train_val_stratify = None if stratify is None else stratify[idx_train_val]
idx_train, idx_val = train_test_split(
idx_train_val,
test_size=val_ratio_in_train_val,
random_state=random_seed,
shuffle=True,
stratify=train_val_stratify,
)
return idx_train, idx_val, idx_test, "solverType" if stratify is not None else "random"
def _ensure_2d(name: str, arr: np.ndarray) -> np.ndarray:
"""保证数组至少为二维,避免单样本情况下维度被压扁。"""
arr = np.asarray(arr, dtype=np.float32)
if arr.ndim != 2:
raise ValueError(f"{name} must be a 2D array, got shape={arr.shape}")
return arr
def _decode_attr_list(raw_names) -> list[str] | None:
"""解码 HDF5 属性中保存的 JSON 字符串或字节串列表。"""
if raw_names is None:
return None
return [
item.decode("utf-8") if isinstance(item, (bytes, np.bytes_)) else str(item)
for item in raw_names
]
def _read_optional_string_dataset(f: h5py.File, name: str) -> np.ndarray | None:
"""读取可选字符串数据集;不存在时返回 None。"""
if name not in f:
return None
return np.asarray(f[name][:]).astype(str)
def _read_optional_numeric_dataset(f: h5py.File, name: str) -> np.ndarray | None:
"""读取可选数值数据集;不存在时返回 None。"""
if name not in f:
return None
return np.asarray(f[name][:])
def _validate_optional_length(name: str, arr: np.ndarray | None, expected_n: int) -> None:
"""检查可选数据集长度是否与主样本数量一致。"""
if arr is not None and len(arr) != expected_n:
raise ValueError(f"{name} row count does not match main arrays: {len(arr)} != {expected_n}")
def _validate_curve_time(
curve_time: np.ndarray,
expected_n: int,
expected_n_time_points: int,
) -> np.ndarray:
"""校验每条曲线对应的物理时间坐标。"""
values = np.asarray(curve_time, dtype=np.float64)
expected_shape = (expected_n, expected_n_time_points)
if values.ndim != 2 or values.shape != expected_shape:
raise ValueError(
f"curve_time must have shape {expected_shape}, got shape={values.shape}"
)
if not np.all(np.isfinite(values)):
raise ValueError("curve_time contains non-finite values")
if np.any(values <= 0.0):
raise ValueError("curve_time values must be positive")
if np.any(np.diff(values, axis=1) <= 0.0):
raise ValueError("each curve_time row must be strictly increasing")
return values
def _validate_curve_valid_mask(
curve_valid_mask: np.ndarray,
expected_n: int,
expected_n_time_points: int,
) -> np.ndarray:
"""校验逐时间点记录物理覆盖范围的有效掩码。"""
values = np.asarray(curve_valid_mask)
expected_shape = (expected_n, expected_n_time_points)
if values.ndim != 2 or values.shape != expected_shape:
raise ValueError(
f"curve_valid_mask must have shape {expected_shape}, got shape={values.shape}"
)
if not np.all(np.isfinite(values)) or not np.all((values == 0) | (values == 1)):
raise ValueError("curve_valid_mask must contain only 0/1 values")
values = values.astype(bool)
if np.any(values.sum(axis=1) < 2):
raise ValueError("each curve_valid_mask row must contain at least two valid points")
return values
def _curve_keep_indices(meta: dict, curve_dim: int) -> np.ndarray:
"""返回旧 processed 曲线中 pressure/derivative 两段的列索引。"""
layout = meta.get("curve_layout") or {}
parts = {str(part["name"]): part for part in layout.get("parts", [])}
if "log_pressure" in parts and "log_derivative" in parts:
indices = []
for name in ("log_pressure", "log_derivative"):
part = parts[name]
indices.extend(range(int(part["start"]), int(part["end"])))
return np.asarray(indices, dtype=np.int64)
if curve_dim % 3 != 0:
raise ValueError(
"processed 数据没有 curve_layout且曲线维度不能按旧版三通道布局识别"
)
n_time_points = curve_dim // 3
return np.arange(2 * n_time_points, dtype=np.int64)
def _slice_standard_scaler(scaler: StandardScaler, keep_indices: np.ndarray) -> StandardScaler:
"""裁剪按列独立拟合的 StandardScaler同时保留其拟合状态。"""
result = copy.deepcopy(scaler)
for name in ("mean_", "scale_", "var_"):
value = getattr(result, name, None)
if value is not None:
setattr(result, name, np.asarray(value)[keep_indices].copy())
result.n_features_in_ = int(keep_indices.size)
return result
def migrate_processed_dataset_without_slope(input_path: Path, output_path: Path) -> None:
"""把旧三通道 processed pkl 精确迁移为 pressure/derivative 双通道数据。"""
if not input_path.exists():
raise FileNotFoundError(f"Input file not found: {input_path}")
source = joblib.load(input_path)
required = ["Y_curve_train", "Y_curve_val", "Y_curve_test", "scaler_curve", "meta"]
missing = [name for name in required if name not in source]
if missing:
raise KeyError(f"processed 数据缺少字段: {missing}")
curve_dim = int(source["Y_curve_train"].shape[1])
keep_indices = _curve_keep_indices(source.get("meta", {}), curve_dim)
if keep_indices.size == curve_dim:
raise ValueError("输入 processed 数据已经不包含 slope无需迁移")
if keep_indices.size % 2 != 0:
raise ValueError(f"保留后的曲线维度必须能被 2 整除,当前为 {keep_indices.size}")
migration_n_time_points = int(keep_indices.size // 2)
migration_meta = dict(source.get("meta", {}))
if migration_meta.get("curve_time_mode") != "fixed":
raise ValueError(
"processed migration requires meta.curve_time_mode='fixed'; "
"regenerate the H5 dataset on a fixed time grid and preprocess it again"
)
prediction_curve_time = np.asarray(
migration_meta.get("prediction_curve_time"),
dtype=np.float64,
)
if prediction_curve_time.shape != (migration_n_time_points,):
raise ValueError(
"processed migration requires meta.prediction_curve_time with shape "
f"({migration_n_time_points},); regenerate the H5 dataset and preprocess it again"
)
if (
not np.all(np.isfinite(prediction_curve_time))
or np.any(prediction_curve_time <= 0.0)
or np.any(np.diff(prediction_curve_time) <= 0.0)
):
raise ValueError(
"processed migration requires a finite, positive, strictly increasing "
"meta.prediction_curve_time"
)
for split in ("train", "val", "test"):
mask_key = f"curve_valid_mask_{split}"
if mask_key not in source:
raise ValueError(
f"processed migration requires {mask_key}; regenerate the H5 dataset "
"on a fixed time grid and preprocess it again"
)
_validate_curve_valid_mask(
source[mask_key],
len(source[f"Y_curve_{split}"]),
migration_n_time_points,
)
payload = dict(source)
for split in ("train", "val", "test"):
key = f"Y_curve_{split}"
payload[key] = np.asarray(source[key][:, keep_indices], dtype=np.float32)
payload["scaler_curve"] = _slice_standard_scaler(source["scaler_curve"], keep_indices)
# 这些字段属于已停用的 time-conditioned 路径,不进入旧正演模型训练包。
for key in (
"T_curve_train", "T_curve_val", "T_curve_test",
"X_time_train", "X_time_val", "X_time_test", "scaler_time",
):
payload.pop(key, None)
n_time_points = int(keep_indices.size // 2)
meta = dict(source.get("meta", {}))
meta.update(
{
"curve_dim": int(keep_indices.size),
"raw_curve_dim": curve_dim,
"dropped_legacy_slope": True,
"migrated_from_processed": str(input_path),
"curve_layout": {
"n_time_points": n_time_points,
"parts": [
{"name": "log_pressure", "start": 0, "end": n_time_points},
{"name": "log_derivative", "start": n_time_points, "end": 2 * n_time_points},
],
},
}
)
meta.pop("time_feature_dim", None)
meta.pop("curve_time_source", None)
payload["meta"] = meta
output_path.parent.mkdir(parents=True, exist_ok=True)
joblib.dump(payload, output_path)
print(f"Processed migration complete, saved to: {output_path}")
print(
f"train={len(payload['Y_curve_train'])}, val={len(payload['Y_curve_val'])}, "
f"test={len(payload['Y_curve_test'])}, curve_dim={curve_dim}->{keep_indices.size}"
)
def preprocess_dataset(
input_path: Path,
output_path: Path,
test_size: float = 0.15,
val_size: float = 0.15,
random_seed: int = 42,
use_param_feature_transform: bool = True,
) -> None:
"""将原始 HDF5 样本整理为训练脚本可直接加载的 joblib 数据包。
处理步骤包括:读取 params/schedule/curve 等核心数组,保留可选元数据,构造参数
特征变换,划分 train/val/test并只在训练集上拟合 StandardScaler。输出 payload
同时保存 scaler 和 meta确保后续评估可以恢复原始尺度并正确拆分压力和导数。
"""
if not input_path.exists():
raise FileNotFoundError(f"Input file not found: {input_path}")
with h5py.File(input_path, "r") as f:
required_keys = ["params", "schedule", "curve"]
for key in required_keys:
if key not in f:
raise KeyError(f"HDF5 dataset is missing required key: {key}")
x_params = _ensure_2d("params", f["params"][:])
x_schedule = _ensure_2d("schedule", f["schedule"][:])
y_curve = _ensure_2d("curve", f["curve"][:])
curve_time = _read_optional_numeric_dataset(f, "curve_time")
curve_valid_mask = _read_optional_numeric_dataset(f, "curve_valid_mask")
# 可选元数据会保留下来,供后续分析和排序验证使用;
# 基础正演训练循环不依赖这些字段。
param_names = _decode_attr_list(f.attrs.get("param_names"))
schedule_meta_names = _decode_attr_list(f.attrs.get("schedule_meta_names"))
source_name_vocab = _decode_attr_list(f.attrs.get("source_name_vocab"))
if param_names is None:
param_names = ["k", "skin", "wellboreC", "phi", "h", "Ct", "Cf"]
solver_type = None
categorical_values: dict[str, list[float]] = {}
if "solverType" in param_names:
solver_col = param_names.index("solverType")
solver_type = np.asarray(x_params[:, solver_col], dtype=np.float64)
categories = np.unique(solver_type)
if not np.all(np.isfinite(categories)):
raise ValueError("solverType contains non-finite values")
categorical_values["solverType"] = categories.tolist()
schedule_meta = _ensure_2d("schedule_meta", f["schedule_meta"][:]) if "schedule_meta" in f else None
family_name = _read_optional_string_dataset(f, "family_name")
source_id = _read_optional_numeric_dataset(f, "source_id")
source_name = _read_optional_string_dataset(f, "source_name")
source_row = _read_optional_numeric_dataset(f, "source_row")
extra_numeric_names = [
"group_id",
"is_anchor",
"neighbor_objective",
"neighbor_objective_p",
"neighbor_objective_d",
"neighbor_span_frac",
"section_index",
]
extra_string_names = ["timeQ_json", "q_json"]
extra_numeric = {name: _read_optional_numeric_dataset(f, name) for name in extra_numeric_names}
extra_string = {name: _read_optional_string_dataset(f, name) for name in extra_string_names}
n = len(x_params)
if len(x_schedule) != n or len(y_curve) != n:
raise ValueError("params / schedule / curve row counts do not match")
raw_curve_dim = int(y_curve.shape[1])
if curve_time is not None:
curve_time_shape = np.asarray(curve_time).shape
if len(curve_time_shape) != 2 or curve_time_shape[0] != n:
raise ValueError(
f"curve_time must be a 2D array with {n} rows, got shape={curve_time_shape}"
)
n_time_points = int(curve_time_shape[1])
if raw_curve_dim == 3 * n_time_points:
# 兼容已有的压力/导数/斜率三通道 HDF5训练时丢弃斜率通道。
y_curve = y_curve[:, : 2 * n_time_points]
elif raw_curve_dim != 2 * n_time_points:
raise ValueError(
"curve dimension does not match curve_time: "
f"curve_dim={raw_curve_dim}, n_time_points={n_time_points}; "
"expected 2*N or legacy 3*N"
)
elif raw_curve_dim % 3 == 0:
n_time_points = raw_curve_dim // 3
y_curve = y_curve[:, : 2 * n_time_points]
elif raw_curve_dim % 2 == 0:
n_time_points = raw_curve_dim // 2
else:
raise ValueError(
f"curve dimension {raw_curve_dim} cannot be interpreted as 2*N or legacy 3*N"
)
curve_dim = int(y_curve.shape[1])
curve_time_mode = "missing"
prediction_curve_time = None
if curve_time is not None:
curve_time = _validate_curve_time(curve_time, n, n_time_points)
if np.allclose(curve_time, curve_time[0:1], rtol=1e-6, atol=1e-10):
curve_time_mode = "fixed"
prediction_curve_time = curve_time[0].copy()
else:
curve_time_mode = "per_sample"
if curve_time_mode != "fixed":
raise ValueError(
"forward-surrogate preprocessing requires one fixed curve_time grid; "
f"got curve_time_mode={curve_time_mode!r}"
)
if curve_valid_mask is None:
raise ValueError(
"fixed-time preprocessing requires curve_valid_mask; regenerate the H5 dataset"
)
curve_valid_mask = _validate_curve_valid_mask(
curve_valid_mask,
n,
n_time_points,
)
curve_valid_mask_mode = "explicit"
curve_mask = np.concatenate([curve_valid_mask, curve_valid_mask], axis=1)
param_feature_transform = build_param_feature_transform(
param_names=param_names,
categorical_values=categorical_values,
enabled=bool(use_param_feature_transform),
)
# 物理参数在 StandardScaler 前先做特征变换:
# 例如 k、wellboreC、h 使用 log10skin 使用 asinh。
x_params_features = transform_param_features(x_params, param_feature_transform)
_validate_optional_length("schedule_meta", schedule_meta, n)
_validate_optional_length("family_name", family_name, n)
_validate_optional_length("source_id", source_id, n)
_validate_optional_length("source_name", source_name, n)
_validate_optional_length("source_row", source_row, n)
for name, arr in extra_numeric.items():
_validate_optional_length(name, arr, n)
for name, arr in extra_string.items():
_validate_optional_length(name, arr, n)
# 同一group通常对应相同流量制度在不同solverType下的样本必须整体进入同一划分。
group_id = extra_numeric.get("group_id")
idx_train, idx_val, idx_test, split_strategy = _split_sample_indices(
n_samples=n,
test_size=test_size,
val_size=val_size,
random_seed=random_seed,
group_id=group_id,
solver_type=solver_type,
)
x_params_train = x_params_features[idx_train]
x_schedule_train = x_schedule[idx_train]
y_curve_train = y_curve[idx_train]
curve_valid_mask_train = curve_valid_mask[idx_train]
curve_mask_train = curve_mask[idx_train]
x_params_val = x_params_features[idx_val]
x_schedule_val = x_schedule[idx_val]
y_curve_val = y_curve[idx_val]
curve_valid_mask_val = curve_valid_mask[idx_val]
curve_mask_val = curve_mask[idx_val]
x_params_test = x_params_features[idx_test]
x_schedule_test = x_schedule[idx_test]
y_curve_test = y_curve[idx_test]
curve_valid_mask_test = curve_valid_mask[idx_test]
curve_mask_test = curve_mask[idx_test]
schedule_meta_train = schedule_meta[idx_train] if schedule_meta is not None else None
schedule_meta_val = schedule_meta[idx_val] if schedule_meta is not None else None
schedule_meta_test = schedule_meta[idx_test] if schedule_meta is not None else None
family_name_train = family_name[idx_train] if family_name is not None else None
family_name_val = family_name[idx_val] if family_name is not None else None
family_name_test = family_name[idx_test] if family_name is not None else None
source_id_train = source_id[idx_train] if source_id is not None else None
source_id_val = source_id[idx_val] if source_id is not None else None
source_id_test = source_id[idx_test] if source_id is not None else None
source_name_train = source_name[idx_train] if source_name is not None else None
source_name_val = source_name[idx_val] if source_name is not None else None
source_name_test = source_name[idx_test] if source_name is not None else None
source_row_train = source_row[idx_train] if source_row is not None else None
source_row_val = source_row[idx_val] if source_row is not None else None
source_row_test = source_row[idx_test] if source_row is not None else None
extra_numeric_split = {
name: (
arr[idx_train] if arr is not None else None,
arr[idx_val] if arr is not None else None,
arr[idx_test] if arr is not None else None,
)
for name, arr in extra_numeric.items()
}
extra_string_split = {
name: (
arr[idx_train] if arr is not None else None,
arr[idx_val] if arr is not None else None,
arr[idx_test] if arr is not None else None,
)
for name, arr in extra_string.items()
}
categorical_indices = np.asarray(
param_feature_transform.get("categorical_feature_indices", []),
dtype=np.int64,
)
scaled_param_indices = np.setdiff1d(
np.arange(x_params_features.shape[1], dtype=np.int64),
categorical_indices,
)
scaler_params = SelectiveStandardScaler(scaled_param_indices)
scaler_schedule = StandardScaler()
scaler_curve = MaskedStandardScaler()
valid_count_train = curve_valid_mask_train.sum(axis=0).astype(np.int64)
if prediction_curve_time is not None and np.any(valid_count_train == 0):
missing_columns = np.flatnonzero(valid_count_train == 0).tolist()
raise ValueError(
"fixed curve_time has training columns without physical coverage: "
f"{missing_columns}"
)
# 只在训练集上拟合 scaler避免验证/测试信息泄漏。
x_params_train_s = scaler_params.fit_transform(x_params_train)
x_schedule_train_s = scaler_schedule.fit_transform(x_schedule_train)
scaler_curve.fit(y_curve_train, curve_mask_train)
y_curve_train_s = scaler_curve.transform(y_curve_train)
y_curve_train_s[~curve_mask_train] = 0.0
x_params_val_s = scaler_params.transform(x_params_val)
x_schedule_val_s = scaler_schedule.transform(x_schedule_val)
y_curve_val_s = scaler_curve.transform(y_curve_val)
y_curve_val_s[~curve_mask_val] = 0.0
x_params_test_s = scaler_params.transform(x_params_test)
x_schedule_test_s = scaler_schedule.transform(x_schedule_test)
y_curve_test_s = scaler_curve.transform(y_curve_test)
y_curve_test_s[~curve_mask_test] = 0.0
output_path.parent.mkdir(parents=True, exist_ok=True)
# 保存 curve_layout训练和评估脚本就能按元数据切分拼接曲线
# 避免在各处硬编码索引范围。
meta = {
"n_total": int(n),
"n_train": int(len(idx_train)),
"n_val": int(len(idx_val)),
"n_test": int(len(idx_test)),
"param_dim": int(x_params_features.shape[1]),
"raw_param_dim": int(x_params.shape[1]),
"schedule_dim": int(x_schedule.shape[1]),
"curve_dim": curve_dim,
"raw_curve_dim": raw_curve_dim,
"dropped_legacy_slope": bool(raw_curve_dim != curve_dim),
"curve_time_mode": curve_time_mode,
"curve_valid_mask_mode": curve_valid_mask_mode,
"curve_valid_fraction": {
"train": float(np.mean(curve_valid_mask_train)),
"val": float(np.mean(curve_valid_mask_val)),
"test": float(np.mean(curve_valid_mask_test)),
},
"curve_valid_count_train": valid_count_train.astype(int).tolist(),
"input_h5": str(input_path),
"param_names": param_names,
"param_feature_transform": param_feature_transform,
"param_feature_names": list(param_feature_transform.get("feature_names", [])),
"param_scaled_indices": scaled_param_indices.tolist(),
"param_unscaled_indices": categorical_indices.tolist(),
"split_strategy": split_strategy,
"solver_type_counts": {
"train": _value_counts(solver_type, idx_train),
"val": _value_counts(solver_type, idx_val),
"test": _value_counts(solver_type, idx_test),
},
"group_counts": {
"train": int(np.unique(group_id[idx_train]).size) if group_id is not None else 0,
"val": int(np.unique(group_id[idx_val]).size) if group_id is not None else 0,
"test": int(np.unique(group_id[idx_test]).size) if group_id is not None else 0,
},
"schedule_meta_names": schedule_meta_names,
"source_name_vocab": source_name_vocab,
"extra_numeric_fields": [name for name, arr in extra_numeric.items() if arr is not None],
"extra_string_fields": [name for name, arr in extra_string.items() if arr is not None],
"curve_layout": {
"n_time_points": int(n_time_points),
"parts": [
{"name": "log_pressure", "start": 0, "end": int(n_time_points)},
{"name": "log_derivative", "start": int(n_time_points), "end": int(2 * n_time_points)},
],
},
}
if prediction_curve_time is not None:
meta["prediction_curve_time"] = prediction_curve_time.tolist()
payload = {
"X_params_train": x_params_train_s.astype(np.float32),
"X_schedule_train": x_schedule_train_s.astype(np.float32),
"Y_curve_train": y_curve_train_s.astype(np.float32),
"X_params_val": x_params_val_s.astype(np.float32),
"X_schedule_val": x_schedule_val_s.astype(np.float32),
"Y_curve_val": y_curve_val_s.astype(np.float32),
"X_params_test": x_params_test_s.astype(np.float32),
"X_schedule_test": x_schedule_test_s.astype(np.float32),
"Y_curve_test": y_curve_test_s.astype(np.float32),
"curve_valid_mask_train": curve_valid_mask_train.astype(np.uint8),
"curve_valid_mask_val": curve_valid_mask_val.astype(np.uint8),
"curve_valid_mask_test": curve_valid_mask_test.astype(np.uint8),
"scaler_params": scaler_params,
"scaler_schedule": scaler_schedule,
"scaler_curve": scaler_curve,
"meta": meta,
}
if schedule_meta is not None:
payload["schedule_meta_train"] = np.asarray(schedule_meta_train, dtype=np.float32)
payload["schedule_meta_val"] = np.asarray(schedule_meta_val, dtype=np.float32)
payload["schedule_meta_test"] = np.asarray(schedule_meta_test, dtype=np.float32)
if family_name is not None:
payload["family_name_train"] = np.asarray(family_name_train, dtype=object)
payload["family_name_val"] = np.asarray(family_name_val, dtype=object)
payload["family_name_test"] = np.asarray(family_name_test, dtype=object)
if source_id is not None:
payload["source_id_train"] = np.asarray(source_id_train)
payload["source_id_val"] = np.asarray(source_id_val)
payload["source_id_test"] = np.asarray(source_id_test)
if source_name is not None:
payload["source_name_train"] = np.asarray(source_name_train, dtype=object)
payload["source_name_val"] = np.asarray(source_name_val, dtype=object)
payload["source_name_test"] = np.asarray(source_name_test, dtype=object)
if source_row is not None:
payload["source_row_train"] = np.asarray(source_row_train)
payload["source_row_val"] = np.asarray(source_row_val)
payload["source_row_test"] = np.asarray(source_row_test)
for name, (train_arr, val_arr, test_arr) in extra_numeric_split.items():
if train_arr is None:
continue
payload[f"{name}_train"] = np.asarray(train_arr)
payload[f"{name}_val"] = np.asarray(val_arr)
payload[f"{name}_test"] = np.asarray(test_arr)
for name, (train_arr, val_arr, test_arr) in extra_string_split.items():
if train_arr is None:
continue
payload[f"{name}_train"] = np.asarray(train_arr, dtype=object)
payload[f"{name}_val"] = np.asarray(val_arr, dtype=object)
payload[f"{name}_test"] = np.asarray(test_arr, dtype=object)
joblib.dump(payload, output_path)
print(f"Preprocess complete, saved to: {output_path}")
print(
f"train={len(idx_train)}, val={len(idx_val)}, test={len(idx_test)}, "
f"param_dim={x_params_features.shape[1]}, schedule_dim={x_schedule.shape[1]}, "
f"curve_dim={y_curve.shape[1]}"
)
print(
f"split_strategy={split_strategy}, "
f"solver_type_counts={meta['solver_type_counts']}, "
f"group_counts={meta['group_counts']}, "
f"curve_time_mode={curve_time_mode}"
)
if schedule_meta is not None:
print(f"schedule_meta_dim={schedule_meta.shape[1]}, family_name_saved={family_name is not None}")
if source_name is not None:
print(f"source_name_saved=True, source_vocab={source_name_vocab}")
if meta["extra_numeric_fields"] or meta["extra_string_fields"]:
print(
"extra_fields_saved="
f"{meta['extra_numeric_fields'] + meta['extra_string_fields']}"
)