| 128 | |
| 129 | template <typename T> |
| 130 | void backward( |
| 131 | _megdnn_tensor_in diff, _megdnn_tensor_in rois, _megdnn_tensor_in index, |
| 132 | _megdnn_tensor_out grad, const Param& param) { |
| 133 | using namespace ::megdnn::roi_align; |
| 134 | switch (param.mode) { |
| 135 | case param::ROIAlign::Mode::MAX: |
| 136 | backward_impl<T, BwdMaxPooler<T>>( |
| 137 | diff, rois, index, grad, param.spatial_scale, param.offset, |
| 138 | param.sample_height, param.sample_width); |
| 139 | break; |
| 140 | case param::ROIAlign::Mode::AVERAGE: |
| 141 | backward_impl<T, BwdAveragePooler<T>>( |
| 142 | diff, rois, index, grad, param.spatial_scale, param.offset, |
| 143 | param.sample_height, param.sample_width); |
| 144 | break; |
| 145 | default: |
| 146 | megdnn_assert_internal(false); |
| 147 | } |
| 148 | } |
| 149 | |
| 150 | } // namespace |
| 151 | |