| 1249 | } |
| 1250 | |
| 1251 | void Reduce::KernScheduler::update_ptr( |
| 1252 | const DeviceTensorND& input, const DeviceTensorND& dest, |
| 1253 | const DeviceTensorND& workspace) { |
| 1254 | auto dtype = dest.layout().dtype; |
| 1255 | mgb_assert(dtype.valid()); |
| 1256 | mgb_assert(m_shape_computed); |
| 1257 | |
| 1258 | if (workspace_size()) { |
| 1259 | mgb_assert( |
| 1260 | workspace.layout().dtype == dtype::Byte() && |
| 1261 | workspace.layout().ndim == 1 && |
| 1262 | workspace.shape()[0] >= workspace_size()); |
| 1263 | } |
| 1264 | |
| 1265 | if (m_kern_param.empty()) |
| 1266 | return; |
| 1267 | |
| 1268 | mgb_assert( |
| 1269 | input.layout().total_nr_elems() == |
| 1270 | m_kern_param[0].input.layout.total_nr_elems()); |
| 1271 | mgb_assert( |
| 1272 | dest.shape().total_nr_elems() == |
| 1273 | m_kern_param.back().output.layout.total_nr_elems()); |
| 1274 | auto in_tensor = input.as_megdnn(); |
| 1275 | in_tensor.layout = m_kern_param[0].input.layout; |
| 1276 | m_kern_param[0].input = in_tensor; |
| 1277 | |
| 1278 | dt_byte *workspace_begin = workspace_size() |
| 1279 | ? const_cast<dt_byte*>(workspace.raw_ptr()) |
| 1280 | : nullptr, |
| 1281 | *tmp_reduce_ptr[2] = |
| 1282 | {workspace_begin + m_workspace_spec[0].offset, |
| 1283 | workspace_begin + m_workspace_spec[1].offset}, |
| 1284 | *kern_workspace = workspace_begin + m_workspace_spec[2].offset; |
| 1285 | for (size_t i = 0; i < m_kern_param.size() - 1; ++i) { |
| 1286 | auto optr = tmp_reduce_ptr[i % 2]; |
| 1287 | m_kern_param[i].output.reset_ptr(optr); |
| 1288 | m_kern_param[i + 1].input.reset_ptr(optr); |
| 1289 | } |
| 1290 | for (auto&& i : m_kern_param) |
| 1291 | i.workspace.raw_ptr = kern_workspace; |
| 1292 | auto out_tensor = dest.as_megdnn(); |
| 1293 | out_tensor.layout = m_kern_param.back().output.layout; |
| 1294 | m_kern_param.back().output = out_tensor; |
| 1295 | } |
| 1296 | |
| 1297 | void Reduce::KernScheduler::execute( |
| 1298 | megdnn::Reduce* opr, const DeviceTensorND& input, const DeviceTensorND& dest) { |