What is TensorDict?
mainTensorDict is a dictionary-like class that inherits properties from tensors, such as indexing, shape operations, and casting to device.
Its primary purpose is to increase code readability and modularity by abstracting away tailored operations. This allows you to write generic training loops that can handle highly heterogeneous tasks (e.g., switching between classification and segmentation) because the model, loss module, and optimizer all interact with a unified TensorDict object rather than individual tensors.
# Example of a generic training loop using TensorDict
for i, tensordict in enumerate(dataset):
# the model reads and writes tensordicts
tensordict = model(tensordict)
loss = loss_module(tensordict)
loss.backward()
optimizer.step()
optimizer.zero_grad()