Pandera

Pandera 0.34: schemas for PyTorch TensorDicts

No items found.
A TensorDict batch with a float64 observation, an out-of-range action, and a NaN reward flows into a Pandera Transition schema that either passes it to the training loop or raises a SchemaError; below, the results: validation is 0.3% of a 17M-parameter training step, a whole 4 GB dataset validates in at most 1.8 seconds, and million-row chunks are 31 times faster

Most of the bugs I've chased in training code weren't in the model. They were in the data that went into it: an observation tensor that came out of NumPy as `float64` and quietly doubled the memory of a replay buffer, an action index that fell outside the action space and only blew up as a CUDA device-side assert three layers down, a `NaN` reward that turned a loss curve into a flat line an hour into a run. In every case the tensor had a contract (a dtype, a shape, a range of sensible values) and nothing in the code wrote that contract down, so nothing checked it.

Pandera has spent years writing those contracts down for dataframes. With v0.34.0, it does the same for tensors. The headline feature of this release is a new `pandera.tensordict` module for validating PyTorch TensorDict objects with the same schema, check, and error-reporting machinery you'd use on a pandas or Polars dataframe.

We also measured what those checks cost. Validating a 256-row batch took 0.3 to 0.5 ms, which is 0.3% of a training step for a 17M-parameter model and 4% for a 1M-parameter one. Only a toy 9K-parameter MLP, whose step takes about a millisecond, pays a noticeable price, at 29%. Validating a whole dataset once, before training starts, is cheaper still: 4 GB took between 112 ms and 1.8 s, and checking it in million-row chunks was 31 times faster than checking the same data one training batch at a time.

Why TensorDict

A `TensorDict` is a dictionary of tensors that share leading batch dimensions. It's the core data structure in TorchRL, and it's become a common way to pass structured batches around PyTorch code more generally, because you can index, slice, reshape, and move to a device the whole collection at once instead of key by key.

That structure maps well onto a dataframe schema. The keys play the role of columns, each tensor has a dtype and a shape the way a column has a dtype, and `batch_size` plays the role of the row count. So a schema for a TensorDict has three layers of things to check:

  • the keys that must be present,
  • each tensor's dtype and shape,
  • and the values inside each tensor, such as ranges and allowed sets.

TensorDict itself already enforces one thing for you: every tensor's leading dimensions must match `batch_size`, or the object won't construct. Everything else is up to you, and that's where Pandera comes in.

What validation costs in a training loop

Checking every batch isn't free, so we measured it. The benchmark trains a policy network on batches of 256 transitions, 5% of which are corrupted in the three ways from the intro. Pandera validates every batch and the loop skips the ones that raise `SchemaError`, and every run skipped exactly the corrupt batches. We ran it as a Flyte task on 4 CPU cores.

Validation is a fixed cost per batch, so what matters is how it compares to the training step. To see that, we kept the batches and the schema fixed and grew the model from 9K to 17M parameters, timing each `validate()` call and each training step without validation.

Validation time as a share of the training step falls from 29% for a 9K-parameter model at 1.1 ms per step to 0.3% for a 17M-parameter model at 142 ms per step

Checking a batch took 0.28 to 0.46 ms at every size, while the step grew 135x. On a 9K-parameter MLP with a 1.1 ms step, validation is 29% of the step, which is about as cheap as a real training step gets. At 1.1M parameters it's 4%, and at 17M parameters it's 0.3%. We also measured the end-to-end wall-clock difference against a loop that skips the same batches with a precomputed mask: it matches these shares for the small models, and from about a million parameters up it's smaller than the run-to-run noise on a shared node. These are CPU numbers, and a GPU step has different overheads, but the shape doesn't change: a fixed cost per batch shrinks against whatever the step costs.

The other thing to keep in mind is that most of the cost is the value checks, which scan every element. In a separate measurement on a laptop, dtype, shape, and key checks alone took about a third as long as the full schema. So if the per-step cost matters for your loop, validate where data enters the system, such as on insertion into a replay buffer, rather than on every batch you sample from it.

Validating once, upstream

So what does validating where the data enters actually cost? We ran a second benchmark with four synthetic datasets shaped like common training data: TorchRL-style continuous-control transitions, Atari frame stacks, a tabular classification set with one-hot categoricals, and tokenized fine-tuning sequences with `-100` ignore-index labels. Each schema checks dtypes and shapes on every key, plus a value check on every non-boolean key, including custom checks for valid one-hot rows, blank frames, and label ids. Each dataset sits in memory as one TensorDict and streams through `validate(inplace=True)` in 64 MB chunks.

Whole-dataset validation time grows linearly from 16 MB to 4 GB for all four datasets, topping out at 1.8 seconds for 4 GB of tabular data; validating the same 1 GB in 256-row chunks takes 5.2 seconds versus 165 ms in million-row chunks

Validation time grows linearly with dataset size. At 4 GB it ranged from 112 ms for the Atari frames to 1.8 s for the tabular data, whose one-hot check does the most work per byte. That's a one-time cost: a dataset validated where it's produced can be handed to any number of training jobs, and none of them has to check it again.

The chart on the right is the link back to the per-batch numbers. Validating the same gigabyte 256 rows at a time, the size of a training batch, took 5.2 s, which is 31 times longer than in million-row chunks, because every `validate()` call carries a fixed overhead and small chunks pay it over and over. So the cheapest way to get the guarantee is to validate in large chunks, once, before the data reaches the loop. Interestingly, validating everything in a single call was slower than million-row chunks, so there's no need to hold out for one giant call either.

One tip from writing these benchmarks: `validate()` clones the TensorDict before checking it unless you pass `inplace=True`. For a large dataset that clone doubles peak memory, so pass `inplace=True` whenever you're only checking the data, not transforming it.

What skipping validation costs

Those numbers are the price of checking. The other side of the ledger is what happens when you don't check and bad batches reach the model. To measure that, we trained a REINFORCE policy on CartPole and corrupted a share of the training batches in one of two ways: a `NaN` reward, which turns the weights into `NaN` and crashes the next rollout, and observations 100 times too large, the kind of thing a units or normalization bug produces, which crashes nothing. The unvalidated loop trains on whatever it gets and rolls back to its last good checkpoint when it crashes. The validated loop runs each batch through a Pandera schema and skips it on `SchemaError`. We counted the environment steps each run needed to solve the task, including steps thrown away by rollbacks and skipped batches, over 30 seeds per setting.

Compute to solve CartPole relative to a clean run as the share of corrupt batches rises from 0% to 20%: without validation, NaN rewards cost 2.1x and observations scaled 100x cost more than 10.6x, with half those runs never solving it; with Pandera validation both stay between 1.0x and 1.2x

Validating every batch took 0.5% of each run's wall-clock time, and the validated runs solved the task every time, at 1.0 to 1.2 times the compute of a clean run. The unvalidated runs did fine at 1% corruption, since a run is only about 45 updates and most never saw a bad batch, but the cost climbed from there. The `NaN` reward cost 2.1 times the compute at 20% corrupt batches, mostly in work lost to 25 rollbacks per run. The scaled observations were worse, because nothing stopped the loop from learning from them: twice the compute at 5%, and at 20% half the runs still hadn't solved the task after a million steps, more than 10 times what a clean run needed.

The loud failure turned out to be the cheap one. A crash costs you the work since your last checkpoint, while a silent corruption costs you however long it takes to notice the model isn't learning. CartPole recovers quickly from a bad update and the checkpoints here were ten updates apart, so a larger model with checkpoints an hour apart would likely pay more on both counts.

Validating a batch

Install the `torch` extra, which pulls in `torch` and `tensordict`:

Copied to clipboard!
pip install 'pandera[torch]'

Here's a schema for a batch of RL transitions, written as a class-based `TensorDictModel`. The type annotations are PyTorch dtypes, and `pa.Field` holds the shape and value constraints:

Copied to clipboard!
import torch
from tensordict import TensorDict
import pandera.tensordict as pa


class Transition(pa.TensorDictModel):
    observation: torch.float32 = pa.Field(shape=(None, 8))
    action: torch.int64 = pa.Field(shape=(None,), isin=[0, 1, 2, 3])
    reward: torch.float32 = pa.Field(shape=(None,), ge=-1.0, le=1.0)
    done: torch.bool = pa.Field(shape=(None,))

    class Config:
        batch_size = (None,)

A `None` in `shape` or `batch_size` means "any size along this dimension," so this schema accepts a batch of 64 transitions or 4,096 of them, as long as each observation has 8 features. If your pipeline always produces a fixed batch size, write `batch_size = (64,)` and Pandera will check that too.

Now let's feed it a batch with two of the bugs from the intro. The observations come from NumPy, so they're `float64`, and the actions are sampled from a range that's one too wide:

Copied to clipboard!
import numpy as np

batch = TensorDict(
    {
        "observation": torch.from_numpy(np.random.randn(64, 8)),  # float64
        "action": torch.randint(0, 5, (64,)),                      # 4 is out of range
        "reward": torch.randn(64).clamp(-1, 1),
        "done": torch.zeros(64, dtype=torch.bool),
    },
    batch_size=[64],
)

try:
    Transition.validate(batch, lazy=True)
except pa.SchemaErrors as exc:
    for err in exc.schema_errors:
        print(err.reason_code, "|", err)

With `lazy=True`, Pandera collects every failure instead of stopping at the first one:

Copied to clipboard!
SchemaErrorReason.WRONG_DATATYPE | Key 'observation': expected dtype torch.float32, got torch.float64
SchemaErrorReason.CHECK_ERROR | Check '<Check isin: isin([0, 1, 2, 3])>' failed for key 'action': check failed

The reason codes are the same ones the dataframe backends use (`WRONG_DATATYPE`, `CHECK_ERROR`, `COLUMN_NOT_IN_DATAFRAME` for a missing key), so any error-handling code you've already written around Pandera keeps working. A `NaN` reward fails the `ge`/`le` range check as well, which means the third bug from the intro gets caught at the batch boundary instead of in the loss curve.

When the data is valid, `validate()` returns it, so it drops into a training loop or a collate function as a one-line guard. By default the returned object is a copy, because Pandera clones the TensorDict before checking it; pass `inplace=True` to skip the clone when you don't need it. If you'd rather cast than reject, set `coerce = True` in `Config`, and the `float64` observations come back as `float32` before any other checks run.

If you prefer to build schemas from objects instead of classes, `pa.TensorDictSchema` with `pa.Tensor` components is the equivalent:

Copied to clipboard!
schema = pa.TensorDictSchema(
    keys={
        "observation": pa.Tensor(dtype=torch.float32, shape=(None, 8)),
        "action": pa.Tensor(dtype=torch.int64, shape=(None,)),
    },
    batch_size=(None,),
)

Both forms validate plain `TensorDict`s, lazily stacked TensorDicts from `TensorDict.lazy_stack()`, and `@tensorclass` instances. That last one matters if you've moved to tensorclasses for the type hints, because you can keep them and validate against the same schema.

The rest of the Pandera toolkit, for tensors

The TensorDict backend isn't validation in isolation. The features that make Pandera useful for dataframes beyond `validate()` came along with it.

Schema inference. Point `pa.infer_schema()` at a known-good batch and you get a schema with dtypes, shapes, and min/max range checks for numeric tensors:

Copied to clipboard!
inferred = pa.infer_schema(good_batch)
print(pa.to_yaml(inferred))

The inferred schema pins every dimension to the exact size it saw, including the batch size, and its range checks are the observed min and max. Treat it as a first draft: replace the batch dimension with `None` and widen the ranges to what the data is actually allowed to be, not what one sample happened to contain.

Serialization. Schemas round-trip through YAML and JSON with `pa.to_yaml` / `pa.from_yaml` and `pa.to_json` / `pa.from_json`, so a schema can live in version control next to the training config and be loaded by a different process. There's also `pa.save(schema, td, path)`, which writes a `.pt` file with the schema embedded, and `pa.load(path)`, which loads the TensorDict and validates it against that embedded schema before handing it back. `pa.load` uses `torch.load(..., weights_only=False)` under the hood, so only load files from sources you trust.

Synthetic data for tests. With `pandera[strategies]` installed, `tensordict_strategy(schema)` and `tensorclass_strategy(cls, schema)` turn a schema into a Hypothesis strategy that generates valid batches, with value checks like `isin` and `in_range` constraining what gets generated. That gives you property-based tests for the code that consumes a batch, such as a policy network, a replay buffer, or a normalization step:

Copied to clipboard!
from hypothesis import given, settings
from pandera import Check
from pandera.strategies import tensordict_strategy

obs_schema = pa.TensorDictSchema(
    keys={
        "observation": pa.Tensor(
            dtype=torch.float32, shape=(None, 8), checks=Check.in_range(-100.0, 100.0)
        ),
    },
    batch_size=(32,),
)

policy = torch.nn.Linear(8, 4)

@given(tensordict_strategy(obs_schema))
@settings(max_examples=50)
def test_policy_outputs_are_finite(td):
    logits = policy(td["observation"])
    assert logits.shape == (32, 4)
    assert torch.isfinite(logits).all()

The same schema that guards production batches now generates the test fixtures, so the two can't drift apart.

Limitations

This is the first release of the TensorDict backend, and a few things you might expect aren't there yet:

  • No nested keys. TensorDicts can nest, and TorchRL uses that for things like `("next", "observation")`. Schemas currently cover top-level keys only.
  • Check errors don't point at elements. A failed check tells you which key and which check failed, but not which indices in the tensor violated it, the way a dataframe failure-case report does.
  • Custom checks need the object API. `pa.Tensor(checks=[Check(...)])` works, but `TensorDictModel` fields don't accept custom checks yet, so a schema with a one-off check (a valid one-hot row, an ignore-index label) has to be a `TensorDictSchema`.
  • No `@check_types` integration. The function-decorator workflow with typed annotations (`DataFrame[Schema]`) isn't wired up for TensorDicts, so you call `validate()` explicitly.

If any of these is blocking you, open an issue, since that's the best signal for what to build next.

Also in 0.34

Alongside the TensorDict backend, 0.34.0 includes nested `DataFrameModel` validation for Polars, a batch of Polars fixes (parsers that fail loudly instead of being skipped, `ignore_na=False` honored when a check returns null, missing-column defaults handled correctly), integer coercion overflow reported as an error instead of wrapping silently, and a fix for string-valued `validation_depth` settings that were silently disabling gating. The full list is in the release notes, and most of it came from first-time contributors, so thank you to everyone who sent a PR.

Conclusion

Pandera 0.34 extends the same schema you'd write for a dataframe to the batches that actually go into a PyTorch model: keys, dtypes, shapes, and value checks on a `TensorDict` or `tensorclass`, with lazy error reports, inference, YAML/JSON serialization, and Hypothesis-backed test data. The bugs it catches are the quiet ones that otherwise surface an hour into a training run.

Install it with `pip install -U 'pandera[torch]'` and start with the PyTorch guide.

Source code

Each benchmark in this post is a Flyte script you can run yourself:

Sign up

30 day free trial

Try the devbox
No items found.