PyTorch Tensor Data Validation

PyTorch provides Tensor objects, TensorDict for managing collections of tensors, and tensorclass for typed tensor collections.

Pandera validates them with the same patterns as the other dataframe backends: schema objects, optional Check instances, and global configuration.

Installation

pip install 'pandera[torch]'
pip install tensordict

Quick start

import torch
from tensordict import TensorDict
import pandera.tensordict as pa

schema = pa.TensorDictSchema(
    keys={
        "observation": pa.Tensor(dtype=torch.float32, shape=(None, 10)),
        "action": pa.Tensor(dtype=torch.float32, shape=(None, 5)),
    },
    batch_size=(32,),
)

td = TensorDict(
    {"observation": torch.randn(32, 10), "action": torch.randn(32, 5)},
    batch_size=[32],
)
schema.validate(td)
TensorDict(
    fields={
        action: Tensor(shape=torch.Size([32, 5]), device=cpu, dtype=torch.float32, is_shared=False),
        observation: Tensor(shape=torch.Size([32, 10]), device=cpu, dtype=torch.float32, is_shared=False)},
    batch_size=torch.Size([32]),
    device=None,
    is_shared=False)

Supported objects

TensorDictSchema validates any tensordict tensor collection:

  • TensorDict and any other TensorDictBase subclass, including lazily stacked tensordicts created with lazy_stack().

  • tensorclass instances created with the tensorclass() decorator.

stacked = TensorDict.lazy_stack(
    [
        TensorDict({"observation": torch.randn(16, 10), "action": torch.randn(16, 5)}, batch_size=[16]),
        TensorDict({"observation": torch.randn(16, 10), "action": torch.randn(16, 5)}, batch_size=[16]),
    ],
    dim=0,
)

stacked_schema = pa.TensorDictSchema(
    keys={
        "observation": pa.Tensor(dtype=torch.float32),
        "action": pa.Tensor(dtype=torch.float32),
    },
    batch_size=(2, 16),
)
stacked_schema.validate(stacked)
LazyStackedTensorDict(
    fields={
        action: Tensor(shape=torch.Size([2, 16, 5]), device=cpu, dtype=torch.float32, is_shared=False),
        observation: Tensor(shape=torch.Size([2, 16, 10]), device=cpu, dtype=torch.float32, is_shared=False)},
    exclusive_fields={
    },
    batch_size=torch.Size([2, 16]),
    device=None,
    is_shared=False,
    stack_dim=0)

Type Coercion

Set coerce=True to automatically convert tensor dtypes:

schema_coerce = pa.TensorDictSchema(
    keys={
        "observation": pa.Tensor(dtype=torch.float32, shape=(None, 10)),
    },
    batch_size=(32,),
    coerce=True,
)

# Input with wrong dtype (float64)
td_wrong_dtype = TensorDict(
    {"observation": torch.randn(32, 10).to(torch.float64)},
    batch_size=[32],
)

# Dtype is automatically coerced to float32
validated = schema_coerce.validate(td_wrong_dtype)
assert validated["observation"].dtype == torch.float32

Define a schema with a class-based model

class RL(pa.TensorDictModel):
    """Schema for reinforcement learning data."""

    # Use PyTorch dtypes in type annotations
    observation: torch.float32 = pa.Field(shape=(None, 10))
    action: torch.int64 = pa.Field(shape=(None,))
    reward: torch.float32 = pa.Field()

    class Config:
        batch_size = (32,)

# Validate using the model - schema is built automatically
td = TensorDict(
    {"observation": torch.randn(32, 10), "action": torch.randint(0, 4, (32,)), "reward": torch.randn(32)},
    batch_size=[32],
)
RL.validate(td)
TensorDict(
    fields={
        action: Tensor(shape=torch.Size([32]), device=cpu, dtype=torch.int64, is_shared=False),
        observation: Tensor(shape=torch.Size([32, 10]), device=cpu, dtype=torch.float32, is_shared=False),
        reward: Tensor(shape=torch.Size([32]), device=cpu, dtype=torch.float32, is_shared=False)},
    batch_size=torch.Size([32]),
    device=None,
    is_shared=False)

Note: Type annotations specify the dtype (e.g., torch.float32, torch.int64). Use Field() to define additional constraints like shape and checks.

Guide contents

See also