Skip to content

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

FieldTypeDescription
namestrTensor name.
canonical_namestrCanonical tensor name.
dtypestrTensor element dtype (e.g. "float32").
shapelist[int] | tuple[int, ...]Tensor dimensions, e.g. [batch, channel, height, width]. -1 marks a dynamic dimension.

Raises

ExceptionCondition
TypeErrorname/canonical_name/dtype is not a str, shape is not a list/tuple, or a shape entry is not an int.

Attributes

AttributeTypeDescription
namestrTensor name.
canonical_namestrCanonical tensor name.
dtypestrTensor element dtype.
shapelist[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

MethodDescription
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

OperationBehavior
a == bTrue 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:

python
"""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()