Add train-split robust scaling for continuous values

This commit is contained in:
2026-08-01 11:37:21 +08:00
parent 89dcf4b362
commit dfb22adf2d
6 changed files with 419 additions and 5 deletions

View File

@@ -161,6 +161,9 @@ def build_model_from_dataset(
n_categories=dataset.n_categories, n_categories=dataset.n_categories,
cont_type_ids=dataset.cont_type_ids, cont_type_ids=dataset.cont_type_ids,
n_bins=int(cfg_get(args, cfg, "n_bins", 16)), n_bins=int(cfg_get(args, cfg, "n_bins", 16)),
continuous_value_scaling=str(
cfg_get(args, cfg, "continuous_value_scaling", "none")
),
extra_pool_reduce=str( extra_pool_reduce=str(
cfg_get(args, cfg, "extra_pool_reduce", "mean") cfg_get(args, cfg, "extra_pool_reduce", "mean")
), ),

View File

@@ -37,6 +37,9 @@ class OtherInfoTokenizer(nn.Module):
cont_type_ids: list[int], cont_type_ids: list[int],
n_value_kinds: int = 3, n_value_kinds: int = 3,
n_bins: int = 16, n_bins: int = 16,
continuous_value_scaling: str = "none",
continuous_value_center: torch.Tensor | list[float] | None = None,
continuous_value_scale: torch.Tensor | list[float] | None = None,
): ):
super().__init__() super().__init__()
if len(cont_type_ids) != n_cont_types: if len(cont_type_ids) != n_cont_types:
@@ -54,6 +57,12 @@ class OtherInfoTokenizer(nn.Module):
raise ValueError( raise ValueError(
f"n_value_kinds must be > {self.CATE_KIND}, got {n_value_kinds}" f"n_value_kinds must be > {self.CATE_KIND}, got {n_value_kinds}"
) )
continuous_value_scaling = str(continuous_value_scaling).lower()
if continuous_value_scaling not in {"none", "robust"}:
raise ValueError(
"continuous_value_scaling must be either 'none' or 'robust', "
f"got {continuous_value_scaling!r}"
)
self.type_emb = nn.Embedding(n_types, n_embd, padding_idx=0) self.type_emb = nn.Embedding(n_types, n_embd, padding_idx=0)
self.kind_emb = nn.Embedding(n_value_kinds, n_embd, padding_idx=0) self.kind_emb = nn.Embedding(n_value_kinds, n_embd, padding_idx=0)
@@ -71,6 +80,36 @@ class OtherInfoTokenizer(nn.Module):
n_embd, n_embd,
padding_idx=0, padding_idx=0,
) )
self.continuous_value_scaling = continuous_value_scaling
if continuous_value_scaling == "robust" and n_cont_types > 0:
center = self._coerce_scaler_buffer(
continuous_value_center,
n_cont_types=n_cont_types,
default=0.0,
name="continuous_value_center",
)
scale = self._coerce_scaler_buffer(
continuous_value_scale,
n_cont_types=n_cont_types,
default=1.0,
name="continuous_value_scale",
)
if not torch.isfinite(center).all():
raise ValueError(
"continuous_value_center must contain only finite values"
)
if not torch.isfinite(scale).all() or torch.any(scale <= 0):
raise ValueError(
"continuous_value_scale must be finite and strictly positive"
)
self.register_buffer("continuous_value_center", center)
self.register_buffer("continuous_value_scale", scale)
else:
# ``None`` buffers are omitted from state_dict. This preserves the
# exact checkpoint schema used by models trained before continuous
# value scaling was introduced.
self.register_buffer("continuous_value_center", None)
self.register_buffer("continuous_value_scale", None)
cont_type_index = torch.full((n_types,), -1, dtype=torch.long) cont_type_index = torch.full((n_types,), -1, dtype=torch.long)
for idx, type_id in enumerate(cont_type_ids): for idx, type_id in enumerate(cont_type_ids):
@@ -86,6 +125,23 @@ class OtherInfoTokenizer(nn.Module):
) )
self.reset_parameters() self.reset_parameters()
@staticmethod
def _coerce_scaler_buffer(
value: torch.Tensor | list[float] | None,
*,
n_cont_types: int,
default: float,
name: str,
) -> torch.Tensor:
if value is None:
return torch.full((n_cont_types,), float(default), dtype=torch.float32)
tensor = torch.as_tensor(value, dtype=torch.float32).detach().clone()
if tensor.shape != (n_cont_types,):
raise ValueError(
f"{name} must have shape ({n_cont_types},), got {tuple(tensor.shape)}"
)
return tensor
def reset_parameters(self) -> None: def reset_parameters(self) -> None:
nn.init.normal_(self.type_emb.weight, mean=0.0, std=0.02) nn.init.normal_(self.type_emb.weight, mean=0.0, std=0.02)
nn.init.zeros_(self.type_emb.weight[0]) nn.init.zeros_(self.type_emb.weight[0])
@@ -127,9 +183,19 @@ class OtherInfoTokenizer(nn.Module):
f"type_id={bad_type} is marked continuous but is not in " f"type_id={bad_type} is marked continuous but is not in "
"cont_type_ids" "cont_type_ids"
) )
cont_value = other_value[cont_pos].to(type_emb.dtype)
if self.continuous_value_scaling == "robust":
if (
self.continuous_value_center is None
or self.continuous_value_scale is None
):
raise RuntimeError("Robust continuous-value scaler buffers are missing")
center = self.continuous_value_center[cont_idx].to(type_emb.dtype)
scale = self.continuous_value_scale[cont_idx].to(type_emb.dtype)
cont_value = (cont_value - center) / scale
value_emb[cont_pos] = self.cont_value_encoder( value_emb[cont_pos] = self.cont_value_encoder(
cont_type_idx=cont_idx, cont_type_idx=cont_idx,
value=other_value[cont_pos].to(type_emb.dtype), value=cont_value,
) )
cate_pos = other_valid & (other_value_kind == self.CATE_KIND) cate_pos = other_valid & (other_value_kind == self.CATE_KIND)
@@ -155,6 +221,9 @@ class DeepHealth(nn.Module):
cont_type_ids: list[int], cont_type_ids: list[int],
n_value_kinds: int = 3, n_value_kinds: int = 3,
n_bins: int = 16, n_bins: int = 16,
continuous_value_scaling: str = "none",
continuous_value_center: torch.Tensor | list[float] | None = None,
continuous_value_scale: torch.Tensor | list[float] | None = None,
target_mode: str = "next_token", # "next_token" or "all_future" target_mode: str = "next_token", # "next_token" or "all_future"
time_mode: str = "absolute", # next_token requires absolute time_mode: str = "absolute", # next_token requires absolute
dist_mode: str = "exponential", # "exponential", "weibull" or "mixed" dist_mode: str = "exponential", # "exponential", "weibull" or "mixed"
@@ -193,11 +262,15 @@ class DeepHealth(nn.Module):
cont_type_ids=cont_type_ids, cont_type_ids=cont_type_ids,
n_value_kinds=n_value_kinds, n_value_kinds=n_value_kinds,
n_bins=n_bins, n_bins=n_bins,
continuous_value_scaling=continuous_value_scaling,
continuous_value_center=continuous_value_center,
continuous_value_scale=continuous_value_scale,
) )
self.target_mode = target_mode self.target_mode = target_mode
self.time_mode = time_mode self.time_mode = time_mode
self.dist_mode = dist_mode self.dist_mode = dist_mode
self.extra_pool_reduce = extra_pool_reduce self.extra_pool_reduce = extra_pool_reduce
self.continuous_value_scaling = str(continuous_value_scaling).lower()
self.model_architecture = model_architecture self.model_architecture = model_architecture
self.n_layer = n_layer self.n_layer = n_layer
self.n_embd = n_embd self.n_embd = n_embd

View File

@@ -0,0 +1,142 @@
import unittest
import numpy as np
import torch
import torch.nn as nn
from torch.utils.data import Subset
from models import OtherInfoTokenizer
from train_util import fit_continuous_robust_scaler
class _ToyAllFutureDataset:
def __init__(self):
self.cont_type_ids = [1, 3]
self.n_types = 4
self.patients = [
self._patient([1, 3], [0.0, 10.0]),
self._patient([1, 3], [1.0, 10.0]),
self._patient([1, 3], [2.0, 10.0]),
self._patient([1, 3], [3.0, 10.0]),
self._patient([1, 3], [4.0, 10.0]),
self._patient([1, 3], [1000.0, 999.0]),
]
@staticmethod
def _patient(types, values):
return {
"other_type": np.asarray(types, dtype=np.int64),
"other_value": np.asarray(values, dtype=np.float32),
"other_value_kind": np.ones(len(types), dtype=np.int64),
}
def __len__(self):
return len(self.patients)
class _CaptureContinuousEncoder(nn.Module):
def __init__(self, n_embd):
super().__init__()
self.n_embd = n_embd
self.last_type = None
self.last_value = None
def forward(self, cont_type_idx, value):
self.last_type = cont_type_idx.detach().clone()
self.last_value = value.detach().clone()
return value[:, None].expand(-1, self.n_embd)
class ContinuousValueScalingTests(unittest.TestCase):
def test_fit_uses_only_training_subset_and_handles_constant_features(self):
dataset = _ToyAllFutureDataset()
train_subset = Subset(dataset, np.asarray([0, 1, 2, 3, 4]))
stats = fit_continuous_robust_scaler(dataset, train_subset)
self.assertEqual(stats.cont_type_ids, (1, 3))
np.testing.assert_array_equal(stats.observation_count, np.asarray([5, 5]))
np.testing.assert_allclose(stats.center, np.asarray([2.0, 10.0]))
np.testing.assert_allclose(stats.scale, np.asarray([2.0, 1.0]))
def test_tokenizer_standardizes_only_continuous_values(self):
tokenizer = OtherInfoTokenizer(
n_embd=4,
n_types=4,
n_cont_types=2,
n_categories=3,
cont_type_ids=[1, 3],
continuous_value_scaling="robust",
continuous_value_center=[10.0, 100.0],
continuous_value_scale=[2.0, 20.0],
)
capture = _CaptureContinuousEncoder(n_embd=4)
tokenizer.cont_value_encoder = capture
tokenizer(
other_type=torch.tensor([[1, 2, 3]], dtype=torch.long),
other_value=torch.tensor([[14.0, 1.0, 80.0]]),
other_value_kind=torch.tensor([[1, 2, 1]], dtype=torch.long),
)
torch.testing.assert_close(capture.last_type, torch.tensor([0, 1]))
torch.testing.assert_close(capture.last_value, torch.tensor([2.0, -1.0]))
def test_scaler_buffers_round_trip_in_new_checkpoint(self):
tokenizer = OtherInfoTokenizer(
n_embd=4,
n_types=4,
n_cont_types=2,
n_categories=2,
cont_type_ids=[1, 3],
continuous_value_scaling="robust",
continuous_value_center=[2.0, 10.0],
continuous_value_scale=[1.5, 4.0],
)
state = tokenizer.state_dict()
self.assertIn("continuous_value_center", state)
self.assertIn("continuous_value_scale", state)
restored = OtherInfoTokenizer(
n_embd=4,
n_types=4,
n_cont_types=2,
n_categories=2,
cont_type_ids=[1, 3],
continuous_value_scaling="robust",
)
restored.load_state_dict(state, strict=True)
torch.testing.assert_close(
restored.continuous_value_center,
torch.tensor([2.0, 10.0]),
)
torch.testing.assert_close(
restored.continuous_value_scale,
torch.tensor([1.5, 4.0]),
)
def test_legacy_none_mode_keeps_old_state_dict_schema(self):
tokenizer = OtherInfoTokenizer(
n_embd=4,
n_types=4,
n_cont_types=2,
n_categories=2,
cont_type_ids=[1, 3],
)
state = tokenizer.state_dict()
self.assertNotIn("continuous_value_center", state)
self.assertNotIn("continuous_value_scale", state)
restored = OtherInfoTokenizer(
n_embd=4,
n_types=4,
n_cont_types=2,
n_categories=2,
cont_type_ids=[1, 3],
)
restored.load_state_dict(state, strict=True)
if __name__ == "__main__":
unittest.main()

View File

@@ -39,9 +39,11 @@ from model_architectures import (
from models import DeepHealth from models import DeepHealth
from targets import CHECKUP_IDX, PAD_IDX from targets import CHECKUP_IDX, PAD_IDX
from train_util import ( from train_util import (
ContinuousRobustScalerStats,
configure_torch_for_training, configure_torch_for_training,
create_unique_run_dir, create_unique_run_dir,
format_extra_info_types, format_extra_info_types,
fit_continuous_robust_scaler,
get_lr, get_lr,
get_model_parameter_counts, get_model_parameter_counts,
load_extra_info_types_file, load_extra_info_types_file,
@@ -104,6 +106,16 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--n_head", type=int, default=10) parser.add_argument("--n_head", type=int, default=10)
parser.add_argument("--n_layer", type=int, default=12) parser.add_argument("--n_layer", type=int, default=12)
parser.add_argument("--n_bins", type=int, default=16) parser.add_argument("--n_bins", type=int, default=16)
parser.add_argument(
"--continuous_value_scaling",
type=str,
default="robust",
choices=["none", "robust"],
help=(
"Continuous extra-info scaling. 'robust' fits the median and IQR "
"on the complete training subset and stores them in the checkpoint."
),
)
parser.add_argument("--extra_pool_reduce", type=str, default="mean", parser.add_argument("--extra_pool_reduce", type=str, default="mean",
choices=["mean", "sum"]) choices=["mean", "sum"])
parser.add_argument("--time_mode", type=str, default="relative", parser.add_argument("--time_mode", type=str, default="relative",
@@ -173,7 +185,22 @@ def parse_args() -> argparse.Namespace:
return args return args
def build_model(args: argparse.Namespace, dataset: AllFutureHealthDataset) -> DeepHealth: def build_model(
args: argparse.Namespace,
dataset: AllFutureHealthDataset,
scaler_stats: ContinuousRobustScalerStats | None = None,
) -> DeepHealth:
if (
args.continuous_value_scaling == "robust"
and dataset.n_cont_types > 0
and scaler_stats is None
):
raise ValueError(
"Robust continuous-value scaling requires statistics fitted on the "
"training subset"
)
center = None if scaler_stats is None else scaler_stats.center
scale = None if scaler_stats is None else scaler_stats.scale
return DeepHealth( return DeepHealth(
vocab_size=dataset.vocab_size, vocab_size=dataset.vocab_size,
n_embd=args.n_embd, n_embd=args.n_embd,
@@ -184,6 +211,9 @@ def build_model(args: argparse.Namespace, dataset: AllFutureHealthDataset) -> De
n_categories=dataset.n_categories, n_categories=dataset.n_categories,
cont_type_ids=dataset.cont_type_ids, cont_type_ids=dataset.cont_type_ids,
n_bins=args.n_bins, n_bins=args.n_bins,
continuous_value_scaling=args.continuous_value_scaling,
continuous_value_center=center,
continuous_value_scale=scale,
extra_pool_reduce=args.extra_pool_reduce, extra_pool_reduce=args.extra_pool_reduce,
target_mode="all_future", target_mode="all_future",
time_mode=args.time_mode, time_mode=args.time_mode,
@@ -322,7 +352,16 @@ def build_metadata(
train_subset, train_subset,
val_subset, val_subset,
test_subset, test_subset,
scaler_stats: ContinuousRobustScalerStats | None,
) -> Dict[str, Any]: ) -> Dict[str, Any]:
scaler_metadata: Dict[str, Any]
if scaler_stats is None:
scaler_metadata = {
"method": "none",
"fitted_on": None,
}
else:
scaler_metadata = scaler_stats.as_metadata()
return { return {
"run_name": run_name, "run_name": run_name,
"dataset_class": "AllFutureHealthDataset", "dataset_class": "AllFutureHealthDataset",
@@ -342,6 +381,7 @@ def build_metadata(
else None else None
), ),
"extra_info_types": [int(x) for x in dataset.extra_info_types], "extra_info_types": [int(x) for x in dataset.extra_info_types],
"continuous_value_scaler": scaler_metadata,
"dataset_metadata": { "dataset_metadata": {
"vocab_size": int(dataset.vocab_size), "vocab_size": int(dataset.vocab_size),
"n_types": int(dataset.n_types), "n_types": int(dataset.n_types),
@@ -385,6 +425,7 @@ def main() -> None:
logger.info(f"Model architecture: {args.model_architecture}") logger.info(f"Model architecture: {args.model_architecture}")
logger.info(f"Disease history mode: {args.disease_history_mode}") logger.info(f"Disease history mode: {args.disease_history_mode}")
logger.info(f"extra_info_types: {format_extra_info_types(args.extra_info_types)}") logger.info(f"extra_info_types: {format_extra_info_types(args.extra_info_types)}")
logger.info(f"Continuous value scaling: {args.continuous_value_scaling}")
logger.info("Loading all-future datasets...") logger.info("Loading all-future datasets...")
train_dataset = AllFutureHealthDataset( train_dataset = AllFutureHealthDataset(
@@ -448,6 +489,23 @@ def main() -> None:
f"Patients/queries: train={len(train_subset)}, val={len(val_subset)}, test={len(test_subset)}" f"Patients/queries: train={len(train_subset)}, val={len(val_subset)}, test={len(test_subset)}"
) )
scaler_stats = None
if args.continuous_value_scaling == "robust" and train_dataset.n_cont_types > 0:
logger.info(
"Fitting continuous RobustScaler on the complete training subset: "
f"patients={len(train_subset):,}, features={train_dataset.n_cont_types}"
)
scaler_stats = fit_continuous_robust_scaler(
train_dataset,
train_subset,
)
logger.info(
"Continuous RobustScaler fitted: "
f"observations={int(scaler_stats.observation_count.sum()):,}, "
f"min_per_feature={int(scaler_stats.observation_count.min()):,}, "
f"max_per_feature={int(scaler_stats.observation_count.max()):,}"
)
train_loader = DataLoader( train_loader = DataLoader(
train_subset, train_subset,
batch_size=args.batch_size, batch_size=args.batch_size,
@@ -479,7 +537,7 @@ def main() -> None:
prefetch_factor=2 if args.num_workers > 0 else None, prefetch_factor=2 if args.num_workers > 0 else None,
) )
model = build_model(args, train_dataset).to(device) model = build_model(args, train_dataset, scaler_stats=scaler_stats).to(device)
parameter_counts = get_model_parameter_counts(model) parameter_counts = get_model_parameter_counts(model)
logger.info( logger.info(
"Model parameters: " "Model parameters: "
@@ -496,7 +554,13 @@ def main() -> None:
adaptive_lr = args.base_lr * math.sqrt(args.batch_size / 128) adaptive_lr = args.base_lr * math.sqrt(args.batch_size / 128)
train_metadata = build_metadata( train_metadata = build_metadata(
args, train_dataset, run_name, train_subset, val_subset, test_subset args,
train_dataset,
run_name,
train_subset,
val_subset,
test_subset,
scaler_stats,
) )
train_metadata.update(parameter_counts) train_metadata.update(parameter_counts)
save_config( save_config(

View File

@@ -10,6 +10,7 @@
# - smoking # - smoking
# - alcohol # - alcohol
# - BMI is already included in the assessment variables # - BMI is already included in the assessment variables
# - continuous values use train-split RobustScaler statistics
# #
# A6000 48 GB default: # A6000 48 GB default:
# batch_size=256 # batch_size=256
@@ -34,7 +35,7 @@ SEED_CSV="42,43,44"
NUM_WORKERS=4 NUM_WORKERS=4
BATCH_SIZE=256 BATCH_SIZE=256
PYTHON_BIN="${PYTHON_BIN:-python}" PYTHON_BIN="${PYTHON_BIN:-python}"
CAMPAIGN_NAME="extra_info_assessment_smoking_alcohol_multiseed" CAMPAIGN_NAME="extra_info_assessment_smoking_alcohol_robust_multiseed"
DRY_RUN=0 DRY_RUN=0
ENTRYPOINT="$SCRIPT_DIR/train_all_future.py" ENTRYPOINT="$SCRIPT_DIR/train_all_future.py"
@@ -66,6 +67,7 @@ Fixed experiment settings:
disease history timed disease history timed
sex enabled by the model sex enabled by the model
extra information extra_info_types_assessment_smoking_alcohol.txt extra information extra_info_types_assessment_smoking_alcohol.txt
continuous scaling robust (training-subset median/IQR)
A6000 48 GB default: A6000 48 GB default:
batch_size 256 batch_size 256
@@ -234,6 +236,7 @@ run_job() {
--time_mode relative --time_mode relative
--dist_mode weibull --dist_mode weibull
--disease_history_mode timed --disease_history_mode timed
--continuous_value_scaling robust
--extra_info_types_file "$EXTRA_INFO_TYPES_FILE" --extra_info_types_file "$EXTRA_INFO_TYPES_FILE"
) )

View File

@@ -5,6 +5,7 @@ import logging
import sys import sys
import time import time
import csv import csv
from dataclasses import dataclass
from datetime import datetime from datetime import datetime
import math import math
from pathlib import Path from pathlib import Path
@@ -19,6 +20,134 @@ from dataset import AllFutureHealthDataset, HealthDataset
from models import DeepHealth from models import DeepHealth
@dataclass(frozen=True)
class ContinuousRobustScalerStats:
"""Train-split robust scaling statistics aligned to ``cont_type_ids``."""
cont_type_ids: tuple[int, ...]
center: np.ndarray
scale: np.ndarray
observation_count: np.ndarray
quantile_range: tuple[float, float] = (25.0, 75.0)
def as_metadata(self) -> Dict[str, Any]:
return {
"method": "robust",
"fitted_on": "train_subset",
"quantile_range": [float(x) for x in self.quantile_range],
"cont_type_ids": [int(x) for x in self.cont_type_ids],
"observation_count": [int(x) for x in self.observation_count.tolist()],
"center_buffer": "tokenizer.continuous_value_center",
"scale_buffer": "tokenizer.continuous_value_scale",
}
def fit_continuous_robust_scaler(
dataset: AllFutureHealthDataset,
subset: Subset,
*,
quantile_range: tuple[float, float] = (25.0, 75.0),
scale_epsilon: float = 1e-6,
) -> ContinuousRobustScalerStats:
"""Fit median/IQR statistics using only patients in the training subset.
The prepared arrays remain unchanged. Scaling is performed later inside the
model tokenizer so the fitted center and scale can live in the checkpoint.
"""
if subset.dataset is not dataset:
raise ValueError("subset must reference the dataset used to fit the scaler")
low, high = (float(quantile_range[0]), float(quantile_range[1]))
if not (0.0 <= low < high <= 100.0):
raise ValueError(
"quantile_range must satisfy 0 <= low < high <= 100, got "
f"{quantile_range!r}"
)
if scale_epsilon <= 0:
raise ValueError("scale_epsilon must be > 0")
cont_type_ids = tuple(int(x) for x in dataset.cont_type_ids)
n_cont_types = len(cont_type_ids)
if n_cont_types == 0:
empty = np.zeros(0, dtype=np.float32)
return ContinuousRobustScalerStats(
cont_type_ids=cont_type_ids,
center=empty.copy(),
scale=empty.copy(),
observation_count=np.zeros(0, dtype=np.int64),
quantile_range=(low, high),
)
subset_indices = np.asarray(subset.indices, dtype=np.int64)
if subset_indices.ndim != 1 or subset_indices.size == 0:
raise ValueError("training subset must contain at least one patient")
type_to_column = np.full(int(dataset.n_types), -1, dtype=np.int64)
for column, type_id in enumerate(cont_type_ids):
if type_id <= 0 or type_id >= len(type_to_column):
raise ValueError(
f"continuous type id {type_id} is outside [1, {len(type_to_column)})"
)
type_to_column[type_id] = column
values = np.full(
(int(subset_indices.size), n_cont_types),
np.nan,
dtype=np.float32,
)
for row, patient_index in enumerate(subset_indices.tolist()):
patient = dataset.patients[int(patient_index)]
other_type = np.asarray(patient["other_type"], dtype=np.int64)
other_value = np.asarray(patient["other_value"], dtype=np.float32)
other_kind = np.asarray(patient["other_value_kind"], dtype=np.int64)
continuous = other_kind == 1
if not np.any(continuous):
continue
selected_type = other_type[continuous]
selected_value = other_value[continuous]
valid_type = (selected_type > 0) & (selected_type < len(type_to_column))
columns = np.full(selected_type.shape, -1, dtype=np.int64)
columns[valid_type] = type_to_column[selected_type[valid_type]]
valid = (columns >= 0) & np.isfinite(selected_value)
values[row, columns[valid]] = selected_value[valid]
observation_count = np.isfinite(values).sum(axis=0).astype(np.int64)
missing_types = [
type_id
for type_id, count in zip(cont_type_ids, observation_count.tolist())
if count == 0
]
if missing_types:
raise ValueError(
"Training subset has no finite observations for continuous type ids: "
f"{missing_types}"
)
low_value, center, high_value = np.nanpercentile(
values,
[low, 50.0, high],
axis=0,
)
scale = high_value - low_value
near_constant = (~np.isfinite(scale)) | (np.abs(scale) <= float(scale_epsilon))
scale[near_constant] = 1.0
center = np.asarray(center, dtype=np.float32)
scale = np.asarray(scale, dtype=np.float32)
if not np.isfinite(center).all():
raise RuntimeError("Robust scaler produced non-finite center values")
if not np.isfinite(scale).all() or np.any(scale <= 0):
raise RuntimeError("Robust scaler produced invalid scale values")
return ContinuousRobustScalerStats(
cont_type_ids=cont_type_ids,
center=center,
scale=scale,
observation_count=observation_count,
quantile_range=(low, high),
)
def create_unique_run_dir(name_fn, runs_root: Path = Path("runs")) -> tuple[Path, str]: def create_unique_run_dir(name_fn, runs_root: Path = Path("runs")) -> tuple[Path, str]:
while True: while True:
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S") timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")