MCPcopy Create free account
hub / github.com/LBANN/lbann / copy_tensor

Function copy_tensor

src/utils/rocm.cpp:175–242  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

173
174template <typename TensorDataType>
175void copy_tensor(hipStream_t stream,
176 const std::vector<size_t>& dims,
177 const TensorDataType* input,
178 const std::vector<size_t>& input_strides,
179 TensorDataType* output,
180 const std::vector<size_t>& output_strides)
181{
182
183 // Check inputs
184 if (dims.empty() || dims.size() > 4) {
185 LBANN_ERROR("invalid number of tensor dimensions (", dims.size(), ")");
186 }
187 if (dims.size() != input_strides.size()) {
188 LBANN_ERROR("number of input strides (",
189 input_strides.size(),
190 ") ",
191 "does not match number of tensor dimensions (",
192 dims.size(),
193 ")");
194 }
195 if (dims.size() != output_strides.size()) {
196 LBANN_ERROR("number of output strides (",
197 output_strides.size(),
198 ") ",
199 "does not match number of tensor dimensions (",
200 dims.size(),
201 ")");
202 }
203
204 // Pad tensor dimensions to 4D
205 std::vector<int> rdims(dims.rbegin(), dims.rend()),
206 input_rstrides(input_strides.rbegin(), input_strides.rend()),
207 output_rstrides(output_strides.rbegin(), output_strides.rend());
208 rdims.resize(4, 1);
209 input_rstrides.resize(4, input_rstrides.back());
210 output_rstrides.resize(4, output_rstrides.back());
211
212 // Launch HIP kernel
213 const auto size = get_linear_size(dims);
214 if (size > 0) {
215 constexpr size_t block_size = 64;
216 dim3 block_dims, grid_dims;
217 block_dims.x = block_size;
218 block_dims.y = 1;
219 block_dims.z = 1;
220 grid_dims.x = (rdims[0] + block_dims.x - 1) / block_dims.x;
221 grid_dims.y = (rdims[1] + block_dims.y - 1) / block_dims.y;
222 grid_dims.z = (rdims[2] + block_dims.z - 1) / block_dims.z;
223 grid_dims.y = El::Min(grid_dims.y, 65535);
224 grid_dims.z = El::Min(grid_dims.z, 65535);
225 hipLaunchKernelGGL(copy_4d_kernel,
226 dim3(grid_dims),
227 dim3(block_dims),
228 0,
229 stream,
230 {rdims[3], rdims[2], rdims[1], rdims[0]},
231 input,
232 {input_rstrides[3],

Callers 2

fp_compute_implFunction · 0.85
bp_compute_implFunction · 0.85

Calls 5

get_linear_sizeFunction · 0.85
MinClass · 0.85
emptyMethod · 0.45
sizeMethod · 0.45
resizeMethod · 0.45

Tested by

no test coverage detected