34 size_t& max_temp_size,
35 int texture_alignment,
37 const std::vector<int>& filter_dims,
40 const TFeat* out_importance,
42 const TFeat* inp_features,
43 const TFeat* inp_neighbors_importance_sum,
44 const int64_t* inp_neighbors_prefix_sum,
45 size_t neighbors_index_size,
46 const TIndex* neighbors_index,
47 const TKernelIndex* neighbors_kernel_index,
48 const TFeat* neighbors_importance,
49 const int64_t* neighbors_row_splits,
52 const bool get_temp_size = !temp;
56 temp_size = std::numeric_limits<int64_t>::max();
61 const int in_channels = filter_dims[filter_dims.size() - 2];
62 const int out_channels = filter_dims[filter_dims.size() - 1];
64 int num_kernel_elements = 1;
65 for (
size_t i = 0; i < filter_dims.size() - 2; ++i)
66 num_kernel_elements *= filter_dims[i];
68 const size_t min_num_cols_per_run = std::min(
size_t(num_out),
size_t(32));
69 const size_t max_num_cols_per_run = num_out;
70 const size_t bytes_per_column =
71 sizeof(TFeat) * (num_kernel_elements * in_channels);
72 const size_t min_temp_size_bytes = min_num_cols_per_run * bytes_per_column;
73 const size_t max_temp_size_bytes = max_num_cols_per_run * bytes_per_column;
76 std::pair<char*, size_t> tmp =
77 mem_temp.
Alloc<
char>(min_temp_size_bytes);
80 mem_temp.
Alloc<
char>(max_temp_size_bytes);
81 max_temp_size = mem_temp.
MaxUsed();
87 if (mem_columns.second < min_temp_size_bytes) {
89 ss <<
"temp is too small " << mem_columns.second
90 <<
" bytes. Expected at least " << min_temp_size_bytes <<
" bytes\n";
91 throw std::runtime_error(ss.str());
94 sycl::event out_features_fill_event =
95 queue.fill(out_features, TOut(0),
size_t(num_out) * out_channels);
97 size_t num_cols_per_run =
98 std::min(mem_columns.second / bytes_per_column,
size_t(num_out));
100 TFeat* columns = (TFeat*)mem_columns.first;
102 size_t num_runs = DivUp(num_out, num_cols_per_run);
109 sycl::event prev_gemm_event;
110 for (
size_t run_i = 0; run_i < num_runs; ++run_i) {
111 const TIndex begin_idx = TIndex(run_i * num_cols_per_run);
112 const TIndex end_idx = TIndex(
113 std::min(
size_t(num_out), (run_i + 1) * num_cols_per_run));
114 const size_t num_cols_this_run = end_idx - begin_idx;
118 queue, columns, in_channels, begin_idx, end_idx, num_out,
119 num_inp, inp_features, inp_neighbors_importance_sum,
120 inp_neighbors_prefix_sum, neighbors_index_size, neighbors_index,
121 neighbors_kernel_index, neighbors_importance,
122 neighbors_row_splits, num_kernel_elements, normalize,
123 run_i == 0 ? std::vector<sycl::event>{out_features_fill_event}
124 : std::vector<sycl::event>{prev_gemm_event});
128 const int m = out_channels;
129 const int k = num_kernel_elements * in_channels;
130 const int n =
static_cast<int>(num_cols_this_run);
131 const float alpha = 1;
132 const float*
const A = filter;
134 const float*
const B = columns;
136 const float beta = 1;
137 float* C = out_features + run_i * num_cols_per_run * out_channels;
141 cutlass::layout::ColumnMajor>(
142 queue, m, n, k, alpha, A, lda,
B, ldb, beta, C, ldc, allow_tf32,
143 {fill_column_event});
146 if (out_importance) {
151 out_importance, {prev_gemm_event});
sycl::queue queue
Definition SYCLContext.cpp:88
void SparseConvTransposeComputeFeaturesSYCL(sycl::queue &queue, void *temp, size_t &temp_size, size_t &max_temp_size, int texture_alignment, TOut *out_features, const std::vector< int > &filter_dims, const TFeat *filter, TIndex num_out, const TFeat *out_importance, TIndex num_inp, const TFeat *inp_features, const TFeat *inp_neighbors_importance_sum, const int64_t *inp_neighbors_prefix_sum, size_t neighbors_index_size, const TIndex *neighbors_index, const TKernelIndex *neighbors_kernel_index, const TFeat *neighbors_importance, const int64_t *neighbors_row_splits, bool normalize, bool allow_tf32)
Definition SparseConvTransposeSYCL.h:30
sycl::event FillColumnTransposeSYCL(sycl::queue &queue, TFeat *columns, int in_channels, TIndex begin_idx, TIndex end_idx, TIndex num_out, const TReal *const out_positions, TIndex num_inp, const TReal *const inp_positions, const TFeat *const inp_features, const TFeat *const inp_neighbors_importance_sum, const int64_t *const inp_neighbors_prefix_sum, size_t neighbors_index_size, const TIndex *const neighbors_index, const TFeat *const neighbors_importance, const int64_t *const neighbors_row_splits, const TReal *const extents, const TReal *const offsets, const std::vector< int > &filter_dims, InterpolationMode interpolation, CoordinateMapping coordinate_mapping, bool align_corners, bool individual_extent, bool isotropic_extent, bool normalize, const std::vector< sycl::event > &deps)
Definition ContinuousConvSYCLKernels.cpp:533
sycl::event GemmColumnMajorSYCL(sycl::queue &queue, int m, int n, int k, float alpha, const float *A, int64_t lda, const float *B, int64_t ldb, float beta, float *C, int64_t ldc, bool allow_tf32, const std::vector< sycl::event > &deps)
Definition GemmSYCL.cpp:462