TorchModel
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.forward
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.(N, output_dim) with the endpoint outputs.
