Skip to main content
Base class for server-side exported model loading and inference. Server models load exported .pt2 checkpoints and run inference. Weights are loaded from /opt/ml/model (SageMaker convention) via torch._inductor.aoti_load_package. For models with no custom forward logic (e.g. embedding backbones), BaseServerModel is used directly. Models that need custom input handling (e.g. bulk RNA, segmentation masks) subclass it and override prepare_sample and forward_batch.

Two-stage failure isolation

Stage 1 — Deserialization (per-request, parallel): prepare_sample decodes, transforms, and validates each request in a thread pool. Failures are returned immediately as ValueError (HTTP 400) to the individual caller. Stage 2 — GPU forward (shared batch): forward_batch receives only validated tensors, runs the model, synchronizes CUDA, and transfers results to CPU. Any failure at this stage raises _ForwardError (RuntimeError subclass → HTTP 503) shared by all requests in the batch.

PreparedSample

A fully validated, preprocessed sample ready for GPU batching. Created by BaseServerModel.prepare_sample during the per-request deserialization stage. Only samples that pass all validation enter the batch queue. Attributes:
  • tile - Preprocessed image tensor of shape (C, H, W).
  • extra_tensors - Additional model inputs (e.g. bulk RNA tensor) keyed by input name. Empty for standard embedding models.

BaseServerModel

Loads an exported .pt2 model and runs single-tile inference. For embedding backbones (h0-mini, h0, h1, …) this class is used directly — no subclass needed. Override prepare_sample and forward_batch for models that require custom input handling (e.g. bulk RNA).
The model specification (tile size, normalization, etc.).
Filename of the .pt2 checkpoint.
Torch device string ("cuda" or "cpu").
Directory containing exported checkpoints.

model_spec

The model’s specification (tile size, normalization, etc.).

model_name

Human-readable model identifier.

preprocess_bulk_rna

Applies the model-spec omics transform to raw counts.
Raw bulk RNA counts as a flat list of floats.
Returns: Transformed counts. Returned unchanged when no omics transform is configured.

prepare_sample

Deserializes, transforms, and validates a single request. Called in a thread pool by the batch scheduler. Any exception raised here is returned as a per-request error (HTTP 400) and the request never enters the batch queue. Subclasses override this to add custom validation (e.g. bulk RNA preprocessing for models).
Raw model request with base64 image data.
Returns: A PreparedSample with the preprocessed tile tensor. Raises:
  • ValueError - If the image cannot be decoded, transformed, or has the wrong shape.

forward_batch

Runs batched GPU inference on pre-validated samples. All samples have already passed prepare_sample. The entire GPU execution path (forward + synchronize + CPU transfer) is wrapped in error handling so that CUDA errors cannot escape.
Pre-validated samples with preprocessed tensors.
Original requests (for response metadata).
"prediction", "embedding", or "prediction_with_embedding". The last returns both the prediction and the embedding from the single forward pass.
Returns: One ModelResponse per sample, positionally aligned. Raises:
  • _ForwardError - If the GPU forward pass, CUDA synchronization, or CPU transfer fails.

get_metadata

Returns model metadata. Returns: Dictionary with model name, input/output specs, and config.

forward

Runs the model on a batch of tile requests (legacy path). Used by subclasses with custom forward logic that haven’t migrated to the two-stage pipeline yet.
List of tile requests.
Returns: Model output tensor of shape (N, D).