Skip to content
AITroveRead. Build. Understand.
Make this comfortable

Tensor contracts: shape, dtype, device and mask

Last updated: 5 Oct 20265 min read
tutorial
IntermediateBy AITrove Editorial

A tensor boundary states what every axis means, which numeric type is valid and whether padding represents data.

Name every axis

A batch of receipt images has shape (batch, channel, height, width). A text encoder may expect (batch, token). Write these contracts at the loader boundary and assert them before the model. A tensor with the right total element count can still have transposed axes. Batch construction] must preserve the same contract for a short final batch.

Pin numeric semantics

Image bytes become floating values under a named scaling rule; class targets remain integer indices for an index-based classification loss. Check finite values before training, and record dtype and precision policy. Move inputs, targets and model parameters to the intended device. A mixed CPU and accelerator calculation can fail only on a less common branch.

Represent absence

Padding a variable-length sequence changes storage, not meaning. Carry a mask marking real elements and apply it to attention or aggregation. An all-zero input might be a real dark image, so zero alone is a weak missingness marker. Missingness rules] belong in the dataset contract.

Test before training

Feed one item, an incomplete last batch, a corrupted image and an unknown label. Assert axes, dtype, range and finite values. Route the corrupted item to a visible rejection path rather than converting it to a plausible black image. Save the contract with the model artifact so inference can enforce it.

Implementation

python
def validate_image_batch(images, class_ids, class_count):
    if images.ndim != 4 or images.shape[1] != 3:
        raise ValueError("expected NCHW RGB images")
    if class_ids.shape != (images.shape[0],):
        raise ValueError("one class index per image")
    if images.dtype != torch.float32 or class_ids.dtype != torch.int64:
        raise TypeError("unexpected tensor dtype")
    if not torch.isfinite(images).all():
        raise ValueError("non-finite input")
    if ((class_ids < 0) | (class_ids >= class_count)).any():
        raise ValueError("class index out of range")

Performance and operating cost

Shape and dtype checks cost O(1); scanning all pixels for non-finite values costs O(BCHW) per batch. The scan has a measurable cost, but silent invalid values can poison an entire training run.

Common Mistakes

  • Do not infer axis meaning from shape alone.
  • Do not cast integer class indices to float for an index-based loss.
  • Do not treat padding as an observation.

Read next

Continue the workflow: Unicode and tokenization: preserve meaning at the text boundary.

ai-data
deep-learning
Storage details