11#include < math.h>
22
3- #include < cassert>
3+ #if GOOGLE_CUDA
4+ #include < mutex>
5+ #include < unordered_map>
6+ #endif
47
58#include " device.h"
69#include " tabulate.h"
@@ -1086,6 +1089,36 @@ void launch_tabulate_fusion_se_a(FPTYPE* out,
10861089 is_sorted);
10871090}
10881091
1092+ #if GOOGLE_CUDA
1093+ namespace {
1094+
1095+ struct CudaSharedMemoryLimits {
1096+ size_t standard;
1097+ size_t opt_in;
1098+ };
1099+
1100+ CudaSharedMemoryLimits get_cuda_shared_memory_limits (const int device) {
1101+ static std::mutex cache_mutex;
1102+ static std::unordered_map<int , CudaSharedMemoryLimits> cache;
1103+ std::lock_guard<std::mutex> lock (cache_mutex);
1104+ const auto cached = cache.find (device);
1105+ if (cached != cache.end ()) {
1106+ return cached->second ;
1107+ }
1108+
1109+ cudaDeviceProp properties;
1110+ DPErrcheck (cudaGetDeviceProperties (&properties, device));
1111+ const CudaSharedMemoryLimits limits{
1112+ properties.sharedMemPerBlock ,
1113+ properties.sharedMemPerBlockOptin ,
1114+ };
1115+ cache.emplace (device, limits);
1116+ return limits;
1117+ }
1118+
1119+ } // namespace
1120+ #endif
1121+
10891122template <typename FPTYPE , int MTILE >
10901123void launch_tabulate_fusion_se_a_grad (FPTYPE * dy_dem_x,
10911124 FPTYPE * dy_dem,
@@ -1103,21 +1136,18 @@ void launch_tabulate_fusion_se_a_grad(FPTYPE* dy_dem_x,
11031136#if GOOGLE_CUDA
11041137 const size_t shared_memory = sizeof (FPTYPE ) * MTILE * last_layer_size;
11051138 int device = 0 ;
1106- cudaDeviceProp properties;
11071139 DPErrcheck (cudaGetDevice (&device));
1108- DPErrcheck ( cudaGetDeviceProperties (&properties, device) );
1140+ const CudaSharedMemoryLimits limits = get_cuda_shared_memory_limits ( device);
11091141 const size_t shared_memory_limit =
1110- properties.sharedMemPerBlock > properties.sharedMemPerBlockOptin
1111- ? properties.sharedMemPerBlock
1112- : properties.sharedMemPerBlockOptin ;
1142+ limits.standard > limits.opt_in ? limits.standard : limits.opt_in ;
11131143 if (shared_memory <= shared_memory_limit) {
11141144 auto kernel =
11151145 tabulate_fusion_se_a_grad_fifth_order_polynomial<FPTYPE , MTILE , KK ,
11161146 true >;
1117- if (shared_memory > properties. sharedMemPerBlock ) {
1147+ if (shared_memory > limits. standard ) {
11181148 DPErrcheck (cudaFuncSetAttribute (
11191149 kernel, cudaFuncAttributeMaxDynamicSharedMemorySize,
1120- static_cast <int >(properties. sharedMemPerBlockOptin )));
1150+ static_cast <int >(limits. opt_in )));
11211151 }
11221152 kernel<<<nloc, KK * WARP_SIZE , shared_memory>>> (
11231153 dy_dem_x, dy_dem, dy_dtwo, table, em_x, em, two_embed, dy,
@@ -1181,10 +1211,10 @@ void tabulate_fusion_se_a_gpu(FPTYPE* out,
11811211 const int last_layer_size,
11821212 const bool is_sorted,
11831213 const int ndescrpt) {
1214+ detail::check_se_a_basis_dimension (ndescrpt);
11841215 if (nloc <= 0 ) {
11851216 return ;
11861217 }
1187- assert (ndescrpt == 4 || ndescrpt == 9 || ndescrpt == 16 || ndescrpt == 25 );
11881218 DPErrcheck (gpuGetLastError ());
11891219 DPErrcheck (gpuDeviceSynchronize ());
11901220 if (ndescrpt == 4 ) {
@@ -1223,10 +1253,10 @@ void tabulate_fusion_se_a_grad_gpu(FPTYPE* dy_dem_x,
12231253 const int last_layer_size,
12241254 const bool is_sorted,
12251255 const int ndescrpt) {
1256+ detail::check_se_a_basis_dimension (ndescrpt);
12261257 if (nloc <= 0 ) {
12271258 return ;
12281259 }
1229- assert (ndescrpt == 4 || ndescrpt == 9 || ndescrpt == 16 || ndescrpt == 25 );
12301260 DPErrcheck (gpuGetLastError ());
12311261 DPErrcheck (gpuDeviceSynchronize ());
12321262 DPErrcheck (gpuMemset (dy_dem_x, 0 , sizeof (FPTYPE ) * nloc * nnei));
@@ -1268,10 +1298,10 @@ void tabulate_fusion_se_a_grad_grad_gpu(FPTYPE* dz_dy,
12681298 const int last_layer_size,
12691299 const bool is_sorted,
12701300 const int ndescrpt) {
1301+ detail::check_se_a_basis_dimension (ndescrpt);
12711302 if (nloc <= 0 ) {
12721303 return ;
12731304 }
1274- assert (ndescrpt == 4 || ndescrpt == 9 || ndescrpt == 16 || ndescrpt == 25 );
12751305 DPErrcheck (gpuGetLastError ());
12761306 DPErrcheck (gpuDeviceSynchronize ());
12771307 DPErrcheck (
0 commit comments