11#include " memory_op.h"
22
33#include " source_base/memory_recorder.h"
4+ #include " source_base/tool_quit.h"
45#include " source_base/tool_threading.h"
56#ifdef __DSP
67#include " source_base/kernels/dsp/dsp_connector.h"
@@ -549,35 +550,52 @@ void set_memory(FPTYPE* arr, const int var, const size_t size, base_device::Abac
549550
550551template <typename FPTYPE >
551552void synchronize_memory (FPTYPE * arr_out, const FPTYPE * arr_in, const size_t size, base_device::AbacusDevice_t device_type_out, base_device::AbacusDevice_t device_type_in){
552- if (device_type_out == base_device::AbacusDevice_t::CpuDevice || device_type_in == base_device::AbacusDevice_t::CpuDevice){
553+ // The four source/destination combinations are mutually exclusive, so each
554+ // branch must test BOTH devices. Using `||` here made the first branch match
555+ // whenever either side was the CPU, which routed host<->device transfers to
556+ // the host-to-host specialization.
557+ if (device_type_out == base_device::AbacusDevice_t::CpuDevice && device_type_in == base_device::AbacusDevice_t::CpuDevice){
553558 synchronize_memory_op<FPTYPE , DEVICE_CPU , DEVICE_CPU >()(arr_out, arr_in, size);
554559 }
555- else if (device_type_out == base_device::AbacusDevice_t::CpuDevice || device_type_in == base_device::AbacusDevice_t::GpuDevice){
560+ #if __CUDA || __UT_USE_CUDA || __ROCM || __UT_USE_ROCM
561+ else if (device_type_out == base_device::AbacusDevice_t::CpuDevice && device_type_in == base_device::AbacusDevice_t::GpuDevice){
556562 synchronize_memory_op<FPTYPE , DEVICE_CPU , DEVICE_GPU >()(arr_out, arr_in, size);
557563 }
558- else if (device_type_out == base_device::AbacusDevice_t::GpuDevice || device_type_in == base_device::AbacusDevice_t::CpuDevice){
564+ else if (device_type_out == base_device::AbacusDevice_t::GpuDevice && device_type_in == base_device::AbacusDevice_t::CpuDevice){
559565 synchronize_memory_op<FPTYPE , DEVICE_GPU , DEVICE_CPU >()(arr_out, arr_in, size);
560566 }
561- else if (device_type_out == base_device::AbacusDevice_t::GpuDevice || device_type_in == base_device::AbacusDevice_t::GpuDevice){
567+ else if (device_type_out == base_device::AbacusDevice_t::GpuDevice && device_type_in == base_device::AbacusDevice_t::GpuDevice){
562568 synchronize_memory_op<FPTYPE , DEVICE_GPU , DEVICE_GPU >()(arr_out, arr_in, size);
563569 }
570+ #endif
571+ else {
572+ ModuleBase::WARNING_QUIT (" base_device::memory::synchronize_memory" ,
573+ " unsupported source/destination device combination" );
574+ }
564575}
565576
566577template <typename FPTYPE_out, typename FPTYPE_in>
567578void cast_memory (FPTYPE_out* arr_out, const FPTYPE_in* arr_in, const size_t size, base_device::AbacusDevice_t device_type_out, base_device::AbacusDevice_t device_type_in)
568579{
569- if (device_type_out == base_device::AbacusDevice_t::CpuDevice || device_type_in == base_device::AbacusDevice_t::CpuDevice){
580+ // See synchronize_memory() above: dispatch on the exact (out, in) device pair.
581+ if (device_type_out == base_device::AbacusDevice_t::CpuDevice && device_type_in == base_device::AbacusDevice_t::CpuDevice){
570582 cast_memory_op<FPTYPE_out, FPTYPE_in, DEVICE_CPU , DEVICE_CPU >()(arr_out, arr_in, size);
571583 }
572- else if (device_type_out == base_device::AbacusDevice_t::CpuDevice || device_type_in == base_device::AbacusDevice_t::GpuDevice){
584+ #if __CUDA || __UT_USE_CUDA || __ROCM || __UT_USE_ROCM
585+ else if (device_type_out == base_device::AbacusDevice_t::CpuDevice && device_type_in == base_device::AbacusDevice_t::GpuDevice){
573586 cast_memory_op<FPTYPE_out, FPTYPE_in, DEVICE_CPU , DEVICE_GPU >()(arr_out, arr_in, size);
574587 }
575- else if (device_type_out == base_device::AbacusDevice_t::GpuDevice || device_type_in == base_device::AbacusDevice_t::CpuDevice){
588+ else if (device_type_out == base_device::AbacusDevice_t::GpuDevice && device_type_in == base_device::AbacusDevice_t::CpuDevice){
576589 cast_memory_op<FPTYPE_out, FPTYPE_in, DEVICE_GPU , DEVICE_CPU >()(arr_out, arr_in, size);
577590 }
578- else if (device_type_out == base_device::AbacusDevice_t::GpuDevice || device_type_in == base_device::AbacusDevice_t::GpuDevice){
591+ else if (device_type_out == base_device::AbacusDevice_t::GpuDevice && device_type_in == base_device::AbacusDevice_t::GpuDevice){
579592 cast_memory_op<FPTYPE_out, FPTYPE_in, DEVICE_GPU , DEVICE_GPU >()(arr_out, arr_in, size);
580593 }
594+ #endif
595+ else {
596+ ModuleBase::WARNING_QUIT (" base_device::memory::cast_memory" ,
597+ " unsupported source/destination device combination" );
598+ }
581599}
582600
583601template <typename FPTYPE >
@@ -591,5 +609,24 @@ void delete_memory(FPTYPE* arr, base_device::AbacusDevice_t device_type)
591609 }
592610}
593611
612+ // Explicit instantiations of the runtime-dispatch wrappers, so that the
613+ // declarations in memory_op.h can actually be linked from another translation
614+ // unit (and covered by unit tests). cast_memory is instantiated only for the
615+ // type pairs that cast_memory_op provides for all four device combinations.
616+ template void synchronize_memory<int >(int *, const int *, const size_t , base_device::AbacusDevice_t, base_device::AbacusDevice_t);
617+ template void synchronize_memory<float >(float *, const float *, const size_t , base_device::AbacusDevice_t, base_device::AbacusDevice_t);
618+ template void synchronize_memory<double >(double *, const double *, const size_t , base_device::AbacusDevice_t, base_device::AbacusDevice_t);
619+ template void synchronize_memory<std::complex <float >>(std::complex <float >*, const std::complex <float >*, const size_t , base_device::AbacusDevice_t, base_device::AbacusDevice_t);
620+ template void synchronize_memory<std::complex <double >>(std::complex <double >*, const std::complex <double >*, const size_t , base_device::AbacusDevice_t, base_device::AbacusDevice_t);
621+
622+ template void cast_memory<float , float >(float *, const float *, const size_t , base_device::AbacusDevice_t, base_device::AbacusDevice_t);
623+ template void cast_memory<double , double >(double *, const double *, const size_t , base_device::AbacusDevice_t, base_device::AbacusDevice_t);
624+ template void cast_memory<float , double >(float *, const double *, const size_t , base_device::AbacusDevice_t, base_device::AbacusDevice_t);
625+ template void cast_memory<double , float >(double *, const float *, const size_t , base_device::AbacusDevice_t, base_device::AbacusDevice_t);
626+ template void cast_memory<std::complex <float >, std::complex <float >>(std::complex <float >*, const std::complex <float >*, const size_t , base_device::AbacusDevice_t, base_device::AbacusDevice_t);
627+ template void cast_memory<std::complex <double >, std::complex <double >>(std::complex <double >*, const std::complex <double >*, const size_t , base_device::AbacusDevice_t, base_device::AbacusDevice_t);
628+ template void cast_memory<std::complex <float >, std::complex <double >>(std::complex <float >*, const std::complex <double >*, const size_t , base_device::AbacusDevice_t, base_device::AbacusDevice_t);
629+ template void cast_memory<std::complex <double >, std::complex <float >>(std::complex <double >*, const std::complex <float >*, const size_t , base_device::AbacusDevice_t, base_device::AbacusDevice_t);
630+
594631} // namespace memory
595632} // namespace base_device
0 commit comments