ModelTensorDefinition
Represents one input/output tensor signature for a model.
Not directly serializable
Although registered as a BaseDataType, its to_pyarrow/from_pyarrow are batch-oriented classmethods (a list[ModelTensorDefinition] in, one StructArray out) rather than the standard per-instance contract. ModelTensorDefinition is not meant to be passed directly to datatypes.serialize() -- it's serialized only as an element of ModelDefinitions.model_inputs/model_outputs.
Parameters
| Field | Type | Description |
|---|---|---|
name | str | Tensor name. |
canonical_name | str | Canonical tensor name. |
dtype | str | Tensor element dtype (e.g. "float32"). |
shape | list[int] | tuple[int, ...] | Tensor dimensions, e.g. [batch, channel, height, width]. -1 marks a dynamic dimension. |
Raises
| Exception | Condition |
|---|---|
TypeError | name/canonical_name/dtype is not a str, shape is not a list/tuple, or a shape entry is not an int. |
Attributes
| Attribute | Type | Description |
|---|---|---|
name | str | Tensor name. |
canonical_name | str | Canonical tensor name. |
dtype | str | Tensor element dtype. |
shape | list[int] | Tensor shape. |
Stored as plain __slots__ attributes rather than validated properties — unlike the other datatypes in this package, assigning directly to name/canonical_name/dtype/shape after construction does not re-run validation.
Methods
| Method | Description |
|---|---|
to_dict() | Returns {"name", "canonical_name", "dtype", "shape"} as a plain dict. |
ModelTensorDefinition.from_dict(value) | Classmethod. Builds an instance from a dict with exactly those four keys. Raises TypeError if value isn't a dict or a field has the wrong type; ValueError if a key is missing. |
ModelTensorDefinition.coerce(value, name="value") | Classmethod. Returns value unchanged if it's already a ModelTensorDefinition; converts a dict via from_dict. Raises TypeError otherwise. |
ModelTensorDefinition.arrow_type() | Static method. Returns the pa.StructType (name, canonical_name, dtype, shape: list<int32>) used to serialize a batch. |
ModelTensorDefinition.to_pyarrow(values) | Classmethod. Serializes a list[ModelTensorDefinition] (may be empty) to one pa.StructArray matching arrow_type(). |
ModelTensorDefinition.from_pyarrow(array) | Classmethod. Inverse of to_pyarrow: deserializes a pa.StructArray into list[ModelTensorDefinition], in order. |
Operators
| Operation | Behavior |
|---|---|
a == b | True if other is a ModelTensorDefinition with equal name, canonical_name, dtype, and shape; NotImplemented for any other type. |
Not hashable (__hash__ is None).
Example
ModelTensorDefinition has no standalone example file — it's demonstrated below as the model_inputs/model_outputs entries of ModelDefinitions:
"""Demonstrates the Telekinesis ModelDefinitions datatype."""
import time
from datetime import datetime, timezone
import rerun as rr
from loguru import logger
from telekinesis import datatypes
def model_definitions_example():
"""Demonstrate batch construction (canonical_name input_0/output_N), access, indexing, empty batches, and serialization."""
# ======================= Create ============================================
created_at = datetime(2024, 6, 1, tzinfo=timezone.utc)
updated_at = datetime(2024, 6, 15, tzinfo=timezone.utc)
model_input = datatypes.ModelTensorDefinition(
name="images", canonical_name="input_0", dtype="float32", shape=[1, 3, 224, 224]
)
model_output = datatypes.ModelTensorDefinition(
name="logits", canonical_name="output_0", dtype="float32", shape=[1, 1000]
)
other_output = datatypes.ModelTensorDefinition(
name="logits", canonical_name="output_0", dtype="float32", shape=[1, 1000]
)
definitions = datatypes.ModelDefinitions(
model_names=["model-a", "model-b"],
model_formats=["onnx", "pytorch"],
visibilities=["private", "public"],
model_statuses=["uploaded", "deploying"],
model_descriptions=["first model", None],
model_inputs=[[model_input], None],
model_outputs=[[model_output, other_output], None],
created_ats=[created_at, updated_at],
updated_ats=[created_at, updated_at],
)
# ======================= Visualize =========================================
rr.init("model_definitions_example", spawn=True)
datatypes.visualize(definitions, entity_path="/ModelDefinitions")
# ======================= Inspect ===========================================
logger.info(f"Number of records: {len(definitions)}")
logger.info(f"model_names={definitions.model_names}, model_formats={definitions.model_formats}")
logger.info(f"visibilities={definitions.visibilities}, model_statuses={definitions.model_statuses}")
logger.info(f"model_inputs={definitions.model_inputs}")
logger.info(f"created_ats={definitions.created_ats}, updated_ats={definitions.updated_ats}")
# ======================= Index =============================================
first = definitions[0]
subset = definitions[0:1]
mask = definitions.model_statuses == "uploaded"
uploaded_only = definitions[mask]
logger.info(f"definitions[0] = {first}")
logger.info(f"definitions[0:1] = {len(subset)} record(s), names={subset.model_names}")
logger.info(f"definitions[uploaded mask] = {len(uploaded_only)} record(s)")
# ======================= Empty Batch =======================================
empty = datatypes.ModelDefinitions(
model_names=[], model_formats=[], visibilities=[], model_statuses=[]
)
datatypes.visualize(empty, entity_path="/ModelDefinitions/empty")
# ======================= Serialize / Deserialize ===========================
start = time.perf_counter()
serialized = datatypes.serialize(definitions)
serialization_ms = (time.perf_counter() - start) * 1000
start = time.perf_counter()
deserialized = datatypes.deserialize(serialized)["param_0"]
deserialization_ms = (time.perf_counter() - start) * 1000
logger.info(f"Deserialized ModelDefinitions: {deserialized}")
logger.info(f"Round-trip successful: {deserialized == definitions}")
logger.info(f"Serialization time: {serialization_ms:.3f} ms")
logger.info(f"Deserialization time: {deserialization_ms:.3f} ms")
if __name__ == "__main__":
model_definitions_example()
