toxr / src /toxarrow /materialize /padding.py
vedatonuryilmaz's picture
Upload src/toxarrow/materialize/padding.py
1fc3c8a verified
Raw
History Blame Contribute Delete
1.82 kB
"""Padding utilities — pad variable-length sequences to a common length."""
from __future__ import annotations
import numpy as np
def pad_sequences(
sequences: list[np.ndarray],
max_length: int | None = None,
pad_value: float = 0.0,
dtype: np.dtype = np.float32,
) -> tuple[np.ndarray, np.ndarray]:
"""Pad a list of 2D sequences to a common length.
Args:
sequences: List of (T_i × F) arrays.
max_length: Target sequence length. If None, uses max observed.
pad_value: Fill value for padded positions.
dtype: Output dtype.
Returns:
(padded_values, padding_mask) where:
padded_values: (S × max_length × F) padded tensor.
padding_mask: (S × max_length) boolean mask (True = padded).
"""
if not sequences:
return np.array([], dtype=dtype), np.array([], dtype=bool)
S = len(sequences)
F = sequences[0].shape[1]
max_len = max_length or max(s.shape[0] for s in sequences)
padded = np.full((S, max_len, F), pad_value, dtype=dtype)
mask = np.ones((S, max_len), dtype=bool)
for i, seq in enumerate(sequences):
length = min(seq.shape[0], max_len)
padded[i, :length, :] = seq[:length, :]
mask[i, :length] = False
return padded, mask
def unpad_sequences(
padded_values: np.ndarray,
padding_mask: np.ndarray,
) -> list[np.ndarray]:
"""Extract ragged sequences from a padded tensor.
Args:
padded_values: (S × T_max × F) padded tensor.
padding_mask: (S × T_max) boolean mask (True = padded).
Returns:
List of (T_i × F) arrays.
"""
sequences = []
for s in range(padded_values.shape[0]):
valid = ~padding_mask[s]
sequences.append(padded_values[s, valid, :])
return sequences