|
|
|
|
@ -25,7 +25,7 @@ import joblib
|
|
|
|
|
import numpy as np
|
|
|
|
|
import torch
|
|
|
|
|
from torch import nn
|
|
|
|
|
from torch.utils.data import DataLoader, Dataset
|
|
|
|
|
from torch.utils.data import DataLoader, Dataset, Sampler
|
|
|
|
|
|
|
|
|
|
from src.models.forward_surrogate import ForwardSurrogate, ForwardSurrogateConfig
|
|
|
|
|
|
|
|
|
|
@ -39,6 +39,8 @@ METRIC_KEYS = (
|
|
|
|
|
"loss_derivative_shape",
|
|
|
|
|
"loss_autofit_pressure",
|
|
|
|
|
"loss_autofit_derivative",
|
|
|
|
|
"loss_delta_pressure",
|
|
|
|
|
"loss_delta_derivative",
|
|
|
|
|
"sample_weight_mean",
|
|
|
|
|
"sample_weight_max",
|
|
|
|
|
)
|
|
|
|
|
@ -47,11 +49,57 @@ METRIC_KEYS = (
|
|
|
|
|
class ForwardDataset(Dataset):
|
|
|
|
|
"""把预处理后的参数、流量制度和曲线数组封装成 PyTorch Dataset。"""
|
|
|
|
|
|
|
|
|
|
def __init__(self, params_x: np.ndarray, schedule_x: np.ndarray, curve_y: np.ndarray):
|
|
|
|
|
def __init__(
|
|
|
|
|
self,
|
|
|
|
|
params_x: np.ndarray,
|
|
|
|
|
schedule_x: np.ndarray,
|
|
|
|
|
curve_y: np.ndarray,
|
|
|
|
|
curve_valid_mask: np.ndarray | None = None,
|
|
|
|
|
group_id: np.ndarray | None = None,
|
|
|
|
|
is_anchor: np.ndarray | None = None,
|
|
|
|
|
source_id: np.ndarray | None = None,
|
|
|
|
|
):
|
|
|
|
|
"""把三个 numpy 数组转为 float32 张量,后续 DataLoader 可直接按样本读取。"""
|
|
|
|
|
self.params_x = torch.tensor(params_x, dtype=torch.float32)
|
|
|
|
|
self.schedule_x = torch.tensor(schedule_x, dtype=torch.float32)
|
|
|
|
|
self.curve_y = torch.tensor(curve_y, dtype=torch.float32)
|
|
|
|
|
n_time_points = int(self.curve_y.shape[1] // 2)
|
|
|
|
|
if curve_valid_mask is None:
|
|
|
|
|
curve_valid_mask = np.ones((len(self.curve_y), n_time_points), dtype=np.float32)
|
|
|
|
|
curve_valid_mask = np.asarray(curve_valid_mask, dtype=np.float32)
|
|
|
|
|
expected_shape = (len(self.curve_y), n_time_points)
|
|
|
|
|
if curve_valid_mask.shape != expected_shape:
|
|
|
|
|
raise ValueError(
|
|
|
|
|
f"curve_valid_mask shape mismatch: {curve_valid_mask.shape} != {expected_shape}"
|
|
|
|
|
)
|
|
|
|
|
self.curve_valid_mask = torch.tensor(curve_valid_mask, dtype=torch.float32)
|
|
|
|
|
self.group_id = self._metadata_tensor("group_id", group_id, default_unique=True)
|
|
|
|
|
self.is_anchor = self._metadata_tensor("is_anchor", is_anchor, default_value=-1)
|
|
|
|
|
self.source_id = self._metadata_tensor("source_id", source_id, default_value=-1)
|
|
|
|
|
|
|
|
|
|
def _metadata_tensor(
|
|
|
|
|
self,
|
|
|
|
|
name: str,
|
|
|
|
|
values: np.ndarray | None,
|
|
|
|
|
default_value: int = -1,
|
|
|
|
|
default_unique: bool = False,
|
|
|
|
|
) -> torch.Tensor:
|
|
|
|
|
"""校验可选样本元数据,并以 int64 张量形式保存。"""
|
|
|
|
|
if values is None:
|
|
|
|
|
if default_unique:
|
|
|
|
|
array = np.arange(len(self.params_x), dtype=np.int64)
|
|
|
|
|
else:
|
|
|
|
|
array = np.full(len(self.params_x), default_value, dtype=np.int64)
|
|
|
|
|
else:
|
|
|
|
|
raw = np.asarray(values).reshape(-1)
|
|
|
|
|
if len(raw) != len(self.params_x):
|
|
|
|
|
raise ValueError(
|
|
|
|
|
f"{name} length mismatch: {len(raw)} != {len(self.params_x)}"
|
|
|
|
|
)
|
|
|
|
|
if np.issubdtype(raw.dtype, np.floating) and not np.all(np.isfinite(raw)):
|
|
|
|
|
raise ValueError(f"{name} contains non-finite values")
|
|
|
|
|
array = raw.astype(np.int64, copy=False)
|
|
|
|
|
return torch.tensor(array, dtype=torch.int64)
|
|
|
|
|
|
|
|
|
|
def __len__(self) -> int:
|
|
|
|
|
"""返回数据集或容器中可迭代样本的数量。"""
|
|
|
|
|
@ -59,7 +107,129 @@ class ForwardDataset(Dataset):
|
|
|
|
|
|
|
|
|
|
def __getitem__(self, idx: int):
|
|
|
|
|
"""按索引取出一个训练样本或数据项。"""
|
|
|
|
|
return self.params_x[idx], self.schedule_x[idx], self.curve_y[idx]
|
|
|
|
|
return (
|
|
|
|
|
self.params_x[idx],
|
|
|
|
|
self.schedule_x[idx],
|
|
|
|
|
self.curve_y[idx],
|
|
|
|
|
self.curve_valid_mask[idx],
|
|
|
|
|
self.group_id[idx],
|
|
|
|
|
self.is_anchor[idx],
|
|
|
|
|
self.source_id[idx],
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class CompleteGroupBatchSampler(Sampler[list[int]]):
|
|
|
|
|
"""按完整样本组组织批次,避免同一个锚点邻域组被拆到不同批次。"""
|
|
|
|
|
|
|
|
|
|
def __init__(
|
|
|
|
|
self,
|
|
|
|
|
group_id: np.ndarray,
|
|
|
|
|
source_id: np.ndarray | None,
|
|
|
|
|
batch_size: int,
|
|
|
|
|
shuffle: bool,
|
|
|
|
|
seed: int,
|
|
|
|
|
drop_last: bool = False,
|
|
|
|
|
):
|
|
|
|
|
if batch_size <= 0:
|
|
|
|
|
raise ValueError("batch_size must be positive")
|
|
|
|
|
|
|
|
|
|
group_values = np.asarray(group_id).reshape(-1)
|
|
|
|
|
if source_id is None:
|
|
|
|
|
source_values = np.full(len(group_values), -1, dtype=np.int64)
|
|
|
|
|
else:
|
|
|
|
|
source_values = np.asarray(source_id).reshape(-1)
|
|
|
|
|
if len(source_values) != len(group_values):
|
|
|
|
|
raise ValueError("source_id length does not match group_id")
|
|
|
|
|
|
|
|
|
|
groups: dict[tuple[str, int, int], list[int]] = {}
|
|
|
|
|
for index, (value, source) in enumerate(zip(group_values.tolist(), source_values.tolist())):
|
|
|
|
|
numeric_value = int(value)
|
|
|
|
|
key = (
|
|
|
|
|
("group", int(source), numeric_value)
|
|
|
|
|
if numeric_value >= 0
|
|
|
|
|
else ("row", int(source), index)
|
|
|
|
|
)
|
|
|
|
|
groups.setdefault(key, []).append(index)
|
|
|
|
|
|
|
|
|
|
self.group_units = list(groups.values())
|
|
|
|
|
largest_group = max((len(unit) for unit in self.group_units), default=0)
|
|
|
|
|
if largest_group > batch_size:
|
|
|
|
|
raise ValueError(
|
|
|
|
|
"complete group does not fit in one batch: "
|
|
|
|
|
f"largest_group={largest_group}, batch_size={batch_size}"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
self.batch_size = int(batch_size)
|
|
|
|
|
self.shuffle = bool(shuffle)
|
|
|
|
|
self.seed = int(seed)
|
|
|
|
|
self.drop_last = bool(drop_last)
|
|
|
|
|
self.epoch = 0
|
|
|
|
|
self._next_batches: list[list[int]] | None = None
|
|
|
|
|
|
|
|
|
|
def _build_batches(self) -> list[list[int]]:
|
|
|
|
|
order = np.arange(len(self.group_units), dtype=np.int64)
|
|
|
|
|
if self.shuffle and len(order) > 1:
|
|
|
|
|
rng = np.random.RandomState(self.seed + self.epoch)
|
|
|
|
|
rng.shuffle(order)
|
|
|
|
|
|
|
|
|
|
batches: list[list[int]] = []
|
|
|
|
|
batch: list[int] = []
|
|
|
|
|
for position in order.tolist():
|
|
|
|
|
unit = self.group_units[position]
|
|
|
|
|
if batch and len(batch) + len(unit) > self.batch_size:
|
|
|
|
|
batches.append(batch)
|
|
|
|
|
batch = []
|
|
|
|
|
batch.extend(unit)
|
|
|
|
|
|
|
|
|
|
if batch and (not self.drop_last or len(batch) == self.batch_size):
|
|
|
|
|
batches.append(batch)
|
|
|
|
|
return batches
|
|
|
|
|
|
|
|
|
|
def _get_next_batches(self) -> list[list[int]]:
|
|
|
|
|
if self._next_batches is None:
|
|
|
|
|
self._next_batches = self._build_batches()
|
|
|
|
|
return self._next_batches
|
|
|
|
|
|
|
|
|
|
def __iter__(self):
|
|
|
|
|
batches = self._get_next_batches()
|
|
|
|
|
self._next_batches = None
|
|
|
|
|
self.epoch += 1
|
|
|
|
|
yield from batches
|
|
|
|
|
|
|
|
|
|
def __len__(self) -> int:
|
|
|
|
|
return len(self._get_next_batches())
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def validate_delta_groups(dataset: ForwardDataset, split: str) -> None:
|
|
|
|
|
"""确保每个局部锚点都具有同一来源下的完整邻域样本组。"""
|
|
|
|
|
group_id = dataset.group_id.numpy()
|
|
|
|
|
is_anchor = dataset.is_anchor.numpy()
|
|
|
|
|
source_id = dataset.source_id.numpy()
|
|
|
|
|
schedule_x = dataset.schedule_x.numpy()
|
|
|
|
|
local_mask = (is_anchor == 0) | (is_anchor == 1)
|
|
|
|
|
if np.any(group_id[local_mask] < 0):
|
|
|
|
|
raise ValueError(f"delta groups must use non-negative group_id values in {split} split")
|
|
|
|
|
local_indices = np.flatnonzero(local_mask)
|
|
|
|
|
if len(local_indices) == 0:
|
|
|
|
|
raise ValueError(f"delta loss requires at least one anchor group in {split} split")
|
|
|
|
|
|
|
|
|
|
keys = {
|
|
|
|
|
(int(source_id[index]), int(group_id[index]))
|
|
|
|
|
for index in local_indices.tolist()
|
|
|
|
|
}
|
|
|
|
|
for key in keys:
|
|
|
|
|
same_group = (source_id == key[0]) & (group_id == key[1])
|
|
|
|
|
anchor_indices = np.flatnonzero(same_group & (is_anchor == 1))
|
|
|
|
|
if len(anchor_indices) != 1:
|
|
|
|
|
raise ValueError(
|
|
|
|
|
f"delta group {key} must have exactly one anchor in {split} split"
|
|
|
|
|
)
|
|
|
|
|
neighbor_indices = np.flatnonzero(same_group & (is_anchor == 0))
|
|
|
|
|
if len(neighbor_indices) == 0:
|
|
|
|
|
raise ValueError(f"delta group {key} has no neighbors in {split} split")
|
|
|
|
|
member_indices = np.flatnonzero(same_group & local_mask)
|
|
|
|
|
if not np.all(schedule_x[member_indices] == schedule_x[anchor_indices[0]]):
|
|
|
|
|
raise ValueError(f"delta group {key} has inconsistent schedules in {split} split")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@dataclass
|
|
|
|
|
@ -92,6 +262,8 @@ class LossWeights:
|
|
|
|
|
derivative_shape: float = 0.10
|
|
|
|
|
autofit_pressure: float = 0.0
|
|
|
|
|
autofit_derivative: float = 0.0
|
|
|
|
|
delta_pressure: float = 0.0
|
|
|
|
|
delta_derivative: float = 0.0
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@dataclass
|
|
|
|
|
@ -101,6 +273,7 @@ class LossConfig:
|
|
|
|
|
weights: LossWeights = field(default_factory=LossWeights)
|
|
|
|
|
use_huber: bool = True
|
|
|
|
|
huber_beta: float = 0.05
|
|
|
|
|
delta_huber_beta: float = 0.01
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@dataclass
|
|
|
|
|
@ -150,6 +323,8 @@ class LossBatchParts:
|
|
|
|
|
pred_d: torch.Tensor
|
|
|
|
|
true_p: torch.Tensor
|
|
|
|
|
true_d: torch.Tensor
|
|
|
|
|
mask_p: torch.Tensor
|
|
|
|
|
mask_d: torch.Tensor
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@dataclass
|
|
|
|
|
@ -194,6 +369,23 @@ def load_processed_dataset(path: Path) -> dict:
|
|
|
|
|
for key in required_keys:
|
|
|
|
|
if key not in data:
|
|
|
|
|
raise KeyError(f"processed dataset 缺少字段: {key}")
|
|
|
|
|
meta = data.get("meta", {}) or {}
|
|
|
|
|
if meta.get("curve_time_mode") != "fixed":
|
|
|
|
|
raise ValueError("training requires processed curve_time_mode='fixed'")
|
|
|
|
|
curve_dim = int(data["Y_curve_train"].shape[1])
|
|
|
|
|
prediction_time = np.asarray(meta.get("prediction_curve_time", []), dtype=np.float64)
|
|
|
|
|
if prediction_time.shape != (curve_dim // 2,):
|
|
|
|
|
raise ValueError(
|
|
|
|
|
f"prediction_curve_time shape {prediction_time.shape} != {(curve_dim // 2,)}"
|
|
|
|
|
)
|
|
|
|
|
for split in ("train", "val", "test"):
|
|
|
|
|
mask_key = f"curve_valid_mask_{split}"
|
|
|
|
|
if mask_key not in data:
|
|
|
|
|
raise KeyError(f"processed fixed-time dataset is missing {mask_key}")
|
|
|
|
|
expected_shape = (data[f"Y_curve_{split}"].shape[0], curve_dim // 2)
|
|
|
|
|
actual_shape = np.asarray(data[mask_key]).shape
|
|
|
|
|
if actual_shape != expected_shape:
|
|
|
|
|
raise ValueError(f"{mask_key} shape {actual_shape} != {expected_shape}")
|
|
|
|
|
return data
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@ -234,32 +426,61 @@ def get_part_slices(curve_layout: dict) -> dict[str, slice]:
|
|
|
|
|
return out
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def smooth_l1_per_sample(pred: torch.Tensor, target: torch.Tensor, beta: float) -> torch.Tensor:
|
|
|
|
|
def masked_mean_per_sample(values: torch.Tensor, valid_mask: torch.Tensor | None) -> torch.Tensor:
|
|
|
|
|
"""逐样本求均值,仅统计掩码标记的有效时间点。"""
|
|
|
|
|
if valid_mask is None:
|
|
|
|
|
return values.mean(dim=1)
|
|
|
|
|
mask = valid_mask.to(dtype=values.dtype)
|
|
|
|
|
if mask.shape != values.shape:
|
|
|
|
|
raise ValueError(f"loss mask shape mismatch: {mask.shape} != {values.shape}")
|
|
|
|
|
return (values * mask).sum(dim=1) / mask.sum(dim=1).clamp_min(1.0)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def smooth_l1_per_sample(
|
|
|
|
|
pred: torch.Tensor,
|
|
|
|
|
target: torch.Tensor,
|
|
|
|
|
beta: float,
|
|
|
|
|
valid_mask: torch.Tensor | None = None,
|
|
|
|
|
) -> torch.Tensor:
|
|
|
|
|
"""按样本计算 Smooth L1 损失,返回每个样本一个损失值。"""
|
|
|
|
|
diff = torch.abs(pred - target)
|
|
|
|
|
loss = torch.where(diff < beta, 0.5 * diff * diff / beta, diff - 0.5 * beta)
|
|
|
|
|
return loss.mean(dim=1)
|
|
|
|
|
return masked_mean_per_sample(loss, valid_mask)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def mse_per_sample(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
|
|
|
|
|
def mse_per_sample(
|
|
|
|
|
pred: torch.Tensor,
|
|
|
|
|
target: torch.Tensor,
|
|
|
|
|
valid_mask: torch.Tensor | None = None,
|
|
|
|
|
) -> torch.Tensor:
|
|
|
|
|
"""按样本计算均方误差。"""
|
|
|
|
|
return ((pred - target) ** 2).mean(dim=1)
|
|
|
|
|
return masked_mean_per_sample((pred - target) ** 2, valid_mask)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def l1_per_sample(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
|
|
|
|
|
def l1_per_sample(
|
|
|
|
|
pred: torch.Tensor,
|
|
|
|
|
target: torch.Tensor,
|
|
|
|
|
valid_mask: torch.Tensor | None = None,
|
|
|
|
|
) -> torch.Tensor:
|
|
|
|
|
"""按样本计算平均绝对误差。"""
|
|
|
|
|
return torch.abs(pred - target).mean(dim=1)
|
|
|
|
|
return masked_mean_per_sample(torch.abs(pred - target), valid_mask)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def regression_per_sample(
|
|
|
|
|
pred: torch.Tensor,
|
|
|
|
|
target: torch.Tensor,
|
|
|
|
|
loss_cfg: LossConfig,
|
|
|
|
|
valid_mask: torch.Tensor | None = None,
|
|
|
|
|
) -> torch.Tensor:
|
|
|
|
|
"""按配置在 Smooth L1 和 MSE 之间切换点值损失。"""
|
|
|
|
|
if loss_cfg.use_huber:
|
|
|
|
|
return smooth_l1_per_sample(pred, target, beta=float(loss_cfg.huber_beta))
|
|
|
|
|
return mse_per_sample(pred, target)
|
|
|
|
|
return smooth_l1_per_sample(
|
|
|
|
|
pred,
|
|
|
|
|
target,
|
|
|
|
|
beta=float(loss_cfg.huber_beta),
|
|
|
|
|
valid_mask=valid_mask,
|
|
|
|
|
)
|
|
|
|
|
return mse_per_sample(pred, target, valid_mask=valid_mask)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def first_diff(x: torch.Tensor) -> torch.Tensor:
|
|
|
|
|
@ -272,30 +493,29 @@ def affine_restore(x_scaled: torch.Tensor, mean: torch.Tensor, scale: torch.Tens
|
|
|
|
|
return x_scaled * scale + mean
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def autofit_curve_objective_per_sample(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor:
|
|
|
|
|
"""用 torch 计算自动拟合风格的曲线误差,作为训练附加目标。"""
|
|
|
|
|
weight_factor = torch.clamp(torch.abs(target) * 0.01, max=100.0)
|
|
|
|
|
weight = 1.0 / (1.0 + weight_factor)
|
|
|
|
|
scale = torch.maximum(
|
|
|
|
|
torch.maximum(torch.abs(target), torch.abs(pred)),
|
|
|
|
|
torch.full_like(target, 1e-12),
|
|
|
|
|
)
|
|
|
|
|
relative_error = torch.abs(target - pred) / scale
|
|
|
|
|
absolute_error = torch.abs(target - pred)
|
|
|
|
|
point_error = 0.7 * relative_error + 0.3 * absolute_error
|
|
|
|
|
weighted_mse = (weight * (point_error**2)).sum(dim=1)
|
|
|
|
|
weighted_mse = weighted_mse / torch.clamp(weight.sum(dim=1), min=1e-12)
|
|
|
|
|
return torch.sqrt(weighted_mse)
|
|
|
|
|
def autofit_curve_objective_per_sample(
|
|
|
|
|
pred: torch.Tensor,
|
|
|
|
|
target: torch.Tensor,
|
|
|
|
|
valid_mask: torch.Tensor | None = None,
|
|
|
|
|
) -> torch.Tensor:
|
|
|
|
|
"""用 torch 计算与 C++ 当前真实误差公式等价的曲线误差。"""
|
|
|
|
|
log_error = torch.abs(target - pred)
|
|
|
|
|
relative_error = -torch.expm1(-log_error)
|
|
|
|
|
point_error = 0.7 * log_error + 0.3 * relative_error
|
|
|
|
|
squared_mean = masked_mean_per_sample(point_error**2, valid_mask)
|
|
|
|
|
return torch.sqrt(torch.clamp(squared_mean, min=1.0e-12))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def build_sample_weight(
|
|
|
|
|
true_p: torch.Tensor,
|
|
|
|
|
true_d: torch.Tensor,
|
|
|
|
|
reweight_cfg: SampleReweightConfig,
|
|
|
|
|
mask_p: torch.Tensor | None = None,
|
|
|
|
|
mask_d: torch.Tensor | None = None,
|
|
|
|
|
) -> torch.Tensor:
|
|
|
|
|
"""根据真实曲线幅值构造样本权重,让高幅值样本训练时权重更高。"""
|
|
|
|
|
p_level = true_p.abs().mean(dim=1)
|
|
|
|
|
d_level = true_d.abs().mean(dim=1)
|
|
|
|
|
p_level = masked_mean_per_sample(true_p.abs(), mask_p)
|
|
|
|
|
d_level = masked_mean_per_sample(true_d.abs(), mask_d)
|
|
|
|
|
|
|
|
|
|
p_norm = p_level / (p_level.mean().detach() + 1e-6)
|
|
|
|
|
d_norm = d_level / (d_level.mean().detach() + 1e-6)
|
|
|
|
|
@ -308,6 +528,7 @@ def build_sample_weight(
|
|
|
|
|
def split_curve_parts(
|
|
|
|
|
pred: torch.Tensor,
|
|
|
|
|
target: torch.Tensor,
|
|
|
|
|
valid_mask: torch.Tensor,
|
|
|
|
|
slices: dict[str, slice],
|
|
|
|
|
) -> LossBatchParts:
|
|
|
|
|
"""把拼接曲线拆成压力和导数两段。"""
|
|
|
|
|
@ -316,6 +537,8 @@ def split_curve_parts(
|
|
|
|
|
pred_d=pred[:, slices["log_derivative"]],
|
|
|
|
|
true_p=target[:, slices["log_pressure"]],
|
|
|
|
|
true_d=target[:, slices["log_derivative"]],
|
|
|
|
|
mask_p=valid_mask,
|
|
|
|
|
mask_d=valid_mask,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@ -325,20 +548,31 @@ def compute_basic_loss_vectors(
|
|
|
|
|
) -> dict[str, torch.Tensor]:
|
|
|
|
|
"""计算标准化空间中的基础点值、偏置和导数形状损失。"""
|
|
|
|
|
return {
|
|
|
|
|
"loss_pressure": regression_per_sample(parts.pred_p, parts.true_p, loss_cfg),
|
|
|
|
|
"loss_derivative": regression_per_sample(parts.pred_d, parts.true_d, loss_cfg),
|
|
|
|
|
"loss_pressure": regression_per_sample(
|
|
|
|
|
parts.pred_p,
|
|
|
|
|
parts.true_p,
|
|
|
|
|
loss_cfg,
|
|
|
|
|
valid_mask=parts.mask_p,
|
|
|
|
|
),
|
|
|
|
|
"loss_derivative": regression_per_sample(
|
|
|
|
|
parts.pred_d,
|
|
|
|
|
parts.true_d,
|
|
|
|
|
loss_cfg,
|
|
|
|
|
valid_mask=parts.mask_d,
|
|
|
|
|
),
|
|
|
|
|
"loss_bias_pressure": l1_per_sample(
|
|
|
|
|
parts.pred_p.mean(dim=1, keepdim=True),
|
|
|
|
|
parts.true_p.mean(dim=1, keepdim=True),
|
|
|
|
|
masked_mean_per_sample(parts.pred_p, parts.mask_p).unsqueeze(1),
|
|
|
|
|
masked_mean_per_sample(parts.true_p, parts.mask_p).unsqueeze(1),
|
|
|
|
|
),
|
|
|
|
|
"loss_bias_derivative": l1_per_sample(
|
|
|
|
|
parts.pred_d.mean(dim=1, keepdim=True),
|
|
|
|
|
parts.true_d.mean(dim=1, keepdim=True),
|
|
|
|
|
masked_mean_per_sample(parts.pred_d, parts.mask_d).unsqueeze(1),
|
|
|
|
|
masked_mean_per_sample(parts.true_d, parts.mask_d).unsqueeze(1),
|
|
|
|
|
),
|
|
|
|
|
"loss_derivative_shape": regression_per_sample(
|
|
|
|
|
first_diff(parts.pred_d),
|
|
|
|
|
first_diff(parts.true_d),
|
|
|
|
|
loss_cfg,
|
|
|
|
|
valid_mask=parts.mask_d[:, 1:] * parts.mask_d[:, :-1],
|
|
|
|
|
),
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
@ -360,14 +594,118 @@ def compute_autofit_loss_vectors(
|
|
|
|
|
"loss_autofit_pressure": autofit_curve_objective_per_sample(
|
|
|
|
|
affine_restore(parts.pred_p, mean_p, scale_p),
|
|
|
|
|
affine_restore(parts.true_p, mean_p, scale_p),
|
|
|
|
|
valid_mask=parts.mask_p,
|
|
|
|
|
),
|
|
|
|
|
"loss_autofit_derivative": autofit_curve_objective_per_sample(
|
|
|
|
|
affine_restore(parts.pred_d, mean_d, scale_d),
|
|
|
|
|
affine_restore(parts.true_d, mean_d, scale_d),
|
|
|
|
|
valid_mask=parts.mask_d,
|
|
|
|
|
),
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def grouped_channel_delta_loss(
|
|
|
|
|
pred_raw: torch.Tensor,
|
|
|
|
|
target_raw: torch.Tensor,
|
|
|
|
|
valid_mask: torch.Tensor,
|
|
|
|
|
group_id: torch.Tensor,
|
|
|
|
|
is_anchor: torch.Tensor,
|
|
|
|
|
source_id: torch.Tensor,
|
|
|
|
|
beta: float,
|
|
|
|
|
) -> tuple[torch.Tensor, int]:
|
|
|
|
|
"""比较锚点与邻域样本的曲线变化,并让每个样本组具有相同权重。"""
|
|
|
|
|
if beta <= 0.0:
|
|
|
|
|
raise ValueError("delta_huber_beta must be positive")
|
|
|
|
|
if group_id.ndim != 1 or is_anchor.ndim != 1 or source_id.ndim != 1:
|
|
|
|
|
raise ValueError("group_id, is_anchor and source_id must be one-dimensional")
|
|
|
|
|
if (
|
|
|
|
|
len(group_id) != len(pred_raw)
|
|
|
|
|
or len(is_anchor) != len(pred_raw)
|
|
|
|
|
or len(source_id) != len(pred_raw)
|
|
|
|
|
):
|
|
|
|
|
raise ValueError("group metadata length does not match the curve batch")
|
|
|
|
|
|
|
|
|
|
zero = pred_raw.sum() * 0.0
|
|
|
|
|
group_losses: list[torch.Tensor] = []
|
|
|
|
|
anchor_indices = torch.nonzero(is_anchor == 1, as_tuple=False).flatten()
|
|
|
|
|
for anchor_index in anchor_indices:
|
|
|
|
|
current_group = group_id[anchor_index]
|
|
|
|
|
current_source = source_id[anchor_index]
|
|
|
|
|
same_group = (group_id == current_group) & (source_id == current_source)
|
|
|
|
|
anchor_mask = same_group & (is_anchor == 1)
|
|
|
|
|
if int(anchor_mask.sum().item()) != 1:
|
|
|
|
|
raise ValueError(
|
|
|
|
|
f"group_id={int(current_group.item())} must contain exactly one anchor"
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
neighbor_indices = torch.nonzero(
|
|
|
|
|
same_group & (is_anchor == 0),
|
|
|
|
|
as_tuple=False,
|
|
|
|
|
).flatten()
|
|
|
|
|
if neighbor_indices.numel() == 0:
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
common_mask = valid_mask[neighbor_indices] * valid_mask[anchor_index].unsqueeze(0)
|
|
|
|
|
usable = common_mask.sum(dim=1) > 0
|
|
|
|
|
if not bool(torch.any(usable).item()):
|
|
|
|
|
continue
|
|
|
|
|
|
|
|
|
|
neighbor_indices = neighbor_indices[usable]
|
|
|
|
|
common_mask = common_mask[usable]
|
|
|
|
|
pred_delta = pred_raw[neighbor_indices] - pred_raw[anchor_index].unsqueeze(0)
|
|
|
|
|
target_delta = target_raw[neighbor_indices] - target_raw[anchor_index].unsqueeze(0)
|
|
|
|
|
neighbor_losses = smooth_l1_per_sample(
|
|
|
|
|
pred_delta,
|
|
|
|
|
target_delta,
|
|
|
|
|
beta=beta,
|
|
|
|
|
valid_mask=common_mask,
|
|
|
|
|
)
|
|
|
|
|
group_losses.append(neighbor_losses.mean())
|
|
|
|
|
|
|
|
|
|
if not group_losses:
|
|
|
|
|
return zero, 0
|
|
|
|
|
return torch.stack(group_losses).mean(), len(group_losses)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def compute_group_delta_losses(
|
|
|
|
|
parts: LossBatchParts,
|
|
|
|
|
group_id: torch.Tensor,
|
|
|
|
|
is_anchor: torch.Tensor,
|
|
|
|
|
source_id: torch.Tensor,
|
|
|
|
|
context: LossContext,
|
|
|
|
|
) -> tuple[torch.Tensor, torch.Tensor, int]:
|
|
|
|
|
"""在反标准化后的原始对数尺度上计算压力和导数差分损失。"""
|
|
|
|
|
pressure_slice = context.slices["log_pressure"]
|
|
|
|
|
derivative_slice = context.slices["log_derivative"]
|
|
|
|
|
mean_p = context.curve_stats.mean_raw[pressure_slice].unsqueeze(0)
|
|
|
|
|
scale_p = context.curve_stats.scale_raw[pressure_slice].unsqueeze(0)
|
|
|
|
|
mean_d = context.curve_stats.mean_raw[derivative_slice].unsqueeze(0)
|
|
|
|
|
scale_d = context.curve_stats.scale_raw[derivative_slice].unsqueeze(0)
|
|
|
|
|
beta = float(context.loss_cfg.delta_huber_beta)
|
|
|
|
|
|
|
|
|
|
pressure_loss, pressure_group_count = grouped_channel_delta_loss(
|
|
|
|
|
pred_raw=affine_restore(parts.pred_p, mean_p, scale_p),
|
|
|
|
|
target_raw=affine_restore(parts.true_p, mean_p, scale_p),
|
|
|
|
|
valid_mask=parts.mask_p,
|
|
|
|
|
group_id=group_id,
|
|
|
|
|
is_anchor=is_anchor,
|
|
|
|
|
source_id=source_id,
|
|
|
|
|
beta=beta,
|
|
|
|
|
)
|
|
|
|
|
derivative_loss, derivative_group_count = grouped_channel_delta_loss(
|
|
|
|
|
pred_raw=affine_restore(parts.pred_d, mean_d, scale_d),
|
|
|
|
|
target_raw=affine_restore(parts.true_d, mean_d, scale_d),
|
|
|
|
|
valid_mask=parts.mask_d,
|
|
|
|
|
group_id=group_id,
|
|
|
|
|
is_anchor=is_anchor,
|
|
|
|
|
source_id=source_id,
|
|
|
|
|
beta=beta,
|
|
|
|
|
)
|
|
|
|
|
if pressure_group_count != derivative_group_count:
|
|
|
|
|
raise RuntimeError("pressure and derivative delta group counts do not match")
|
|
|
|
|
return pressure_loss, derivative_loss, pressure_group_count
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def weighted_total_vector(
|
|
|
|
|
loss_vectors: dict[str, torch.Tensor],
|
|
|
|
|
weights: LossWeights,
|
|
|
|
|
@ -387,21 +725,59 @@ def weighted_total_vector(
|
|
|
|
|
def compute_weighted_loss(
|
|
|
|
|
pred: torch.Tensor,
|
|
|
|
|
target: torch.Tensor,
|
|
|
|
|
valid_mask: torch.Tensor,
|
|
|
|
|
context: LossContext,
|
|
|
|
|
group_id: torch.Tensor | None = None,
|
|
|
|
|
is_anchor: torch.Tensor | None = None,
|
|
|
|
|
source_id: torch.Tensor | None = None,
|
|
|
|
|
) -> dict[str, torch.Tensor]:
|
|
|
|
|
"""计算正演代理模型的复合训练损失。"""
|
|
|
|
|
parts = split_curve_parts(pred, target, context.slices)
|
|
|
|
|
parts = split_curve_parts(pred, target, valid_mask, context.slices)
|
|
|
|
|
loss_vectors = compute_basic_loss_vectors(parts, context.loss_cfg)
|
|
|
|
|
loss_vectors.update(compute_autofit_loss_vectors(parts, context))
|
|
|
|
|
|
|
|
|
|
total_vec = weighted_total_vector(loss_vectors, context.loss_cfg.weights)
|
|
|
|
|
if context.reweight_cfg.enabled:
|
|
|
|
|
sample_weight = build_sample_weight(parts.true_p, parts.true_d, context.reweight_cfg)
|
|
|
|
|
sample_weight = build_sample_weight(
|
|
|
|
|
parts.true_p,
|
|
|
|
|
parts.true_d,
|
|
|
|
|
context.reweight_cfg,
|
|
|
|
|
mask_p=parts.mask_p,
|
|
|
|
|
mask_d=parts.mask_d,
|
|
|
|
|
)
|
|
|
|
|
else:
|
|
|
|
|
sample_weight = torch.ones_like(total_vec)
|
|
|
|
|
|
|
|
|
|
metrics = {key: value.mean() for key, value in loss_vectors.items()}
|
|
|
|
|
metrics["loss"] = (total_vec * sample_weight).mean()
|
|
|
|
|
base_loss = (total_vec * sample_weight).mean()
|
|
|
|
|
metrics["loss"] = base_loss
|
|
|
|
|
weights = context.loss_cfg.weights
|
|
|
|
|
delta_enabled = bool(weights.delta_pressure != 0.0 or weights.delta_derivative != 0.0)
|
|
|
|
|
if delta_enabled:
|
|
|
|
|
if group_id is None or is_anchor is None:
|
|
|
|
|
raise ValueError("delta loss requires group_id and is_anchor metadata")
|
|
|
|
|
if source_id is None:
|
|
|
|
|
source_id = torch.full_like(group_id, -1)
|
|
|
|
|
delta_pressure, delta_derivative, delta_group_count = compute_group_delta_losses(
|
|
|
|
|
parts=parts,
|
|
|
|
|
group_id=group_id,
|
|
|
|
|
is_anchor=is_anchor,
|
|
|
|
|
source_id=source_id,
|
|
|
|
|
context=context,
|
|
|
|
|
)
|
|
|
|
|
metrics["loss"] = (
|
|
|
|
|
metrics["loss"]
|
|
|
|
|
+ weights.delta_pressure * delta_pressure
|
|
|
|
|
+ weights.delta_derivative * delta_derivative
|
|
|
|
|
)
|
|
|
|
|
else:
|
|
|
|
|
delta_pressure = pred.sum() * 0.0
|
|
|
|
|
delta_derivative = pred.sum() * 0.0
|
|
|
|
|
delta_group_count = 0
|
|
|
|
|
metrics["loss_base"] = base_loss
|
|
|
|
|
metrics["loss_delta_pressure"] = delta_pressure
|
|
|
|
|
metrics["loss_delta_derivative"] = delta_derivative
|
|
|
|
|
metrics["delta_group_count"] = pred.new_tensor(float(delta_group_count))
|
|
|
|
|
metrics["sample_weight_mean"] = sample_weight.mean()
|
|
|
|
|
metrics["sample_weight_max"] = sample_weight.max()
|
|
|
|
|
return metrics
|
|
|
|
|
@ -454,19 +830,43 @@ def run_loader_epoch(
|
|
|
|
|
|
|
|
|
|
total = init_metric_accumulator()
|
|
|
|
|
total_n = 0
|
|
|
|
|
base_loss_total = 0.0
|
|
|
|
|
delta_pressure_total = 0.0
|
|
|
|
|
delta_derivative_total = 0.0
|
|
|
|
|
delta_group_total = 0
|
|
|
|
|
grad_context = torch.enable_grad() if is_train else torch.no_grad()
|
|
|
|
|
|
|
|
|
|
with grad_context:
|
|
|
|
|
for params_x, schedule_x, curve_y in loader:
|
|
|
|
|
for (
|
|
|
|
|
params_x,
|
|
|
|
|
schedule_x,
|
|
|
|
|
curve_y,
|
|
|
|
|
curve_valid_mask,
|
|
|
|
|
group_id,
|
|
|
|
|
is_anchor,
|
|
|
|
|
source_id,
|
|
|
|
|
) in loader:
|
|
|
|
|
params_x = params_x.to(device)
|
|
|
|
|
schedule_x = schedule_x.to(device)
|
|
|
|
|
curve_y = curve_y.to(device)
|
|
|
|
|
curve_valid_mask = curve_valid_mask.to(device)
|
|
|
|
|
group_id = group_id.to(device)
|
|
|
|
|
is_anchor = is_anchor.to(device)
|
|
|
|
|
source_id = source_id.to(device)
|
|
|
|
|
|
|
|
|
|
if is_train:
|
|
|
|
|
optimizer.zero_grad()
|
|
|
|
|
|
|
|
|
|
pred = model_forward(model, params_x, schedule_x, use_schedule)
|
|
|
|
|
losses = compute_weighted_loss(pred=pred, target=curve_y, context=context)
|
|
|
|
|
losses = compute_weighted_loss(
|
|
|
|
|
pred=pred,
|
|
|
|
|
target=curve_y,
|
|
|
|
|
valid_mask=curve_valid_mask,
|
|
|
|
|
context=context,
|
|
|
|
|
group_id=group_id,
|
|
|
|
|
is_anchor=is_anchor,
|
|
|
|
|
source_id=source_id,
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
if is_train:
|
|
|
|
|
losses["loss"].backward()
|
|
|
|
|
@ -474,9 +874,27 @@ def run_loader_epoch(
|
|
|
|
|
|
|
|
|
|
batch_size = params_x.size(0)
|
|
|
|
|
accumulate_metrics(total, losses, batch_size)
|
|
|
|
|
base_loss_total += losses["loss_base"].item() * batch_size
|
|
|
|
|
delta_group_count = int(losses["delta_group_count"].item())
|
|
|
|
|
delta_pressure_total += losses["loss_delta_pressure"].item() * delta_group_count
|
|
|
|
|
delta_derivative_total += losses["loss_delta_derivative"].item() * delta_group_count
|
|
|
|
|
delta_group_total += delta_group_count
|
|
|
|
|
total_n += batch_size
|
|
|
|
|
|
|
|
|
|
return average_metrics(total, total_n)
|
|
|
|
|
metrics = average_metrics(total, total_n)
|
|
|
|
|
if delta_group_total > 0:
|
|
|
|
|
metrics["loss_delta_pressure"] = delta_pressure_total / delta_group_total
|
|
|
|
|
metrics["loss_delta_derivative"] = delta_derivative_total / delta_group_total
|
|
|
|
|
else:
|
|
|
|
|
metrics["loss_delta_pressure"] = 0.0
|
|
|
|
|
metrics["loss_delta_derivative"] = 0.0
|
|
|
|
|
weights = context.loss_cfg.weights
|
|
|
|
|
metrics["loss"] = (
|
|
|
|
|
base_loss_total / max(total_n, 1)
|
|
|
|
|
+ weights.delta_pressure * metrics["loss_delta_pressure"]
|
|
|
|
|
+ weights.delta_derivative * metrics["loss_delta_derivative"]
|
|
|
|
|
)
|
|
|
|
|
return metrics
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
def evaluate(
|
|
|
|
|
@ -510,21 +928,97 @@ def build_curve_stats(data: dict, device: str) -> CurveStats:
|
|
|
|
|
|
|
|
|
|
def build_dataloaders(data: dict, cfg: TrainConfig) -> DatasetBundle:
|
|
|
|
|
"""根据预处理数组构造训练、验证、测试 DataLoader。"""
|
|
|
|
|
train_ds = ForwardDataset(data["X_params_train"], data["X_schedule_train"], data["Y_curve_train"])
|
|
|
|
|
val_ds = ForwardDataset(data["X_params_val"], data["X_schedule_val"], data["Y_curve_val"])
|
|
|
|
|
test_ds = ForwardDataset(data["X_params_test"], data["X_schedule_test"], data["Y_curve_test"])
|
|
|
|
|
|
|
|
|
|
loader_generator = torch.Generator()
|
|
|
|
|
loader_generator.manual_seed(int(cfg.runtime.seed))
|
|
|
|
|
|
|
|
|
|
train_loader = DataLoader(
|
|
|
|
|
train_ds,
|
|
|
|
|
batch_size=cfg.optim.batch_size,
|
|
|
|
|
shuffle=True,
|
|
|
|
|
generator=loader_generator,
|
|
|
|
|
delta_enabled = bool(
|
|
|
|
|
cfg.loss.weights.delta_pressure != 0.0
|
|
|
|
|
or cfg.loss.weights.delta_derivative != 0.0
|
|
|
|
|
)
|
|
|
|
|
val_loader = DataLoader(val_ds, batch_size=cfg.optim.batch_size, shuffle=False)
|
|
|
|
|
test_loader = DataLoader(test_ds, batch_size=cfg.optim.batch_size, shuffle=False)
|
|
|
|
|
if cfg.loss.weights.delta_pressure < 0.0 or cfg.loss.weights.delta_derivative < 0.0:
|
|
|
|
|
raise ValueError("delta loss weights must be non-negative")
|
|
|
|
|
if delta_enabled and cfg.loss.delta_huber_beta <= 0.0:
|
|
|
|
|
raise ValueError("delta_huber_beta must be positive when delta loss is enabled")
|
|
|
|
|
if delta_enabled:
|
|
|
|
|
for split in ("train", "val", "test"):
|
|
|
|
|
for field_name in ("group_id", "is_anchor"):
|
|
|
|
|
key = f"{field_name}_{split}"
|
|
|
|
|
if key not in data:
|
|
|
|
|
raise KeyError(f"delta loss requires processed dataset field: {key}")
|
|
|
|
|
|
|
|
|
|
train_ds = ForwardDataset(
|
|
|
|
|
data["X_params_train"],
|
|
|
|
|
data["X_schedule_train"],
|
|
|
|
|
data["Y_curve_train"],
|
|
|
|
|
data.get("curve_valid_mask_train"),
|
|
|
|
|
data.get("group_id_train"),
|
|
|
|
|
data.get("is_anchor_train"),
|
|
|
|
|
data.get("source_id_train"),
|
|
|
|
|
)
|
|
|
|
|
val_ds = ForwardDataset(
|
|
|
|
|
data["X_params_val"],
|
|
|
|
|
data["X_schedule_val"],
|
|
|
|
|
data["Y_curve_val"],
|
|
|
|
|
data.get("curve_valid_mask_val"),
|
|
|
|
|
data.get("group_id_val"),
|
|
|
|
|
data.get("is_anchor_val"),
|
|
|
|
|
data.get("source_id_val"),
|
|
|
|
|
)
|
|
|
|
|
test_ds = ForwardDataset(
|
|
|
|
|
data["X_params_test"],
|
|
|
|
|
data["X_schedule_test"],
|
|
|
|
|
data["Y_curve_test"],
|
|
|
|
|
data.get("curve_valid_mask_test"),
|
|
|
|
|
data.get("group_id_test"),
|
|
|
|
|
data.get("is_anchor_test"),
|
|
|
|
|
data.get("source_id_test"),
|
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
if delta_enabled:
|
|
|
|
|
validate_delta_groups(train_ds, "train")
|
|
|
|
|
validate_delta_groups(val_ds, "val")
|
|
|
|
|
validate_delta_groups(test_ds, "test")
|
|
|
|
|
train_loader = DataLoader(
|
|
|
|
|
train_ds,
|
|
|
|
|
batch_sampler=CompleteGroupBatchSampler(
|
|
|
|
|
train_ds.group_id.numpy(),
|
|
|
|
|
train_ds.source_id.numpy(),
|
|
|
|
|
batch_size=cfg.optim.batch_size,
|
|
|
|
|
shuffle=True,
|
|
|
|
|
seed=cfg.runtime.seed,
|
|
|
|
|
drop_last=False,
|
|
|
|
|
),
|
|
|
|
|
)
|
|
|
|
|
val_loader = DataLoader(
|
|
|
|
|
val_ds,
|
|
|
|
|
batch_sampler=CompleteGroupBatchSampler(
|
|
|
|
|
val_ds.group_id.numpy(),
|
|
|
|
|
val_ds.source_id.numpy(),
|
|
|
|
|
batch_size=cfg.optim.batch_size,
|
|
|
|
|
shuffle=False,
|
|
|
|
|
seed=cfg.runtime.seed + 1,
|
|
|
|
|
drop_last=False,
|
|
|
|
|
),
|
|
|
|
|
)
|
|
|
|
|
test_loader = DataLoader(
|
|
|
|
|
test_ds,
|
|
|
|
|
batch_sampler=CompleteGroupBatchSampler(
|
|
|
|
|
test_ds.group_id.numpy(),
|
|
|
|
|
test_ds.source_id.numpy(),
|
|
|
|
|
batch_size=cfg.optim.batch_size,
|
|
|
|
|
shuffle=False,
|
|
|
|
|
seed=cfg.runtime.seed + 2,
|
|
|
|
|
drop_last=False,
|
|
|
|
|
),
|
|
|
|
|
)
|
|
|
|
|
else:
|
|
|
|
|
loader_generator = torch.Generator()
|
|
|
|
|
loader_generator.manual_seed(int(cfg.runtime.seed))
|
|
|
|
|
train_loader = DataLoader(
|
|
|
|
|
train_ds,
|
|
|
|
|
batch_size=cfg.optim.batch_size,
|
|
|
|
|
shuffle=True,
|
|
|
|
|
generator=loader_generator,
|
|
|
|
|
)
|
|
|
|
|
val_loader = DataLoader(val_ds, batch_size=cfg.optim.batch_size, shuffle=False)
|
|
|
|
|
test_loader = DataLoader(test_ds, batch_size=cfg.optim.batch_size, shuffle=False)
|
|
|
|
|
|
|
|
|
|
return DatasetBundle(
|
|
|
|
|
train_loader=train_loader,
|
|
|
|
|
@ -581,7 +1075,9 @@ def print_training_config(cfg: TrainConfig, curve_layout: dict) -> None:
|
|
|
|
|
f" weights: pressure={weights.pressure}, derivative={weights.derivative}, "
|
|
|
|
|
f"bias_p={weights.bias_pressure}, "
|
|
|
|
|
f"bias_d={weights.bias_derivative}, d_shape={weights.derivative_shape}, "
|
|
|
|
|
f"autofit_p={weights.autofit_pressure}, autofit_d={weights.autofit_derivative}"
|
|
|
|
|
f"autofit_p={weights.autofit_pressure}, autofit_d={weights.autofit_derivative}, "
|
|
|
|
|
f"delta_p={weights.delta_pressure}, delta_d={weights.delta_derivative}, "
|
|
|
|
|
f"delta_beta={cfg.loss.delta_huber_beta}"
|
|
|
|
|
)
|
|
|
|
|
print(
|
|
|
|
|
f" sample_reweight={reweight.enabled}, alpha={reweight.alpha}, "
|
|
|
|
|
@ -603,6 +1099,8 @@ def format_metric_line(epoch: int, train_metrics: dict[str, float], val_metrics:
|
|
|
|
|
f"ds={train_metrics['loss_derivative_shape']:.6f}, "
|
|
|
|
|
f"ap={train_metrics['loss_autofit_pressure']:.6f}, "
|
|
|
|
|
f"ad={train_metrics['loss_autofit_derivative']:.6f}, "
|
|
|
|
|
f"dp={train_metrics['loss_delta_pressure']:.6f}, "
|
|
|
|
|
f"dd={train_metrics['loss_delta_derivative']:.6f}, "
|
|
|
|
|
f"wmean={train_metrics['sample_weight_mean']:.4f}, "
|
|
|
|
|
f"wmax={train_metrics['sample_weight_max']:.4f}) "
|
|
|
|
|
f"val={val_metrics['loss']:.6f} "
|
|
|
|
|
@ -613,6 +1111,8 @@ def format_metric_line(epoch: int, train_metrics: dict[str, float], val_metrics:
|
|
|
|
|
f"ds={val_metrics['loss_derivative_shape']:.6f}, "
|
|
|
|
|
f"ap={val_metrics['loss_autofit_pressure']:.6f}, "
|
|
|
|
|
f"ad={val_metrics['loss_autofit_derivative']:.6f}, "
|
|
|
|
|
f"dp={val_metrics['loss_delta_pressure']:.6f}, "
|
|
|
|
|
f"dd={val_metrics['loss_delta_derivative']:.6f}, "
|
|
|
|
|
f"wmean={val_metrics['sample_weight_mean']:.4f}, "
|
|
|
|
|
f"wmax={val_metrics['sample_weight_max']:.4f})"
|
|
|
|
|
)
|
|
|
|
|
@ -629,6 +1129,8 @@ def format_final_line(test_metrics: dict[str, float]) -> str:
|
|
|
|
|
f"ds={test_metrics['loss_derivative_shape']:.6f}, "
|
|
|
|
|
f"ap={test_metrics['loss_autofit_pressure']:.6f}, "
|
|
|
|
|
f"ad={test_metrics['loss_autofit_derivative']:.6f}, "
|
|
|
|
|
f"dp={test_metrics['loss_delta_pressure']:.6f}, "
|
|
|
|
|
f"dd={test_metrics['loss_delta_derivative']:.6f}, "
|
|
|
|
|
f"wmean={test_metrics['sample_weight_mean']:.4f}, "
|
|
|
|
|
f"wmax={test_metrics['sample_weight_max']:.4f})"
|
|
|
|
|
)
|
|
|
|
|
@ -652,6 +1154,7 @@ def build_checkpoint_payload(
|
|
|
|
|
"seed": int(cfg.runtime.seed),
|
|
|
|
|
"curve_layout": curve_layout,
|
|
|
|
|
"loss_weights": asdict(cfg.loss.weights),
|
|
|
|
|
"delta_huber_beta": float(cfg.loss.delta_huber_beta),
|
|
|
|
|
"sample_reweight": asdict(cfg.sample_reweight),
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
@ -748,6 +1251,7 @@ def build_metrics_payload(
|
|
|
|
|
"use_schedule": cfg.model.use_schedule,
|
|
|
|
|
"seed": int(cfg.runtime.seed),
|
|
|
|
|
"loss_weights": asdict(cfg.loss.weights),
|
|
|
|
|
"delta_huber_beta": float(cfg.loss.delta_huber_beta),
|
|
|
|
|
"sample_reweight": asdict(cfg.sample_reweight),
|
|
|
|
|
"curve_layout": curve_layout,
|
|
|
|
|
}
|
|
|
|
|
|