Skip to main content
PyTorch adapter for Bioptimus endpoint models.

TorchModel

Wraps a SageMaker backbone endpoint with a PyTorch forward API. Pass a model name plus the AWS endpoint details and the wrapper builds the SageMaker client for you. On each forward call it serializes the incoming image batch into endpoint requests, calls the endpoint, and returns a stacked tensor of outputs. This keeps loops shaped like features = model(images) while the model runs behind a SageMaker endpoint. Keep ToTensor in the dataloader but drop Normalize (and any other scaling/stain transforms): the endpoint applies the model’s own preprocessing server-side, so tiles must be passed as un-normalized RGB. forward also accepts numpy arrays or PIL images if you prefer not to convert to tensors at all. Notes: This is a thin remote adapter, not a real local model. Inference runs server-side behind the SageMaker endpoint, and the wrapper registers no parameters or buffers. The usual Module controls are therefore inert no-ops: .to(device), .cuda(), .half(), .eval(), .train(), and .parameters() do nothing meaningful, and there is no local device or dtype to configure. forward always returns a freshly built CPU float32 tensor regardless of any device the caller set. The Module base class exists only so the wrapper drops into loops that expect a callable model(images); do not rely on device, dtype, or eval/train state to change where or how the endpoint runs.
Model to build the endpoint for, as a Models member (e.g. Models.M_OPTIMUS) or its string name (e.g. "m-optimus").
Name of the deployed SageMaker endpoint to invoke.
AWS region of the endpoint. Falls back to the default boto3 resolution (env/config) when omitted.
"embedding" to call EndpointModel.embed (tile features), or "prediction" to call EndpointModel.predict (e.g. M-Optimus gene expression).
Per-request timeout in seconds.
Maximum number of tiles sent to the endpoint concurrently per forward call. Each tile is one request, so this caps how many in-flight requests a single batch produces. Lower it for single-GPU endpoints that cannot absorb many concurrent requests.
Example:

forward

Runs a batch of raw tiles through the wrapped endpoint. Inference runs server-side, so the result is always a CPU float32 tensor regardless of any device set on the wrapper (see the class note on the inert Module controls).
Un-normalized tiles, passed as-is without transforms. May be a tensor batch shaped (N, C, H, W) or (N, H, W, C), a numpy array, a single PIL image, or a list of PIL images / tensors / numpy arrays.
Returns: CPU tensor shaped (N, output_dim) with the endpoint outputs.