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:
TensorDictand any otherTensorDictBasesubclass, including lazily stacked tensordicts created withlazy_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¶
TensorDictSchema — validating a
TensorDictwithTensorcomponentsTensorDictModel — class-based
TensorDictModelChecks and Lazy Validation — checks, parsers, and lazy validation
Schema Inference — infer schemas from data automatically
Serialization and Deserialization — save/load schemas with YAML/JSON
Hypothesis Strategies — generate synthetic data with Hypothesis
Error Reporting —
SchemaError/SchemaErrors, lazy validation, and failure cases
See also¶
Supported DataFrame Libraries — other backends
Validating with Checks — general
CheckbehaviourLazy Validation —
lazy=TrueandSchemaErrorsConfiguration —
ValidationDepthand environment variables