import json
import os
import warnings
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Optional, Union
import numpy as np
import onnxruntime
import polars as pl
import torch
from beartype import beartype
from beartype.typing import Iterator
from loguru import logger
from sequifier.config.infer_config import InfererModel, load_inferer_config
from sequifier.config.train_config import load_train_config
from sequifier.helpers import (
PANDAS_TO_TORCH_TYPES,
configure_determinism,
configured_model_window_stride,
construct_index_maps,
normalize_path,
numpy_storage_to_pytorch,
numpy_to_pytorch,
resolve_unified_polars_numeric_dtype,
resolve_window_sampling_plan,
resolve_window_view,
subset_to_input_columns,
validate_stored_window_width,
write_data,
)
from sequifier.objectives import get_objective_class
from sequifier.special_tokens import (
ONNX_CATEGORICAL_TARGET_CODECS_KEY,
SPECIAL_TOKEN_IDS,
resolve_categorical_decoder_ids,
)
from sequifier.train import (
infer_with_embedding_model,
infer_with_generative_model,
load_inference_model,
)
ONNX_NUMPY_DTYPES = {
"tensor(float16)": np.float16,
"tensor(float)": np.float32,
"tensor(double)": np.float64,
"tensor(int8)": np.int8,
"tensor(int16)": np.int16,
"tensor(int32)": np.int32,
"tensor(int64)": np.int64,
"tensor(uint8)": np.uint8,
"tensor(uint16)": np.uint16,
"tensor(uint32)": np.uint32,
"tensor(uint64)": np.uint64,
"tensor(bool)": np.bool_,
}
[docs]@beartype
def load_onnx_target_decoder_ids(
session: onnxruntime.InferenceSession,
target_columns: list[str],
target_column_types: dict[str, str],
model_type: str,
) -> dict[str, list[int]]:
"""Load and validate categorical decoder-index mappings from ONNX metadata."""
raw_codecs = session.get_modelmeta().custom_metadata_map.get(
ONNX_CATEGORICAL_TARGET_CODECS_KEY
)
if raw_codecs is None:
raise ValueError(
"ONNX model is missing required metadata "
f"{ONNX_CATEGORICAL_TARGET_CODECS_KEY!r}."
)
try:
loaded_codecs = json.loads(raw_codecs)
except (TypeError, json.JSONDecodeError) as exc:
raise ValueError(
"ONNX categorical target codec metadata must be valid JSON."
) from exc
if not isinstance(loaded_codecs, dict):
raise ValueError("ONNX categorical target codec metadata must be an object.")
categorical_targets = {
column
for column in target_columns
if target_column_types[column] == "categorical"
}
if set(loaded_codecs) != categorical_targets:
raise ValueError(
"ONNX categorical target codecs must match the configured categorical "
f"targets: expected {sorted(categorical_targets)}, "
f"found {sorted(loaded_codecs)}."
)
target_decoder_ids: dict[str, list[int]] = {}
for column, decoder_ids in loaded_codecs.items():
if not isinstance(decoder_ids, list) or not decoder_ids:
raise ValueError(
f"ONNX categorical target codec for {column!r} must be a non-empty list."
)
if not all(
isinstance(decoder_id, int) and not isinstance(decoder_id, bool)
for decoder_id in decoder_ids
):
raise ValueError(
f"ONNX categorical target codec for {column!r} must contain integers."
)
if any(decoder_id < 0 for decoder_id in decoder_ids):
raise ValueError(
f"ONNX categorical target codec for {column!r} contains a negative ID."
)
if len(decoder_ids) != len(set(decoder_ids)):
raise ValueError(
f"ONNX categorical target codec for {column!r} contains duplicate IDs."
)
target_decoder_ids[column] = decoder_ids
if model_type == "generative":
outputs_by_name = {output.name: output for output in session.get_outputs()}
for column, decoder_ids in target_decoder_ids.items():
output = outputs_by_name.get(column)
if output is None:
output = outputs_by_name.get(f"{column}_out")
if output is None:
raise ValueError(
f"ONNX model has no output for categorical target {column!r}."
)
output_width = output.shape[-1]
if isinstance(output_width, int) and output_width != len(decoder_ids):
raise ValueError(
f"ONNX output for {column!r} has width {output_width}, but its "
f"categorical target codec contains {len(decoder_ids)} IDs."
)
return target_decoder_ids
[docs]@beartype
def infer(args: Any, args_config: dict[str, Any]) -> None:
"""Load inference config and dispatch the worker."""
logger.info("--- Starting Inference ---")
config_path = (
args.config_path if args.config_path is not None else "configs/infer.yaml"
)
skip_metadata = args_config.get("skip_metadata", False)
config = load_inferer_config(config_path, args_config, skip_metadata)
if config.map_to_id or (len(config.real_columns) > 0):
metadata = config.dataset_metadata
if metadata is None:
raise ValueError(
"Resolved inference metadata is required for ID mapping and "
"real-column normalization."
)
id_maps = metadata.id_maps
selected_columns_statistics = metadata.selected_columns_statistics
normalize_real_columns = metadata.normalize_real_columns
else:
id_maps = None
selected_columns_statistics = {}
normalize_real_columns = True
configure_determinism(config.seed, config.enforce_deterministic_inference)
infer_worker(
config,
args_config,
id_maps,
selected_columns_statistics,
(0.0, 100.0),
normalize_real_columns=normalize_real_columns,
)
[docs]@beartype
def load_pt_dataset(data_path: str, start_pct: float, end_pct: float) -> Iterator[Any]:
"""Yield a percentage slice of sorted top-level PT files."""
pt_files = sorted(Path(data_path).glob("*.pt"))
total = len(pt_files)
start_idx = int(total * start_pct / 100)
end_idx = int(total * end_pct / 100)
for pt_file in pt_files[start_idx:end_idx]:
yield torch.load(pt_file, weights_only=False)
[docs]@beartype
def load_parquet_folder_dataset(
data_path: str, start_pct: float, end_pct: float
) -> Iterator[Any]:
"""Yield a percentage slice of sorted top-level Parquet files."""
parquet_files = sorted(Path(data_path).glob("*.parquet"))
total = len(parquet_files)
start_idx = int(total * start_pct / 100)
end_idx = int(total * end_pct / 100)
for parquet_file in parquet_files[start_idx:end_idx]:
yield pl.read_parquet(parquet_file)
@beartype
def _torch_column_types(config: InfererModel) -> dict[str, torch.dtype]:
return {
col: PANDAS_TO_TORCH_TYPES[config.column_data_types[col]]
for col in config.column_data_types
}
@beartype
def _sequence_position_columns(config: InfererModel, data: pl.DataFrame) -> list[str]:
return [
str(i)
for i in range(config.storage_layout.stored_context_width - 1, -1, -1)
if str(i) in data.columns
]
@beartype
def _configured_types_for_loaded_rows(
config: InfererModel, data: pl.DataFrame
) -> dict[str, str]:
if "inputCol" not in data.columns:
return {
column: config.column_data_types[column]
for column in config.input_columns
if column in config.column_data_types
}
loaded_columns = [
column
for column in data.get_column("inputCol").unique().to_list()
if column in config.column_data_types
]
return {column: config.column_data_types[column] for column in loaded_columns}
[docs]@beartype
def apply_inference_column_types(
data: pl.DataFrame, config: InfererModel
) -> pl.DataFrame:
"""Cast loaded long-format sequence values to the configured unified dtype."""
sequence_columns = _sequence_position_columns(config, data)
if not sequence_columns:
return data
configured_types = _configured_types_for_loaded_rows(config, data)
if not configured_types:
return data
unified_dtype = resolve_unified_polars_numeric_dtype(configured_types)
casts = [
pl.col(column).cast(unified_dtype)
for column in sequence_columns
if data.schema[column] != unified_dtype
]
if not casts:
return data
return data.with_columns(casts)
[docs]@beartype
def apply_inference_tensor_types(
sequences_dict: dict[str, torch.Tensor],
column_data_types: dict[str, torch.dtype],
) -> dict[str, torch.Tensor]:
"""Cast loaded PT feature tensors to the configured per-column dtype."""
return {
column: tensor.to(dtype=column_data_types[column])
if column in column_data_types and tensor.dtype != column_data_types[column]
else tensor
for column, tensor in sequences_dict.items()
}
[docs]@dataclass(frozen=True)
class WindowedInferenceBatch:
"""Model-facing windows plus physical identities and adjusted starts."""
inputs: dict[str, torch.Tensor]
metadata: dict[str, torch.Tensor]
sequence_ids: torch.Tensor
subsequence_ids: torch.Tensor
model_start_positions: torch.Tensor
window_start_offsets: torch.Tensor
@beartype
def _windowed_inference_batch_from_storage(
config: InfererModel,
sequences: dict[str, torch.Tensor],
sequence_ids: torch.Tensor,
subsequence_ids: torch.Tensor,
start_positions: torch.Tensor,
left_pad_lengths: torch.Tensor,
) -> WindowedInferenceBatch:
plan = resolve_window_sampling_plan(
config.storage_layout,
config.window_view,
configured_model_window_stride(config),
)
sample_index = plan.build_index(left_pad_lengths)
if len(sample_index) == 0:
raise ValueError("No usable model windows were found for inference.")
logical_indices = torch.arange(len(sample_index), dtype=torch.int64)
stored_rows, input_starts = sample_index.resolve(logical_indices)
inputs = {
column: plan.gather(
sequences[column],
stored_rows,
input_starts,
)
for column in config.input_columns
}
metadata = plan.build_masks(
left_pad_lengths[stored_rows],
input_starts,
)
return WindowedInferenceBatch(
inputs=inputs,
metadata=metadata,
sequence_ids=sequence_ids[stored_rows],
subsequence_ids=subsequence_ids[stored_rows],
model_start_positions=start_positions[stored_rows] + input_starts,
window_start_offsets=input_starts,
)
@beartype
def _windowed_inference_batch_from_dataframe(
config: InfererModel,
data: pl.DataFrame,
column_data_types: dict[str, torch.dtype],
) -> WindowedInferenceBatch:
if config.input_columns is not None:
subset = subset_to_input_columns(data, config.input_columns)
if not isinstance(subset, pl.DataFrame):
raise TypeError("Expected eager preprocessed inference data")
data = subset
data = apply_inference_column_types(data, config)
sequences, left_pad_lengths = numpy_storage_to_pytorch(
data,
column_data_types,
config.input_columns,
config.storage_layout.stored_context_width,
sort_rows=False,
)
identities = data.group_by(
["sequenceId", "subsequenceId"], maintain_order=True
).agg(pl.col("startItemPosition").first().alias("startItemPosition"))
return _windowed_inference_batch_from_storage(
config,
sequences,
torch.tensor(
identities.get_column("sequenceId").to_numpy(),
dtype=torch.int64,
),
torch.tensor(
identities.get_column("subsequenceId").to_numpy(),
dtype=torch.int64,
),
torch.tensor(
identities.get_column("startItemPosition").to_numpy(),
dtype=torch.int64,
),
left_pad_lengths,
)
@beartype
def _windowed_inference_batch_from_pt(
config: InfererModel,
data: tuple,
column_data_types: dict[str, torch.dtype],
) -> WindowedInferenceBatch:
(
sequences,
sequence_ids,
subsequence_ids,
start_positions,
left_pad_lengths,
) = data
sequences = apply_inference_tensor_types(sequences, column_data_types)
for tensor in sequences.values():
validate_stored_window_width(
tensor,
config.storage_layout.stored_context_width,
)
return _windowed_inference_batch_from_storage(
config,
sequences,
sequence_ids,
subsequence_ids,
start_positions,
left_pad_lengths,
)
@beartype
def _windowed_inference_batch(
config: InfererModel,
data: Any,
column_data_types: dict[str, torch.dtype],
) -> WindowedInferenceBatch:
if isinstance(data, pl.DataFrame):
return _windowed_inference_batch_from_dataframe(
config,
data,
column_data_types,
)
if isinstance(data, tuple):
return _windowed_inference_batch_from_pt(
config,
data,
column_data_types,
)
raise TypeError(f"Unsupported preprocessed inference chunk: {type(data).__name__}")
[docs]@beartype
def infer_worker(
config: Any,
args_config: dict[str, Any],
id_maps: Optional[dict[str, dict[str | int, int]]],
selected_columns_statistics: dict[str, dict[str, float]],
percentage_limits: Optional[tuple[float, float]],
normalize_real_columns: bool,
):
"""Load data, instantiate models, and run the configured inference mode."""
logger.info(f"[INFO] Reading data from '{config.data_path}'...")
is_folder_input = os.path.isdir(
normalize_path(config.data_path, config.project_root)
)
dataset = None
if not is_folder_input:
# Standalone Single-File Path Execution
if config.read_format == "parquet":
dataset = [pl.read_parquet(config.data_path)]
elif config.read_format == "csv":
dataset = [pl.read_csv(config.data_path)]
model_paths = (
config.model_path
if isinstance(config.model_path, list)
else [config.model_path]
)
for model_path in model_paths:
target_decoder_ids = None
if model_path.lower().endswith(".pt"):
if config.training_config_path is None:
raise ValueError("training_config_path is required for PyTorch models")
training_config = load_train_config(
config.training_config_path,
{
key: value
for key, value in args_config.items()
if key not in ["model_path", "data_path"]
},
args_config.get("skip_metadata", False),
)
target_column_types = training_config.target_column_types
if target_column_types is None:
raise ValueError("target_column_types must be provided or derived")
target_decoder_ids = resolve_categorical_decoder_ids(
training_config.target_columns,
target_column_types,
training_config.n_classes,
training_config.categorical_decoder_special_tokens,
)
if is_folder_input:
if percentage_limits is None:
raise ValueError(
"percentage_limits must be provided for folder-based read formats"
)
start_pct, end_pct = percentage_limits
# Direct folders to their respective lazy loaders based on file format
if config.read_format == "pt":
dataset = load_pt_dataset(config.data_path, start_pct, end_pct)
elif config.read_format == "parquet":
dataset = load_parquet_folder_dataset(
config.data_path, start_pct, end_pct
)
if dataset is None:
raise Exception(
f"Unsupported input type or read format: {config.read_format}"
)
objective_class = get_objective_class(config.training_objective)
prediction_length = (
config.prediction_length
if config.prediction_length is not None
else objective_class.default_prediction_length(
config.window_view.context_length
)
)
inferer = Inferer(
config.model_type,
model_path,
config.project_root,
id_maps,
selected_columns_statistics,
config.map_to_id,
config.categorical_columns,
config.real_columns,
config.input_columns,
config.target_columns,
config.target_column_types,
config.sample_from_distribution_columns,
config.infer_with_dropout,
prediction_length,
config.inference_batch_size,
config.device,
args_config=args_config,
training_config_path=config.training_config_path,
training_objective=config.training_objective,
normalize_real_columns=normalize_real_columns,
target_decoder_ids=target_decoder_ids,
)
column_data_types = _torch_column_types(config)
model_id = os.path.split(model_path)[1].replace(
f".{inferer.inference_model_type}", ""
)
logger.info(f"[INFO] Inferring for {model_id}")
if config.model_type == "generative":
infer_generative(config, inferer, model_id, dataset, column_data_types)
if config.model_type == "embedding":
infer_embedding(config, inferer, model_id, dataset, column_data_types)
logger.info("--- Inference Complete ---")
[docs]def calculate_item_positions(
start_positions: np.ndarray,
context_length: int,
prediction_length: int,
training_objective: str,
target_offset: int = 1,
) -> np.ndarray:
"""Return flattened absolute item positions for inference outputs."""
objective_class = get_objective_class(training_objective)
positions = objective_class.item_positions(
start_positions,
context_length,
prediction_length,
)
if objective_class.forward_looking:
positions = positions + target_offset - 1
return positions
@beartype
def _flatten_valid_mask(
config: InfererModel,
metadata: dict[str, Any],
prediction_length: int,
mask_key: str = "target_valid_mask",
) -> np.ndarray:
valid_mask = metadata.get(mask_key)
if isinstance(valid_mask, torch.Tensor):
valid_mask = valid_mask.detach().cpu().numpy() # type: ignore
valid_mask = np.asarray(valid_mask, dtype=bool)
if valid_mask.ndim != 2:
raise ValueError(f"{mask_key} must be 2D, got shape {valid_mask.shape}.")
return valid_mask[:, -prediction_length:].reshape(-1)
@beartype
def _bert_reference_column(config: InfererModel, data_columns: set[str]) -> str:
preferred_columns = (
[col for col in config.target_columns if col in config.categorical_columns]
+ [col for col in config.input_columns if col in config.categorical_columns]
+ [col for col in config.target_columns if col in data_columns]
+ [col for col in config.input_columns if col in data_columns]
)
for column_name in preferred_columns:
if column_name in data_columns:
return column_name
raise ValueError("Could not find a reference column for BERT padding metadata.")
@beartype
def _valid_mask_from_preprocessed_data(
config: InfererModel,
data: pl.DataFrame,
prediction_length: int,
mask_key: str = "target_valid_mask",
) -> np.ndarray:
data_columns = set(data.get_column("inputCol").unique())
column_name = _bert_reference_column(config, data_columns)
reference_rows = data.filter(pl.col("inputCol") == column_name)
left_pad_lengths = torch.tensor(
reference_rows.get_column("leftPadLength").to_numpy(), dtype=torch.int64
)
resolved_view = resolve_window_view(config.storage_layout, config.window_view)
metadata = resolved_view.build_masks(left_pad_lengths)
return _flatten_valid_mask(config, metadata, prediction_length, mask_key)
@beartype
def _apply_valid_prediction_mask(
values: np.ndarray,
valid_prediction_mask: np.ndarray,
label: str,
) -> np.ndarray:
values = np.asarray(values)
if values.shape[0] != valid_prediction_mask.shape[0]:
raise ValueError(
f"{label} has {values.shape[0]} rows, but valid_prediction_mask has "
f"{valid_prediction_mask.shape[0]} rows."
)
return values[valid_prediction_mask]
@beartype
def _apply_valid_prediction_mask_to_dict(
values: Optional[dict[str, np.ndarray]],
valid_prediction_mask: np.ndarray,
label: str,
) -> Optional[dict[str, np.ndarray]]:
if values is None:
return None
return {
key: _apply_valid_prediction_mask(
value, valid_prediction_mask, f"{label}.{key}"
)
for key, value in values.items()
}
@beartype
def _autoregression_seed_dataframe(
config: InfererModel,
data: pl.DataFrame,
) -> pl.DataFrame:
"""Keep the first physical subsequence for each autoregressive sequence."""
verify_variable_order(data)
selected = subset_to_input_columns(data, config.input_columns)
if not isinstance(selected, pl.DataFrame):
raise TypeError("Expected eager preprocessed autoregression data")
seed_data = selected.filter(
pl.col("subsequenceId") == pl.col("subsequenceId").first().over("sequenceId")
)
expected_columns = set(config.input_columns)
for sequence_id, sequence_data in seed_data.group_by(
"sequenceId",
maintain_order=True,
):
found_columns = set(sequence_data.get_column("inputCol").to_list())
if found_columns != expected_columns:
raise ValueError(
"The first autoregression subsequence must contain every input "
f"column exactly once for sequenceId={sequence_id!r}; expected "
f"{sorted(expected_columns)}, found {sorted(found_columns)}."
)
if sequence_data.height != len(expected_columns):
raise ValueError(
"The first autoregression subsequence contains duplicate input "
f"rows for sequenceId={sequence_id!r}."
)
return seed_data
[docs]@beartype
def infer_embedding(
config: "InfererModel",
inferer: "Inferer",
model_id: str,
dataset: Union[list[Any], Iterator[Any]],
column_data_types: dict[str, torch.dtype],
) -> None:
"""Write embeddings for each dataset chunk."""
data_path = config.data_path
if data_path is None:
raise ValueError("data_path must be provided or resolved from metadata")
for data_id, data in enumerate(dataset):
prediction_length = inferer.prediction_length
is_folder_input = os.path.isdir(normalize_path(data_path, config.project_root))
windowed = _windowed_inference_batch(config, data, column_data_types)
if (
isinstance(data, pl.DataFrame)
and configured_model_window_stride(config) is None
):
embeddings = get_embeddings(config, inferer, data, column_data_types)
else:
embeddings = inferer.infer_embedding(
{key: value.numpy() for key, value in windowed.inputs.items()},
metadata={
key: value.numpy() for key, value in windowed.metadata.items()
},
column_data_types=column_data_types,
)
valid_prediction_mask = _flatten_valid_mask(
config,
windowed.metadata,
prediction_length,
mask_key="attention_valid_mask",
)
base_offsets = np.arange(
config.window_view.context_length - prediction_length,
config.window_view.context_length,
)
item_positions_for_preds_base = windowed.model_start_positions.numpy()
base_positions_repeated = np.repeat(
item_positions_for_preds_base, prediction_length
)
final_positions = base_positions_repeated + np.tile(
base_offsets,
len(item_positions_for_preds_base),
)
sequence_ids_repeated = np.repeat(
windowed.sequence_ids.numpy(),
prediction_length,
)
subsequence_ids_repeated = np.repeat(
windowed.subsequence_ids.numpy(),
prediction_length,
)
window_offsets_repeated = np.repeat(
windowed.window_start_offsets.numpy(),
prediction_length,
)
embeddings = _apply_valid_prediction_mask(
embeddings, valid_prediction_mask, "embeddings"
)
final_positions = _apply_valid_prediction_mask(
final_positions, valid_prediction_mask, "itemPosition"
)
sequence_ids_repeated = _apply_valid_prediction_mask(
sequence_ids_repeated, valid_prediction_mask, "sequenceId"
)
subsequence_ids_repeated = _apply_valid_prediction_mask(
subsequence_ids_repeated, valid_prediction_mask, "subsequenceId"
)
window_offsets_repeated = _apply_valid_prediction_mask(
window_offsets_repeated,
valid_prediction_mask,
"windowStartOffset",
)
embeddings_df = pl.DataFrame(
{
"sequenceId": sequence_ids_repeated,
"subsequenceId": subsequence_ids_repeated,
"windowStartOffset": window_offsets_repeated,
"itemPosition": final_positions,
**dict(
zip(
[str(v) for v in range(embeddings.shape[1])],
[embeddings[:, i] for i in range(embeddings.shape[1])],
)
),
}
)
os.makedirs(
os.path.join(config.project_root, "outputs", "embeddings"),
exist_ok=True,
)
if not is_folder_input:
file_name = f"{model_id}-embeddings.{config.write_format}"
else:
dirname = f"{model_id}-embeddings"
file_name = os.path.join(
dirname,
f"{model_id}-{data_id}-embeddings.{config.write_format}",
)
dir_path = os.path.join(
config.project_root, "outputs", "embeddings", dirname
)
os.makedirs(dir_path, exist_ok=True)
embeddings_path = os.path.join(
config.project_root, "outputs", "embeddings", file_name
)
logger.info(f"[INFO] Writing predictions to '{embeddings_path}'")
write_data(
embeddings_df,
embeddings_path,
config.write_format,
)
[docs]def infer_generative(
config: "InfererModel",
inferer: "Inferer",
model_id: str,
dataset: Union[list[Any], Iterator[Any]],
column_data_types: dict[str, torch.dtype],
):
"""Write generative predictions/probabilities for each dataset chunk."""
data_path = config.data_path
if data_path is None:
raise ValueError("data_path must be provided or resolved from metadata")
for data_id, data in enumerate(dataset):
is_folder_input = os.path.isdir(normalize_path(data_path, config.project_root))
if config.autoregression and isinstance(data, pl.DataFrame):
data = _autoregression_seed_dataframe(config, data)
windowed = _windowed_inference_batch(config, data, column_data_types)
if config.autoregression and inferer.prediction_length != 1:
raise ValueError(
"prediction_length must be 1 for autoregression, "
f"got {inferer.prediction_length}"
)
total_steps = (
config.autoregression_total_steps
if config.autoregression and config.autoregression_total_steps is not None
else 1
)
if (
isinstance(data, pl.DataFrame)
and total_steps == 1
and configured_model_window_stride(config) is None
):
probs, preds = get_probs_preds_from_df(
config,
inferer,
data,
column_data_types,
)
else:
probs, preds = get_probs_preds_from_dict(
config,
inferer,
windowed.inputs,
windowed.metadata,
column_data_types,
total_steps,
)
if total_steps == 1:
output_count_per_window = inferer.prediction_length
valid_prediction_mask = _flatten_valid_mask(
config,
windowed.metadata,
inferer.prediction_length,
)
item_positions_for_preds = calculate_item_positions(
windowed.model_start_positions.numpy(),
config.window_view.context_length,
inferer.prediction_length,
config.training_objective,
config.window_view.target_offset,
)
else:
output_count_per_window = total_steps
valid_prediction_mask = np.repeat(
_flatten_valid_mask(config, windowed.metadata, 1),
total_steps,
)
first_positions = (
windowed.model_start_positions.numpy()
+ config.window_view.context_length
+ config.window_view.target_offset
- 1
)
item_positions_for_preds = np.concatenate(
[np.arange(start, start + total_steps) for start in first_positions]
)
sequence_ids_for_preds = np.repeat(
windowed.sequence_ids.numpy(),
output_count_per_window,
)
subsequence_ids_for_preds = np.repeat(
windowed.subsequence_ids.numpy(),
output_count_per_window,
)
window_offsets_for_preds = np.repeat(
windowed.window_start_offsets.numpy(),
output_count_per_window,
)
if inferer.map_to_id:
for target_column, predictions in preds.items():
if target_column in inferer.index_map:
preds[target_column] = np.array(
[inferer.index_map[target_column][i] for i in predictions]
)
for target_column, predictions in preds.items():
if inferer.target_column_types[target_column] == "real":
preds[target_column] = inferer.invert_normalization(
predictions, target_column
)
assert valid_prediction_mask is not None
sequence_ids_for_preds = _apply_valid_prediction_mask(
sequence_ids_for_preds, valid_prediction_mask, "sequenceId"
)
subsequence_ids_for_preds = _apply_valid_prediction_mask(
subsequence_ids_for_preds,
valid_prediction_mask,
"subsequenceId",
)
window_offsets_for_preds = _apply_valid_prediction_mask(
window_offsets_for_preds,
valid_prediction_mask,
"windowStartOffset",
)
item_positions_for_preds = _apply_valid_prediction_mask(
item_positions_for_preds, valid_prediction_mask, "itemPosition"
)
preds = _apply_valid_prediction_mask_to_dict(
preds, valid_prediction_mask, "preds"
)
probs = _apply_valid_prediction_mask_to_dict(
probs, valid_prediction_mask, "probs"
)
os.makedirs(
os.path.join(config.project_root, "outputs", "predictions"),
exist_ok=True,
)
if config.output_probabilities:
assert probs is not None
os.makedirs(
os.path.join(config.project_root, "outputs", "probabilities"),
exist_ok=True,
)
for target_column in inferer.target_columns:
if not is_folder_input:
file_name = f"{model_id}-{target_column}-probabilities.{config.write_format}"
else:
dirname = f"{model_id}-{target_column}-probabilities"
file_name = os.path.join(
dirname,
f"{model_id}-{data_id}-probabilities.{config.write_format}",
)
dir_path = os.path.join(
config.project_root, "outputs", "probabilities", dirname
)
os.makedirs(dir_path, exist_ok=True)
if inferer.target_column_types[target_column] == "categorical":
probabilities_path = os.path.join(
config.project_root, "outputs", "probabilities", file_name
)
logger.info(
f"[INFO] Writing probabilities to '{probabilities_path}'"
)
# Step 5: Finalize Output and I/O (write_data now handles Polars DF)
write_data(
pl.DataFrame(
probs[target_column],
schema=[
str(inferer.index_map[target_column][global_id])
for global_id in inferer.target_decoder_ids.get(
target_column,
range(probs[target_column].shape[1]),
)
],
),
probabilities_path,
config.write_format,
)
assert preds is not None
predictions = pl.DataFrame(
{
"sequenceId": sequence_ids_for_preds,
"subsequenceId": subsequence_ids_for_preds,
"windowStartOffset": window_offsets_for_preds,
"itemPosition": item_positions_for_preds,
**{
target_column: preds[target_column].flatten()
for target_column in inferer.target_columns
},
}
)
if not is_folder_input:
file_name = f"{model_id}-predictions.{config.write_format}"
else:
dirname = f"{model_id}-predictions"
file_name = os.path.join(
dirname, f"{model_id}-{data_id}-predictions.{config.write_format}"
)
dir_path = os.path.join(
config.project_root, "outputs", "predictions", dirname
)
os.makedirs(dir_path, exist_ok=True)
predictions_path = os.path.join(
config.project_root, "outputs", "predictions", file_name
)
logger.info(f"[INFO] Writing predictions to '{predictions_path}'")
write_data(
predictions,
predictions_path,
config.write_format,
)
[docs]@beartype
def get_embeddings_pt(
config: Any,
inferer: "Inferer",
data: dict[str, torch.Tensor],
metadata: dict[str, torch.Tensor],
column_data_types: dict[str, torch.dtype],
) -> np.ndarray:
"""Infer embeddings from PT tensors."""
resolved_view = resolve_window_view(config.storage_layout, config.window_view)
for tensor in data.values():
validate_stored_window_width(tensor, config.storage_layout.stored_context_width)
X = {
key: val[:, resolved_view.input_slice].numpy()
for key, val in data.items()
if key in config.input_columns
}
metadata_np = {key: val.numpy() for key, val in metadata.items()}
embeddings = inferer.infer_embedding(
X, metadata=metadata_np, column_data_types=column_data_types
)
return embeddings
[docs]@beartype
def get_probs_preds_from_dict(
config: Any,
inferer: "Inferer",
data: dict[str, torch.Tensor],
metadata: dict[str, torch.Tensor],
column_data_types: dict[str, torch.dtype],
total_steps: int = 1,
) -> tuple[Optional[dict[str, np.ndarray]], dict[str, np.ndarray]]:
"""Infer PT predictions, flattened sample-major across autoregressive steps."""
target_cols = inferer.target_columns
X = {
key: tensor.numpy()
for key, tensor in data.items()
if key in config.input_columns
}
metadata_np = {key: tensor.numpy() for key, tensor in metadata.items()}
all_probs_list = {col: [] for col in target_cols}
all_preds_list = {col: [] for col in target_cols}
metadata_for_step = metadata_np
for i in range(total_steps):
if config.output_probabilities:
probs_for_step = inferer.infer_generative(
X,
metadata_for_step,
column_data_types=column_data_types,
return_probs=True,
)
preds_for_step = inferer.infer_generative(
None, metadata_for_step, probs_for_step
)
for col in target_cols:
all_probs_list[col].append(probs_for_step[col])
else:
preds_for_step = inferer.infer_generative(
X,
metadata_for_step,
column_data_types=column_data_types,
return_probs=False,
)
for col in target_cols:
all_preds_list[col].append(preds_for_step[col])
if i == (total_steps - 1):
break
X_next = {}
for col in X.keys():
shifted_input = X[col][:, 1:]
new_value = preds_for_step[col].reshape(-1, 1).astype(shifted_input.dtype)
X_next[col] = np.concatenate([shifted_input, new_value], axis=1)
X = X_next
if (
metadata_for_step is not None
and "attention_valid_mask" in metadata_for_step
):
shifted_mask = metadata_for_step["attention_valid_mask"][:, 1:]
appended_valid = np.ones(
(shifted_mask.shape[0], 1), dtype=shifted_mask.dtype
)
metadata_for_step["attention_valid_mask"] = np.concatenate(
[shifted_mask, appended_valid], axis=1
)
final_preds = {
col: np.array(preds_list).T.reshape(-1, 1).flatten()
for col, preds_list in all_preds_list.items()
}
if config.output_probabilities:
final_probs = {
col: np.array(probs_list)
.transpose((1, 0, 2))
.reshape(-1, probs_list[0].shape[1])
for col, probs_list in all_probs_list.items()
}
else:
final_probs = None
return (final_probs, final_preds)
[docs]@beartype
def get_embeddings(
config: Any,
inferer: "Inferer",
data: pl.DataFrame,
column_data_types: dict[str, torch.dtype],
) -> np.ndarray:
"""Infer embeddings from a Polars chunk."""
windowed = _windowed_inference_batch_from_dataframe(
config,
data,
column_data_types,
)
embeddings = inferer.infer_embedding(
{key: value.numpy() for key, value in windowed.inputs.items()},
metadata={key: value.numpy() for key, value in windowed.metadata.items()},
column_data_types=column_data_types,
)
return embeddings
[docs]@beartype
def get_probs_preds_from_df(
config: Any,
inferer: "Inferer",
data: pl.DataFrame,
column_data_types: dict[str, torch.dtype],
) -> tuple[Optional[dict[str, np.ndarray]], dict[str, np.ndarray]]:
"""Infer non-autoregressive predictions from a Polars chunk."""
windowed = _windowed_inference_batch_from_dataframe(
config,
data,
column_data_types,
)
X = {key: value.numpy() for key, value in windowed.inputs.items()}
metadata_np = {key: value.numpy() for key, value in windowed.metadata.items()}
if config.output_probabilities:
probs = inferer.infer_generative(
X, metadata_np, column_data_types=column_data_types, return_probs=True
)
preds = inferer.infer_generative(None, metadata_np, probs)
else:
probs = None
preds = inferer.infer_generative(
X, metadata_np, column_data_types=column_data_types
)
return (probs, preds)
[docs]@beartype
def fill_number(number: Union[int, float], max_length: int) -> str:
"""Left-pad a number for sortable string keys."""
number_str = str(number)
return f"{'0' * (max_length - len(number_str))}{number_str}"
[docs]@beartype
def verify_variable_order(data: pl.DataFrame) -> None:
"""Require sequenceId order and in-sequence subsequenceId order."""
is_globally_sorted = data.select(
(pl.col("sequenceId").diff().fill_null(0) >= 0).all()
).item()
if not is_globally_sorted:
raise ValueError("sequenceId must be in ascending order for autoregression")
is_group_sorted = (
data.select(
(pl.col("subsequenceId").diff().fill_null(0) >= 0)
.all()
.over("sequenceId")
.alias("is_sorted")
)
.get_column("is_sorted")
.all()
)
if not is_group_sorted:
raise ValueError("subsequenceId must be sorted within sequenceId groups")
[docs]@beartype
def get_probs_preds_autoregression(
config: Any,
inferer: "Inferer",
data: pl.DataFrame,
column_data_types: dict[str, torch.dtype],
context_length: int,
) -> tuple[
Optional[dict[str, np.ndarray]],
dict[str, np.ndarray],
np.ndarray,
np.ndarray,
np.ndarray,
]:
"""Infer autoregressive predictions with sequence IDs, positions, and mask."""
verify_variable_order(data)
distinct_cols = len(np.unique(data["inputCol"].to_numpy()))
head_data_df = data.group_by("sequenceId", maintain_order=True).head(distinct_cols)
aligned_sequence_ids = (
head_data_df.get_column("sequenceId").unique(maintain_order=True).to_numpy()
)
resolved_view = resolve_window_view(config.storage_layout, config.window_view)
input_start = resolved_view.input_slice.start
if input_start is None:
raise ValueError("Resolved input slice must have a concrete start")
aligned_start_positions = (
head_data_df.group_by("sequenceId", maintain_order=True)
.agg(pl.col("startItemPosition").max())
.get_column("startItemPosition")
.to_numpy()
+ input_start
+ context_length
+ config.window_view.target_offset
- 1
)
head_data, metadata = numpy_to_pytorch(
head_data_df,
column_data_types,
config.input_columns,
resolved_view,
)
probs, preds = get_probs_preds_from_dict(
config,
inferer,
head_data,
metadata,
column_data_types,
config.autoregression_total_steps,
)
item_positions_for_preds = np.concatenate(
[
np.arange(start_pos, start_pos + config.autoregression_total_steps)
for start_pos in aligned_start_positions
],
axis=0,
)
sequence_ids_for_preds = np.repeat(
aligned_sequence_ids, config.autoregression_total_steps
)
base_mask = _flatten_valid_mask(config, metadata, 1)
valid_prediction_mask = np.repeat(base_mask, config.autoregression_total_steps)
return (
probs,
preds,
sequence_ids_for_preds,
item_positions_for_preds,
valid_prediction_mask,
)
[docs]class Inferer:
"""Inference runtime for PT/ONNX sequifier models."""
[docs] @beartype
def __init__(
self,
model_type: str,
model_path: str,
project_root: str,
id_maps: Optional[dict[str, dict[Union[str, int], int]]],
selected_columns_statistics: dict[str, dict[str, float]],
map_to_id: bool,
categorical_columns: list[str],
real_columns: list[str],
input_columns: Optional[list[str]],
target_columns: list[str],
target_column_types: dict[str, str],
sample_from_distribution_columns: Optional[list[str]],
infer_with_dropout: bool,
prediction_length: int,
inference_batch_size: int,
device: str,
args_config: dict[str, Any],
training_config_path: Optional[str],
training_objective: Optional[str] = None,
normalize_real_columns: bool = True,
target_decoder_ids: Optional[dict[str, list[int]]] = None,
):
"""Load a PT or ONNX backend and postprocessing state."""
self.model_type = model_type
self.training_objective = training_objective
self.map_to_id = map_to_id
self.selected_columns_statistics = selected_columns_statistics
self.normalize_real_columns = normalize_real_columns
target_columns_index_map = [
c for c in target_columns if target_column_types[c] == "categorical"
]
self.index_map = construct_index_maps(
id_maps, target_columns_index_map, map_to_id
)
self.device = device
self.categorical_columns = categorical_columns
self.real_columns = real_columns
self.input_columns = input_columns
self.target_columns = target_columns
self.target_column_types = target_column_types
self.sample_from_distribution_columns = sample_from_distribution_columns
self.infer_with_dropout = infer_with_dropout
self.prediction_length = prediction_length
self.inference_batch_size = inference_batch_size
self.inference_model_type = model_path.split(".")[-1]
self.args_config = args_config
self.training_config_path = training_config_path
self.target_decoder_ids = target_decoder_ids or {}
if self.inference_model_type == "onnx":
execution_providers = [
"CUDAExecutionProvider" if device == "cuda" else "CPUExecutionProvider"
]
kwargs = {}
if self.infer_with_dropout:
kwargs["disabled_optimizers"] = ["EliminateDropout"]
warnings.warn(
"For inference with onnx, 'infer_with_dropout==True' is only effective if 'export_with_dropout==True' in training"
)
self.ort_session = onnxruntime.InferenceSession(
normalize_path(model_path, project_root),
providers=execution_providers,
**kwargs,
)
self.target_decoder_ids = load_onnx_target_decoder_ids(
self.ort_session,
self.target_columns,
self.target_column_types,
self.model_type,
)
if self.inference_model_type == "pt":
if self.training_config_path is None:
raise ValueError("training_config_path is required for PyTorch models")
self.inference_model = load_inference_model(
self.model_type,
normalize_path(model_path, project_root),
self.training_config_path,
self.args_config,
self.device,
self.infer_with_dropout,
)
[docs] @beartype
def invert_normalization(
self, values: np.ndarray, target_column: str
) -> np.ndarray:
"""Invert target-column Z-score normalization."""
if not self.normalize_real_columns:
return values
std = self.selected_columns_statistics[target_column]["std"]
mean = self.selected_columns_statistics[target_column]["mean"]
return (values * (std + 1e-9)) + mean
def _exclude_mask_token(
self,
values: np.ndarray,
*,
is_log_probabilities: bool,
target_column: Optional[str] = None,
) -> np.ndarray:
"""Remove the input-only BERT mask class and renormalize rows."""
if self.training_objective != "bert":
return values
decoder_ids = self.target_decoder_ids.get(
target_column or "", list(range(values.shape[1]))
)
if SPECIAL_TOKEN_IDS.mask not in decoder_ids:
return values
mask_id = decoder_ids.index(SPECIAL_TOKEN_IDS.mask)
adjusted = np.array(values, copy=True)
if is_log_probabilities:
adjusted[:, mask_id] = -np.inf
row_max = np.max(adjusted, axis=1, keepdims=True)
if not np.isfinite(row_max).all():
raise ValueError(
"BERT categorical outputs contain no generatable classes."
)
log_normalizers = row_max + np.log(
np.exp(adjusted - row_max).sum(axis=1, keepdims=True)
)
return adjusted - log_normalizers
adjusted[:, mask_id] = 0.0
row_sums = adjusted.sum(axis=1, keepdims=True)
if not np.isfinite(row_sums).all() or np.any(row_sums <= 0):
raise ValueError(
"BERT categorical probabilities contain no generatable classes."
)
return adjusted / row_sums
[docs] @beartype
def infer_embedding(
self,
x: dict[str, np.ndarray],
metadata: dict[str, np.ndarray],
column_data_types: dict[str, torch.dtype],
) -> np.ndarray:
"""Return embeddings for a feature-array batch."""
assert x is not None
size = x[list(x.keys())[0]].shape[0]
embedding = self.adjust_and_infer_embedding(
x, size, metadata, column_data_types
)
return embedding
[docs] @beartype
def infer_generative(
self,
x: Optional[dict[str, np.ndarray]],
metadata: dict[str, np.ndarray],
probs: Optional[dict[str, np.ndarray]] = None,
return_probs: bool = False,
column_data_types: Optional[dict[str, torch.dtype]] = None,
) -> dict[str, np.ndarray]:
"""Return target probabilities or decoded predictions."""
if probs is None or (
x is not None and len(set(x.keys()).difference(set(probs.keys()))) > 0
): # type: ignore
assert x is not None
size = x[list(x.keys())[0]].shape[0]
if (
probs is not None
and len(set(x.keys()).difference(set(probs.keys()))) > 0
): # type: ignore
assert x is not None
warnings.warn(
f"Not all keys in x are in probs - {x.keys() = } != {probs.keys() = }. Full inference is executed."
)
outs = self.adjust_and_infer_generative(
x, size, metadata, column_data_types or {}
)
for target_column, target_outs in outs.items():
if np.any(target_outs == np.inf):
raise ValueError(
f"Inference resulted in infinite values: {target_outs}"
)
if return_probs:
preds = {
target_column: outputs
for target_column, outputs in outs.items()
if self.target_column_types[target_column] != "categorical"
}
logits = {
target_column: self._exclude_mask_token(
outputs,
is_log_probabilities=True,
target_column=target_column,
)
for target_column, outputs in outs.items()
if self.target_column_types[target_column] == "categorical"
}
return {**preds, **normalize(logits)}
else:
outs = dict(probs)
for target_column in self.target_columns:
if self.target_column_types[target_column] == "categorical":
outs[target_column] = self._exclude_mask_token(
outs[target_column],
is_log_probabilities=probs is None,
target_column=target_column,
)
if (
self.sample_from_distribution_columns is None
or target_column not in self.sample_from_distribution_columns
):
outs[target_column] = outs[target_column].argmax(1)
else:
outs[target_column] = sample_with_cumsum(
outs[target_column], is_log_probs=(probs is None)
)
if target_column in self.target_decoder_ids:
outs[target_column] = np.asarray(
self.target_decoder_ids[target_column]
)[outs[target_column]]
return outs
[docs] @beartype
def adjust_and_infer_embedding(
self,
x: dict[str, np.ndarray],
size: int,
metadata: dict[str, np.ndarray],
column_data_types: dict[str, torch.dtype],
):
"""Batch embedding inference across the active backend."""
if self.inference_model_type == "onnx":
assert x is not None
x_adjusted = self.prepare_inference_batches(x, pad_to_batch_size=True)
metadata_adjusted = self.prepare_inference_batches(
metadata, pad_to_batch_size=True
)
inference_batch_embeddings = [
self.infer_pure(x_sub, metadata_sub)[0]
for x_sub, metadata_sub in zip(x_adjusted, metadata_adjusted)
]
embeddings = np.concatenate(inference_batch_embeddings, axis=0)[
: size * self.prediction_length
]
elif self.inference_model_type == "pt":
x_adjusted = self.prepare_inference_batches(x, pad_to_batch_size=False)
metadata_adjusted = self.prepare_inference_batches(
metadata, pad_to_batch_size=False
)
embeddings = infer_with_embedding_model(
self.inference_model,
x_adjusted,
self.device,
size,
self.target_columns,
metadata=metadata_adjusted,
column_data_types=column_data_types,
)
else:
assert False, "not possible"
return embeddings
[docs] @beartype
def adjust_and_infer_generative(
self,
x: dict[str, np.ndarray],
size: int,
metadata: dict[str, np.ndarray],
column_data_types: dict[str, torch.dtype],
):
"""Batch generative inference across the active backend."""
if self.inference_model_type == "onnx":
assert x is not None
x_adjusted = self.prepare_inference_batches(x, pad_to_batch_size=True)
metadata_adjusted = self.prepare_inference_batches(
metadata, pad_to_batch_size=True
)
out_subs = [
dict(zip(self.target_columns, self.infer_pure(x_sub, metadata_sub)))
for x_sub, metadata_sub in zip(x_adjusted, metadata_adjusted)
]
outs = {
target_column: np.concatenate(
[out_sub[target_column] for out_sub in out_subs], axis=0
)[: size * self.prediction_length, :]
for target_column in self.target_columns
}
elif self.inference_model_type == "pt":
assert x is not None
x_adjusted = self.prepare_inference_batches(x, pad_to_batch_size=False)
metadata_adjusted = self.prepare_inference_batches(
metadata, pad_to_batch_size=False
)
outs = infer_with_generative_model(
self.inference_model,
x_adjusted,
self.device,
size * self.prediction_length,
self.target_columns,
metadata=metadata_adjusted,
column_data_types=column_data_types,
)
else:
assert False
outs = {} # for type checking
return outs
[docs] @beartype
def prepare_inference_batches(
self, x: dict[str, np.ndarray], pad_to_batch_size: bool
) -> list[dict[str, np.ndarray]]:
"""Split feature arrays into backend-sized batches."""
size = x[list(x.keys())[0]].shape[0]
if size == self.inference_batch_size:
return [x]
elif size < self.inference_batch_size:
if pad_to_batch_size:
x_expanded = {
col: self.expand_to_batch_size(x_col) for col, x_col in x.items()
}
return [x_expanded]
else:
return [x]
else:
starts = range(0, size, self.inference_batch_size)
ends = range(
self.inference_batch_size,
size + self.inference_batch_size,
self.inference_batch_size,
)
xs = [
{col: x_col[start:end, :] for col, x_col in x.items()}
for start, end in zip(starts, ends)
]
return xs
[docs] @beartype
def infer_pure(
self,
x: dict[str, np.ndarray],
metadata: dict[str, np.ndarray],
) -> list[np.ndarray]:
"""Run one ONNX batch and flatten sequence-major outputs."""
metadata = metadata or {}
ort_inputs = {}
for session_input in self.ort_session.get_inputs():
input_name = session_input.name
if input_name in metadata:
value = metadata[input_name]
elif input_name.endswith("_in") and input_name[:-3] in x:
feature_column = input_name[:-3]
value = x[feature_column]
elif input_name in x:
value = x[input_name]
else:
raise ValueError(
f"Could not map ONNX input '{input_name}' to a feature or metadata array."
)
expected_dtype = ONNX_NUMPY_DTYPES.get(session_input.type)
if expected_dtype is not None and value.dtype != expected_dtype:
value = value.astype(expected_dtype, copy=False)
ort_inputs[input_name] = self.expand_to_batch_size(value)
ort_outs = self.ort_session.run(None, ort_inputs)
return [
oo.transpose(1, 0, 2).reshape(oo.shape[0] * oo.shape[1], oo.shape[2])
for oo in ort_outs
]
[docs] @beartype
def expand_to_batch_size(self, x: np.ndarray) -> np.ndarray:
"""Repeat leading samples until the ONNX batch size is met."""
repetitions = self.inference_batch_size // x.shape[0]
filler = self.inference_batch_size % x.shape[0]
return np.concatenate(([x] * repetitions) + [x[0:filler, :]], axis=0)
[docs]@beartype
def normalize(outs: dict[str, np.ndarray]) -> dict[str, np.ndarray]:
"""Softmax logits by target column."""
shifted_values = {
target_column: target_values - np.max(target_values, axis=1, keepdims=True)
for target_column, target_values in outs.items()
}
exp_values = {
target_column: np.exp(target_values)
for target_column, target_values in shifted_values.items()
}
probs = {
target_column: target_values / np.sum(target_values, axis=1, keepdims=True)
for target_column, target_values in exp_values.items()
}
return probs
[docs]@beartype
def sample_with_cumsum(probs: np.ndarray, is_log_probs: bool = True) -> np.ndarray:
"""Sample class indices from log-probabilities or probabilities."""
if is_log_probs:
sampling_probs = np.exp(probs)
else:
sampling_probs = probs
if not np.isfinite(sampling_probs).all():
raise ValueError("Sampling probabilities must be finite.")
if np.any(sampling_probs < 0):
raise ValueError("Sampling probabilities must be non-negative.")
row_sums = sampling_probs.sum(axis=1)
if not np.allclose(row_sums, 1.0):
raise ValueError("Sampling probabilities must sum to 1.0 for each row.")
cumulative_probs = np.cumsum(sampling_probs, axis=1)
cumulative_probs[:, -1] = 1.0
random_threshold = np.random.rand(cumulative_probs.shape[0], 1)
random_threshold = np.repeat(random_threshold, probs.shape[1], axis=1)
return (random_threshold < cumulative_probs).argmax(axis=1)