diff --git a/src/fvdb/detail/ops/gsplat/FusedSSIM.cu b/src/fvdb/detail/ops/gsplat/FusedSSIM.cu index bbcff57d..812148ec 100644 --- a/src/fvdb/detail/ops/gsplat/FusedSSIM.cu +++ b/src/fvdb/detail/ops/gsplat/FusedSSIM.cu @@ -26,16 +26,17 @@ // SPDX-License-Identifier: Apache-2.0 // #include +#include #include -#include - +#include #include #include #include #include +#include namespace fvdb { @@ -582,45 +583,93 @@ fusedSSIMBackwardCUDA(double C1, namespace { +struct ImageBlockChunk { + size_t tileRowOffset; + size_t tileRowCount; + int blockOffset; + int blockCount; +}; + +ImageBlockChunk +imageBlockChunk(int B, int H, int W, int deviceId) { + // Blocks are flattened with x varying fastest. Split only at full-width tile-row + // boundaries so each device's NCHW working set can be expressed as compact row ranges. + const size_t blocksPerTileRow = (W + BLOCK_X - 1) / BLOCK_X; + const size_t tileRowsPerImage = (H + BLOCK_Y - 1) / BLOCK_Y; + const size_t globalTileRows = B * tileRowsPerImage; + + size_t localTileRowOffset, localTileRowCount; + std::tie(localTileRowOffset, localTileRowCount) = deviceChunk(globalTileRows, deviceId); + + return {localTileRowOffset, + localTileRowCount, + static_cast(localTileRowOffset * blocksPerTileRow), + static_cast(localTileRowCount * blocksPerTileRow)}; +} + void -imagePrefetchBatchAsync(const torch::TensorList &tensors, - int localElementOffset, - int localElementCount, - int deviceId, - cudaStream_t stream) { - TORCH_CHECK(stream, "cudaMemPrefetchBatchAsync does not support the default stream"); -#if (CUDART_VERSION < 13000) - for (size_t i = 0; i < tensors.size(); ++i) { - const auto &tensor = tensors[i]; - TORCH_CHECK(tensor.is_contiguous(), "Tensor to prefetch is not contiguous"); - C10_CUDA_CHECK( - nanovdb::util::cuda::memPrefetchAsync(tensor.data_ptr() + localElementOffset, - localElementCount * sizeof(float), - deviceId, - stream)); +appendImagePrefetchRanges(std::vector &prefetchPointers, + std::vector &prefetchSizes, + const torch::TensorList &tensors, + size_t tileRowOffset, + size_t tileRowCount, + int B, + int CH, + int H, + int W) { + if (!tileRowCount) { + return; } -#else - std::vector prefetchPointers; - std::vector prefetchSizes; - cudaMemLocation location = {cudaMemLocationTypeDevice, deviceId}; - std::vector prefetchLocations = {location}; - std::vector prefetchLocationIndices = {0}; - - for (size_t i = 0; i < tensors.size(); ++i) { - const auto &tensor = tensors[i]; + + const size_t tileRowsPerImage = (H + BLOCK_Y - 1) / BLOCK_Y; + const size_t tileRowEnd = tileRowOffset + tileRowCount; + + TORCH_CHECK(tileRowEnd <= B * tileRowsPerImage, "Invalid image tile-row range"); + + for (const auto &tensor: tensors) { TORCH_CHECK(tensor.is_contiguous(), "Tensor to prefetch is not contiguous"); - prefetchPointers.emplace_back(tensor.data_ptr() + localElementOffset); - prefetchSizes.emplace_back(localElementCount * sizeof(float)); + TORCH_CHECK(tensor.dim() == 4 && tensor.size(0) == B && tensor.size(1) == CH && + tensor.size(2) == H && tensor.size(3) == W, + "Tensor to prefetch does not match the input image shape"); + + const size_t firstTensorRange = prefetchPointers.size(); + const size_t scalarSize = c10::elementSize(tensor.scalar_type()); + auto *tensorData = static_cast(tensor.data_ptr()); + + const size_t firstBatch = tileRowOffset / tileRowsPerImage; + const size_t lastBatch = (tileRowEnd - 1) / tileRowsPerImage; + for (size_t batch = firstBatch; batch <= lastBatch; ++batch) { + const size_t batchTileRowOffset = batch * tileRowsPerImage; + const size_t firstTileRow = + std::max(tileRowOffset, batchTileRowOffset) - batchTileRowOffset; + const size_t lastTileRow = + std::min(tileRowEnd, batchTileRowOffset + tileRowsPerImage) - batchTileRowOffset; + + // Prefetch only the rows owned by this device. Halo reads remain demand-driven so + // adjacent devices never issue prefetches for overlapping logical ranges. + const size_t firstRow = firstTileRow * BLOCK_Y; + const size_t lastRow = std::min(lastTileRow * BLOCK_Y, static_cast(H)); + const size_t rowCount = lastRow - firstRow; + + for (int channel = 0; channel < CH; ++channel) { + const size_t elementOffset = batch * tensor.stride(0) + channel * tensor.stride(1) + + firstRow * tensor.stride(2); + auto *pointer = tensorData + elementOffset * scalarSize; + const size_t byteCount = rowCount * static_cast(W) * scalarSize; + + if (prefetchPointers.size() > firstTensorRange) { + auto *previousEnd = + static_cast(prefetchPointers.back()) + prefetchSizes.back(); + if (previousEnd == pointer) { + prefetchSizes.back() += byteCount; + continue; + } + } + prefetchPointers.emplace_back(pointer); + prefetchSizes.emplace_back(byteCount); + } + } } - C10_CUDA_CHECK(cudaMemPrefetchBatchAsync(prefetchPointers.data(), - prefetchSizes.data(), - prefetchPointers.size(), - prefetchLocations.data(), - prefetchLocationIndices.data(), - prefetchLocations.size(), - 0, - stream)); -#endif } } // namespace @@ -654,6 +703,13 @@ fusedSSIMPrivateUse1( auto img1_ = img1.contiguous(); auto img2_ = img2.contiguous(); + std::vector imageTensors = {img1_, img2_, ssim_map}; + if (train) { + imageTensors.emplace_back(dm_dmu1); + imageTensors.emplace_back(dm_dsigma1_sq); + imageTensors.emplace_back(dm_dsigma12); + } + std::vector events(c10::cuda::device_count()); for (const auto deviceId: c10::irange(c10::cuda::device_count())) { C10_CUDA_CHECK(cudaSetDevice(deviceId)); @@ -662,34 +718,25 @@ fusedSSIMPrivateUse1( C10_CUDA_CHECK(cudaEventRecord(events[deviceId], stream)); } - const auto globalBlockCount = ((W + BLOCK_X - 1) / BLOCK_X) * ((H + BLOCK_Y - 1) / BLOCK_Y) * B; for (const auto deviceId: c10::irange(c10::cuda::device_count())) { C10_CUDA_CHECK(cudaSetDevice(deviceId)); auto stream = c10::cuda::getStreamFromPool(false, deviceId); C10_CUDA_CHECK(cudaStreamWaitEvent(stream, events[deviceId])); - constexpr size_t kAlignment = kPageSize / (sizeof(float) * BLOCK_X * BLOCK_Y); - int localBlockOffset, localBlockCount; - std::tie(localBlockOffset, localBlockCount) = - deviceAlignedChunk(kAlignment, globalBlockCount, deviceId); - - if (localBlockCount) { - auto localElementOffset = localBlockOffset * BLOCK_X * BLOCK_Y * CH; - auto localElementCount = localBlockCount * BLOCK_X * BLOCK_Y * CH; - if (localElementOffset + localElementCount > img1_.numel()) { - localElementOffset = std::min(localElementOffset, static_cast(img1_.numel())); - localElementCount = std::min(localElementCount, - static_cast(img1_.numel()) - localElementOffset); - } - - std::vector tensors = {img1_, img2_, ssim_map}; - if (train) { - tensors.emplace_back(dm_dmu1); - tensors.emplace_back(dm_dsigma1_sq); - tensors.emplace_back(dm_dsigma12); - } - imagePrefetchBatchAsync( - tensors, localElementOffset, localElementCount, deviceId, stream); + const auto chunk = imageBlockChunk(B, H, W, deviceId); + if (chunk.blockCount) { + std::vector prefetchPointers; + std::vector prefetchSizes; + appendImagePrefetchRanges(prefetchPointers, + prefetchSizes, + imageTensors, + chunk.tileRowOffset, + chunk.tileRowCount, + B, + CH, + H, + W); + memPrefetchBatchAsync(prefetchPointers, prefetchSizes, deviceId, stream); } C10_CUDA_CHECK(cudaEventRecord(events[deviceId], stream)); } @@ -700,26 +747,14 @@ fusedSSIMPrivateUse1( C10_CUDA_CHECK(cudaStreamWaitEvent(stream, events[deviceId])); C10_CUDA_CHECK(cudaEventDestroy(events[deviceId])); - constexpr size_t kAlignment = kPageSize / (sizeof(float) * BLOCK_X * BLOCK_Y); - int localBlockOffset, localBlockCount; - std::tie(localBlockOffset, localBlockCount) = - deviceAlignedChunk(kAlignment, globalBlockCount, deviceId); - - if (localBlockCount) { - auto localElementOffset = localBlockOffset * BLOCK_X * BLOCK_Y * CH; - auto localElementCount = localBlockCount * BLOCK_X * BLOCK_Y * CH; - if (localElementOffset + localElementCount > img1_.numel()) { - localElementOffset = std::min(localElementOffset, static_cast(img1_.numel())); - localElementCount = std::min(localElementCount, - static_cast(img1_.numel()) - localElementOffset); - } - + const auto chunk = imageBlockChunk(B, H, W, deviceId); + if (chunk.blockCount) { // Launch config - dim3 grid(localBlockCount); + dim3 grid(chunk.blockCount); dim3 block(BLOCK_X, BLOCK_Y); fusedSSIMKernel<<>>( - localBlockOffset, + chunk.blockOffset, B, H, W, @@ -776,6 +811,9 @@ fusedSSIMBackwardPrivateUse1(double C1, auto dm_dsigma1_sq_ = dm_dsigma1_sq.contiguous(); auto dm_dsigma12_ = dm_dsigma12.contiguous(); + std::vector imageTensors = { + img1_, img2_, dL_dmap_, dm_dmu1_, dm_dsigma1_sq_, dm_dsigma12_, dL_dimg1}; + std::vector events(c10::cuda::device_count()); for (const auto deviceId: c10::irange(c10::cuda::device_count())) { C10_CUDA_CHECK(cudaSetDevice(deviceId)); @@ -784,29 +822,25 @@ fusedSSIMBackwardPrivateUse1(double C1, C10_CUDA_CHECK(cudaEventRecord(events[deviceId], stream)); } - const auto globalBlockCount = ((W + BLOCK_X - 1) / BLOCK_X) * ((H + BLOCK_Y - 1) / BLOCK_Y) * B; for (const auto deviceId: c10::irange(c10::cuda::device_count())) { C10_CUDA_CHECK(cudaSetDevice(deviceId)); auto stream = c10::cuda::getStreamFromPool(false, deviceId); C10_CUDA_CHECK(cudaStreamWaitEvent(stream, events[deviceId])); - constexpr size_t kAlignment = kPageSize / (sizeof(float) * BLOCK_X * BLOCK_Y); - int localBlockOffset, localBlockCount; - std::tie(localBlockOffset, localBlockCount) = - deviceAlignedChunk(kAlignment, globalBlockCount, deviceId); - - if (localBlockCount) { - auto localElementOffset = localBlockOffset * BLOCK_X * BLOCK_Y * CH; - auto localElementCount = localBlockCount * BLOCK_X * BLOCK_Y * CH; - if (localElementOffset + localElementCount > img1_.numel()) { - localElementOffset = std::min(localElementOffset, static_cast(img1_.numel())); - localElementCount = std::min(localElementCount, - static_cast(img1_.numel()) - localElementOffset); - } - - std::vector tensors = {dL_dimg1}; - imagePrefetchBatchAsync( - tensors, localElementOffset, localElementCount, deviceId, stream); + const auto chunk = imageBlockChunk(B, H, W, deviceId); + if (chunk.blockCount) { + std::vector prefetchPointers; + std::vector prefetchSizes; + appendImagePrefetchRanges(prefetchPointers, + prefetchSizes, + imageTensors, + chunk.tileRowOffset, + chunk.tileRowCount, + B, + CH, + H, + W); + memPrefetchBatchAsync(prefetchPointers, prefetchSizes, deviceId, stream); } C10_CUDA_CHECK(cudaEventRecord(events[deviceId], stream)); } @@ -817,26 +851,14 @@ fusedSSIMBackwardPrivateUse1(double C1, C10_CUDA_CHECK(cudaStreamWaitEvent(stream, events[deviceId])); C10_CUDA_CHECK(cudaEventDestroy(events[deviceId])); - constexpr size_t kAlignment = kPageSize / (sizeof(float) * BLOCK_X * BLOCK_Y); - int localBlockOffset, localBlockCount; - std::tie(localBlockOffset, localBlockCount) = - deviceAlignedChunk(kAlignment, globalBlockCount, deviceId); - - if (localBlockCount) { - auto localElementOffset = localBlockOffset * BLOCK_X * BLOCK_Y * CH; - auto localElementCount = localBlockCount * BLOCK_X * BLOCK_Y * CH; - if (localElementOffset + localElementCount > img1_.numel()) { - localElementOffset = std::min(localElementOffset, static_cast(img1_.numel())); - localElementCount = std::min(localElementCount, - static_cast(img1_.numel()) - localElementOffset); - } - + const auto chunk = imageBlockChunk(B, H, W, deviceId); + if (chunk.blockCount) { // Launch config - dim3 grid(localBlockCount); + dim3 grid(chunk.blockCount); dim3 block(BLOCK_X, BLOCK_Y); fusedSSIMBackwardKernel<<>>( - localBlockOffset, + chunk.blockOffset, B, H, W,