[FEA][Java][JNI] Support multi-input aggregations without requiring a struct-column wrapper
- Dominant language
- C++
- Stars
- 9.8k
- Forks
- 1.1k
- Avg merge
- 3d 6m
- Merged PRs (30d)
- 278
Description
### 🎯 Problem
`cuDF Java` attaches one column per aggregation via `GroupByAggregation.onColumn(int)` at [`GroupByAggregation.java:29-31`](https://github.com/rapidsai/cudf/blob/main/java/src/main/java/ai/rapids/cudf/GroupByAggregation.java#L29-L31). Multi-input aggregations therefore pack inputs into a struct column via `ColumnVector.makeStruct(...)` and unpack afterward — a full-size copy per batch or merge. `libcudf` has no such limit (see [`AggregationJni.cpp:26-107`](https://github.com/rapidsai/cudf/blob/main/java/src/main/native/src/AggregationJni.cpp#L26-L107)); the limitation lives in **`cuDF Java` + `cuDF JNI`** and in the `spark-rapids` aggregation architecture (one `CudfAggregate` → one column, flowed through `GpuAggregateExec`) that drives them.
### 💡 Proposed solution
Express a multi-input aggregation as **N role-tagged `Aggregation` instances sharing a multi-input id**. The `cuDF JNI` layer buckets them by id and dispatches directly to `libcudf`. Additive; no `libcudf` changes.
### 🔍 Implementation example — `min_by`
**1. Extend `Kind` ([`Aggregation.java:27-69`](https://github.com/rapidsai/cudf/blob/main/java/src/main/java/ai/rapids/cudf/Aggregation.java#L27-L69)):**
```java
ORDERING_FOR_MIN_BY(37),
VALUE_FOR_MIN_BY(38),
```
**2. Role-tagged subclasses sharing an id** — the id survives Java→JNI (step 3) and lets C++ (step 6) pair siblings when multiple `min_by` calls coexist (`SELECT min_by(a, b), min_by(c, d) …`):
```java
static final class OrderingForMinByAgg extends Aggregation {
private final long multiInputId;
OrderingForMinByAgg(long id) { super(Kind.ORDERING_FOR_MIN_BY); this.multiInputId = id; }
@Override long createNativeInstance() { return createMultiInputAgg(kind.nativeId, multiInputId); }
}
static final class ValueForMinByAgg extends Aggregation { /* same shape, VALUE_FOR_MIN_BY */ }
```
**3. Native factory (mirrors [`Aggregation.java:1054-1119`](https://github.com/rapidsai/cudf/blob/main/java/src/main/java/ai/rapids/cudf/Aggregation.java#L1054-L1119)):**
```java
private static native long createMultiInputAgg(int kind, long multiInputId);
```
C++ side in [`AggregationJni.cpp`](https://github.com/rapidsai/cudf/blob/main/java/src/main/native/src/AggregationJni.cpp) (boilerplate omitted); holder derives from `cudf::groupby_aggregation` so (a) the existing close path at [`:13-24`](https://github.com/rapidsai/cudf/blob/main/java/src/main/native/src/AggregationJni.cpp#L13-L24) works unchanged, and (b) the `dynamic_cast` check at [`TableJni.cpp:3904`](https://github.com/rapidsai/cudf/blob/main/java/src/main/native/src/TableJni.cpp#L3904) still succeeds:
```cpp
enum class multi_input_role : int32_t {
ORDERING_FOR_MIN_BY = 37, // must equal Aggregation.Kind.ORDERING_FOR_MIN_BY.nativeId
VALUE_FOR_MIN_BY = 38,
// future: ORDERING_FOR_MAX_BY, VALUE_FOR_MAX_BY, {two roles per multi-input aggregation}…
};
class multi_input_aggregation : public cudf::groupby_aggregation {
public:
multi_input_role const role;
int64_t const multi_input_id;
multi_input_aggregation(multi_input_role r, int64_t id)
: cudf::groupby_aggregation{/* JNI-private sentinel Kind; never dispatched to libcudf */},
role(r), multi_input_id(id) {}
std::unique_ptr clone() const override { /* … */ }
};
// JNI factory body
auto agg = std::make_unique(
static_cast(kind), multi_input_id);
return reinterpret_cast(agg.release());
```
**4. `ai.rapids.cudf.MultiInputIds`:**
```java
package ai.rapids.cudf;
import java.util.concurrent.atomic.AtomicLong;
public final class MultiInputIds {
private static final AtomicLong counter = new AtomicLong();
public static long next() { return counter.incrementAndGet(); }
private MultiInputIds() {}
}
```
**5. Public factories on `GroupByAggregation` (mirrors `min()` at [`GroupByAggregation.java:103-105`](https://github.com/rapidsai/cudf/blob/main/java/src/main/java/ai/rapids/cudf/GroupByAggregation.java#L103-L105)):**
```java
public static GroupByAggregation orderingForMinBy(long multiInputId) {
return new GroupByAggregation(new Aggregation.OrderingForMinByAgg(multiInputId));
}
public static GroupByAggregation valueForMinBy(long multiInputId) {
return new GroupByAggregation(new Aggregation.ValueForMinByAgg(multiInputId));
}
```
`GpuMinBy` collapses — no `GpuCreateNamedStruct`, no `GpuGetStructField`, `postUpdate` / `preMerge` / `postMerge` fall back to trait defaults ([`aggregateBase.scala:128-129`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/aggregateBase.scala#L128-L129), [`:140`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/aggregateBase.scala#L140), [`:163-164`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/aggregateBase.scala#L163-L164)):
```scala
case class GpuMinBy(valueExpr: Expression, orderingExpr: Expression) extends GpuAggregateFunction {
private val id = MultiInputIds.next()
private val cudfValue = new CudfValueForMinBy(id, valueExpr.dataType)
private val cudfOrdering = new CudfOrderingForMinBy(id, orderingExpr.dataType)
override lazy val inputProjection = Seq(valueExpr, orderingExpr)
override lazy val updateAggregates = Seq(cudfValue, cudfOrdering)
override lazy val mergeAggregates = Seq(cudfValue, cudfOrdering)
}
```
vs. today's `GpuMaxMinByBase` at [`aggregateFunctions.scala:2272-2283`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/aggregateFunctions.scala#L2272-L2283).
**6. `cuDF JNI` regroup + dispatch** — augment the groupby JNI path at [`TableJni.cpp:3885-3932`](https://github.com/rapidsai/cudf/blob/main/java/src/main/native/src/TableJni.cpp#L3885-L3932):
```cpp
// 1. Bucket multi-input aggs by id; remember each entry's slot.
// (The existing loop at TableJni.cpp:3900-3918 must also skip any `multi_input_aggregation`
// so it is not pushed into `requests` twice or handed to libcudf directly.)
std::unordered_map>> groups;
std::unordered_map argmin_req_idx;
for (int i = 0; i < n_values.size(); ++i) {
if (auto* c = dynamic_cast(n_agg_instances[i]))
groups[c->multi_input_id].emplace_back(c->role, n_input_table->column(n_values[i]), i);
}
// 2. One argmin request per group, merged into the existing requests vector.
for (auto const& [id, entries] : groups) {
cudf::column_view ord;
for (auto const& [role, col, _] : entries)
if (role == multi_input_role::ORDERING_FOR_MIN_BY) ord = col;
cudf::groupby::aggregation_request req;
req.values = ord;
req.aggregations.push_back(cudf::make_argmin_aggregation());
argmin_req_idx[id] = requests.size();
requests.push_back(std::move(req));
}
auto result = grouper.aggregate(requests); // single batched libcudf call
// 3. Gather {value, ordering} per group in one call; drop into slots.
std::vector> result_columns(n_values.size());
for (auto const& [id, entries] : groups) {
cudf::column_view ord, val;
int ord_slot = 0, val_slot = 0;
for (auto const& [role, col, slot] : entries) {
if (role == multi_input_role::ORDERING_FOR_MIN_BY) { ord = col; ord_slot = slot; }
else { val = col; val_slot = slot; }
}
auto indices = result.second[argmin_req_idx[id]].results.front()->view();
auto gathered = cudf::gather(cudf::table_view{{val, ord}}, indices);
auto cols = gathered->release();
result_columns[val_slot] = std::move(cols[0]);
result_columns[ord_slot] = std::move(cols[1]);
}
return convert_table_for_return(env, result.first, std::move(result_columns));
```
### 🔍 Extension — N-ary homogeneous roles (e.g. `GpuHyperLogLogPlusPlus`, ~52 long slots)
Generalize the role-tagged aggregation with a **slot index** alongside the correlation id. In JNI regroup, the N children sharing a correlation id compose into a **zero-copy `struct` `column_view`** and pass to libcudf's existing HLL merge — no data copy, no libcudf change.
```java
WORD_FOR_HLLPP(39), // one Kind per N-ary role, not per slot
static final class WordForHllppAgg extends Aggregation {
private final long multiInputId;
private final int slotIndex; // 0..N-1
WordForHllppAgg(long multiInputId, int slotIndex) { ... }
@Override long createNativeInstance() {
return createMultiInputAggWithSlot(kind.nativeId, multiInputId, slotIndex);
}
}
public static GroupByAggregation wordForHllpp(long multiInputId, int slotIndex) { ... }
private static native long createMultiInputAggWithSlot(int kind, long multiInputId, int slotIndex);
```
```cpp
// JNI regroup — per correlated group:
std::sort(group.entries, by slot_index);
req.values = cudf::column_view(STRUCT, rows, ..., /* children */ group_columns); // zero-copy
req.aggregations.push_back(cudf::make_merge_hllpp_aggregation<...>(precision));
```
```scala
// spark-rapids caller
val id = MultiInputIds.next()
override lazy val mergeAggregates = Seq.tabulate(numWords)(i =>
new CudfWordForHllppMerge(id, slotIndex = i, precision))
```
Binary roles (`min_by` in steps 1–6) ignore `slotIndex`; it's only meaningful for N-ary homogeneous roles.
---
### 🚀 Beneficiaries in `spark-rapids`
#### A. Collapses existing struct-wrap hot paths
| class | File:Line | Struct today | Built |
|---|---|---|---|
| `GpuMaxBy` / `GpuMinBy` | [`aggregateFunctions.scala:2297-2313`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/aggregateFunctions.scala#L2297-L2313) (struct at [`:2258-2263`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/aggregateFunctions.scala#L2258-L2263)) | `struct<_key_ordering, _key_value>` | **per batch** + preMerge |
| `GpuStddevPop` / `GpuStddevSamp` / `GpuVariancePop` / `GpuVarianceSamp` (via `GpuM2`) | [`:2086`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/aggregateFunctions.scala#L2086), [`:2119`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/aggregateFunctions.scala#L2119), [`:2152`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/aggregateFunctions.scala#L2152), [`:2166`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/aggregateFunctions.scala#L2166); preMerge [`:2064-2070`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/aggregateFunctions.scala#L2064-L2070) | `struct` | **per merge** |
| `GpuHyperLogLogPlusPlus` | [`GpuHyperLogLogPlusPlus.scala:151-154`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/GpuHyperLogLogPlusPlus.scala#L151-L154), [`:166`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/GpuHyperLogLogPlusPlus.scala#L166), [`:179`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/aggregate/GpuHyperLogLogPlusPlus.scala#L179) | `struct` | **per merge + eval** |
Copy site: `ColumnVector.makeStruct(...)` at [`complexTypeCreator.scala:235`](https://github.com/NVIDIA/spark-rapids/blob/main/sql-plugin/src/main/scala/org/apache/spark/sql/rapids/complexTypeCreator.scala#L235).
#### B. Unblocks Spark aggregates not GPU-accelerated today (verified by grep — all CPU-fallback)
| function | Inputs | Notes |
|---|---|---|
| `corr(y, x)` | 2 | Pearson correlation |
| `covar_pop(y, x)` / `covar_samp(y, x)` | 2 each | Population / sample covariance |
| `regr_count` / `regr_slope` / `regr_intercept` / `regr_r2` | 2 each | Spark regression family (SPARK-23907) |
| `regr_sxx` / `regr_syy` / `regr_sxy` | 2 each | |
| `regr_avgx` / `regr_avgy` | 2 each | |
| `percentile(col, pct, frequency)` | 2 | frequency variant |
Contributor guide
Assessment
This issue has not been assessed yet.