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.typing import Iterator
from loguru import logger
from sequifier.config.infer_config import InfererModel, load_inferer_config
from sequifier.config.train_config import ResolvedSequifierConfig as TrainModel
from sequifier.helpers import (
PANDAS_TO_TORCH_TYPES,
configure_determinism,
configured_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,
)
from sequifier.train import (
infer_with_embedding_model,
infer_with_generative_model,
load_inference_model,
)
from sequifier.typechecking import beartype
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":
from sequifier.model.execution_schema import (
EXECUTION_SCHEMA_KEY,
ExecutionSchema,
)
raw_schema = session.get_modelmeta().custom_metadata_map.get(
EXECUTION_SCHEMA_KEY
)
output_names = (
{
item["key"]: item["name"]
for item in ExecutionSchema.from_dict(json.loads(raw_schema)).outputs
}
if raw_schema is not None
else None
)
outputs_by_name = {output.name: output for output in session.get_outputs()}
for column, decoder_ids in target_decoder_ids.items():
if output_names is not None:
output = outputs_by_name.get(output_names.get(column))
else:
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.decode_categories 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.deterministic)
# Seed once for this CLI run, before creating any ORT sessions.
onnxruntime.set_seed(config.seed)
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]:
from sequifier.config.depth_layout import DepthLayoutRegistryModel
from sequifier.io.pt_payload import load_pt_payload
metadata_path = Path(data_path) / "metadata.json"
folder_metadata = (
json.loads(metadata_path.read_text()) if metadata_path.exists() else {}
)
yield load_pt_payload(
pt_file,
layouts=DepthLayoutRegistryModel.model_validate(
folder_metadata.get("depth_layouts", {})
),
n_classes=folder_metadata.get("n_classes"),
)
[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.window_length - 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,
depth_valid_masks: Optional[dict[str, torch.Tensor]] = None,
) -> WindowedInferenceBatch:
plan = resolve_window_sampling_plan(
config.storage_layout,
config.window_view,
configured_window_stride(config),
)
sample_index = plan.build_index(left_pad_lengths)
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,
)
from sequifier.config.depth_layout import depth_mask_metadata_key
for name, mask in (depth_valid_masks or {}).items():
metadata[depth_mask_metadata_key(name)] = plan.gather(
mask, 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.window_length,
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: Any,
column_data_types: dict[str, torch.dtype],
) -> WindowedInferenceBatch:
(
sequences,
sequence_ids,
subsequence_ids,
start_positions,
left_pad_lengths,
) = data
from sequifier.config.depth_layout import DepthLayoutRegistryModel
from sequifier.io.pt_payload import StoredTensorBatch
masks = data.depth_valid_masks if isinstance(data, StoredTensorBatch) else {}
if isinstance(data, StoredTensorBatch) and config.dataset_metadata is not None:
selected = config.dataset_metadata.depth_layouts.relevant_layouts(
config.input_columns
)
# Complete payload validation occurred at load; selected-view validation
# happens before graph execution using the artifact contract.
masks = {name: masks[name] for name in selected.root}
if config.dataset_metadata is not None:
from sequifier.io.pt_payload import validate_tensor_inputs
selected_layouts = config.dataset_metadata.depth_layouts.relevant_layouts(
config.input_columns
)
validate_tensor_inputs(
{c: sequences[c] for c in config.input_columns},
masks,
selected_layouts,
n_classes=config.dataset_metadata.n_classes,
)
# Sanitize legal masked values before an input dtype override narrows them.
layouts = (
config.dataset_metadata.depth_layouts
if config.dataset_metadata
else DepthLayoutRegistryModel()
)
for column, tensor in sequences.items():
name = layouts.column_to_layout.get(column)
if name is not None and name in masks:
sequences[column] = torch.where(
masks[name], tensor, torch.zeros_like(tensor)
)
sequences = apply_inference_tensor_types(sequences, column_data_types)
for tensor in sequences.values():
validate_stored_window_width(
tensor,
config.storage_layout.window_length,
)
return _windowed_inference_batch_from_storage(
config,
sequences,
sequence_ids,
subsequence_ids,
start_positions,
left_pad_lengths,
depth_valid_masks=masks,
)
@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,
)
from sequifier.io.pt_payload import StoredTensorBatch
if isinstance(data, (tuple, StoredTensorBatch)):
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"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)]
if not is_folder_input and config.read_format == "pt":
from sequifier.io.pt_payload import load_pt_payload
dataset = [
load_pt_payload(
config.data_path,
layouts=config.dataset_metadata.depth_layouts
if config.dataset_metadata
else None,
n_classes=config.dataset_metadata.n_classes
if config.dataset_metadata
else None,
)
]
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
route_args = dict(args_config)
for key in ("dataset", "part", "model_interface"):
value = getattr(config, key, None)
if value is not None:
route_args[key] = value
if model_path.lower().endswith(".pt"):
model_state = torch.load(
normalize_path(model_path, config.project_root),
map_location="cpu",
weights_only=False,
)
embedded_model_config = model_state.get("model_config")
if embedded_model_config is not None:
interface_name = route_args.get("model_interface")
if (
interface_name is None
and len(embedded_model_config.get("interfaces", {})) == 1
):
interface_name = next(iter(embedded_model_config["interfaces"]))
if interface_name is not None:
interface = embedded_model_config["interfaces"].get(interface_name)
if interface is None:
raise ValueError(
f"Unknown PT model interface {interface_name!r}"
)
target_decoder_ids = interface.get("target_decoder_ids", {})
elif model_state.get("training_config") is not None:
training_config = TrainModel.model_validate(
model_state["training_config"]
)
selected_dataset = route_args.get("dataset")
selected_interface = route_args.get("model_interface")
datasets = training_config.dataset_training
if selected_dataset is not None:
if selected_dataset not in datasets:
raise ValueError(
f"Unknown inference dataset {selected_dataset!r}"
)
dataset_config = datasets[selected_dataset]
if (
selected_interface is not None
and dataset_config.model_interface != selected_interface
):
raise ValueError(
f"Dataset {selected_dataset!r} maps to interface "
f"{dataset_config.model_interface!r}, not "
f"{selected_interface!r}"
)
elif selected_interface is not None:
dataset_config = next(
(
dataset
for dataset in datasets.values()
if dataset.model_interface == selected_interface
),
None,
)
if dataset_config is None:
raise ValueError(
"No execution route for model interface "
f"{selected_interface!r}"
)
elif len(datasets) == 1:
dataset_config = next(iter(datasets.values()))
else:
raise ValueError(
"A dataset or model_interface selection is required for "
"multi-dataset checkpoint inference"
)
target_decoder_ids = dict(dataset_config.interface.target_decoder_ids)
else:
raise ValueError(
"PyTorch artifact has neither model_config nor a resolved "
"training_config."
)
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.decode_categories,
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=route_args,
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 = Path(model_path).stem
logger.info(f"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]@beartype
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 _autoregressive_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 autoregressive 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 autoregressive 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 autoregressive subsequence contains duplicate input "
f"rows for sequenceId={sequence_id!r}."
)
return seed_data
[docs]@beartype
def inference_output_path(
project_root: str,
write_format: str,
artifact_type: str,
model_id: str,
data_id: int,
target_column: Optional[str] = None,
) -> str:
"""Return a canonical inference output path and create its directory."""
output_dir = Path(project_root) / "outputs" / artifact_type / model_id
if target_column is not None:
output_dir /= target_column
output_dir.mkdir(parents=True, exist_ok=True)
return str(output_dir / f"part-{data_id:03d}.{write_format}")
[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
windowed = _windowed_inference_batch(config, data, column_data_types)
if windowed.sequence_ids.numel() == 0:
continue
if isinstance(data, pl.DataFrame) and configured_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])],
)
),
}
)
embeddings_path = inference_output_path(
config.project_root,
config.write_format,
"embeddings",
model_id,
data_id,
)
logger.info(f"Writing embeddings to '{embeddings_path}'")
write_data(
embeddings_df,
embeddings_path,
config.write_format,
)
[docs]@beartype
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):
if config.autoregressive and isinstance(data, pl.DataFrame):
data = _autoregressive_seed_dataframe(config, data)
windowed = _windowed_inference_batch(config, data, column_data_types)
if windowed.sequence_ids.numel() == 0:
continue
if config.autoregressive and inferer.prediction_length != 1:
raise ValueError(
"prediction_length must be 1 for autoregressive inference, "
f"got {inferer.prediction_length}"
)
total_steps = (
config.generation_steps
if config.autoregressive and config.generation_steps is not None
else 1
)
if (
isinstance(data, pl.DataFrame)
and total_steps == 1
and configured_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.decode_categories:
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"
)
if config.output_probabilities:
assert probs is not None
for target_column in inferer.target_columns:
if inferer.target_column_types[target_column] == "categorical":
probabilities_path = inference_output_path(
config.project_root,
config.write_format,
"probabilities",
model_id,
data_id,
target_column,
)
logger.info(f"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
},
}
)
predictions_path = inference_output_path(
config.project_root,
config.write_format,
"predictions",
model_id,
data_id,
)
logger.info(f"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.window_length)
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 autoregressive inference"
)
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_autoregressive(
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.generation_steps,
)
item_positions_for_preds = np.concatenate(
[
np.arange(start_pos, start_pos + config.generation_steps)
for start_pos in aligned_start_positions
],
axis=0,
)
sequence_ids_for_preds = np.repeat(aligned_sequence_ids, config.generation_steps)
base_mask = _flatten_valid_mask(config, metadata, 1)
valid_prediction_mask = np.repeat(base_mask, config.generation_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]],
decode_categories: 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.decode_categories = decode_categories
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, decode_categories
)
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.startswith("cuda")
else "CPUExecutionProvider"
]
if any(
provider not in onnxruntime.get_available_providers()
for provider in execution_providers
):
raise ValueError(
f"Requested ONNX providers are unavailable: {execution_providers}"
)
kwargs = {}
if self.infer_with_dropout:
session_options = onnxruntime.SessionOptions()
session_options.graph_optimization_level = (
onnxruntime.GraphOptimizationLevel.ORT_DISABLE_ALL
)
kwargs["sess_options"] = session_options
kwargs["disabled_optimizers"] = ["EliminateDropout"]
self.ort_session = onnxruntime.InferenceSession(
normalize_path(model_path, project_root),
providers=execution_providers,
**kwargs,
)
from sequifier.model.execution_schema import (
DROPOUT_MODE_KEY,
EXECUTION_SCHEMA_KEY,
ExecutionSchema,
)
properties = self.ort_session.get_modelmeta().custom_metadata_map
mode = properties.get(DROPOUT_MODE_KEY)
requested = "stochastic" if self.infer_with_dropout else "evaluation"
if mode is not None and mode != requested:
raise ValueError(
f"ONNX graph provides {mode} dropout mode, but inference requests {requested}; export a graph for the requested mode"
)
if mode is None and self.infer_with_dropout:
warnings.warn(
"Legacy ONNX graph has no dropout capability metadata; runtime dropout behavior cannot be guaranteed"
)
self.execution_schema = (
ExecutionSchema.from_dict(json.loads(properties[EXECUTION_SCHEMA_KEY]))
if EXECUTION_SCHEMA_KEY in properties
else None
)
if self.execution_schema is not None:
import onnx
from sequifier.export.onnx import validate_graph_contract
graph = onnx.load(
normalize_path(model_path, project_root),
load_external_data=False,
)
validate_graph_contract(graph, self.execution_schema)
elif any(len(value.shape) == 3 for value in self.ort_session.get_inputs()):
raise ValueError("Depth ONNX inputs require execution schema metadata")
self.symbolic_batch = all(
not isinstance(value.shape[0], int)
for value in self.ort_session.get_inputs()
)
if not self.symbolic_batch:
batches = {
value.shape[0]
for value in self.ort_session.get_inputs()
if isinstance(value.shape[0], int)
}
if len(batches) != 1:
raise ValueError("ONNX graph inputs disagree on static batch size")
self.inference_batch_size = next(iter(batches))
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":
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,
)
route_model = getattr(
self.inference_model,
"transformer_model",
self.inference_model,
)
self.target_decoder_ids = dict(route_model.target_decoder_ids)
[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
@beartype
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=not self.symbolic_batch
)
metadata_adjusted = self.prepare_inference_batches(
metadata, pad_to_batch_size=not self.symbolic_batch
)
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=not self.symbolic_batch
)
metadata_adjusted = self.prepare_inference_batches(
metadata, pad_to_batch_size=not self.symbolic_batch
)
out_subs = [
dict(
zip(
[item["key"] for item in self.execution_schema.outputs]
if self.execution_schema
else 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."""
if not x:
return []
size = next(iter(x.values())).shape[0]
if any(value.shape[0] != size for value in x.values()):
raise ValueError("Inference arrays disagree on batch size")
result = []
for start in range(0, size, self.inference_batch_size):
batch = {
key: value[start : start + self.inference_batch_size]
for key, value in x.items()
}
if (
pad_to_batch_size
and next(iter(batch.values())).shape[0] < self.inference_batch_size
):
batch = {
key: self.expand_to_batch_size(value)
for key, value in batch.items()
}
result.append(batch)
return result
[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 {}
schema = getattr(self, "execution_schema", None)
if schema is not None:
schema.validate(
{key: torch.from_numpy(value) for key, value in x.items()},
{key: torch.from_numpy(value) for key, value in metadata.items()},
)
descriptors = {item.name: item for item in schema.inputs} if schema else {}
ort_inputs = {}
sizes = set()
for session_input in self.ort_session.get_inputs():
input_name = session_input.name
descriptor = descriptors.get(input_name)
if descriptor is not None:
value = (x if descriptor.role == "feature" else metadata)[
descriptor.key
]
if descriptor.role == "feature" and schema is not None:
for name, layout in schema.depth_layouts.items():
if descriptor.key in layout["columns"]:
value = np.where(
metadata[f"depth_valid_mask:{name}"], value, 0
)
elif input_name in metadata:
value = metadata[input_name]
elif input_name.endswith("_in") and input_name[:-3] in x:
value = x[input_name[:-3]]
elif input_name in x:
value = x[input_name]
else:
raise ValueError(f"Could not bind ONNX input {input_name!r}")
if len(value.shape) != len(session_input.shape):
raise ValueError(f"ONNX input {input_name!r} has incorrect rank")
for axis, expected in enumerate(session_input.shape):
if axis and isinstance(expected, int) and value.shape[axis] != expected:
raise ValueError(
f"ONNX input {input_name!r} has incorrect capacity"
)
expected_dtype = ONNX_NUMPY_DTYPES.get(session_input.type)
if expected_dtype is None:
raise ValueError(f"Unsupported ONNX input dtype {session_input.type}")
if expected_dtype == np.bool_ and value.dtype != np.bool_:
raise ValueError("ONNX validity masks must have boolean dtype")
value = value.astype(expected_dtype, copy=False)
if np.issubdtype(value.dtype, np.floating) and not np.isfinite(value).all():
raise ValueError(f"{input_name}: input conversion overflowed")
sizes.add(value.shape[0])
ort_inputs[input_name] = (
value if self.symbolic_batch else self.expand_to_batch_size(value)
)
if len(sizes) != 1 or 0 in sizes:
raise ValueError("ONNX inputs need one common non-empty batch size")
output_names = [item["name"] for item in schema.outputs] if schema else None
ort_outs = self.ort_session.run(output_names, 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)