Source code for sequifier.infer

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)