| 95 | } |
| 96 | |
| 97 | void MHAForwardProxyBase::layout_refill(MHA_PROXY_FORWARD_LAYOUT_CONST_PARAM) { |
| 98 | MEGDNN_MARK_USED_VAR(handle); |
| 99 | MEGDNN_MARK_USED_VAR(attn_mask); |
| 100 | MEGDNN_MARK_USED_VAR(bias_k); |
| 101 | MEGDNN_MARK_USED_VAR(bias_v); |
| 102 | MEGDNN_MARK_USED_VAR(out); |
| 103 | MEGDNN_MARK_USED_VAR(attn_weight); |
| 104 | MEGDNN_MARK_USED_VAR(mask_reservespace); |
| 105 | MEGDNN_MARK_USED_VAR(othr_reservespace); |
| 106 | |
| 107 | m_heads = param.num_heads; |
| 108 | m_embed_size = param.embeding_size; |
| 109 | m_ksize = param.k_size; |
| 110 | m_vsize = param.v_size; |
| 111 | m_qproj_size = param.qproj_size; |
| 112 | m_kproj_size = param.kproj_size; |
| 113 | m_vproj_size = param.vproj_size; |
| 114 | m_oproj_size = param.oproj_size; |
| 115 | m_qbias = param.qbias; |
| 116 | m_kbias = param.kbias; |
| 117 | m_vbias = param.vbias; |
| 118 | m_obias = param.obias; |
| 119 | auto cal_type = qkvo_weight_bias.dtype; |
| 120 | TensorLayout placeholder_layout; |
| 121 | |
| 122 | auto reflash_dtype = [&](DType dtype) { |
| 123 | m_q_layout.dtype = dtype; |
| 124 | m_k_layout.dtype = dtype; |
| 125 | m_v_layout.dtype = dtype; |
| 126 | m_nq_layout.dtype = dtype; |
| 127 | m_nk_layout.dtype = dtype; |
| 128 | m_nv_layout.dtype = dtype; |
| 129 | m_nx_layout.dtype = dtype; |
| 130 | m_mask1_layout.dtype = dtype; |
| 131 | m_nz_layout.dtype = dtype; |
| 132 | m_z_layout.dtype = dtype; |
| 133 | m_out_layout.dtype = dtype; |
| 134 | m_mask2_layout.dtype = dtype; |
| 135 | }; |
| 136 | reflash_dtype(queries.dtype); |
| 137 | m_datatype = queries.dtype.enumv(); |
| 138 | #define cb(DType) \ |
| 139 | if (queries.dtype.enumv() == DTypeTrait<DType>::enumv) { \ |
| 140 | m_sizeof_datatype = sizeof(DTypeTrait<DType>::ctype); \ |
| 141 | } |
| 142 | MEGDNN_FOREACH_COMPUTING_DTYPE(cb) |
| 143 | #undef cb |
| 144 | |
| 145 | // proxy opr |
| 146 | m_matmul_opr->param().format = param::MatrixMul::Format::DEFAULT; |
| 147 | m_bmatmul_opr->param().format = param::MatrixMul::Format::DEFAULT; |
| 148 | m_softmax_opr->param().axis = -1; |
| 149 | m_dropout_opr->param().seed = param.seed; |
| 150 | |
| 151 | // wq/wk/wv/wo |
| 152 | m_wq_layout = TensorLayout{{m_embed_size, m_qproj_size}, cal_type}; |
| 153 | m_wk_layout = TensorLayout{{m_ksize, m_kproj_size}, cal_type}; |
| 154 | m_wv_layout = TensorLayout{{m_vsize, m_vproj_size}, cal_type}; |
nothing calls this directly
no test coverage detected