Skip to content
Merged
Changes from 3 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
241 changes: 131 additions & 110 deletions src/fvdb/detail/ops/gsplat/FusedSSIM.cu
Original file line number Diff line number Diff line change
Expand Up @@ -26,16 +26,17 @@
// SPDX-License-Identifier: Apache-2.0
//
#include <fvdb/detail/ops/gsplat/FusedSSIM.h>
#include <fvdb/detail/utils/cuda/Prefetch.h>
#include <fvdb/detail/utils/cuda/Utils.cuh>

#include <nanovdb/util/cuda/Util.h>

#include <c10/core/ScalarType.h>
#include <c10/cuda/CUDAGuard.h>
#include <torch/types.h>

#include <cooperative_groups.h>

#include <algorithm>
#include <cstdint>

namespace fvdb {

Expand Down Expand Up @@ -582,45 +583,92 @@ 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<int>(localTileRowOffset * blocksPerTileRow),
static_cast<int>(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<float>() + localElementOffset,
localElementCount * sizeof(float),
deviceId,
stream));
appendImagePrefetchRanges(std::vector<void *> &prefetchPointers,
std::vector<size_t> &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<void *> prefetchPointers;
std::vector<size_t> prefetchSizes;
cudaMemLocation location = {cudaMemLocationTypeDevice, deviceId};
std::vector<cudaMemLocation> prefetchLocations = {location};
std::vector<size_t> 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<float>() + 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<uint8_t *>(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<size_t>(H));

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 = (lastRow - firstRow) * tensor.stride(2) * scalarSize;
Comment thread
matthewdcong marked this conversation as resolved.
Outdated

if (prefetchPointers.size() > firstTensorRange) {
auto *previousEnd =
static_cast<uint8_t *>(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
Expand Down Expand Up @@ -654,6 +702,13 @@ fusedSSIMPrivateUse1(
auto img1_ = img1.contiguous();
auto img2_ = img2.contiguous();

std::vector<torch::Tensor> 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<cudaEvent_t> events(c10::cuda::device_count());
for (const auto deviceId: c10::irange(c10::cuda::device_count())) {
C10_CUDA_CHECK(cudaSetDevice(deviceId));
Expand All @@ -662,34 +717,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<int>(img1_.numel()));
localElementCount = std::min(localElementCount,
static_cast<int>(img1_.numel()) - localElementOffset);
}

std::vector<torch::Tensor> 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<void *> prefetchPointers;
std::vector<size_t> 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));
}
Expand All @@ -700,26 +746,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<int>(img1_.numel()));
localElementCount = std::min(localElementCount,
static_cast<int>(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<<<grid, block, 0, stream>>>(
localBlockOffset,
chunk.blockOffset,
B,
H,
W,
Expand Down Expand Up @@ -776,6 +810,9 @@ fusedSSIMBackwardPrivateUse1(double C1,
auto dm_dsigma1_sq_ = dm_dsigma1_sq.contiguous();
auto dm_dsigma12_ = dm_dsigma12.contiguous();

std::vector<torch::Tensor> imageTensors = {
img1_, img2_, dL_dmap_, dm_dmu1_, dm_dsigma1_sq_, dm_dsigma12_, dL_dimg1};

std::vector<cudaEvent_t> events(c10::cuda::device_count());
for (const auto deviceId: c10::irange(c10::cuda::device_count())) {
C10_CUDA_CHECK(cudaSetDevice(deviceId));
Expand All @@ -784,29 +821,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<int>(img1_.numel()));
localElementCount = std::min(localElementCount,
static_cast<int>(img1_.numel()) - localElementOffset);
}

std::vector<torch::Tensor> tensors = {dL_dimg1};
imagePrefetchBatchAsync(
tensors, localElementOffset, localElementCount, deviceId, stream);
const auto chunk = imageBlockChunk(B, H, W, deviceId);
if (chunk.blockCount) {
std::vector<void *> prefetchPointers;
std::vector<size_t> 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));
}
Expand All @@ -817,26 +850,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<int>(img1_.numel()));
localElementCount = std::min(localElementCount,
static_cast<int>(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<<<grid, block, 0, stream>>>(
localBlockOffset,
chunk.blockOffset,
B,
H,
W,
Expand Down
Loading