| 74 | } |
| 75 | |
| 76 | void MHABackwardProxyBase::layout_refill(MHA_PROXY_BACKWARD_LAYOUT_CONST_PARAM) { |
| 77 | MEGDNN_MARK_USED_VAR(attn_weight); |
| 78 | MEGDNN_MARK_USED_VAR(mask_reservespace); |
| 79 | MEGDNN_MARK_USED_VAR(handle); |
| 80 | MEGDNN_MARK_USED_VAR(diff); |
| 81 | MEGDNN_MARK_USED_VAR(attn_mask); |
| 82 | MEGDNN_MARK_USED_VAR(othr_reservespace); |
| 83 | MEGDNN_MARK_USED_VAR(dqueries); |
| 84 | MEGDNN_MARK_USED_VAR(dkeys); |
| 85 | MEGDNN_MARK_USED_VAR(dvalues); |
| 86 | MEGDNN_MARK_USED_VAR(dqkvo_weight_bias); |
| 87 | MEGDNN_MARK_USED_VAR(dbias_k); |
| 88 | MEGDNN_MARK_USED_VAR(dbias_v); |
| 89 | // proxy opr |
| 90 | m_softmaxbw_opr->param().axis = -1; |
| 91 | m_matmul_opr->param().format = param::MatrixMul::Format::DEFAULT; |
| 92 | m_bmatmul_opr->param().format = param::MatrixMul::Format::DEFAULT; |
| 93 | m_dropoutbw_opr->param().seed = param.seed; |
| 94 | m_dropout_opr->param().seed = param.seed; |
| 95 | m_reduce_opr->param().mode = param::Reduce::Mode::SUM; |
| 96 | m_reduce_opr->param().data_type = param::Reduce::DataType::DEFAULT; |
| 97 | |
| 98 | m_head = param.num_heads; |
| 99 | m_embed_size = param.embeding_size; |
| 100 | m_ksize = param.k_size; |
| 101 | m_vsize = param.v_size; |
| 102 | m_qproj_size = param.qproj_size; |
| 103 | m_kproj_size = param.kproj_size; |
| 104 | m_vproj_size = param.vproj_size; |
| 105 | m_oproj_size = param.oproj_size; |
| 106 | m_qbias = param.qbias; |
| 107 | m_kbias = param.kbias; |
| 108 | m_vbias = param.vbias; |
| 109 | m_obias = param.obias; |
| 110 | auto cal_type = qkvo_weight_bias.dtype; |
| 111 | m_grad_qin_layout = queries; |
| 112 | m_grad_kin_layout = keys; |
| 113 | m_grad_vin_layout = values; |
| 114 | |
| 115 | auto reflash_dtype = [&](DType dtype) { |
| 116 | m_grad_drop2_layout.dtype = dtype; |
| 117 | m_grad_out_layout.dtype = dtype; |
| 118 | m_grad_z_layout.dtype = dtype; |
| 119 | m_grad_wo_layout.dtype = dtype; |
| 120 | m_grad_bo_layout.dtype = dtype; |
| 121 | m_grad_nz_layout.dtype = dtype; |
| 122 | m_grad_nv_layout.dtype = dtype; |
| 123 | m_grad_ny_layout.dtype = dtype; |
| 124 | m_grad_drop1_layout.dtype = dtype; |
| 125 | m_grad_nx_layout.dtype = dtype; |
| 126 | m_grad_nq_layout.dtype = dtype; |
| 127 | m_grad_nk_layout.dtype = dtype; |
| 128 | m_grad_q_layout.dtype = dtype; |
| 129 | m_grad_k_layout.dtype = dtype; |
| 130 | m_grad_v_layout.dtype = dtype; |
| 131 | m_grad_qin_layout.dtype = dtype; |
| 132 | m_grad_wq_layout.dtype = dtype; |
| 133 | m_grad_bq_layout.dtype = dtype; |
nothing calls this directly
no test coverage detected