michelangelo.lib.trainer.torch.data_collate_functions
Collate helpers for Ray Data / PyTorch training.
This module exposes small building blocks so callers can compose custom collate functions:
DEFAULT_COLLATE_NUMPY_DTYPE/DEFAULT_COLLATE_TORCH_DTYPE— default dtypes (float32unless overridden via function orLiteralEvalFloat32Collatekwargs).pad_ragged_lists— pad nested Python lists to a dense array of numpy_dtype.cell_is_nested_subsequence/row_is_list_of_nested_cells— structure checks.collate_value_to_float32_numpy— one feature column →numpy.ndarray.collate_value_to_float32_tensor— one feature column →torch.Tensor.collate_batch_to_float32_tensors— full batch dict → tensors.
The default literal_eval_data_collate_function is implemented on top of these.
LiteralEvalFloat32Collate wraps the same behavior for subclassing (custom device, hooks).
Example:
Using the default collate with a PyTorch DataLoader:
from torch.utils.data import DataLoader
from michelangelo.lib.trainer.torch.data_collate_functions import (
literal_eval_data_collate_function,
)
loader = DataLoader(
dataset,
batch_size=32,
collate_fn=literal_eval_data_collate_function,
)
cell_is_nested_subsequence
def cell_is_nested_subsequence(cell) -> bool
Return True if cell is a vector-valued slot (list/tuple or ndarray with ndim >= 1).
Scalars and 0-D ndarrays are leaves for the 2-D-ragged path (one flat vector per row).
row_is_list_of_nested_cells
def row_is_list_of_nested_cells(flat0: list | np.ndarray) -> bool
Return True when flat0 is a row of cells where at least one cell is a sub-sequence (3-D path).
Uses every cell, not only flat0[0], so a leading scalar with later list cells still
selects the 3-D normalization branch.
pad_ragged_lists
def pad_ragged_lists(items: list,
pad_value: float | None = None,
*,
numpy_dtype: np.dtype | None = None) -> np.ndarray
Pad nested lists to a rectangular array of numpy_dtype (default: DEFAULT_COLLATE_NUMPY_DTYPE).
collate_value_to_float32_numpy
def collate_value_to_float32_numpy(
value,
*,
reshape_1d_features: bool = True,
parse_string_with_literal_eval: bool = True,
numpy_dtype: np.dtype | None = None) -> np.ndarray
Convert a single batch column value to a numpy.ndarray of numpy_dtype.
collate_value_to_float32_tensor
def collate_value_to_float32_tensor(
value,
*,
device: str | torch.device = "cpu",
reshape_1d_features: bool = True,
parse_string_with_literal_eval: bool = True,
numpy_dtype: np.dtype | None = None) -> torch.Tensor
Convert one column value to torch.Tensor on device (see collate_value_to_float32_numpy).
collate_batch_to_float32_tensors
def collate_batch_to_float32_tensors(
batch_data: dict,
*,
device: str | torch.device = "cpu",
reshape_1d_features: bool = True,
parse_string_with_literal_eval: bool = True,
numpy_dtype: np.dtype | None = None) -> dict[str, torch.Tensor]
Map a batch dict of Python / NumPy values to tensors (default element dtype: float32).
LiteralEvalFloat32Collate Objects
class LiteralEvalFloat32Collate()
Default collate with ast.literal_eval for stringified arrays.
__init__
def __init__(*,
device: str | torch.device = "cpu",
reshape_1d_features: bool = True,
parse_string_with_literal_eval: bool = True,
numpy_dtype: np.dtype | None = None) -> None
Initialize the collate.
Arguments:
device- Target device for emitted tensors.reshape_1d_features- If True, scalar features are reshaped to(N, 1).parse_string_with_literal_eval- If True, string-encoded arrays are decoded viaast.literal_eval.numpy_dtype- Optional numpy dtype to cast numeric values to before tensor conversion; defaults toDEFAULT_COLLATE_NUMPY_DTYPE.
collate_value_to_numpy
def collate_value_to_numpy(value) -> np.ndarray
Convert one column value to numpy.ndarray (override in subclasses).
collate_value_to_tensor
def collate_value_to_tensor(value) -> torch.Tensor
Convert one column value to torch.Tensor on device.
collate_batch
def collate_batch(batch_data: dict) -> dict[str, torch.Tensor]
Map a batch dict to tensors (override for per-key routing).
__call__
def __call__(batch_data: dict) -> dict[str, torch.Tensor]
Delegate to collate_batch.
literal_eval_data_collate_function
def literal_eval_data_collate_function(
batch_data: dict) -> dict[str, torch.Tensor]
Convert processed batch data to tensors (default training collate).