MCPcopy Create free account
hub / github.com/arrayfire/arrayfire / matmul

Function matmul

src/backend/opencl/cpu/cpu_sparse_blas.cpp:194–265  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

192////////////////////////////////////////////////////////////////////////////////
193template<typename T>
194Array<T> matmul(const common::SparseArray<T> lhs, const Array<T> rhs,
195 af_mat_prop optLhs, af_mat_prop optRhs) {
196 // MKL: CSRMM Does not support optRhs
197 UNUSED(optRhs);
198
199 lhs.eval();
200 rhs.eval();
201
202 // Similar Operations to GEMM
203 sparse_operation_t lOpts = toSparseTranspose(optLhs);
204
205 int lRowDim = (lOpts == SPARSE_OPERATION_NON_TRANSPOSE) ? 0 : 1;
206 // int lColDim = (lOpts == SPARSE_OPERATION_NON_TRANSPOSE) ? 1 : 0;
207
208 // Unsupported : (rOpts == SPARSE_OPERATION_NON_TRANSPOSE;) ? 1 : 0;
209 static const int rColDim = 1;
210
211 dim4 lDims = lhs.dims();
212 dim4 rDims = rhs.dims();
213 int M = lDims[lRowDim];
214 int N = rDims[rColDim];
215 // int K = lDims[lColDim];
216
217 Array<T> out = createValueArray<T>(af::dim4(M, N, 1, 1), scalar<T>(0));
218 out.eval();
219
220 auto alpha = getScale<T, 1>();
221 auto beta = getScale<T, 0>();
222
223 int ldb = rhs.strides()[1];
224 int ldc = out.strides()[1];
225
226 // get host pointers from mapped memory
227 mapped_ptr<T> rhsPtr = rhs.getMappedPtr(CL_MAP_READ);
228 mapped_ptr<T> outPtr = out.getMappedPtr();
229
230 Array<T> values = lhs.getValues();
231 Array<int> rowIdx = lhs.getRowIdx();
232 Array<int> colIdx = lhs.getColIdx();
233
234 mapped_ptr<T> vPtr = values.getMappedPtr();
235 mapped_ptr<int> rPtr = rowIdx.getMappedPtr();
236 mapped_ptr<int> cPtr = colIdx.getMappedPtr();
237 int *pB = rPtr.get();
238 int *pE = rPtr.get() + 1;
239
240 sparse_matrix_t csrLhs;
241 create_csr_func<T>()(&csrLhs, SPARSE_INDEX_BASE_ZERO, lhs.dims()[0],
242 lhs.dims()[1], pB, pE, cPtr.get(),
243 reinterpret_cast<ptr_type<T>>(vPtr.get()));
244
245 struct matrix_descr descrLhs {};
246 descrLhs.type = SPARSE_MATRIX_TYPE_GENERAL;
247
248 mkl_sparse_optimize(csrLhs);
249
250 if (rDims[rColDim] == 1) {
251 mkl_sparse_set_mv_hint(csrLhs, lOpts, descrLhs, 1);

Callers

nothing calls this directly

Calls 7

toSparseTransposeFunction · 0.70
dim4Class · 0.50
evalMethod · 0.45
dimsMethod · 0.45
stridesMethod · 0.45
getMappedPtrMethod · 0.45
getMethod · 0.45

Tested by

no test coverage detected