diff --git a/velox/experimental/cudf/exec/DecimalAggregationDevice.cu b/velox/experimental/cudf/exec/DecimalAggregationDevice.cu index a24c2b71171..ea56801d86e 100644 --- a/velox/experimental/cudf/exec/DecimalAggregationDevice.cu +++ b/velox/experimental/cudf/exec/DecimalAggregationDevice.cu @@ -234,133 +234,132 @@ std::pair buildStateValidityMaskImpl( namespace detail { -void fillOffsetsForDecimalSumState( - bool use64BitOffsets, - void* offsetsMutable, - int32_t numRows, - rmm::cuda_stream_view stream) { - // The offsets buffer holds numRows + 1 entries. - auto const n = static_cast(numRows) + 1; - if (use64BitOffsets) { - launchFillOffsets( - cuda::std::span{static_cast(offsetsMutable), n}, - stream); - } else { - launchFillOffsets( - cuda::std::span{static_cast(offsetsMutable), n}, - stream); - } +template + requires OffsetStorageType +void fillOffsetsForDecimalSumState::operator()( + cudf::mutable_column_view offsetsView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const { + launchFillOffsets( + cuda::std::span{ + offsetsView.data(), static_cast(numRows) + 1}, + stream); } -void packDecimalSumState( - cudf::type_id sumType, - bool use64BitOffsets, - const void* sumPtr, - const int64_t* countPtr, - const void* offsetsPtr, - uint8_t* chars, - int32_t numRows, - rmm::cuda_stream_view stream) { +template void fillOffsetsForDecimalSumState::operator()( + cudf::mutable_column_view offsetsView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const; +template void fillOffsetsForDecimalSumState::operator()( + cudf::mutable_column_view offsetsView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const; + +template + requires OffsetStorageType +void unpackDecimalSumState::operator()( + cudf::column_view offsetsView, + const uint8_t* chars, + cudf::mutable_column_view sumView, + cudf::mutable_column_view countView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const { auto const n = static_cast(numRows); - cuda::std::span counts{countPtr, n}; - if (use64BitOffsets) { - cuda::std::span offsets{ - static_cast(offsetsPtr), n}; - if (sumType == cudf::type_id::DECIMAL64) { - launchPackState( - cuda::std::span{ - static_cast(sumPtr), n}, - counts, - offsets, - chars, - stream); - } else { - launchPackState( - cuda::std::span{ - static_cast(sumPtr), n}, - counts, - offsets, - chars, - stream); - } - } else { - cuda::std::span offsets{ - static_cast(offsetsPtr), n}; - if (sumType == cudf::type_id::DECIMAL64) { - launchPackState( - cuda::std::span{ - static_cast(sumPtr), n}, - counts, - offsets, - chars, - stream); - } else { - launchPackState( - cuda::std::span{ - static_cast(sumPtr), n}, - counts, - offsets, - chars, - stream); - } - } + launchUnpackState( + cuda::std::span{offsetsView.data(), n}, + chars, + cuda::std::span<__int128_t>{sumView.data<__int128_t>(), n}, + cuda::std::span{countView.data(), n}, + stream); } -void unpackDecimalSumState( - bool offsets64, - const void* offsetsPtr, +template void unpackDecimalSumState::operator()( + cudf::column_view offsetsView, const uint8_t* chars, - __int128_t* sums, - int64_t* counts, - int32_t numRows, - rmm::cuda_stream_view stream) { + cudf::mutable_column_view sumView, + cudf::mutable_column_view countView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const; +template void unpackDecimalSumState::operator()( + cudf::column_view offsetsView, + const uint8_t* chars, + cudf::mutable_column_view sumView, + cudf::mutable_column_view countView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const; + +template + requires DecimalSumStorageType +void packDecimalSumState::operator()( + cudf::column_view sumCol, + const int64_t* counts, + cudf::column_view offsetsView, + uint8_t* chars, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const { auto const n = static_cast(numRows); - cuda::std::span<__int128_t> sumsSpan{sums, n}; - cuda::std::span countsSpan{counts, n}; - if (offsets64) { - launchUnpackState( - cuda::std::span{ - static_cast(offsetsPtr), n}, + auto const sums = sumCol.data(); + if (offsetsView.type().id() == cudf::type_id::INT32) { + launchPackState( + cuda::std::span{sums, n}, + cuda::std::span{counts, n}, + cuda::std::span{offsetsView.data(), n}, chars, - sumsSpan, - countsSpan, stream); } else { - launchUnpackState( - cuda::std::span{ - static_cast(offsetsPtr), n}, + launchPackState( + cuda::std::span{sums, n}, + cuda::std::span{counts, n}, + cuda::std::span{offsetsView.data(), n}, chars, - sumsSpan, - countsSpan, stream); } } -void averageRoundDecimalSum( - cudf::type_id sumType, - const void* sums, +template void packDecimalSumState::operator()( + cudf::column_view sumCol, const int64_t* counts, - void* out, - int32_t numRows, - rmm::cuda_stream_view stream) { + cudf::column_view offsetsView, + uint8_t* chars, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const; +template void packDecimalSumState::operator()<__int128_t>( + cudf::column_view sumCol, + const int64_t* counts, + cudf::column_view offsetsView, + uint8_t* chars, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const; + +template + requires DecimalSumStorageType +void averageRoundDecimalSum::operator()( + cudf::column_view sumCol, + const int64_t* counts, + cudf::mutable_column_view outView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const { auto const n = static_cast(numRows); - cuda::std::span countsSpan{counts, n}; - if (sumType == cudf::type_id::DECIMAL64) { - launchAvgRound( - cuda::std::span{static_cast(sums), n}, - countsSpan, - cuda::std::span{static_cast(out), n}, - stream); - } else { - launchAvgRound( - cuda::std::span{ - static_cast(sums), n}, - countsSpan, - cuda::std::span<__int128_t>{static_cast<__int128_t*>(out), n}, - stream); - } + launchAvgRound( + cuda::std::span{sumCol.data(), n}, + cuda::std::span{counts, n}, + cuda::std::span{outView.data(), n}, + stream); } +template void averageRoundDecimalSum::operator()( + cudf::column_view sumCol, + const int64_t* counts, + cudf::mutable_column_view outView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const; +template void averageRoundDecimalSum::operator()<__int128_t>( + cudf::column_view sumCol, + const int64_t* counts, + cudf::mutable_column_view outView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const; + std::pair buildStateValidityMask( const cudf::column_view& sumCol, const cudf::column_view& countCol, diff --git a/velox/experimental/cudf/exec/DecimalAggregationDevice.h b/velox/experimental/cudf/exec/DecimalAggregationDevice.h index 25411ca1a2b..152b3652501 100644 --- a/velox/experimental/cudf/exec/DecimalAggregationDevice.h +++ b/velox/experimental/cudf/exec/DecimalAggregationDevice.h @@ -22,11 +22,21 @@ #include #include +#include +#include #include #include namespace facebook::velox::cudf_velox::detail { +template +concept OffsetStorageType = + std::same_as || std::same_as; + +template +concept DecimalSumStorageType = + std::same_as || std::same_as; + // Size in bytes of each row's packed decimal SUM intermediate state in the // strings payload (count, overflow placeholder, and 128-bit sum split into // words). @@ -35,78 +45,120 @@ constexpr size_t kDecimalSumStateSize = 32; /** * Writes strings-style prefix offsets: offset[i] == i * kDecimalSumStateSize. * - * @param use64BitOffsets whether offsets are INT64 (else INT32). - * @param offsetsMutable output buffer of numRows + 1 offset elements. + * @param offsetsView output offsets column of numRows + 1 elements. * @param numRows number of payload rows. * @param stream CUDA stream for the launch. */ -void fillOffsetsForDecimalSumState( - bool use64BitOffsets, - void* offsetsMutable, - cudf::size_type numRows, - rmm::cuda_stream_view stream); +struct fillOffsetsForDecimalSumState { + template + requires OffsetStorageType + void operator()( + cudf::mutable_column_view offsetsView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const; + + template + requires(!OffsetStorageType) + void operator()( + cudf::mutable_column_view offsetsView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const {} +}; /** * Encodes each row's partial sum and count into the fixed-width device layout * used for VARBINARY interchange. * - * @param sumType element type of sumPtr (DECIMAL64 or DECIMAL128). - * @param use64BitOffsets whether offsetsPtr is INT64 (else INT32). - * @param sumPtr per-row sums. - * @param countPtr per-row int64 counts. - * @param offsetsPtr per-row byte offsets into chars. + * @param sumCol per-row sums. + * @param counts per-row int64 counts. + * @param offsetsView per-row byte offsets into chars. * @param chars output payload buffer. * @param numRows number of rows. * @param stream CUDA stream for the launch. */ -void packDecimalSumState( - cudf::type_id sumType, - bool use64BitOffsets, - const void* sumPtr, - const int64_t* countPtr, - const void* offsetsPtr, - uint8_t* chars, - cudf::size_type numRows, - rmm::cuda_stream_view stream); +struct packDecimalSumState { + template + requires DecimalSumStorageType + void operator()( + cudf::column_view sumCol, + const int64_t* counts, + cudf::column_view offsetsView, + uint8_t* chars, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const; + + template + requires(!DecimalSumStorageType) + void operator()( + cudf::column_view sumCol, + const int64_t* counts, + cudf::column_view offsetsView, + uint8_t* chars, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const {} +}; /** * Inverse of packDecimalSumState. * - * @param offsets64 whether offsetsPtr is INT64 (else INT32). - * @param offsetsPtr per-row byte offsets into chars. + * @param offsetsView per-row byte offsets into chars. * @param chars packed payload buffer. - * @param sums output per-row DECIMAL128 sums. - * @param counts output per-row counts. + * @param sumView output per-row DECIMAL128 sums. + * @param countView output per-row counts. * @param numRows number of rows. * @param stream CUDA stream for the launch. */ -void unpackDecimalSumState( - bool offsets64, - const void* offsetsPtr, - const uint8_t* chars, - __int128_t* sums, - int64_t* counts, - cudf::size_type numRows, - rmm::cuda_stream_view stream); +struct unpackDecimalSumState { + template + requires OffsetStorageType + void operator()( + cudf::column_view offsetsView, + const uint8_t* chars, + cudf::mutable_column_view sumView, + cudf::mutable_column_view countView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const; + + template + requires(!OffsetStorageType) + void operator()( + cudf::column_view offsetsView, + const uint8_t* chars, + cudf::mutable_column_view sumView, + cudf::mutable_column_view countView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const {} +}; /** * Per-row half-up integer divide of sum by count; count == 0 writes zero * (validity is applied separately). * - * @param sumType element type of sums/out (DECIMAL64 or DECIMAL128). - * @param sums per-row sums. + * @param sumCol per-row sums. * @param counts per-row counts. - * @param out output per-row averages. + * @param outView output per-row averages. * @param numRows number of rows. * @param stream CUDA stream for the launch. */ -void averageRoundDecimalSum( - cudf::type_id sumType, - const void* sums, - const int64_t* counts, - void* out, - cudf::size_type numRows, - rmm::cuda_stream_view stream); +struct averageRoundDecimalSum { + template + requires DecimalSumStorageType + void operator()( + cudf::column_view sumCol, + const int64_t* counts, + cudf::mutable_column_view outView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const; + + template + requires(!DecimalSumStorageType) + void operator()( + cudf::column_view sumCol, + const int64_t* counts, + cudf::mutable_column_view outView, + cudf::size_type numRows, + rmm::cuda_stream_view stream) const {} +}; /** * Builds a null mask for rows where sum and count are both valid and count is diff --git a/velox/experimental/cudf/exec/DecimalAggregationState.cpp b/velox/experimental/cudf/exec/DecimalAggregationState.cpp index a0756d82a7c..6238bb3a885 100644 --- a/velox/experimental/cudf/exec/DecimalAggregationState.cpp +++ b/velox/experimental/cudf/exec/DecimalAggregationState.cpp @@ -25,6 +25,7 @@ #include #include #include +#include #include @@ -81,7 +82,6 @@ DecimalSumStateColumns deserializeDecimalSumState( cudf::strings_column_view strings(stateCol); auto offsetsView = strings.offsets(); - auto offsetsType = offsetsView.type().id(); auto charsPtr = reinterpret_cast(strings.chars_begin(stream)); auto sumCol = cudf::make_fixed_width_column( @@ -101,18 +101,19 @@ DecimalSumStateColumns deserializeDecimalSumState( auto countView = countCol->mutable_view(); // numRows is guaranteed positive here - const bool offsets64 = (offsetsType == cudf::type_id::INT64); + auto const offsetsType = offsetsView.type().id(); VELOX_CHECK( - offsets64 || offsetsType == cudf::type_id::INT32, + offsetsType == cudf::type_id::INT32 || + offsetsType == cudf::type_id::INT64, "Decimal sum state requires INT32 or INT64 offsets (offset type is {})", - static_cast(offsetsType)); - detail::unpackDecimalSumState( - offsets64, - offsets64 ? static_cast(offsetsView.data()) - : static_cast(offsetsView.data()), + cudf::type_to_name(offsetsView.type())); + cudf::type_dispatcher( + offsetsView.type(), + detail::unpackDecimalSumState{}, + offsetsView, charsPtr, - sumView.data<__int128_t>(), - countView.data(), + sumView, + countView, numRows, stream); @@ -183,32 +184,25 @@ std::unique_ptr serializeDecimalSumState( rmm::device_buffer charsBuf( static_cast(numRows) * detail::kDecimalSumStateSize, stream, mr); - detail::fillOffsetsForDecimalSumState( - useLargeOffsets, - useLargeOffsets ? static_cast(offsetsView.data()) - : static_cast(offsetsView.data()), + auto charsPtr = reinterpret_cast(charsBuf.data()); + cudf::type_dispatcher( + offsetsView.type(), + detail::fillOffsetsForDecimalSumState{}, + offsetsView, rowCount, stream); - auto charsPtr = reinterpret_cast(charsBuf.data()); - const void* offsetsPtr = useLargeOffsets - ? static_cast(offsetsView.data()) - : static_cast(offsetsView.data()); - const auto sumType = sumCol.type().id(); VELOX_CHECK( - sumType == cudf::type_id::DECIMAL64 || - sumType == cudf::type_id::DECIMAL128, + sumCol.type().id() == cudf::type_id::DECIMAL64 || + sumCol.type().id() == cudf::type_id::DECIMAL128, "Unsupported decimal sum column type (type is {})", cudf::type_to_name(sumCol.type())); - const void* sumPtr = sumType == cudf::type_id::DECIMAL64 - ? static_cast(sumCol.data()) - : static_cast(sumCol.data<__int128_t>()); - detail::packDecimalSumState( - sumType, - useLargeOffsets, - sumPtr, + cudf::type_dispatcher( + sumCol.type(), + detail::packDecimalSumState{}, + sumCol, countCol.data(), - offsetsPtr, + offsetsView, charsPtr, rowCount, stream); @@ -250,15 +244,14 @@ std::unique_ptr computeDecimalAverage( if (numRows > 0) { auto const rowCount = static_cast(numRows); - const auto sumType = sumCol.type().id(); - const void* sumsPtr = sumType == cudf::type_id::DECIMAL64 - ? static_cast(sumCol.data()) - : static_cast(sumCol.data<__int128_t>()); - void* outPtr = sumType == cudf::type_id::DECIMAL64 - ? static_cast(out->mutable_view().data()) - : static_cast(out->mutable_view().data<__int128_t>()); - detail::averageRoundDecimalSum( - sumType, sumsPtr, countCol.data(), outPtr, rowCount, stream); + cudf::type_dispatcher( + sumCol.type(), + detail::averageRoundDecimalSum{}, + sumCol, + countCol.data(), + out->mutable_view(), + rowCount, + stream); } auto [nullMask, nullCount] =