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
4041namespace 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-
597593ImageBlockChunk
598594imageBlockChunk (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
614610void
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