MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / exec

Method exec

dnn/src/cuda/matrix_mul/conv1x1.cpp:120–169  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

118}
119
120void MatrixMulForwardImpl::AlgoConv1X1CUDNN::exec(const ExecArgs& args) const {
121 SizeArgs size_args(args.opr, args.layout_a, args.layout_b, args.layout_c);
122
123 auto conv_opr_ptr = prepare_conv_opr(size_args);
124
125 size_t m, k, n;
126 std::tie(m, k, n) = gen_matrixmul_shape(size_args);
127
128 auto bundle = get_workspace_bundle(args.workspace.raw_ptr, size_args);
129 auto A_dst_tensor = args.tensor_a;
130 auto B_dst_tensor = args.tensor_b;
131 if (args.opr->param().transposeA || args.opr->param().transposeB) {
132 auto trans = args.opr->handle()->create_operator<RelayoutForward>();
133
134 auto trans_tensor = [&](size_t workspace_pos, const TensorND& ori_tensor,
135 TensorND& dst_tensor) {
136 TensorLayout dst_layout(
137 {ori_tensor.layout.shape[1], ori_tensor.layout.shape[0]},
138 ori_tensor.layout.dtype);
139 dst_tensor = TensorND(bundle.get(workspace_pos), dst_layout);
140 TensorND src_tensor(ori_tensor.raw_ptr(), dst_layout);
141 src_tensor.layout.stride[0] = ori_tensor.layout.stride[1];
142 src_tensor.layout.stride[1] = ori_tensor.layout.stride[0];
143
144 trans->exec(src_tensor, dst_tensor, args.opr->handle());
145 };
146
147 if (args.opr->param().transposeA) {
148 trans_tensor(1, args.tensor_a, A_dst_tensor);
149 }
150 if (args.opr->param().transposeB) {
151 trans_tensor(bundle.nr_workspace() - 1, args.tensor_b, B_dst_tensor);
152 }
153 }
154
155 TensorLayout src_layout({1, k, 1, n}, args.layout_b.dtype);
156 TensorLayout filter_layout({m, k, 1, 1}, args.layout_a.dtype);
157 TensorLayout dst_layout({1, m, 1, n}, args.layout_c.dtype);
158
159 TensorND src(B_dst_tensor.raw_ptr(), src_layout);
160 TensorND filter(A_dst_tensor.raw_ptr(), filter_layout);
161 TensorND z(nullptr, TensorLayout(src_layout.dtype));
162 TensorND bias(nullptr, TensorLayout(src_layout.dtype));
163 TensorND dst(args.tensor_c.raw_ptr(), dst_layout);
164
165 ConvBiasForwardImpl::AlgoBase::ExecArgs conv_exec_args(
166 static_cast<ConvBiasForwardImpl*>(conv_opr_ptr.get()), src, filter, bias, z,
167 dst, bundle.get_workspace(0));
168 m_impl->exec(conv_exec_args);
169}

Callers

nothing calls this directly

Calls 11

prepare_conv_oprFunction · 0.85
gen_matrixmul_shapeFunction · 0.85
TensorLayoutClass · 0.85
nr_workspaceMethod · 0.80
get_workspace_bundleFunction · 0.50
TensorNDClass · 0.50
paramMethod · 0.45
handleMethod · 0.45
getMethod · 0.45
raw_ptrMethod · 0.45
get_workspaceMethod · 0.45

Tested by

no test coverage detected