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 653ac61283 [python] Group write rows by partition/bucket in Arrow to 
avoid GIL-bound per-row extraction (#9047)
653ac61283 is described below

commit 653ac61283576ff9142da479cdfe0c087f8aa05b
Author: Yu Gan <[email protected]>
AuthorDate: Thu Aug 6 12:12:09 2026 +0800

    [python] Group write rows by partition/bucket in Arrow to avoid GIL-bound 
per-row extraction (#9047)
---
 .../pypaimon/tests/write/table_write_test.py       | 201 ++++++++++++++++++++-
 paimon-python/pypaimon/write/row_key_extractor.py  | 130 +++++++++++++
 paimon-python/pypaimon/write/table_write.py        |  31 ++--
 3 files changed, 346 insertions(+), 16 deletions(-)

diff --git a/paimon-python/pypaimon/tests/write/table_write_test.py 
b/paimon-python/pypaimon/tests/write/table_write_test.py
index 1189974b2e..a640b8d31f 100644
--- a/paimon-python/pypaimon/tests/write/table_write_test.py
+++ b/paimon-python/pypaimon/tests/write/table_write_test.py
@@ -141,14 +141,30 @@ class TableWriteTest(unittest.TestCase):
             write.close()
             commit.close()
 
+    @staticmethod
+    def _groups_from_rows(partitions, buckets, num_rows):
+        """Mirror RowKeyExtractor.extract_partition_bucket_groups' contract 
from
+        per-row (partition, bucket) values: list of (partition, bucket,
+        row_indices) with row_indices None when the whole batch is one 
group."""
+        grouped = {}
+        for i in range(num_rows):
+            grouped.setdefault((tuple(partitions[i]), buckets[i]), 
[]).append(i)
+        out = []
+        for (partition, bucket), idxs in grouped.items():
+            row_indices = None if len(idxs) == num_rows \
+                else pa.array(idxs, type=pa.int64())
+            out.append((partition, bucket, row_indices))
+        return out
+
     @staticmethod
     def _mock_table_write(partitions, buckets):
         table_write = object.__new__(TableWrite)
         table_write._validate_pyarrow_schema = Mock()
         table_write.row_key_extractor = Mock()
         table_write.file_store_write = Mock()
-        
table_write.row_key_extractor.extract_partition_bucket_batch.return_value = (
-            partitions, buckets)
+        
table_write.row_key_extractor.extract_partition_bucket_groups.side_effect = (
+            lambda data: TableWriteTest._groups_from_rows(
+                partitions, buckets, data.num_rows))
         return table_write
 
     def test_write_arrow_batch_reuses_full_batch(self):
@@ -207,6 +223,187 @@ class TableWriteTest(unittest.TestCase):
         self.assertEqual({'id': [1, 3], 'payload': [b'b', b'd']},
                          calls[1][0][2].to_pydict())
 
+    def test_write_arrow_batch_noncontiguous_group_uses_take(self):
+        # The extractor delivers each group's indices in ascending input order.
+        # A non-contiguous group (endpoints span more than len(group)) must be
+        # gathered via take and keep input order. Here ('p1',) = rows [0, 2, 
3]:
+        # endpoints 0..3 span 4 != 3 == len -> take, not a contiguous slice.
+        data = pa.RecordBatch.from_pydict({
+            'id': [0, 1, 2, 3],
+            'payload': [b'a', b'b', b'c', b'd'],
+        })
+        table_write = object.__new__(TableWrite)
+        table_write._validate_pyarrow_schema = Mock()
+        table_write.file_store_write = Mock()
+        table_write.row_key_extractor = Mock()
+        
table_write.row_key_extractor.extract_partition_bucket_groups.return_value = [
+            (('p1',), 0, pa.array([0, 2, 3], type=pa.int64())),
+            (('p2',), 0, pa.array([1], type=pa.int64())),
+        ]
+
+        table_write.write_arrow_batch(data)
+
+        calls = table_write.file_store_write.write.call_args_list
+        self.assertEqual(2, len(calls))
+        # Order-sensitive: sequence numbers are assigned in the delivered 
order.
+        self.assertEqual([0, 2, 3], calls[0][0][2].column('id').to_pylist())
+        self.assertEqual([1], calls[1][0][2].column('id').to_pylist())
+
+    def test_write_arrow_batch_contiguous_group_uses_zero_copy_slice(self):
+        # A contiguous group (ascending endpoints span exactly len(group)) 
takes
+        # the zero-copy slice path instead of allocating a copy via take.
+        data = pa.RecordBatch.from_pydict({
+            'id': [0, 1, 2, 3],
+            'payload': [b'a', b'b', b'c', b'd'],
+        })
+        table_write = object.__new__(TableWrite)
+        table_write._validate_pyarrow_schema = Mock()
+        table_write.file_store_write = Mock()
+        table_write.row_key_extractor = Mock()
+        
table_write.row_key_extractor.extract_partition_bucket_groups.return_value = [
+            (('p0',), 0, pa.array([0], type=pa.int64())),
+            (('p1',), 0, pa.array([1, 2, 3], type=pa.int64())),
+        ]
+
+        with patch.object(pa.compute, 'take', wraps=pa.compute.take) as take:
+            table_write.write_arrow_batch(data)
+
+        take.assert_not_called()
+        calls = table_write.file_store_write.write.call_args_list
+        self.assertEqual([1, 2, 3], calls[1][0][2].column('id').to_pylist())
+
+    def _unaware_partitioned_extractor(self, name, partition_keys):
+        schema = Schema.from_pyarrow_schema(
+            self.pa_schema, partition_keys=partition_keys,
+            options={'bucket': '-1'})
+        self.catalog.create_table(name, schema, False)
+        return self.catalog.get_table(name).create_row_key_extractor()
+
+    def test_extract_partition_bucket_groups_multi_partition(self):
+        ex = self._unaware_partitioned_extractor(
+            'default.t_groups_multi', ['dt'])
+        data = pa.RecordBatch.from_pydict({
+            'user_id': [1, 2, 3, 4],
+            'item_id': [1, 2, 3, 4],
+            'behavior': ['a', 'b', 'c', 'd'],
+            'dt': ['p1', 'p2', 'p1', 'p2'],   # interleaved -> non-contiguous
+        }, schema=self.pa_schema)
+
+        groups = ex.extract_partition_bucket_groups(data)
+        # Within-group indices must be in ascending input order (the extractor
+        # sorts Arrow's threaded aggregation output) so the writer's
+        # sequence-number assignment stays latest-wins correct.
+        self.assertEqual(
+            {(('p1',), 0): [0, 2], (('p2',), 0): [1, 3]},
+            {(p, b): idx.to_pylist() for p, b, idx in groups})
+
+    def test_extract_partition_bucket_groups_single_group_is_none(self):
+        ex = self._unaware_partitioned_extractor(
+            'default.t_groups_single', ['dt'])
+        data = pa.RecordBatch.from_pydict({
+            'user_id': [1, 2, 3],
+            'item_id': [1, 2, 3],
+            'behavior': ['a', 'b', 'c'],
+            'dt': ['p1', 'p1', 'p1'],   # one partition -> whole batch
+        }, schema=self.pa_schema)
+
+        groups = ex.extract_partition_bucket_groups(data)
+        self.assertEqual(1, len(groups))
+        partition, bucket, row_indices = groups[0]
+        self.assertEqual(('p1',), partition)
+        self.assertEqual(0, bucket)
+        # None signals "reuse the original batch" (no BLOB copy via take).
+        self.assertIsNone(row_indices)
+
+    def test_extract_partition_bucket_groups_arrow_matches_fallback(self):
+        ex = self._unaware_partitioned_extractor(
+            'default.t_groups_equiv', ['behavior', 'dt'])
+        data = pa.RecordBatch.from_pydict({
+            'user_id': [1, 2, 3, 4, 5],
+            'item_id': [1, 2, 3, 4, 5],
+            'behavior': ['a', 'b', 'a', 'b', 'a'],
+            'dt': ['p1', 'p1', 'p2', 'p1', 'p2'],
+        }, schema=self.pa_schema)
+        buckets = ex._extract_buckets_batch(data)
+
+        def norm(groups):
+            # Both paths must return within-group indices in ascending input
+            # order (arrow sorts its threaded aggregation; python appends in
+            # row order), so compare order-sensitively.
+            return {
+                (p, b): (None if idx is None else idx.to_pylist())
+                for p, b, idx in groups
+            }
+
+        self.assertEqual(
+            norm(ex._group_indices_arrow(data, buckets)),
+            norm(ex._group_indices_python(data, buckets)))
+
+    def test_extract_partition_bucket_groups_without_arrow_group_by(self):
+        # pyarrow < 7.0.0 (e.g. 6.0.1 on the Python 3.6 lane) has no
+        # Table.group_by; extract_partition_bucket_groups must transparently 
use
+        # the per-row fallback instead of raising AttributeError.
+        ex = self._unaware_partitioned_extractor(
+            'default.t_groups_no_group_by', ['dt'])
+        data = pa.RecordBatch.from_pydict({
+            'user_id': [1, 2, 3, 4],
+            'item_id': [1, 2, 3, 4],
+            'behavior': ['a', 'b', 'c', 'd'],
+            'dt': ['p1', 'p2', 'p1', 'p2'],
+        }, schema=self.pa_schema)
+
+        with 
patch('pypaimon.write.row_key_extractor._ARROW_GROUP_BY_SUPPORTED',
+                   False):
+            groups = ex.extract_partition_bucket_groups(data)
+
+        self.assertEqual(
+            {(('p1',), 0): [0, 2], (('p2',), 0): [1, 3]},
+            {(p, b): sorted(idx.to_pylist()) for p, b, idx in groups})
+
+    def test_group_indices_arrow_sorts_unordered_aggregation(self):
+        # Arrow's threaded hash_list can return a group's row indices out of
+        # input order (reproduced upstream: a group whose last input index was
+        # 2,999,988 ended at 1,048,575). Out-of-order indices would make the
+        # writer assign sequence numbers in the wrong order, letting an earlier
+        # input row with a repeated primary key wrongly win latest-wins dedup.
+        # _group_indices_arrow must sort each group back to ascending input 
order.
+        ex = self._unaware_partitioned_extractor(
+            'default.t_groups_sorted', ['dt'])
+        data = pa.RecordBatch.from_pydict({
+            'user_id': [1, 2, 3, 4, 5],
+            'item_id': [1, 2, 3, 4, 5],
+            'behavior': ['a', 'b', 'c', 'd', 'e'],
+            'dt': ['p1', 'p2', 'p1', 'p2', 'p1'],
+        }, schema=self.pa_schema)
+        buckets = ex._extract_buckets_batch(data)
+        # Force Arrow's aggregation to report each group's indices out of 
order.
+        unordered = pa.table({
+            '__p0': pa.array(['p1', 'p2']),
+            '__bucket': pa.array([0, 0], type=pa.int32()),
+            '__idx_list': pa.array([[4, 0, 2], [3, 1]],
+                                   type=pa.list_(pa.int64())),
+        })
+        with patch.object(pa.TableGroupBy, 'aggregate', 
return_value=unordered):
+            groups = ex._group_indices_arrow(data, buckets)
+
+        by_part = {p: idx.to_pylist() for p, b, idx in groups}
+        self.assertEqual([0, 2, 4], by_part[('p1',)])
+        self.assertEqual([1, 3], by_part[('p2',)])
+
+    def test_probe_arrow_group_by_false_when_hash_list_missing(self):
+        # pyarrow 7 has Table.group_by but not the hash_list aggregate kernel
+        # (added in Arrow 8); it raises ArrowKeyError, and pyarrow>=7,<13 is
+        # still allowed on the Python 3.7 lane. The capability probe must treat
+        # that as unsupported so writes fall back instead of failing every
+        # write_arrow_batch.
+        from pypaimon.write import row_key_extractor as rk
+
+        def raise_missing_kernel(self, *args, **kwargs):
+            raise pa.ArrowKeyError("No function registered with name: 
hash_list")
+
+        with patch.object(pa.TableGroupBy, 'aggregate', raise_missing_kernel):
+            self.assertFalse(rk._probe_arrow_group_by())
+
     def test_write_snapshot(self):
         schema = Schema.from_pyarrow_schema(self.pa_schema, 
partition_keys=['dt'])
         self.catalog.create_table('default.test_write_snapshot', schema, False)
diff --git a/paimon-python/pypaimon/write/row_key_extractor.py 
b/paimon-python/pypaimon/write/row_key_extractor.py
index f97933fb83..3313620d7d 100644
--- a/paimon-python/pypaimon/write/row_key_extractor.py
+++ b/paimon-python/pypaimon/write/row_key_extractor.py
@@ -15,12 +15,14 @@
 # specific language governing permissions and limitations
 # under the License.
 
+import logging
 import math
 import random
 import struct
 from abc import ABC, abstractmethod
 from typing import Any, Dict, List, Optional, Tuple
 
+import numpy as np
 import pyarrow as pa
 
 from pypaimon.common.options.core_options import CoreOptions
@@ -30,6 +32,38 @@ from pypaimon.table.bucket_mode import BucketMode
 from pypaimon.table.row.generic_row import GenericRow, GenericRowSerializer
 from pypaimon.table.row.internal_row import RowKind
 
+logger = logging.getLogger(__name__)
+
+
+def _probe_arrow_group_by() -> bool:
+    """Return True only if this pyarrow can run the write path's group-by.
+
+    Two versions matter, and a plain ``hasattr(pa.Table, "group_by")`` 
conflates
+    them: ``Table.group_by`` (Acero) landed in pyarrow 7.0.0, but the
+    ``hash_list`` aggregate kernel this path relies on only landed in 8.0.0.
+    pyarrow 7 therefore *has* ``group_by`` yet raises ``ArrowKeyError`` for
+    ``hash_list`` -- and the Python 3.7 dependency range still permits
+    ``pyarrow>=7,<13``. Probe the actual aggregate once at import so both
+    pyarrow<7 (no ``group_by``) and pyarrow 7 (no ``hash_list``) fall through 
to
+    the per-row grouping instead of failing every ``write_arrow_batch``.
+    """
+    if not hasattr(pa.Table, "group_by"):
+        return False
+    try:
+        probe = pa.table({
+            "__k": pa.array([0], type=pa.int32()),
+            "__idx": pa.array([0], type=pa.int64()),
+        })
+        probe.group_by(["__k"]).aggregate([("__idx", "list")])
+    except Exception:  # any failure here means "use the fallback"
+        return False
+    return True
+
+
+# pyarrow < 7.0.0 has no ``Table.group_by`` and pyarrow 7 has no ``hash_list``
+# aggregate kernel; on either the write path must use per-row grouping.
+_ARROW_GROUP_BY_SUPPORTED = _probe_arrow_group_by()
+
 _MURMUR_C1 = 0xCC9E2D51
 _MURMUR_C2 = 0x1B873593
 _DEFAULT_SEED = 42
@@ -95,6 +129,102 @@ class RowKeyExtractor(ABC):
         """Return partition tuples without calculating bucket hashes."""
         return self._extract_partitions_batch(data)
 
+    def extract_partition_bucket_groups(
+            self, data: pa.RecordBatch) -> List[Tuple[Tuple, int, 
Optional[pa.Array]]]:
+        """Group row indices by (partition, bucket) for the write path.
+
+        Returns a list of ``(partition, bucket, row_indices)`` where
+        ``row_indices`` is an Arrow ``int64`` array of the rows belonging to 
the
+        group, or ``None`` when the whole batch is a single group (so callers 
can
+        pass the original batch through without copying large values, e.g. 
BLOBs).
+
+        The grouping is done in Arrow so only the distinct group keys are
+        materialized into Python objects, instead of one ``.as_py()`` scalar 
per
+        row. The old per-row loop held the GIL for the entire batch, which
+        serialized multi-threaded writers down to ~1 core.
+
+        Buckets are computed once here, in row order, via 
``_extract_buckets_batch``
+        so stateful extractors (dynamic bucket) keep their exact assignment
+        sequence and side effects regardless of which grouping path runs.
+        """
+        buckets = self._extract_buckets_batch(data)
+        if _ARROW_GROUP_BY_SUPPORTED:
+            try:
+                return self._group_indices_arrow(data, buckets)
+            except (pa.ArrowNotImplementedError, pa.ArrowInvalid):
+                # Only Arrow's own "can't group this column type" errors fall
+                # back to the legacy per-row grouping; any other exception is a
+                # real bug and must propagate rather than silently degrade to 
the
+                # GIL-bound path. `buckets` is reused (never recomputed) so
+                # stateful extractors are not double-notified. Log so the
+                # (GIL-bound) fallback is visible.
+                logger.warning(
+                    "Arrow group_by could not handle the partition/bucket key "
+                    "types; falling back to per-row grouping (GIL-bound).",
+                    exc_info=True)
+        # pyarrow < 7.0.0 has no group_by; use the per-row grouping directly.
+        return self._group_indices_python(data, buckets)
+
+    def _group_indices_arrow(
+            self, data: pa.RecordBatch,
+            buckets: List[int]) -> List[Tuple[Tuple, int, Optional[pa.Array]]]:
+        num_rows = data.num_rows
+        columns = {}
+        key_names = []
+        for k, pi in enumerate(self.partition_indices):
+            name = f"__p{k}"
+            columns[name] = data.column(pi)
+            key_names.append(name)
+        columns["__bucket"] = pa.array(buckets, type=pa.int32())
+        key_names.append("__bucket")
+        # Build the row index with numpy (C speed, releases the GIL). Using a
+        # Python range() here makes pyarrow iterate it element by element under
+        # the GIL, which dominates this method and kills multi-thread scaling.
+        columns["__idx"] = pa.array(np.arange(num_rows, dtype=np.int64))
+
+        grouped = pa.table(columns).group_by(key_names).aggregate([("__idx", 
"list")])
+        num_groups = grouped.num_rows
+
+        num_part = len(self.partition_indices)
+        part_values = [grouped.column(f"__p{k}").to_pylist() for k in 
range(num_part)]
+        bucket_values = grouped.column("__bucket").to_pylist()
+        idx_lists = grouped.column("__idx_list")
+
+        groups = []
+        for gi in range(num_groups):
+            partition = tuple(part_values[k][gi] for k in range(num_part))
+            if num_groups == 1:
+                row_indices = None
+            else:
+                # Arrow's threaded ``hash_list`` may return a group's indices 
out
+                # of input order. The writer assigns sequence numbers in the
+                # order it receives rows, so unordered indices let an earlier
+                # input row (with a repeated primary key) win latest-wins
+                # deduplication / partial update. Sort back to ascending input
+                # order; np.sort runs in C (releases the GIL) so multi-threaded
+                # scaling is preserved.
+                row_indices = pa.array(
+                    
np.sort(idx_lists[gi].values.to_numpy(zero_copy_only=False)))
+            groups.append((partition, bucket_values[gi], row_indices))
+        return groups
+
+    def _group_indices_python(
+            self, data: pa.RecordBatch,
+            buckets: List[int]) -> List[Tuple[Tuple, int, Optional[pa.Array]]]:
+        partitions = self._extract_partitions_batch(data)
+        num_rows = data.num_rows
+        partition_bucket_groups = {}
+        for i in range(num_rows):
+            partition_bucket_groups.setdefault(
+                (tuple(partitions[i]), buckets[i]), []).append(i)
+
+        groups = []
+        for (partition, bucket), row_indices in 
partition_bucket_groups.items():
+            indices = None if len(row_indices) == num_rows \
+                else pa.array(row_indices, type=pa.int64())
+            groups.append((partition, bucket, indices))
+        return groups
+
     def extract_partition_bucket_row(
             self, values_by_name: Dict[str, Any]) -> Tuple[Tuple, int]:
         partition = tuple(
diff --git a/paimon-python/pypaimon/write/table_write.py 
b/paimon-python/pypaimon/write/table_write.py
index 493f3ba72e..f0a68bcf6f 100644
--- a/paimon-python/pypaimon/write/table_write.py
+++ b/paimon-python/pypaimon/write/table_write.py
@@ -15,7 +15,6 @@
 # specific language governing permissions and limitations
 # under the License.
 
-from collections import defaultdict
 from typing import TYPE_CHECKING, Any, Dict, List, Optional
 
 import pyarrow as pa
@@ -61,25 +60,29 @@ class TableWrite:
 
     def write_arrow_batch(self, data: pa.RecordBatch):
         self._validate_pyarrow_schema(data.schema)
-        partitions, buckets = 
self.row_key_extractor.extract_partition_bucket_batch(data)
 
-        partition_bucket_groups = defaultdict(list)
-        for i in range(data.num_rows):
-            partition_bucket_groups[(tuple(partitions[i]), 
buckets[i])].append(i)
-
-        for (partition, bucket), row_indices in 
partition_bucket_groups.items():
-            if len(row_indices) == data.num_rows:
+        for partition, bucket, row_indices in \
+                self.row_key_extractor.extract_partition_bucket_groups(data):
+            if row_indices is None:
                 # Every input row belongs to the same partition/bucket. 
Passing the
                 # original batch through avoids copying large BLOB values 
through
                 # Arrow take before the dedicated BLOB writer consumes them.
                 sub_table = data
-            elif row_indices[-1] - row_indices[0] + 1 == len(row_indices):
-                # Contiguous groups can share the original Arrow buffers 
instead of
-                # gathering their rows into newly allocated buffers with take.
-                sub_table = data.slice(row_indices[0], len(row_indices))
             else:
-                indices_array = pa.array(row_indices, type=pa.int64())
-                sub_table = pa.compute.take(data, indices_array)
+                # row_indices is an int64 array of this group's rows in
+                # ascending input order (the extractor sorts grouped indices so
+                # sequence-number assignment stays latest-wins correct), so the
+                # span is just first..last.
+                lo = row_indices[0].as_py()
+                hi = row_indices[-1].as_py()
+                count = len(row_indices)
+                if hi - lo + 1 == count:
+                    # Distinct row indices spanning exactly `count` values are
+                    # contiguous, so share the original Arrow buffers instead 
of
+                    # gathering their rows into newly allocated buffers with 
take.
+                    sub_table = data.slice(lo, count)
+                else:
+                    sub_table = pa.compute.take(data, row_indices)
             self._write_partition_bucket_batch(partition, bucket, sub_table)
 
     def _write_partition_bucket_batch(self, partition, bucket, data):

Reply via email to