toxr / src /toxarrow /materialize /materializer.py
vedatonuryilmaz's picture
Upload src/toxarrow/materialize/materializer.py
0ec92f0 verified
Raw
History Blame Contribute Delete
2.7 kB
"""Core materializer — dispatch TensorSpec to the correct materialization function."""
from __future__ import annotations
from typing import Any
import pyarrow as pa
from toxarrow.compile.arrow_store import ArrowStudyStore
from toxarrow.schemas.tensor import (
DenseGridSpec,
GraphBundleSpec,
MaskedGridSpec,
MultiViewSpec,
SequenceSpec,
TensorSpec,
)
class TensorMaterializer:
"""Materialize tensor batches from an ArrowStudyStore.
Usage:
store = ArrowStudyStore.from_records(records)
spec = DenseGridSpec(...)
materializer = TensorMaterializer(store)
batch = materializer.materialize(spec)
xt = batch["values"] # numpy array
"""
def __init__(self, store: ArrowStudyStore):
"""Initialize with an ArrowStudyStore.
Args:
store: ArrowStudyStore containing canonicalized study tables.
"""
self.store = store
def materialize(self, spec: TensorSpec) -> dict[str, Any]:
"""Materialize a tensor batch according to the given spec.
Dispatches to the appropriate materialization function based on spec type.
Args:
spec: A TensorSpec instance (DenseGridSpec, SequenceSpec, etc.)
Returns:
Dict with tensor arrays, coordinate metadata, and masks.
Raises:
ValueError: If the observations table is not in the store.
"""
obs = self.store.observations
if obs is None:
raise ValueError("ArrowStudyStore has no observations table. Run from_records() first.")
units = self.store.units
if isinstance(spec, DenseGridSpec):
# Check if it's actually a MaskedGridSpec (subclass)
if isinstance(spec, MaskedGridSpec):
from toxarrow.specs.masked_grid import materialize_masked_grid
return materialize_masked_grid(obs, spec, units)
from toxarrow.specs.dense_grid import materialize_dense_grid
return materialize_dense_grid(obs, spec, units)
elif isinstance(spec, SequenceSpec):
from toxarrow.specs.sequence import materialize_sequence
return materialize_sequence(obs, spec, units)
elif isinstance(spec, GraphBundleSpec):
from toxarrow.specs.graph_bundle import materialize_graph_bundle
return materialize_graph_bundle(obs, spec, units)
elif isinstance(spec, MultiViewSpec):
from toxarrow.specs.multiview import materialize_multiview
return materialize_multiview(obs, spec, units)
else:
raise TypeError(f"Unsupported TensorSpec type: {type(spec).__name__}")