1use sparse_ir::gemm::{
31 Dgemm64FnPtr, DgemmFnPtr, ExternalBlas64Backend, ExternalBlasBackend, GemmBackendHandle,
32 Zgemm64FnPtr, ZgemmFnPtr,
33};
34
35#[repr(C)]
46pub struct spir_gemm_backend {
47 pub(crate) _private: *const std::ffi::c_void,
48}
49
50impl spir_gemm_backend {
51 pub(crate) fn inner(&self) -> &GemmBackendHandle {
53 unsafe { &*(self._private as *const GemmBackendHandle) }
54 }
55
56 pub(crate) fn new(handle: GemmBackendHandle) -> Self {
57 Self {
58 _private: Box::into_raw(Box::new(handle)) as *const std::ffi::c_void,
59 }
60 }
61}
62
63impl Drop for spir_gemm_backend {
64 fn drop(&mut self) {
65 if !self._private.is_null() {
66 unsafe {
67 let _ = Box::from_raw(
68 self._private as *const GemmBackendHandle as *mut GemmBackendHandle,
69 );
70 }
71 }
72 }
73}
74
75impl Clone for spir_gemm_backend {
76 fn clone(&self) -> Self {
77 let inner = self.inner().clone();
79 Self::new(inner)
80 }
81}
82
83#[unsafe(no_mangle)]
104pub extern "C" fn spir_gemm_backend_new_from_fblas_lp64(
105 dgemm: *const libc::c_void,
106 zgemm: *const libc::c_void,
107) -> *mut spir_gemm_backend {
108 if dgemm.is_null() || zgemm.is_null() {
110 return std::ptr::null_mut();
111 }
112
113 let result = std::panic::catch_unwind(|| {
115 let dgemm_fn: DgemmFnPtr = unsafe { std::mem::transmute(dgemm) };
117 let zgemm_fn: ZgemmFnPtr = unsafe { std::mem::transmute(zgemm) };
118
119 let backend = ExternalBlasBackend::new(dgemm_fn, zgemm_fn);
121
122 let handle = GemmBackendHandle::new(Box::new(backend));
124 Box::into_raw(Box::new(spir_gemm_backend::new(handle)))
125 });
126
127 result.unwrap_or(std::ptr::null_mut())
128}
129
130#[unsafe(no_mangle)]
151pub extern "C" fn spir_gemm_backend_new_from_fblas_ilp64(
152 dgemm64: *const libc::c_void,
153 zgemm64: *const libc::c_void,
154) -> *mut spir_gemm_backend {
155 if dgemm64.is_null() || zgemm64.is_null() {
157 return std::ptr::null_mut();
158 }
159
160 let result = std::panic::catch_unwind(|| {
162 let dgemm64_fn: Dgemm64FnPtr = unsafe { std::mem::transmute(dgemm64) };
164 let zgemm64_fn: Zgemm64FnPtr = unsafe { std::mem::transmute(zgemm64) };
165
166 let backend = ExternalBlas64Backend::new(dgemm64_fn, zgemm64_fn);
168
169 let handle = GemmBackendHandle::new(Box::new(backend));
171 Box::into_raw(Box::new(spir_gemm_backend::new(handle)))
172 });
173
174 result.unwrap_or(std::ptr::null_mut())
175}
176
177#[unsafe(no_mangle)]
189pub extern "C" fn spir_gemm_backend_release(backend: *mut spir_gemm_backend) {
190 if !backend.is_null() {
191 unsafe {
192 let _ = Box::from_raw(backend);
193 }
194 }
195}
196
197pub(crate) unsafe fn get_backend_handle<'a>(
204 backend: *const spir_gemm_backend,
205) -> Option<&'a GemmBackendHandle> {
206 if backend.is_null() {
207 None
208 } else {
209 unsafe { Some((*backend).inner()) }
210 }
211}
212
213#[cfg(test)]
214mod tests {
215 use super::*;
216
217 unsafe extern "C" fn mock_dgemm(
219 _transa: *const libc::c_char,
220 _transb: *const libc::c_char,
221 _m: *const libc::c_int,
222 _n: *const libc::c_int,
223 _k: *const libc::c_int,
224 _alpha: *const libc::c_double,
225 _a: *const libc::c_double,
226 _lda: *const libc::c_int,
227 _b: *const libc::c_double,
228 _ldb: *const libc::c_int,
229 _beta: *const libc::c_double,
230 _c: *mut libc::c_double,
231 _ldc: *const libc::c_int,
232 ) {
233 }
235
236 unsafe extern "C" fn mock_zgemm(
237 _transa: *const libc::c_char,
238 _transb: *const libc::c_char,
239 _m: *const libc::c_int,
240 _n: *const libc::c_int,
241 _k: *const libc::c_int,
242 _alpha: *const num_complex::Complex<f64>,
243 _a: *const num_complex::Complex<f64>,
244 _lda: *const libc::c_int,
245 _b: *const num_complex::Complex<f64>,
246 _ldb: *const libc::c_int,
247 _beta: *const num_complex::Complex<f64>,
248 _c: *mut num_complex::Complex<f64>,
249 _ldc: *const libc::c_int,
250 ) {
251 }
253
254 unsafe extern "C" fn mock_dgemm64(
255 _transa: *const libc::c_char,
256 _transb: *const libc::c_char,
257 _m: *const i64,
258 _n: *const i64,
259 _k: *const i64,
260 _alpha: *const libc::c_double,
261 _a: *const libc::c_double,
262 _lda: *const i64,
263 _b: *const libc::c_double,
264 _ldb: *const i64,
265 _beta: *const libc::c_double,
266 _c: *mut libc::c_double,
267 _ldc: *const i64,
268 ) {
269 }
271
272 unsafe extern "C" fn mock_zgemm64(
273 _transa: *const libc::c_char,
274 _transb: *const libc::c_char,
275 _m: *const i64,
276 _n: *const i64,
277 _k: *const i64,
278 _alpha: *const num_complex::Complex<f64>,
279 _a: *const num_complex::Complex<f64>,
280 _lda: *const i64,
281 _b: *const num_complex::Complex<f64>,
282 _ldb: *const i64,
283 _beta: *const num_complex::Complex<f64>,
284 _c: *mut num_complex::Complex<f64>,
285 _ldc: *const i64,
286 ) {
287 }
289
290 #[test]
291 fn test_backend_new_from_fblas_lp64_success() {
292 unsafe {
293 let backend = spir_gemm_backend_new_from_fblas_lp64(
294 mock_dgemm as *const _,
295 mock_zgemm as *const _,
296 );
297 assert!(!backend.is_null(), "Backend should not be null");
298 spir_gemm_backend_release(backend);
299 }
300 }
301
302 #[test]
303 fn test_backend_new_from_fblas_ilp64_success() {
304 unsafe {
305 let backend = spir_gemm_backend_new_from_fblas_ilp64(
306 mock_dgemm64 as *const _,
307 mock_zgemm64 as *const _,
308 );
309 assert!(!backend.is_null(), "Backend should not be null");
310 spir_gemm_backend_release(backend);
311 }
312 }
313
314 #[test]
315 fn test_backend_new_from_fblas_lp64_null_dgemm() {
316 unsafe {
317 let backend =
318 spir_gemm_backend_new_from_fblas_lp64(std::ptr::null(), mock_zgemm as *const _);
319 assert!(
320 backend.is_null(),
321 "Backend should be null when dgemm is null"
322 );
323 }
324 }
325
326 #[test]
327 fn test_backend_new_from_fblas_lp64_null_zgemm() {
328 unsafe {
329 let backend =
330 spir_gemm_backend_new_from_fblas_lp64(mock_dgemm as *const _, std::ptr::null());
331 assert!(
332 backend.is_null(),
333 "Backend should be null when zgemm is null"
334 );
335 }
336 }
337
338 #[test]
339 fn test_backend_new_from_fblas_ilp64_null_pointers() {
340 unsafe {
341 let backend =
342 spir_gemm_backend_new_from_fblas_ilp64(std::ptr::null(), std::ptr::null());
343 assert!(
344 backend.is_null(),
345 "Backend should be null when pointers are null"
346 );
347 }
348 }
349
350 #[test]
351 fn test_backend_release_null() {
352 unsafe {
353 spir_gemm_backend_release(std::ptr::null_mut());
355 }
356 }
357
358 #[cfg(all(test, feature = "system-blas"))]
360 mod system_blas_tests {
361 use super::*;
362 use blas_sys::{dgemm_, zgemm_};
363 use mdarray::tensor;
364 use sparse_ir::gemm::matmul_par;
365
366 unsafe fn create_blas_backend() -> *mut spir_gemm_backend {
368 unsafe {
369 spir_gemm_backend_new_from_fblas_lp64(
370 dgemm_ as *const _,
371 unsafe {
373 std::mem::transmute::<
374 unsafe extern "C" fn(
375 *const libc::c_char,
376 *const libc::c_char,
377 *const libc::c_int,
378 *const libc::c_int,
379 *const libc::c_int,
380 *const blas_sys::c_double_complex,
381 *const blas_sys::c_double_complex,
382 *const libc::c_int,
383 *const blas_sys::c_double_complex,
384 *const libc::c_int,
385 *const blas_sys::c_double_complex,
386 *mut blas_sys::c_double_complex,
387 *const libc::c_int,
388 ),
389 sparse_ir::gemm::ZgemmFnPtr,
390 >(zgemm_)
391 } as *const _,
392 )
393 }
394 }
395
396 #[test]
397 fn test_default_backend_matrix_multiplication_f64() {
398 unsafe {
399 let backend = std::ptr::null();
401
402 let a: mdarray::DTensor<f64, 2> = tensor![[1.0, 2.0], [3.0, 4.0]];
407 let b: mdarray::DTensor<f64, 2> = tensor![[5.0, 6.0], [7.0, 8.0]];
408 let backend_handle = get_backend_handle(backend);
409 let c = matmul_par(&a, &b, backend_handle);
410
411 assert!(
413 (c[[0, 0]] - 19.0).abs() < 1e-10,
414 "c[0,0] should be 19.0, got {}",
415 c[[0, 0]]
416 );
417 assert!((c[[0, 1]] - 22.0).abs() < 1e-10, "c[0,1] should be 22.0");
418 assert!((c[[1, 0]] - 43.0).abs() < 1e-10, "c[1,0] should be 43.0");
419 assert!((c[[1, 1]] - 50.0).abs() < 1e-10, "c[1,1] should be 50.0");
420 }
421 }
422
423 #[test]
424 fn test_lp64_backend_matrix_multiplication_f64() {
425 unsafe {
426 let backend = create_blas_backend();
428 assert!(!backend.is_null());
429
430 let a: mdarray::DTensor<f64, 2> = tensor![[1.0, 2.0], [3.0, 4.0]];
435 let b: mdarray::DTensor<f64, 2> = tensor![[5.0, 6.0], [7.0, 8.0]];
436 let backend_handle = get_backend_handle(backend);
437 let c = matmul_par(&a, &b, backend_handle);
438
439 assert!(
441 (c[[0, 0]] - 19.0).abs() < 1e-10,
442 "c[0,0] should be 19.0, got {}",
443 c[[0, 0]]
444 );
445 assert!((c[[0, 1]] - 22.0).abs() < 1e-10, "c[0,1] should be 22.0");
446 assert!((c[[1, 0]] - 43.0).abs() < 1e-10, "c[1,0] should be 43.0");
447 assert!((c[[1, 1]] - 50.0).abs() < 1e-10, "c[1,1] should be 50.0");
448
449 spir_gemm_backend_release(backend);
451 }
452 }
453
454 #[test]
455 fn test_default_backend_matrix_multiplication_complex() {
456 unsafe {
457 let backend = std::ptr::null();
459
460 let a: mdarray::DTensor<num_complex::Complex<f64>, 2> = tensor![
462 [
463 num_complex::Complex::new(1.0, 0.0),
464 num_complex::Complex::new(2.0, 0.0)
465 ],
466 [
467 num_complex::Complex::new(3.0, 0.0),
468 num_complex::Complex::new(4.0, 0.0)
469 ]
470 ];
471 let b: mdarray::DTensor<num_complex::Complex<f64>, 2> = tensor![
472 [
473 num_complex::Complex::new(5.0, 0.0),
474 num_complex::Complex::new(6.0, 0.0)
475 ],
476 [
477 num_complex::Complex::new(7.0, 0.0),
478 num_complex::Complex::new(8.0, 0.0)
479 ]
480 ];
481 let backend_handle = get_backend_handle(backend);
482 let c = matmul_par(&a, &b, backend_handle);
483
484 assert!((c[[0, 0]].re - 19.0).abs() < 1e-10);
486 assert!((c[[0, 1]].re - 22.0).abs() < 1e-10);
487 assert!((c[[1, 0]].re - 43.0).abs() < 1e-10);
488 assert!((c[[1, 1]].re - 50.0).abs() < 1e-10);
489 assert!(c[[0, 0]].im.abs() < 1e-10);
490 }
491 }
492
493 #[test]
494 fn test_lp64_backend_matrix_multiplication_complex() {
495 unsafe {
496 let backend = create_blas_backend();
498 assert!(!backend.is_null());
499
500 let a: mdarray::DTensor<num_complex::Complex<f64>, 2> = tensor![
502 [
503 num_complex::Complex::new(1.0, 0.0),
504 num_complex::Complex::new(2.0, 0.0)
505 ],
506 [
507 num_complex::Complex::new(3.0, 0.0),
508 num_complex::Complex::new(4.0, 0.0)
509 ]
510 ];
511 let b: mdarray::DTensor<num_complex::Complex<f64>, 2> = tensor![
512 [
513 num_complex::Complex::new(5.0, 0.0),
514 num_complex::Complex::new(6.0, 0.0)
515 ],
516 [
517 num_complex::Complex::new(7.0, 0.0),
518 num_complex::Complex::new(8.0, 0.0)
519 ]
520 ];
521 let backend_handle = get_backend_handle(backend);
522 let c = matmul_par(&a, &b, backend_handle);
523
524 assert!((c[[0, 0]].re - 19.0).abs() < 1e-10);
526 assert!((c[[0, 1]].re - 22.0).abs() < 1e-10);
527 assert!((c[[1, 0]].re - 43.0).abs() < 1e-10);
528 assert!((c[[1, 1]].re - 50.0).abs() < 1e-10);
529 assert!(c[[0, 0]].im.abs() < 1e-10);
530
531 spir_gemm_backend_release(backend);
533 }
534 }
535
536 #[test]
537 fn test_default_backend_larger_matrix() {
538 unsafe {
539 let backend = std::ptr::null();
541
542 let a: mdarray::DTensor<f64, 2> = tensor![[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]];
544 let b: mdarray::DTensor<f64, 2> =
545 tensor![[7.0, 8.0, 9.0, 10.0], [11.0, 12.0, 13.0, 14.0]];
546 let backend_handle = get_backend_handle(backend);
547 let c = matmul_par(&a, &b, backend_handle);
548
549 assert!((c[[0, 0]] - 29.0).abs() < 1e-10);
552 assert!((c[[0, 1]] - 32.0).abs() < 1e-10);
553 assert!((c[[0, 2]] - 35.0).abs() < 1e-10);
554 assert!((c[[0, 3]] - 38.0).abs() < 1e-10);
555 }
556 }
557
558 #[test]
559 fn test_lp64_backend_larger_matrix() {
560 unsafe {
561 let backend = create_blas_backend();
563 assert!(!backend.is_null());
564
565 let a: mdarray::DTensor<f64, 2> = tensor![[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]];
567 let b: mdarray::DTensor<f64, 2> =
568 tensor![[7.0, 8.0, 9.0, 10.0], [11.0, 12.0, 13.0, 14.0]];
569 let backend_handle = get_backend_handle(backend);
570 let c = matmul_par(&a, &b, backend_handle);
571
572 assert!((c[[0, 0]] - 29.0).abs() < 1e-10);
575 assert!((c[[0, 1]] - 32.0).abs() < 1e-10);
576 assert!((c[[0, 2]] - 35.0).abs() < 1e-10);
577 assert!((c[[0, 3]] - 38.0).abs() < 1e-10);
578
579 spir_gemm_backend_release(backend);
581 }
582 }
583 }
584}