| 499 | #undef INST |
| 500 | |
| 501 | void ResizeBackwardImpl::exec( |
| 502 | _megdnn_tensor_in diff, _megdnn_tensor_out grad, _megdnn_workspace workspace) { |
| 503 | check_exec(diff.layout, grad.layout, workspace.size); |
| 504 | megdnn_assert( |
| 505 | param().format == param::Resize::Format::NCHW || |
| 506 | param().format == param::Resize::Format::NHWC, |
| 507 | "invalid resize format"); |
| 508 | size_t N, C, IH, IW, OH, OW; |
| 509 | bool is_nhwc = param().format == param::Resize::Format::NHWC; |
| 510 | if (is_nhwc) { |
| 511 | if (param().imode != Param::InterpolationMode::LINEAR && |
| 512 | is_nhwc_contig_wc(grad.layout)) { |
| 513 | megdnn_assert( |
| 514 | 0, |
| 515 | "unsupport mode in resizeBackward, only support param().imode = " |
| 516 | "LINEAR"); |
| 517 | } |
| 518 | N = grad.layout.shape[0]; |
| 519 | C = grad.layout.shape[3]; |
| 520 | IH = grad.layout.shape[1]; |
| 521 | IW = grad.layout.shape[2]; |
| 522 | OH = diff.layout.shape[1]; |
| 523 | OW = diff.layout.shape[2]; |
| 524 | } else { |
| 525 | N = grad.layout.shape[0], C = grad.layout.shape[1], IH = grad.layout.shape[2], |
| 526 | IW = grad.layout.shape[3]; |
| 527 | OH = diff.layout.shape[2], OW = diff.layout.shape[3]; |
| 528 | } |
| 529 | switch (grad.layout.dtype.enumv()) { |
| 530 | #define cb(_t) \ |
| 531 | case DTypeTrait<_t>::enumv: { \ |
| 532 | typedef DTypeTrait<_t>::ctype ct; \ |
| 533 | ct* diff_ptr = diff.ptr<ct>(); \ |
| 534 | ct* grad_ptr = grad.ptr<ct>(); \ |
| 535 | ResizeBackwardImpl::kern_naive( \ |
| 536 | is_nhwc, param().imode, diff_ptr, grad_ptr, N, C, IH, IW, OH, OW); \ |
| 537 | break; \ |
| 538 | } |
| 539 | cb(megdnn::dtype::Float32); |
| 540 | DNN_INC_FLOAT16(cb(megdnn::dtype::Float16)); |
| 541 | default: |
| 542 | megdnn_throw(ssprintf( |
| 543 | "unsupported dtype: %s in resize backward", |
| 544 | grad.layout.dtype.name())); |
| 545 | } |
| 546 | } |
| 547 | |
| 548 | template <typename ctype> |
| 549 | void Resize3DImpl::kern_naive( |
nothing calls this directly
no test coverage detected