Skip to content

Commit 111b939

Browse files
committed
Refactor FusedSSIM prefetch handling
Signed-off-by: Matthew Cong <mcong@nvidia.com>
1 parent a44cd7e commit 111b939

1 file changed

Lines changed: 54 additions & 89 deletions

File tree

src/fvdb/detail/ops/gsplat/FusedSSIM.cu

Lines changed: 54 additions & 89 deletions
Original file line numberDiff line numberDiff line change
@@ -26,16 +26,17 @@
2626
// SPDX-License-Identifier: Apache-2.0
2727
//
2828
#include <fvdb/detail/ops/gsplat/FusedSSIM.h>
29+
#include <fvdb/detail/utils/cuda/Prefetch.h>
2930
#include <fvdb/detail/utils/cuda/Utils.cuh>
3031

31-
#include <nanovdb/util/cuda/Util.h>
32-
32+
#include <c10/core/ScalarType.h>
3333
#include <c10/cuda/CUDAGuard.h>
3434
#include <torch/types.h>
3535

3636
#include <cooperative_groups.h>
3737

3838
#include <algorithm>
39+
#include <cstdint>
3940

4041
namespace fvdb {
4142

@@ -589,11 +590,6 @@ struct ImageBlockChunk {
589590
int blockCount;
590591
};
591592

592-
struct ImagePrefetchRange {
593-
void *pointer;
594-
size_t byteCount;
595-
};
596-
597593
ImageBlockChunk
598594
imageBlockChunk(int B, int H, int W, int deviceId) {
599595
// Blocks are flattened with x varying fastest. Split only at full-width tile-row
@@ -612,7 +608,8 @@ imageBlockChunk(int B, int H, int W, int deviceId) {
612608
}
613609

614610
void
615-
appendImagePrefetchRanges(std::vector<ImagePrefetchRange> &ranges,
611+
appendImagePrefetchRanges(std::vector<void *> &prefetchPointers,
612+
std::vector<size_t> &prefetchSizes,
616613
const torch::TensorList &tensors,
617614
size_t tileRowOffset,
618615
size_t tileRowCount,
@@ -635,8 +632,9 @@ appendImagePrefetchRanges(std::vector<ImagePrefetchRange> &ranges,
635632
tensor.size(2) == H && tensor.size(3) == W,
636633
"Tensor to prefetch does not match the input image shape");
637634

638-
const size_t firstTensorRange = ranges.size();
639-
auto *tensorData = tensor.data_ptr<float>();
635+
const size_t firstTensorRange = prefetchPointers.size();
636+
const size_t scalarSize = c10::elementSize(tensor.scalar_type());
637+
auto *tensorData = static_cast<uint8_t *>(tensor.data_ptr());
640638

641639
const size_t firstBatch = tileRowOffset / tileRowsPerImage;
642640
const size_t lastBatch = (tileRowEnd - 1) / tileRowsPerImage;
@@ -649,68 +647,30 @@ appendImagePrefetchRanges(std::vector<ImagePrefetchRange> &ranges,
649647

650648
// Prefetch only the rows owned by this device. Halo reads remain demand-driven so
651649
// adjacent devices never issue prefetches for overlapping logical ranges.
652-
const size_t firstRow = firstTileRow * BLOCK_Y;
653-
const size_t lastRow = std::min(lastTileRow * BLOCK_Y, static_cast<size_t>(H));
654-
const size_t elementCount = (lastRow - firstRow) * W;
650+
const size_t firstRow = firstTileRow * BLOCK_Y;
651+
const size_t lastRow = std::min(lastTileRow * BLOCK_Y, static_cast<size_t>(H));
655652

656653
for (int channel = 0; channel < CH; ++channel) {
657-
const size_t elementOffset = ((batch * CH + channel) * H + firstRow) * W;
658-
auto *pointer = tensorData + elementOffset;
659-
const size_t byteCount = elementCount * sizeof(float);
654+
const size_t elementOffset = batch * tensor.stride(0) + channel * tensor.stride(1) +
655+
firstRow * tensor.stride(2);
656+
auto *pointer = tensorData + elementOffset * scalarSize;
657+
const size_t byteCount = (lastRow - firstRow) * tensor.stride(2) * scalarSize;
660658

661-
if (ranges.size() > firstTensorRange) {
662-
auto &previousRange = ranges.back();
659+
if (prefetchPointers.size() > firstTensorRange) {
663660
auto *previousEnd =
664-
static_cast<char *>(previousRange.pointer) + previousRange.byteCount;
665-
if (previousEnd == reinterpret_cast<char *>(pointer)) {
666-
previousRange.byteCount += byteCount;
661+
static_cast<uint8_t *>(prefetchPointers.back()) + prefetchSizes.back();
662+
if (previousEnd == pointer) {
663+
prefetchSizes.back() += byteCount;
667664
continue;
668665
}
669666
}
670-
ranges.emplace_back(ImagePrefetchRange{pointer, byteCount});
667+
prefetchPointers.emplace_back(pointer);
668+
prefetchSizes.emplace_back(byteCount);
671669
}
672670
}
673671
}
674672
}
675673

676-
void
677-
imagePrefetchBatchAsync(const std::vector<ImagePrefetchRange> &ranges,
678-
int deviceId,
679-
cudaStream_t stream) {
680-
if (ranges.empty()) {
681-
return;
682-
}
683-
684-
TORCH_CHECK(stream, "cudaMemPrefetchBatchAsync does not support the default stream");
685-
#if (CUDART_VERSION < 13000)
686-
for (const auto &range: ranges) {
687-
C10_CUDA_CHECK(nanovdb::util::cuda::memPrefetchAsync(
688-
range.pointer, range.byteCount, deviceId, stream));
689-
}
690-
#else
691-
std::vector<void *> prefetchPointers;
692-
std::vector<size_t> prefetchSizes;
693-
cudaMemLocation location = {cudaMemLocationTypeDevice, deviceId};
694-
std::vector<cudaMemLocation> prefetchLocations = {location};
695-
std::vector<size_t> prefetchLocationIndices = {0};
696-
697-
prefetchPointers.reserve(ranges.size());
698-
prefetchSizes.reserve(ranges.size());
699-
for (const auto &range: ranges) {
700-
prefetchPointers.emplace_back(range.pointer);
701-
prefetchSizes.emplace_back(range.byteCount);
702-
}
703-
C10_CUDA_CHECK(cudaMemPrefetchBatchAsync(prefetchPointers.data(),
704-
prefetchSizes.data(),
705-
prefetchPointers.size(),
706-
prefetchLocations.data(),
707-
prefetchLocationIndices.data(),
708-
prefetchLocations.size(),
709-
0,
710-
stream));
711-
#endif
712-
}
713-
714674
} // namespace
715675

716676
// ------------------------------------------
@@ -742,6 +702,13 @@ fusedSSIMPrivateUse1(
742702
auto img1_ = img1.contiguous();
743703
auto img2_ = img2.contiguous();
744704

705+
std::vector<torch::Tensor> imageTensors = {img1_, img2_, ssim_map};
706+
if (train) {
707+
imageTensors.emplace_back(dm_dmu1);
708+
imageTensors.emplace_back(dm_dsigma1_sq);
709+
imageTensors.emplace_back(dm_dsigma12);
710+
}
711+
745712
std::vector<cudaEvent_t> events(c10::cuda::device_count());
746713
for (const auto deviceId: c10::irange(c10::cuda::device_count())) {
747714
C10_CUDA_CHECK(cudaSetDevice(deviceId));
@@ -757,20 +724,18 @@ fusedSSIMPrivateUse1(
757724

758725
const auto chunk = imageBlockChunk(B, H, W, deviceId);
759726
if (chunk.blockCount) {
760-
std::vector<ImagePrefetchRange> ranges;
761-
std::vector<torch::Tensor> inputTensors = {img1_, img2_};
762-
appendImagePrefetchRanges(
763-
ranges, inputTensors, chunk.tileRowOffset, chunk.tileRowCount, B, CH, H, W);
764-
765-
std::vector<torch::Tensor> outputTensors = {ssim_map};
766-
if (train) {
767-
outputTensors.emplace_back(dm_dmu1);
768-
outputTensors.emplace_back(dm_dsigma1_sq);
769-
outputTensors.emplace_back(dm_dsigma12);
770-
}
771-
appendImagePrefetchRanges(
772-
ranges, outputTensors, chunk.tileRowOffset, chunk.tileRowCount, B, CH, H, W);
773-
imagePrefetchBatchAsync(ranges, deviceId, stream);
727+
std::vector<void *> prefetchPointers;
728+
std::vector<size_t> prefetchSizes;
729+
appendImagePrefetchRanges(prefetchPointers,
730+
prefetchSizes,
731+
imageTensors,
732+
chunk.tileRowOffset,
733+
chunk.tileRowCount,
734+
B,
735+
CH,
736+
H,
737+
W);
738+
memPrefetchBatchAsync(prefetchPointers, prefetchSizes, deviceId, stream);
774739
}
775740
C10_CUDA_CHECK(cudaEventRecord(events[deviceId], stream));
776741
}
@@ -845,6 +810,9 @@ fusedSSIMBackwardPrivateUse1(double C1,
845810
auto dm_dsigma1_sq_ = dm_dsigma1_sq.contiguous();
846811
auto dm_dsigma12_ = dm_dsigma12.contiguous();
847812

813+
std::vector<torch::Tensor> imageTensors = {
814+
img1_, img2_, dL_dmap_, dm_dmu1_, dm_dsigma1_sq_, dm_dsigma12_, dL_dimg1};
815+
848816
std::vector<cudaEvent_t> events(c10::cuda::device_count());
849817
for (const auto deviceId: c10::irange(c10::cuda::device_count())) {
850818
C10_CUDA_CHECK(cudaSetDevice(deviceId));
@@ -860,21 +828,18 @@ fusedSSIMBackwardPrivateUse1(double C1,
860828

861829
const auto chunk = imageBlockChunk(B, H, W, deviceId);
862830
if (chunk.blockCount) {
863-
std::vector<ImagePrefetchRange> ranges;
864-
865-
std::vector<torch::Tensor> imageTensors = {img1_, img2_};
866-
appendImagePrefetchRanges(
867-
ranges, imageTensors, chunk.tileRowOffset, chunk.tileRowCount, B, CH, H, W);
868-
869-
std::vector<torch::Tensor> derivativeTensors = {
870-
dL_dmap_, dm_dmu1_, dm_dsigma1_sq_, dm_dsigma12_};
871-
appendImagePrefetchRanges(
872-
ranges, derivativeTensors, chunk.tileRowOffset, chunk.tileRowCount, B, CH, H, W);
873-
874-
std::vector<torch::Tensor> outputTensors = {dL_dimg1};
875-
appendImagePrefetchRanges(
876-
ranges, outputTensors, chunk.tileRowOffset, chunk.tileRowCount, B, CH, H, W);
877-
imagePrefetchBatchAsync(ranges, deviceId, stream);
831+
std::vector<void *> prefetchPointers;
832+
std::vector<size_t> prefetchSizes;
833+
appendImagePrefetchRanges(prefetchPointers,
834+
prefetchSizes,
835+
imageTensors,
836+
chunk.tileRowOffset,
837+
chunk.tileRowCount,
838+
B,
839+
CH,
840+
H,
841+
W);
842+
memPrefetchBatchAsync(prefetchPointers, prefetchSizes, deviceId, stream);
878843
}
879844
C10_CUDA_CHECK(cudaEventRecord(events[deviceId], stream));
880845
}

0 commit comments

Comments
 (0)