This is an automated email from the ASF dual-hosted git repository.
JingsongLi pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/paimon.git
The following commit(s) were added to refs/heads/master by this push:
new d5d42d4551 [python] Support window joins in temporal alignment (#9760)
d5d42d4551 is described below
commit d5d42d4551f52e2e3417be03e13d1811f8ae11fa
Author: XiaoHongbo <[email protected]>
AuthorDate: Sun Sep 13 16:29:27 2026 +0800
[python] Support window joins in temporal alignment (#9760)
---
docs/docs/pypaimon/multimodal-reading.md | 23 +
paimon-python/pypaimon/multimodal/__init__.py | 2 +
paimon-python/pypaimon/multimodal/temporal.py | 328 +++++++++++-
.../pypaimon/tests/multimodal_temporal_test.py | 573 +++++++++++++++++++++
4 files changed, 910 insertions(+), 16 deletions(-)
diff --git a/docs/docs/pypaimon/multimodal-reading.md
b/docs/docs/pypaimon/multimodal-reading.md
index adb32b75b8..234dd4f6e2 100644
--- a/docs/docs/pypaimon/multimodal-reading.md
+++ b/docs/docs/pypaimon/multimodal-reading.md
@@ -102,6 +102,29 @@ states_at_steps = aligned.interpolate(
)
```
+Use `join_window(left, right, ...)` directly or chain
+`aligned.join_window(...)`. Unlike single-table `rolling().agg()`, it joins
+each left row to right rows in the same `by` group and
+`[left - preceding, left + following]` time window, then aggregates them.
+`following` defaults to zero. `mean`, `min`, and `max` require integer or
+floating-point scalars; `count`, `first`, and `last` also accept non-numeric
+values. `closed` is `both`, `left`, `right`, or `neither` and refers to the
+interval endpoints. Nulls are skipped; empty windows return null, except
+`count` returns zero. Use `(source, operation)` pairs to aggregate one source
+column more than once; `{"value": "mean"}` remains valid shorthand.
+
+```python
+steps_with_imu = aligned.join_window(
+ imu.scan().select("acceleration"),
+ preceding=timedelta(milliseconds=50),
+ following=timedelta(milliseconds=50),
+ aggregations={
+ "acceleration_mean": ("acceleration", "mean"),
+ "acceleration_max": ("acceleration", "max"),
+ },
+)
+```
+
### Reading BLOB columns
`scan().read_blobs(column)` bulk-fetches a BLOB column's bytes for the filtered
diff --git a/paimon-python/pypaimon/multimodal/__init__.py
b/paimon-python/pypaimon/multimodal/__init__.py
index 1807ba4d6f..00d59d5c71 100644
--- a/paimon-python/pypaimon/multimodal/__init__.py
+++ b/paimon-python/pypaimon/multimodal/__init__.py
@@ -46,6 +46,7 @@ from pypaimon.multimodal.temporal import (
TemporalAlignment,
interpolate,
join_asof,
+ join_window,
)
from pypaimon.multimodal.video import VideoFrameCollator
from pypaimon.table.row.blob import Blob, BlobDescriptor, VideoFrameDescriptor
@@ -79,6 +80,7 @@ __all__ = [
"connect",
"interpolate",
"join_asof",
+ "join_window",
"lit",
"source_col",
"target_col",
diff --git a/paimon-python/pypaimon/multimodal/temporal.py
b/paimon-python/pypaimon/multimodal/temporal.py
index 43de0a8ed1..bb42053ef3 100644
--- a/paimon-python/pypaimon/multimodal/temporal.py
+++ b/paimon-python/pypaimon/multimodal/temporal.py
@@ -18,6 +18,7 @@
from bisect import bisect_left, bisect_right
from datetime import timedelta
+from fractions import Fraction
import json
import math
from numbers import Integral, Real
@@ -74,6 +75,21 @@ def interpolate(left, right, *, on, by, tolerance=None,
)
+def join_window(left, right, *, on, by, preceding, aggregations,
+ following=None, closed="both", right_on=None,
+ suffix="_right") -> "TemporalAlignment":
+ """Join and aggregate right values in each left row's time window."""
+ return TemporalAlignment(left, on=on, by=by).join_window(
+ right,
+ preceding=preceding,
+ following=following,
+ aggregations=aggregations,
+ closed=closed,
+ right_on=right_on,
+ suffix=suffix,
+ )
+
+
def _normalize_temporal_keys(on, by):
if not isinstance(on, str) or not on:
raise ValueError("on must be a non-empty column name.")
@@ -140,6 +156,25 @@ class TemporalAlignment:
)
return self._append(source)
+ def join_window(self, right, *, preceding, aggregations,
+ following=None, closed="both", right_on=None,
+ suffix="_right") -> "TemporalAlignment":
+ """Append a right-side window join with aggregation."""
+ position = len(self._sources) + 1
+ source = _WindowJoinRight(
+ "right source %d" % position,
+ right,
+ self._on,
+ self._by,
+ preceding,
+ following,
+ aggregations,
+ closed,
+ right_on,
+ suffix,
+ )
+ return self._append(source)
+
def _append(self, source):
result = object.__new__(TemporalAlignment)
result._anchor = self._anchor
@@ -168,7 +203,8 @@ class TemporalAlignment:
for source in self._sources:
source.plan()
source_fetchers.append(
- _RowIdFetcher(source.query, row_group_cache))
+ _RowIdFetcher(
+ source.query, row_group_cache, source._fetch_names))
schema = self._output_schema(anchor_fetcher.schema, source_fetchers)
self.schema = schema
@@ -250,11 +286,9 @@ class TemporalAlignment:
source.payload_schema if source_fetchers is None
else source_fetchers[position].schema
)
- for name in source.payload_schema.names:
- field = source.output_field(
- payload_schema.field(name),
- effective=source_fetchers is not None,
- )
+ for field in source.output_fields(
+ payload_schema,
+ effective=source_fetchers is not None):
output_name = field.name
if output_name in names:
output_name += source.suffix
@@ -333,6 +367,7 @@ class _AsOfJoinRight:
field for field, path in zip(schema, paths)
if tuple(path) not in excluded
])
+ self._fetch_names = None
self._index = None
def plan(self):
@@ -373,6 +408,12 @@ class _AsOfJoinRight:
def output_field(field, effective=True):
return field
+ def output_fields(self, payload_schema, effective=True):
+ return [
+ self.output_field(payload_schema.field(name), effective)
+ for name in self.payload_schema.names
+ ]
+
def build_arrays(self, anchor_rows, fetcher):
matches = [self.match(row) for row in anchor_rows]
matched_ids = [match for match in matches if match is not None]
@@ -484,6 +525,235 @@ class _LinearInterpolationRight(_AsOfJoinRight):
return arrays
+class _WindowJoinRight(_AsOfJoinRight):
+
+ _SUPPORTED_AGGREGATIONS = {
+ "count", "first", "last", "max", "mean", "min",
+ }
+
+ def __init__(self, label, query, anchor_on, by, preceding, following,
+ aggregations, closed, right_on, suffix):
+ super().__init__(
+ label, query, anchor_on, by, "nearest", None,
+ right_on, suffix)
+ self._preceding_key = _window_bound_key(
+ "preceding", preceding, self.time_type)
+ if following is None:
+ following = (
+ timedelta(0) if pa.types.is_timestamp(self.time_type) else 0)
+ self._following_key = _window_bound_key(
+ "following", following, self.time_type)
+ if closed not in ("both", "left", "neither", "right"):
+ raise ValueError(
+ "closed must be 'both', 'left', 'right', or 'neither'.")
+ self.closed = closed
+ self.aggregations = _normalize_aggregations(
+ aggregations, self.payload_schema, self.label,
+ self._SUPPORTED_AGGREGATIONS)
+ source_names = {
+ specification[1] for specification in self.aggregations
+ }
+ self._fetch_names = tuple(
+ field.name for field in self.payload_schema
+ if field.name in source_names)
+ self.payload_schema = pa.schema([
+ field for field in self.payload_schema
+ if field.name in source_names
+ ], metadata=self.payload_schema.metadata)
+
+ def output_fields(self, payload_schema, effective=True):
+ fields = []
+ for output_name, source_name, aggregation in self.aggregations:
+ source = payload_schema.field(source_name)
+ try:
+ output_type = _aggregate_output_type(
+ source.type, aggregation)
+ except TypeError:
+ if effective:
+ raise
+ output_type = source.type
+ fields.append(pa.field(
+ output_name, output_type, nullable=True,
+ metadata=source.metadata))
+ return fields
+
+ def match(self, anchor_row):
+ key = tuple(anchor_row[name] for name in self.by)
+ bounds = self._index.get(key)
+ if bounds is None:
+ return []
+ target = anchor_row[_TIME_KEY]
+ start, end = bounds
+ left = target - self._preceding_key
+ right = target + self._following_key
+ if (pa.types.is_integer(self.time_type)
+ or pa.types.is_timestamp(self.time_type)):
+ first_key = (
+ math.ceil(left)
+ if self.closed in ("both", "left")
+ else math.floor(left) + 1
+ )
+ last_key = (
+ math.floor(right)
+ if self.closed in ("both", "right")
+ else math.ceil(right) - 1
+ )
+ # Avoid comparing NumPy keys with out-of-range Python integers.
+ first_key = max(first_key, _python_scalar(self._time_keys[start]))
+ last_key = min(last_key, _python_scalar(self._time_keys[end - 1]))
+ if first_key > last_key:
+ return []
+ first = bisect_left(
+ self._time_keys, first_key, start, end)
+ last = bisect_right(
+ self._time_keys, last_key, first, end)
+ else:
+ first = (
+ bisect_left(self._time_keys, left, start, end)
+ if self.closed in ("both", "left")
+ else bisect_right(self._time_keys, left, start, end)
+ )
+ last = (
+ bisect_right(self._time_keys, right, first, end)
+ if self.closed in ("both", "right")
+ else bisect_left(self._time_keys, right, first, end)
+ )
+ return [self._row_ids[index].as_py()
+ for index in range(first, last)]
+
+ def build_arrays(self, anchor_rows, fetcher):
+ matches = [self.match(row) for row in anchor_rows]
+ unique_ids = list(dict.fromkeys(
+ row_id for match in matches for row_id in match))
+ values = fetcher.fetch(unique_ids)
+ positions = {
+ row_id: index for index, row_id in enumerate(unique_ids)
+ }
+ indices = [
+ [positions[row_id] for row_id in match]
+ for match in matches
+ ]
+ arrays = []
+ for _, source_name, aggregation in self.aggregations:
+ effective = fetcher.schema.field(source_name)
+ output_type = _aggregate_output_type(
+ effective.type, aggregation)
+ arrays.append(pa.array([
+ _aggregate_values(
+ values[source_name], row_indices, aggregation)
+ for row_indices in indices
+ ], type=output_type))
+ return arrays
+
+
+def _normalize_aggregations(aggregations, schema, label, supported):
+ if not isinstance(aggregations, dict) or not aggregations:
+ raise ValueError("aggregations must be a non-empty dict.")
+ normalized = []
+ missing = []
+ for output_name, specification in aggregations.items():
+ if not isinstance(output_name, str) or not output_name:
+ raise ValueError(
+ "Aggregation output names must be non-empty strings.")
+ if isinstance(specification, str):
+ source_name = output_name
+ aggregation = specification
+ elif isinstance(specification, tuple) and len(specification) == 2:
+ source_name, aggregation = specification
+ else:
+ raise ValueError(
+ "Aggregation %r must be an operation or a "
+ "(source column, operation) pair." % output_name)
+ if not isinstance(source_name, str) or not source_name:
+ raise ValueError(
+ "Aggregation source columns must be non-empty strings.")
+ if source_name not in schema.names:
+ missing.append(source_name)
+ if not isinstance(aggregation, str) or aggregation not in supported:
+ raise ValueError(
+ "Unsupported aggregation %r for output %r; expected one of "
+ "%r." % (aggregation, output_name, sorted(supported)))
+ normalized.append((output_name, source_name, aggregation))
+ if missing:
+ raise ValueError(
+ "%s is missing aggregation columns %r." % (label, missing))
+ return tuple(normalized)
+
+
+def _aggregate_output_type(data_type, aggregation):
+ if aggregation == "count":
+ return pa.int64()
+ if aggregation in ("first", "last"):
+ return data_type
+ if not (pa.types.is_integer(data_type)
+ or pa.types.is_floating(data_type)):
+ raise TypeError(
+ "Window %s aggregation requires an integer or floating-point "
+ "scalar column; got %s." % (aggregation, data_type))
+ if aggregation == "mean":
+ return pa.float64()
+ return data_type
+
+
+def _aggregate_values(values, indices, aggregation):
+ if not indices:
+ return 0 if aggregation == "count" else None
+ selected = pc.take(values, pa.array(indices, type=pa.int64()))
+ if aggregation == "count":
+ return pc.count(selected).as_py()
+ if aggregation == "mean":
+ items = [item for item in selected.to_pylist()
+ if item is not None]
+ if not items:
+ return None
+ if pa.types.is_integer(values.type):
+ return sum(items) / len(items)
+ if not all(math.isfinite(item) for item in items):
+ return pc.mean(selected).as_py()
+ try:
+ return math.fsum(items) / len(items)
+ except OverflowError:
+ pass
+ scale = max(abs(item) for item in items)
+ if scale == 0:
+ return 0.0
+ return (math.fsum(item / scale for item in items) / len(items)) * scale
+ if aggregation == "min":
+ return pc.min(selected).as_py()
+ if aggregation == "max":
+ return pc.max(selected).as_py()
+ items = selected.to_pylist()
+ if aggregation == "first":
+ return next((item for item in items if item is not None), None)
+ return next((item for item in reversed(items) if item is not None), None)
+
+
+def _window_bound_key(name, value, data_type):
+ if isinstance(value, bool) or not isinstance(value, (Real, timedelta)):
+ raise TypeError(
+ "%s must be numeric or datetime.timedelta." % name)
+ if isinstance(value, Real):
+ value = _python_scalar(value)
+ if isinstance(value, Integral):
+ value = int(value)
+ if (isinstance(value, Real) and not isinstance(value, Integral)
+ and not math.isfinite(value)):
+ raise ValueError("%s must be finite." % name)
+ zero = timedelta(0) if isinstance(value, timedelta) else 0
+ if value < zero:
+ raise ValueError("%s must be non-negative." % name)
+ _validate_tolerance(value, data_type)
+ if pa.types.is_integer(data_type) and not isinstance(value, Integral):
+ try:
+ exact = Fraction(value)
+ except TypeError:
+ exact = Fraction(float(value))
+ if exact.denominator == 1:
+ return exact.numerator
+ return exact
+ return _time_tolerance_key(value, data_type)
+
+
def _validate_join_options(direction, tolerance, right_on, suffix):
if direction not in ("backward", "forward", "nearest"):
raise ValueError(
@@ -769,13 +1039,29 @@ def _validate_metadata(query, metadata, key_columns):
class _RowIdFetcher:
- def __init__(self, query, row_group_cache):
+ def __init__(self, query, row_group_cache, output_names=None):
_validate_pinned_tag(query)
- self._schema = _query_schema(query)
+ query_schema, query_paths = _query_schema_and_paths(query)
+ visible_projection = query._effective_projection()
+ if output_names is None:
+ self._schema = query_schema
+ visible_paths = query_paths
+ else:
+ output_names = set(output_names)
+ selected = [
+ (field, path)
+ for field, path in zip(query_schema, query_paths)
+ if field.name in output_names
+ ]
+ self._schema = pa.schema(
+ [field for field, unused in selected],
+ metadata=query_schema.metadata,
+ )
+ visible_paths = [path for unused, path in selected]
+ visible_projection = [".".join(path) for path in visible_paths]
table = query._table.copy_without_time_travel({
CoreOptions.BLOB_AS_DESCRIPTOR.key(): "true",
})
- visible_projection = query._effective_projection()
plan_builder = table.new_read_builder()
if visible_projection is not None:
plan_projection = visible_projection
@@ -804,8 +1090,10 @@ class _RowIdFetcher:
projected_builder.read_type())
projected_paths = projected_builder._nested_name_paths()
if projected_paths is not None:
+ table_names = set(_table_schema(query).names)
for field, path in zip(projected_schema, projected_paths):
- if field.name in masking and field.name != path[0]:
+ if (field.name in masking and field.name != path[0]
+ and field.name not in table_names):
raise ValueError(
"Temporal alignment cannot safely apply column "
"masking to nested projection %r."
@@ -847,10 +1135,6 @@ class _RowIdFetcher:
builder.read_type())
effective_schema = _effective_masked_schema(
physical_schema, masking)
- visible_paths = (
- None if self._name_paths is None
- else self._name_paths[:len(self._schema)]
- )
self._schema = _project_effective_schema(
self._schema, visible_paths, effective_schema, masking)
self._fetch_schema = _project_effective_schema(
@@ -936,7 +1220,11 @@ class _RowIdFetcher:
)
take = pa.array(
[positions[row_id] for row_id in row_ids], type=pa.int64())
- return arrow.select(self._schema.names).take(take)
+ visible = pa.Table.from_arrays(
+ [arrow.column(index) for index in range(len(self._schema))],
+ schema=self._schema,
+ )
+ return visible.take(take)
def _find_splits(self, ranges):
split_indices = set()
@@ -1197,7 +1485,15 @@ def _time_search_keys(values, data_type):
def _time_tolerance_key(tolerance, data_type):
if tolerance is None or not pa.types.is_timestamp(data_type):
return tolerance
- return pa.scalar(tolerance, type=pa.duration(data_type.unit)).value
+ microseconds = (
+ (tolerance.days * 24 * 60 * 60 + tolerance.seconds) * 1_000_000
+ + tolerance.microseconds
+ )
+ divisors = {"s": 1_000_000, "ms": 1_000, "us": 1}
+ if data_type.unit == "ns":
+ return microseconds * 1_000
+ exact = Fraction(microseconds, divisors[data_type.unit])
+ return exact.numerator if exact.denominator == 1 else exact
def _python_scalar(value):
diff --git a/paimon-python/pypaimon/tests/multimodal_temporal_test.py
b/paimon-python/pypaimon/tests/multimodal_temporal_test.py
index 174352bf42..5ee34556d4 100644
--- a/paimon-python/pypaimon/tests/multimodal_temporal_test.py
+++ b/paimon-python/pypaimon/tests/multimodal_temporal_test.py
@@ -23,6 +23,7 @@ import unittest
from datetime import datetime, timedelta
from unittest import mock
+import numpy as np
import pyarrow as pa
import pypaimon.multimodal as pmm
from pypaimon.multimodal import temporal
@@ -188,6 +189,578 @@ class MultimodalTemporalTest(unittest.TestCase):
row["value"] for row in rows
])
+ def test_window_join_stays_in_group_and_skips_nulls(self):
+ anchors = self._table("window_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ })
+ samples = self._table("window_samples", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ "value": pa.int32(),
+ "label": pa.string(),
+ })
+ anchors.add([
+ {"episode_id": 1, "event_time": 10},
+ {"episode_id": 2, "event_time": 10},
+ {"episode_id": 3, "event_time": 10},
+ ])
+ samples.add([
+ {"episode_id": 1, "event_time": 5,
+ "value": 1, "label": "a"},
+ {"episode_id": 1, "event_time": 10,
+ "value": None, "label": None},
+ {"episode_id": 1, "event_time": 15,
+ "value": 5, "label": "c"},
+ {"episode_id": 2, "event_time": 10,
+ "value": 100, "label": "z"},
+ ])
+
+ result = pmm.join_window(
+ anchors.scan(),
+ samples.scan().select(["value", "label"]),
+ on="event_time",
+ by="episode_id",
+ preceding=5,
+ following=5,
+ aggregations={
+ "average": ("value", "mean"),
+ "minimum": ("value", "min"),
+ "maximum": ("value", "max"),
+ "first_label": ("label", "first"),
+ "last_label": ("label", "last"),
+ "valid_count": ("label", "count"),
+ },
+ )
+ rows = sorted(result.to_list(), key=lambda row: row["episode_id"])
+
+ self.assertIsInstance(result, pmm.TemporalAlignment)
+ self.assertEqual(pa.float64(), result.schema.field("average").type)
+ self.assertEqual(pa.string(), result.schema.field("first_label").type)
+ self.assertEqual(pa.int64(), result.schema.field("valid_count").type)
+ self.assertEqual(
+ (3.0, 1, 5, "a", "c", 2),
+ tuple(rows[0][name] for name in (
+ "average", "minimum", "maximum", "first_label",
+ "last_label", "valid_count")),
+ )
+ self.assertEqual(100.0, rows[1]["average"])
+ self.assertIsNone(rows[2]["average"])
+ self.assertEqual(0, rows[2]["valid_count"])
+
+ def test_window_join_supports_asymmetric_timestamp_bounds(self):
+ anchors = self._table("window_timestamp_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.timestamp("ms"),
+ })
+ samples = self._table("window_timestamp_samples", {
+ "episode_id": pa.int32(),
+ "captured_at": pa.timestamp("ms"),
+ "value": pa.float32(),
+ })
+ anchor = datetime(2026, 9, 1, 12, 0, 0)
+ anchors.add([{"episode_id": 1, "event_time": anchor}])
+ samples.add([
+ {"episode_id": 1,
+ "captured_at": anchor + timedelta(milliseconds=offset),
+ "value": value}
+ for offset, value in ((-11, 100.0), (-10, 1.0),
+ (0, 2.0), (5, 3.0), (6, 100.0))
+ ])
+
+ row = pmm.join_window(
+ anchors.scan(), samples.scan().select("value"),
+ on="event_time", right_on="captured_at", by="episode_id",
+ preceding=timedelta(milliseconds=10),
+ following=timedelta(milliseconds=5),
+ aggregations={"value": "mean"},
+ ).to_list()[0]
+
+ self.assertEqual(2.0, row["value"])
+
+ right_closed = pmm.join_window(
+ anchors.scan(), samples.scan().select("value"),
+ on="event_time", right_on="captured_at", by="episode_id",
+ preceding=timedelta(milliseconds=10),
+ following=timedelta(milliseconds=5),
+ aggregations={"value": "mean"},
+ closed="right",
+ ).to_list()[0]
+ self.assertEqual(2.5, right_closed["value"])
+
+ def test_window_join_preserves_subunit_timestamp_bounds(self):
+ anchors = self._table("window_subunit_timestamp_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.timestamp("ms"),
+ })
+ samples = self._table("window_subunit_timestamp_samples", {
+ "episode_id": pa.int32(),
+ "event_time": pa.timestamp("ms"),
+ "value": pa.int32(),
+ })
+ anchor = datetime(2026, 9, 1, 12, 0, 0)
+ anchors.add([{"episode_id": 1, "event_time": anchor}])
+ samples.add([
+ {"episode_id": 1,
+ "event_time": anchor + timedelta(milliseconds=offset),
+ "value": offset}
+ for offset in (-1, 0, 1)
+ ])
+
+ exact = pmm.join_window(
+ anchors.scan(), samples.scan().select("value"),
+ on="event_time", by="episode_id",
+ preceding=timedelta(microseconds=500),
+ following=timedelta(0), closed="right",
+ aggregations={"matches": ("value", "count")},
+ ).to_list()[0]
+ open_window = pmm.join_window(
+ anchors.scan(), samples.scan().select("value"),
+ on="event_time", by="episode_id",
+ preceding=timedelta(microseconds=1_500),
+ following=timedelta(microseconds=1_500), closed="neither",
+ aggregations={"matches": ("value", "count")},
+ ).to_list()[0]
+
+ self.assertEqual(1, exact["matches"])
+ self.assertEqual(3, open_window["matches"])
+
+ def test_window_join_normalizes_numpy_float_bounds(self):
+ anchors = self._table("window_numpy_float_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.float64(),
+ })
+ samples = self._table("window_numpy_float_samples", {
+ "episode_id": pa.int32(),
+ "event_time": pa.float64(),
+ "value": pa.int32(),
+ })
+ timestamp = 1_700_000_000.001
+ anchors.add([{"episode_id": 1, "event_time": timestamp}])
+ samples.add([{
+ "episode_id": 1, "event_time": timestamp, "value": 7,
+ }])
+
+ row = pmm.join_window(
+ anchors.scan(), samples.scan().select("value"),
+ on="event_time", by="episode_id",
+ preceding=0, following=np.float32(0),
+ aggregations={"matches": ("value", "count")},
+ ).to_list()[0]
+
+ self.assertEqual(1, row["matches"])
+
+ def test_window_join_keeps_numeric_bounds_exact_for_integer_time(self):
+ anchors = self._table("window_integer_bound_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ })
+ samples = self._table("window_integer_bound_samples", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ "value": pa.int32(),
+ })
+ timestamp = 1_700_000_000_000_000_001
+ anchors.add(pa.Table.from_pydict({
+ "episode_id": [1], "event_time": [timestamp],
+ }))
+ samples.add(pa.Table.from_pydict({
+ "episode_id": [1, 1, 1],
+ "event_time": [timestamp - 1, timestamp, timestamp + 1],
+ "value": [1, 2, 3],
+ }))
+
+ for preceding, following in (
+ (0.0, 0.0), (0.125, 0.125),
+ (np.int64(0), 0), (0, np.int64(0)),
+ (np.int64(0), np.int64(0))):
+ with self.subTest(preceding=preceding, following=following):
+ row = pmm.join_window(
+ anchors.scan(), samples.scan().select("value"),
+ on="event_time", by="episode_id",
+ preceding=preceding, following=following,
+ aggregations={
+ "matches": ("value", "count"),
+ "first_value": ("value", "first"),
+ },
+ ).to_list()[0]
+
+ self.assertEqual(1, row["matches"])
+ self.assertEqual(2, row["first_value"])
+
+ def test_window_join_clips_bounds_to_integer_time_range(self):
+ anchors = self._table("window_range_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ })
+ samples = self._table("window_range_samples", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ "value": pa.int32(),
+ })
+ minimum, maximum = -(1 << 63), (1 << 63) - 1
+ anchors.add(pa.Table.from_pydict({
+ "episode_id": [1, 2], "event_time": [minimum, maximum],
+ }))
+ times = {1: [minimum, minimum + 1], 2: [maximum - 1, maximum]}
+ samples.add(pa.Table.from_pydict({
+ "episode_id": [1, 1, 2, 2],
+ "event_time": times[1] + times[2],
+ "value": [1, 2, 3, 4],
+ }))
+
+ for preceding, following in ((0, 0), (0, 1), (1, 0), (1, 1)):
+ for closed in ("both", "left", "right", "neither"):
+ with self.subTest(
+ preceding=preceding, following=following,
+ closed=closed):
+ rows = pmm.join_window(
+ anchors.scan(), samples.scan().select("value"),
+ on="event_time", by="episode_id",
+ preceding=preceding, following=following,
+ closed=closed,
+ aggregations={"matches": ("value", "count")},
+ ).to_list()
+ for row in rows:
+ left = row["event_time"] - preceding
+ right = row["event_time"] + following
+ expected = sum(
+ (time >= left if closed in ("both", "left")
+ else time > left)
+ and (time <= right if closed in ("both", "right")
+ else time < right)
+ for time in times[row["episode_id"]]
+ )
+ self.assertEqual(expected, row["matches"])
+
+ def test_window_join_prunes_unaggregated_right_columns(self):
+ anchors = self._table("window_projection_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ })
+ samples = self._table("window_projection_samples", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ "payload": pa.struct([
+ pa.field("unused_nested", pa.int32()),
+ pa.field("value", pa.int32()),
+ ]),
+ "unused": pa.string(),
+ })
+ anchors.add([{"episode_id": 1, "event_time": 10}])
+ samples.add([{
+ "episode_id": 1,
+ "event_time": 10,
+ "payload": {"value": 7, "unused_nested": 8},
+ "unused": "not read",
+ }])
+ payload_reads = []
+ original = FormatPyArrowReader._read_parquet_row_group_batches
+
+ def tracked(reader, row_group, columns):
+ if columns is not None and "payload" in columns:
+ payload_reads.append(tuple(columns))
+ yield from original(reader, row_group, columns)
+
+ with mock.patch.object(
+ FormatPyArrowReader,
+ "_read_parquet_row_group_batches", tracked):
+ row = pmm.join_window(
+ anchors.scan(),
+ samples.scan().select(["payload.value", "unused"]),
+ on="event_time", by="episode_id", preceding=0,
+ aggregations={"payload_value": "mean"},
+ ).to_list()[0]
+
+ self.assertEqual(7.0, row["payload_value"])
+ self.assertTrue(payload_reads)
+ self.assertNotIn("unused", {
+ name for columns in payload_reads for name in columns
+ })
+
+ def test_window_join_preserves_nested_projection_aliases_when_pruned(self):
+ anchors = self._table("window_alias_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ })
+ samples = self._table("window_alias_samples", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ "a_b": pa.int32(),
+ "a": pa.struct([pa.field("b", pa.int32())]),
+ })
+ anchors.add([{"episode_id": 1, "event_time": 10}])
+ samples.add([{
+ "episode_id": 1,
+ "event_time": 10,
+ "a_b": 3,
+ "a": {"b": 7},
+ }])
+
+ row = pmm.join_window(
+ anchors.scan(), samples.scan().select(["a_b", "a.b"]),
+ on="event_time", by="episode_id", preceding=0,
+ aggregations={"nested_mean": ("a_b__0", "mean")},
+ ).to_list()[0]
+
+ self.assertEqual(7.0, row["nested_mean"])
+
+ def test_window_join_preserves_masked_alias_types_when_pruned(self):
+ anchors = self._table("window_masked_alias_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ })
+ samples = self._table("window_masked_alias_samples", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ "a_b": pa.int32(),
+ "a": pa.struct([pa.field("b", pa.int32())]),
+ })
+ anchors.add(pa.Table.from_pydict({
+ "episode_id": [1, 1], "event_time": [10, 20],
+ }))
+ samples.add(pa.Table.from_pydict({
+ "episode_id": [1], "event_time": [10],
+ "a_b": [3],
+ "a": pa.array([{"b": 7}],
+ type=pa.struct([pa.field("b", pa.int32())])),
+ }))
+ auth = TableQueryAuthResult(
+ filter=None,
+ column_masking={"a_b": json.dumps({
+ "name": "CAST",
+ "fieldRef": {"index": 2, "name": "a_b", "type": "INT"},
+ "type": "STRING",
+ })},
+ )
+ samples.raw_table.catalog_environment.table_query_auth = (
+ lambda options, identifier: lambda select: auth)
+
+ for projection, source in (
+ (["a_b"], "a_b"), (["a.b", "a_b"], "a_b__0")):
+ with self.subTest(projection=projection):
+ result = pmm.join_window(
+ anchors.scan(), samples.scan().select(projection),
+ on="event_time", by="episode_id", preceding=0,
+ aggregations={
+ "first_value": (source, "first"),
+ "last_value": (source, "last"),
+ },
+ ).to_arrow()
+ for name in ("first_value", "last_value"):
+ self.assertEqual(pa.string(), result[name].type)
+ self.assertEqual(["3", None], result[name].to_pylist())
+
+ for operation in ("mean", "min", "max"):
+ with self.subTest(projection=projection, operation=operation):
+ with self.assertRaisesRegex(
+ TypeError, "requires an integer or floating"):
+ pmm.join_window(
+ anchors.scan(), samples.scan().select(projection),
+ on="event_time", by="episode_id", preceding=0,
+ aggregations={"value": (source, operation)},
+ ).to_arrow()
+
+ def test_window_join_matches_masks_by_original_nested_path(self):
+ anchors = self._table("window_nested_mask_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ })
+ samples = self._table("window_nested_mask_samples", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ "a_b": pa.int32(),
+ "a": pa.struct([pa.field("b", pa.int32())]),
+ })
+ anchors.add([{"episode_id": 1, "event_time": 10}])
+ samples.add([{
+ "episode_id": 1, "event_time": 10,
+ "a_b": 3, "a": {"b": 7},
+ }])
+ auth = TableQueryAuthResult(
+ filter=None,
+ column_masking={"a_b": json.dumps({"name": "NULL"})},
+ )
+ samples.raw_table.catalog_environment.table_query_auth = (
+ lambda options, identifier: lambda select: auth)
+
+ row = pmm.join_window(
+ anchors.scan(), samples.scan().select(["a_b", "a.b"]),
+ on="event_time", by="episode_id", preceding=0,
+ aggregations={"nested_mean": ("a_b__0", "mean")},
+ ).to_list()[0]
+
+ self.assertEqual(7.0, row["nested_mean"])
+
+ def test_window_mean_uses_masked_numeric_alias_type_when_pruned(self):
+ anchors = self._table("window_numeric_alias_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ })
+ samples = self._table("window_numeric_alias_samples", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ "a_b": pa.string(),
+ "a": pa.struct([pa.field("b", pa.int32())]),
+ })
+ anchors.add(pa.Table.from_pydict({
+ "episode_id": [1, 1], "event_time": [10, 20],
+ }))
+ samples.add(pa.Table.from_pydict({
+ "episode_id": [1, 1], "event_time": [9, 11],
+ "a_b": ["6", "10"],
+ "a": pa.array([{"b": 1}, {"b": 2}],
+ type=pa.struct([pa.field("b", pa.int32())])),
+ }))
+ auth = TableQueryAuthResult(
+ filter=None,
+ column_masking={"a_b": json.dumps({
+ "name": "CAST",
+ "fieldRef": {"index": 2, "name": "a_b", "type": "STRING"},
+ "type": "DOUBLE",
+ })},
+ )
+ samples.raw_table.catalog_environment.table_query_auth = (
+ lambda options, identifier: lambda select: auth)
+
+ for projection, source in (
+ (["a_b"], "a_b"), (["a.b", "a_b"], "a_b__0")):
+ with self.subTest(projection=projection):
+ result = pmm.join_window(
+ anchors.scan(), samples.scan().select(projection),
+ on="event_time", by="episode_id", preceding=1, following=1,
+ aggregations={"value": (source, "mean")},
+ ).to_arrow()
+ self.assertEqual(pa.float64(), result["value"].type)
+ self.assertEqual([8.0, None], result["value"].to_pylist())
+
+ def test_window_mean_avoids_numeric_overflow_and_integer_rounding(self):
+ anchors = self._table("window_numeric_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ })
+ samples = self._table("window_numeric_samples", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ "integer_value": pa.int64(),
+ "float_value": pa.float64(),
+ })
+ anchors.add([{"episode_id": 1, "event_time": 1}])
+ samples.add([
+ {"episode_id": 1, "event_time": 0,
+ "integer_value": -(1 << 63),
+ "float_value": sys.float_info.max},
+ {"episode_id": 1, "event_time": 2,
+ "integer_value": (1 << 63) - 1,
+ "float_value": sys.float_info.max},
+ ])
+
+ row = pmm.join_window(
+ anchors.scan(),
+ samples.scan().select(["integer_value", "float_value"]),
+ on="event_time", by="episode_id", preceding=1, following=1,
+ aggregations={
+ "integer_value": "mean",
+ "float_value": "mean",
+ },
+ ).to_list()[0]
+
+ self.assertEqual(-0.5, row["integer_value"])
+ self.assertEqual(sys.float_info.max, row["float_value"])
+
+ def test_window_mean_preserves_finite_float_cancellation(self):
+ anchors = self._table("window_float_mean_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ })
+ samples = self._table("window_float_mean_samples", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ "value": pa.float64(),
+ })
+ anchors.add([{"episode_id": 1, "event_time": 1}])
+ samples.add([
+ {"episode_id": 1, "event_time": 0, "value": 1e16},
+ {"episode_id": 1, "event_time": 2, "value": -1e16 + 2},
+ ])
+
+ row = pmm.join_window(
+ anchors.scan(), samples.scan().select("value"),
+ on="event_time", by="episode_id", preceding=1, following=1,
+ aggregations={"value": "mean"},
+ ).to_list()[0]
+
+ self.assertEqual(1.0, row["value"])
+
+ def test_window_join_can_follow_an_asof_join(self):
+ anchors = self._table("window_chain_anchors", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ })
+ images = self._table("window_chain_images", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ "image": pa.string(),
+ })
+ imu = self._table("window_chain_imu", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ "acceleration": pa.float32(),
+ })
+ anchors.add([{"episode_id": 1, "event_time": 10}])
+ images.add([{"episode_id": 1, "event_time": 9, "image": "frame"}])
+ imu.add([
+ {"episode_id": 1, "event_time": 8, "acceleration": 1.0},
+ {"episode_id": 1, "event_time": 10, "acceleration": 3.0},
+ ])
+
+ row = pmm.join_asof(
+ anchors.scan(), images.scan().select("image"),
+ on="event_time", by="episode_id", direction="nearest",
+ ).join_window(
+ imu.scan().select("acceleration"),
+ preceding=2,
+ aggregations={"acceleration": "mean"},
+ ).to_list()[0]
+
+ self.assertEqual("frame", row["image"])
+ self.assertEqual(2.0, row["acceleration"])
+
+ def test_window_join_validates_options(self):
+ table = self._table("window_validation", {
+ "episode_id": pa.int32(),
+ "event_time": pa.int64(),
+ "value": pa.int32(),
+ "text": pa.string(),
+ })
+
+ def scan():
+ return table.scan().select("value")
+
+ with self.assertRaisesRegex(ValueError, "non-negative"):
+ pmm.join_window(
+ scan(), scan(), on="event_time", by="episode_id",
+ preceding=-1, aggregations={"value": "mean"})
+ with self.assertRaisesRegex(ValueError, "Unsupported aggregation"):
+ pmm.join_window(
+ scan(), scan(), on="event_time", by="episode_id",
+ preceding=1, aggregations={"value": "median"})
+ with self.assertRaisesRegex(ValueError, "closed must be"):
+ pmm.join_window(
+ scan(), scan(), on="event_time", by="episode_id",
+ preceding=1, aggregations={"value": "mean"},
+ closed="middle")
+ with self.assertRaisesRegex(ValueError, "missing aggregation columns"):
+ pmm.join_window(
+ scan(), scan(), on="event_time", by="episode_id",
+ preceding=1, aggregations={"missing": "mean"})
+ with self.assertRaisesRegex(
+ TypeError, "requires an integer or floating"):
+ pmm.join_window(
+ scan(), table.scan().select("text"),
+ on="event_time", by="episode_id", preceding=1,
+ aggregations={"text": "mean"}).to_arrow()
+
def test_linear_interpolation_preserves_an_exact_infinite_float(self):
anchors = self._table("linear_exact_anchors", {
"episode_id": pa.int32(),