| 192 | //////////////////////////////////////////////////////////////////////////////// |
| 193 | template<typename T> |
| 194 | Array<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); |
nothing calls this directly
no test coverage detected