| 41 | } |
| 42 | |
| 43 | void DropoutForwardImpl::exec( |
| 44 | _megdnn_tensor_in inp, _megdnn_tensor_out oup, _megdnn_tensor_out mask, |
| 45 | _megdnn_workspace workspace) { |
| 46 | check_exec(inp.layout, oup.layout, mask.layout, workspace.size); |
| 47 | uint64_t seed = param().seed; |
| 48 | float drop_prob = param().drop_prob; |
| 49 | |
| 50 | if (!dropout_status.initialized()) { |
| 51 | dropout_status.set(cudnn_handle(this->handle()), seed, drop_prob); |
| 52 | } |
| 53 | if (dropout_status.drop_prob != drop_prob) { |
| 54 | dropout_status.drop_prob = drop_prob; |
| 55 | dropout_status.restore_desc(cudnn_handle(this->handle())); |
| 56 | } |
| 57 | megdnn_assert(dropout_status.seed == seed); |
| 58 | |
| 59 | DropoutTensorDesc inp_desc(inp.layout), oup_desc(oup.layout); |
| 60 | auto&& op_desc = dropout_status.desc; |
| 61 | |
| 62 | cudnn_check(cudnnDropoutForward( |
| 63 | cudnn_handle(this->handle()), op_desc.desc, inp_desc.desc, inp.raw_ptr(), |
| 64 | oup_desc.desc, oup.raw_ptr(), mask.raw_ptr(), |
| 65 | mask.layout.total_nr_elems())); |
| 66 | } |
| 67 | |
| 68 | void DropoutBackwardImpl::exec( |
| 69 | _megdnn_tensor_in doup, _megdnn_tensor_in mask, _megdnn_tensor_out dinp, |
nothing calls this directly
no test coverage detected