diff --git a/src/lib/dbcsr_internal_operations.F b/src/lib/dbcsr_internal_operations.F index acaec07891..cc8d74a3dd 100644 --- a/src/lib/dbcsr_internal_operations.F +++ b/src/lib/dbcsr_internal_operations.F @@ -16,7 +16,6 @@ MODULE dbcsr_internal_operations USE array_types, ONLY: array_data USE dbcsr_block_access, ONLY: dbcsr_put_block,& dbcsr_reserve_blocks - USE dbcsr_block_operations, ONLY: dbcsr_data_clear USE dbcsr_data_methods, ONLY: & dbcsr_data_clear_pointer, dbcsr_data_ensure_size, dbcsr_data_get_size, & dbcsr_data_get_size_referenced, dbcsr_data_init, dbcsr_data_new, & @@ -90,6 +89,13 @@ MODULE dbcsr_internal_operations REAL, PARAMETER :: default_resize_factor = 1.618034 +#if defined (__INTERNAL_GEMM) + LOGICAL, PARAMETER :: internal_gemm = .TRUE. +#else + LOGICAL, PARAMETER :: internal_gemm = .FALSE. +#endif + + PUBLIC :: dbcsr_mult_m_e_e PUBLIC :: dbcsr_insert_blocks @@ -1699,7 +1705,8 @@ CONTAINS routineP = moduleN//':'//routineN REAL, PARAMETER :: resize_factor = 1.618034 - INTEGER :: error_handler + INTEGER :: error_handler, maxs, sp, sp_c + LOGICAL :: do_resize DEBUG_HEADER ! --------------------------------------------------------------------------- @@ -1709,26 +1716,66 @@ CONTAINS !CALL print_dgemm_parameters (params(1:stack_size),& ! params_a, params_b, params_c) !WRITE(*,*)routineN//"========== END of multiplies" + ! Increase product data area size if necessary. + do_resize = .FALSE. + !if (remote_memory) then + ! min_size_a = dbcsr_data_get_size(left_data_area) + ! min_size_b = dbcsr_data_get_size(right_data_area) + ! min_size_c = dbcsr_data_get_size(product_data_area) + ! max_size_a = 1 + ! max_size_b = 1 + ! max_size_c = 1 + !endif + maxs = dbcsr_data_get_size(product_data_area) + DO sp = 1, stack_size + sp_c = params(sp)%p_c + IF (params_c(sp_c)%offset+params_c(sp_c)%nze-1 .GT. maxs) THEN + maxs = params_c(sp_c)%offset+params_c(sp_c)%nze-1 + do_resize = .TRUE. + ENDIF + !if (remote_memory) then + ! sp_a = params(sp)%p_a + ! min_size_a = MIN (min_size_a, params_a(sp_a)%offset) + ! max_size_a = MAX (max_size_a, params_a(sp_a)%offset& + ! +params_a(sp_a)%nze-1) + ! sp_b = params(sp)%p_b + ! min_size_b = MIN (min_size_b, params_b(sp_b)%offset) + ! max_size_b = MAX (max_size_b, params_b(sp_b)%offset& + ! +params_b(sp_b)%nze-1) + ! sp_c = params(sp)%p_c + ! min_size_c = MIN (min_size_c, params_c(sp_c)%offset) + ! max_size_c = MAX (max_size_c, params_c(sp_c)%offset& + ! +params_c(sp_c)%nze-1) + !endif + ENDDO + IF (maxs .GT. dbcsr_data_get_size_referenced (product_data_area)) THEN + CALL dbcsr_data_set_size_referenced (product_data_area, maxs) + ENDIF + IF (do_resize) THEN + CALL dbcsr_data_ensure_size (product_data_area,& + maxs, factor=resize_factor, error=error) + ENDIF + ! SELECT CASE (product_data_area%d%data_type) CASE (dbcsr_type_real_4) CALL process_dgemm_stack_s (params, params_a, params_b, params_c,& stack_size, stack_size_a, stack_size_b, stack_size_c,& - left_data_area%d%r_sp, right_data_area%d%r_sp, product_data_area,& + left_data_area%d%r_sp, right_data_area%d%r_sp, product_data_area%d%r_sp,& use_plasma=use_plasma, lflop=lflop, ltime=ltime, error=error) CASE (dbcsr_type_real_8) CALL process_dgemm_stack_d (params, params_a, params_b, params_c,& stack_size, stack_size_a, stack_size_b, stack_size_c,& - left_data_area%d%r_dp, right_data_area%d%r_dp, product_data_area,& + left_data_area%d%r_dp, right_data_area%d%r_dp, product_data_area%d%r_dp,& use_plasma=use_plasma, lflop=lflop, ltime=ltime, error=error) CASE (dbcsr_type_complex_4) CALL process_dgemm_stack_c (params, params_a, params_b, params_c,& stack_size, stack_size_a, stack_size_b, stack_size_c,& - left_data_area%d%c_sp, right_data_area%d%c_sp, product_data_area,& + left_data_area%d%c_sp, right_data_area%d%c_sp, product_data_area%d%c_sp,& use_plasma=use_plasma, lflop=lflop, ltime=ltime, error=error) CASE (dbcsr_type_complex_8) CALL process_dgemm_stack_z (params, params_a, params_b, params_c,& stack_size, stack_size_a, stack_size_b, stack_size_c,& - left_data_area%d%c_dp, right_data_area%d%c_dp, product_data_area,& + left_data_area%d%c_dp, right_data_area%d%c_dp, product_data_area%d%c_dp,& use_plasma=use_plasma, lflop=lflop, ltime=ltime, error=error) CASE default CALL dbcsr_assert (.FALSE., dbcsr_fatal_level, dbcsr_caller_error,& diff --git a/src/lib/dbcsr_internal_operations__nametype1_.template b/src/lib/dbcsr_internal_operations__nametype1_.template index d66211d56b..597d126df5 100644 --- a/src/lib/dbcsr_internal_operations__nametype1_.template +++ b/src/lib/dbcsr_internal_operations__nametype1_.template @@ -12,7 +12,7 @@ SUBROUTINE process_dgemm_stack_[nametype1](params,& params_a, params_b, params_c,& stack_size, stack_size_a, stack_size_b, stack_size_c,& - left_data_a, right_data_a, product_data_area, use_plasma, lflop, ltime, error) + left_data_a, right_data_a, product_data_a, use_plasma, lflop, ltime, error) INTEGER, INTENT(IN) :: stack_size, stack_size_a,& stack_size_b, stack_size_c TYPE(dgemm_join), DIMENSION(1:stack_size), & @@ -22,7 +22,7 @@ TYPE(block_parameters), DIMENSION(1:stack_size_c), INTENT(IN) :: params_c [type1], DIMENSION(:), INTENT(IN) :: left_data_a, & right_data_a - TYPE(dbcsr_data_obj), INTENT(INOUT) :: product_data_area + [type1], DIMENSION(:), INTENT(INOUT) :: product_data_a LOGICAL, INTENT(IN), OPTIONAL :: use_plasma INTEGER(KIND=int_8), INTENT(OUT), & OPTIONAL :: lflop @@ -33,10 +33,10 @@ routineP = moduleN//':'//routineN REAL, PARAMETER :: resize_factor = 1.618034 - INTEGER :: c, plasma_info, lda, ldb, ldc, maxs, r, sp,& + INTEGER :: c, plasma_info, lda, ldb, ldc, r, sp,& poff INTEGER :: sp_c, sp_a, sp_b - LOGICAL :: do_resize, my_use_plasma + LOGICAL :: my_use_plasma #ifdef __PLASMA INCLUDE 'plasmaf.h' @@ -49,23 +49,6 @@ IF (PRESENT(use_plasma)) my_use_plasma = use_plasma IF (PRESENT (lflop)) lflop = INT(0, int_8) IF (PRESENT (ltime)) ltime = 0.0_real_8 - ! Increase product data area size if necessary. - do_resize = .FALSE. - maxs = dbcsr_data_get_size(product_data_area) - DO sp = 1, stack_size - sp_c = params(sp)%p_c - IF (params_c(sp_c)%offset+params_c(sp_c)%nze-1 .GT. maxs) THEN - maxs = params_c(sp_c)%offset+params_c(sp_c)%nze-1 - do_resize = .TRUE. - ENDIF - ENDDO - IF (maxs .GT. dbcsr_data_get_size_referenced (product_data_area)) THEN - CALL dbcsr_data_set_size_referenced (product_data_area, maxs) - ENDIF - IF (do_resize) THEN - CALL dbcsr_data_ensure_size (product_data_area,& - maxs, factor=resize_factor, error=error) - ENDIF ! Setup encapsulated data area DO sp = 1, stack_size sp_c = params(sp)%p_c @@ -84,19 +67,25 @@ IF (params(sp)%beta%[base1]_[prec1] .EQ. [zero1]) THEN IF (.NOT. params_c(sp_c)%tr & .OR. params(sp)%last_n .EQ. params_c(sp_c)%logical_cols) THEN - CALL dbcsr_data_clear (product_data_area,& - lb=poff,& - ub=poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) + product_data_a(poff : poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) = [zero1] + !CALL dbcsr_data_clear (product_data_area,& + ! lb=poff,& + ! ub=poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) ELSE ! It must be transposed. + FORALL (r = 1 : params_c(sp_c)%logical_rows,& + c = 1 : params(sp)%last_n) + product_data_a(poff-1+(r-1)*params_c(sp_c)%logical_cols+c) =& + [zero1] + END FORALL !### optimize out the inner loop (turn into a range) - DO r = 1, params_c(sp_c)%logical_rows - DO c = 1, params(sp)%last_n - CALL dbcsr_data_clear(product_data_area,& - lb=poff-1+(r-1)*params_c(sp_c)%logical_cols+c,& - ub=poff-1+(r-1)*params_c(sp_c)%logical_cols+c) - ENDDO - ENDDO + !DO r = 1, params_c(sp_c)%logical_rows + ! DO c = 1, params(sp)%last_n + ! CALL dbcsr_data_clear(product_data_area,& + ! lb=poff-1+(r-1)*params_c(sp_c)%logical_cols+c,& + ! ub=poff-1+(r-1)*params_c(sp_c)%logical_cols+c) + ! ENDDO + !ENDDO !FORALL (r = 1:params(sp)%lrows_c, c = 1:params(sp)%last_n) ! product_data_d((r-1)*params(sp)%lcols_c+c) = 0.0_dp !END FORALL @@ -123,14 +112,25 @@ left_data_a(params_a(sp_a)%offset), lda,& right_data_a(params_b(sp_b)%offset), ldb,& params(sp)%beta%[base1]_[prec1],& - product_data_area%d%[base1]_[prec1](poff), ldc,& + product_data_a(poff), ldc,& plasma_info) - CALL dbcsr_assert( plasma_info==0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& + CALL dbcsr_assert( plasma_info, "EQ", 0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma_gemm failed", __LINE__, error=error) #else CALL dbcsr_assert( .FALSE., dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma badly set", __LINE__, error=error) #endif + ELSEIF (internal_gemm) THEN + CALL internal_[gemmname1](& + params_a(sp_a)%tr,& + params_b(sp_b)%tr,& + params_c(sp_c)%logical_rows, params(sp)%last_n,& !m, n + params(sp)%last_k,& ! k + params(sp)%alpha%[base1]_[prec1],& + left_data_a(params_a(sp_a)%offset:), lda,& + right_data_a(params_b(sp_b)%offset:), ldb,& + params(sp)%beta%[base1]_[prec1],& + product_data_a(poff:), ldc) ELSE CALL [gemmname1](& blas_mat_type(params_a(sp_a)%tr),& @@ -141,7 +141,7 @@ left_data_a(params_a(sp_a)%offset), lda,& right_data_a(params_b(sp_b)%offset), ldb,& params(sp)%beta%[base1]_[prec1],& - product_data_area%d%[base1]_[prec1](poff), ldc) + product_data_a(poff), ldc) ENDIF IF (PRESENT (ltime)) ltime = m_walltime() - ltime IF (PRESENT (lflop)) & @@ -161,14 +161,25 @@ right_data_a(params_b(sp_b)%offset), ldb,& left_data_a(params_a(sp_a)%offset), lda,& params(sp)%beta%[base1]_[prec1],& - product_data_area%d%[base1]_[prec1](poff), ldc,& + product_data_a(poff), ldc,& plasma_info) - CALL dbcsr_assert( plasma_info==0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& + CALL dbcsr_assert( plasma_info, "EQ", 0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma_gemm failed", __LINE__, error=error) #else CALL dbcsr_assert( .FALSE., dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma badly set", __LINE__, error=error) #endif + ELSEIF (internal_gemm) THEN + CALL internal_[gemmname1](& + .NOT.params_b(sp_b)%tr,& + .NOT.params_a(sp_a)%tr,& + params_c(sp_c)%logical_cols, params_c(sp_c)%logical_rows,& !m, n (switched) + params(sp)%last_k,& ! k + params(sp)%alpha%[base1]_[prec1],& + right_data_a(params_b(sp_b)%offset:), ldb,& + left_data_a(params_a(sp_a)%offset:), lda,& + params(sp)%beta%[base1]_[prec1],& + product_data_a(poff:), ldc) ELSE CALL [gemmname1](& blas_mat_type(.NOT.params_b(sp_b)%tr),& @@ -179,7 +190,7 @@ right_data_a(params_b(sp_b)%offset), ldb,& left_data_a(params_a(sp_a)%offset), lda,& params(sp)%beta%[base1]_[prec1],& - product_data_area%d%[base1]_[prec1](poff), ldc) + product_data_a(poff), ldc) ENDIF IF (PRESENT (ltime)) ltime = m_walltime() - ltime IF (PRESENT (lflop)) & @@ -191,3 +202,68 @@ END SUBROUTINE process_dgemm_stack_[nametype1] + PURE SUBROUTINE internal_[gemmname1](& + TRANSA,TRANSB,M,N,K,ALPHA,A,LDA,B,LDB,BETA,C,LDC) + LOGICAL, INTENT(IN) :: TRANSA, TRANSB + INTEGER, INTENT(IN) :: M, N, K + INTEGER, INTENT(IN) :: LDC + [type1], INTENT(INOUT) :: C(LDC,*) + [type1], INTENT(IN) :: BETA + INTEGER, INTENT(IN) :: LDB + [type1], INTENT(IN) :: B(LDB,*) + INTEGER, INTENT(IN) :: LDA + [type1], INTENT(IN) :: A(LDA,*), ALPHA + IF (.NOT. transa .AND. .NOT. transb) THEN + call internal_[gemmname1]_nn(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (.NOT. transa .AND. transb) THEN + call internal_[gemmname1]_nt(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (transa .AND. .NOT. transb) THEN + call internal_[gemmname1]_tn(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (transa .AND. transb) THEN + call internal_[gemmname1]_tt(M,N,K,ALPHA,A,B,BETA,C) + ENDIF + END SUBROUTINE internal_[gemmname1] + + PURE SUBROUTINE internal_[gemmname1]_nn(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + [type1], INTENT(INOUT) :: C(M,N) + [type1], INTENT(IN) :: BETA + [type1], INTENT(IN) :: B(K,N) + [type1], INTENT(IN) :: A(M,K), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (A, B) + END SUBROUTINE internal_[gemmname1]_nn + + PURE SUBROUTINE internal_[gemmname1]_nt(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + [type1], INTENT(INOUT) :: C(M,N) + [type1], INTENT(IN) :: BETA + [type1], INTENT(IN) :: B(N,K) + [type1], INTENT(IN) :: A(M,K), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (A, TRANSPOSE(B)) + END SUBROUTINE internal_[gemmname1]_nt + + PURE SUBROUTINE internal_[gemmname1]_tn(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + [type1], INTENT(INOUT) :: C(M,N) + [type1], INTENT(IN) :: BETA + [type1], INTENT(IN) :: B(K,N) + [type1], INTENT(IN) :: A(K,M), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (TRANSPOSE(A), B) + END SUBROUTINE internal_[gemmname1]_tn + + PURE SUBROUTINE internal_[gemmname1]_tt(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + [type1], INTENT(INOUT) :: C(M,N) + [type1], INTENT(IN) :: BETA + [type1], INTENT(IN) :: B(N,K) + [type1], INTENT(IN) :: A(K,M), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (TRANSPOSE(A), TRANSPOSE(B)) + END SUBROUTINE internal_[gemmname1]_tt diff --git a/src/lib/dbcsr_internal_operations_c.F b/src/lib/dbcsr_internal_operations_c.F index aa7952548a..0c3648594a 100644 --- a/src/lib/dbcsr_internal_operations_c.F +++ b/src/lib/dbcsr_internal_operations_c.F @@ -12,7 +12,7 @@ SUBROUTINE process_dgemm_stack_c(params,& params_a, params_b, params_c,& stack_size, stack_size_a, stack_size_b, stack_size_c,& - left_data_a, right_data_a, product_data_area, use_plasma, lflop, ltime, error) + left_data_a, right_data_a, product_data_a, use_plasma, lflop, ltime, error) INTEGER, INTENT(IN) :: stack_size, stack_size_a,& stack_size_b, stack_size_c TYPE(dgemm_join), DIMENSION(1:stack_size), & @@ -22,7 +22,7 @@ TYPE(block_parameters), DIMENSION(1:stack_size_c), INTENT(IN) :: params_c COMPLEX(kind=real_4), DIMENSION(:), INTENT(IN) :: left_data_a, & right_data_a - TYPE(dbcsr_data_obj), INTENT(INOUT) :: product_data_area + COMPLEX(kind=real_4), DIMENSION(:), INTENT(INOUT) :: product_data_a LOGICAL, INTENT(IN), OPTIONAL :: use_plasma INTEGER(KIND=int_8), INTENT(OUT), & OPTIONAL :: lflop @@ -33,10 +33,10 @@ routineP = moduleN//':'//routineN REAL, PARAMETER :: resize_factor = 1.618034 - INTEGER :: c, plasma_info, lda, ldb, ldc, maxs, r, sp,& + INTEGER :: c, plasma_info, lda, ldb, ldc, r, sp,& poff INTEGER :: sp_c, sp_a, sp_b - LOGICAL :: do_resize, my_use_plasma + LOGICAL :: my_use_plasma #ifdef __PLASMA INCLUDE 'plasmaf.h' @@ -49,23 +49,6 @@ IF (PRESENT(use_plasma)) my_use_plasma = use_plasma IF (PRESENT (lflop)) lflop = INT(0, int_8) IF (PRESENT (ltime)) ltime = 0.0_real_8 - ! Increase product data area size if necessary. - do_resize = .FALSE. - maxs = dbcsr_data_get_size(product_data_area) - DO sp = 1, stack_size - sp_c = params(sp)%p_c - IF (params_c(sp_c)%offset+params_c(sp_c)%nze-1 .GT. maxs) THEN - maxs = params_c(sp_c)%offset+params_c(sp_c)%nze-1 - do_resize = .TRUE. - ENDIF - ENDDO - IF (maxs .GT. dbcsr_data_get_size_referenced (product_data_area)) THEN - CALL dbcsr_data_set_size_referenced (product_data_area, maxs) - ENDIF - IF (do_resize) THEN - CALL dbcsr_data_ensure_size (product_data_area,& - maxs, factor=resize_factor, error=error) - ENDIF ! Setup encapsulated data area DO sp = 1, stack_size sp_c = params(sp)%p_c @@ -84,19 +67,25 @@ IF (params(sp)%beta%c_sp .EQ. CMPLX(0.0, 0.0, real_4)) THEN IF (.NOT. params_c(sp_c)%tr & .OR. params(sp)%last_n .EQ. params_c(sp_c)%logical_cols) THEN - CALL dbcsr_data_clear (product_data_area,& - lb=poff,& - ub=poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) + product_data_a(poff : poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) = CMPLX(0.0, 0.0, real_4) + !CALL dbcsr_data_clear (product_data_area,& + ! lb=poff,& + ! ub=poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) ELSE ! It must be transposed. + FORALL (r = 1 : params_c(sp_c)%logical_rows,& + c = 1 : params(sp)%last_n) + product_data_a(poff-1+(r-1)*params_c(sp_c)%logical_cols+c) =& + CMPLX(0.0, 0.0, real_4) + END FORALL !### optimize out the inner loop (turn into a range) - DO r = 1, params_c(sp_c)%logical_rows - DO c = 1, params(sp)%last_n - CALL dbcsr_data_clear(product_data_area,& - lb=poff-1+(r-1)*params_c(sp_c)%logical_cols+c,& - ub=poff-1+(r-1)*params_c(sp_c)%logical_cols+c) - ENDDO - ENDDO + !DO r = 1, params_c(sp_c)%logical_rows + ! DO c = 1, params(sp)%last_n + ! CALL dbcsr_data_clear(product_data_area,& + ! lb=poff-1+(r-1)*params_c(sp_c)%logical_cols+c,& + ! ub=poff-1+(r-1)*params_c(sp_c)%logical_cols+c) + ! ENDDO + !ENDDO !FORALL (r = 1:params(sp)%lrows_c, c = 1:params(sp)%last_n) ! product_data_d((r-1)*params(sp)%lcols_c+c) = 0.0_dp !END FORALL @@ -123,14 +112,25 @@ left_data_a(params_a(sp_a)%offset), lda,& right_data_a(params_b(sp_b)%offset), ldb,& params(sp)%beta%c_sp,& - product_data_area%d%c_sp(poff), ldc,& + product_data_a(poff), ldc,& plasma_info) - CALL dbcsr_assert( plasma_info==0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& + CALL dbcsr_assert( plasma_info, "EQ", 0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma_gemm failed", __LINE__, error=error) #else CALL dbcsr_assert( .FALSE., dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma badly set", __LINE__, error=error) #endif + ELSEIF (internal_gemm) THEN + CALL internal_CGEMM(& + params_a(sp_a)%tr,& + params_b(sp_b)%tr,& + params_c(sp_c)%logical_rows, params(sp)%last_n,& !m, n + params(sp)%last_k,& ! k + params(sp)%alpha%c_sp,& + left_data_a(params_a(sp_a)%offset:), lda,& + right_data_a(params_b(sp_b)%offset:), ldb,& + params(sp)%beta%c_sp,& + product_data_a(poff:), ldc) ELSE CALL CGEMM(& blas_mat_type(params_a(sp_a)%tr),& @@ -141,7 +141,7 @@ left_data_a(params_a(sp_a)%offset), lda,& right_data_a(params_b(sp_b)%offset), ldb,& params(sp)%beta%c_sp,& - product_data_area%d%c_sp(poff), ldc) + product_data_a(poff), ldc) ENDIF IF (PRESENT (ltime)) ltime = m_walltime() - ltime IF (PRESENT (lflop)) & @@ -161,14 +161,25 @@ right_data_a(params_b(sp_b)%offset), ldb,& left_data_a(params_a(sp_a)%offset), lda,& params(sp)%beta%c_sp,& - product_data_area%d%c_sp(poff), ldc,& + product_data_a(poff), ldc,& plasma_info) - CALL dbcsr_assert( plasma_info==0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& + CALL dbcsr_assert( plasma_info, "EQ", 0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma_gemm failed", __LINE__, error=error) #else CALL dbcsr_assert( .FALSE., dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma badly set", __LINE__, error=error) #endif + ELSEIF (internal_gemm) THEN + CALL internal_CGEMM(& + .NOT.params_b(sp_b)%tr,& + .NOT.params_a(sp_a)%tr,& + params_c(sp_c)%logical_cols, params_c(sp_c)%logical_rows,& !m, n (switched) + params(sp)%last_k,& ! k + params(sp)%alpha%c_sp,& + right_data_a(params_b(sp_b)%offset:), ldb,& + left_data_a(params_a(sp_a)%offset:), lda,& + params(sp)%beta%c_sp,& + product_data_a(poff:), ldc) ELSE CALL CGEMM(& blas_mat_type(.NOT.params_b(sp_b)%tr),& @@ -179,7 +190,7 @@ right_data_a(params_b(sp_b)%offset), ldb,& left_data_a(params_a(sp_a)%offset), lda,& params(sp)%beta%c_sp,& - product_data_area%d%c_sp(poff), ldc) + product_data_a(poff), ldc) ENDIF IF (PRESENT (ltime)) ltime = m_walltime() - ltime IF (PRESENT (lflop)) & @@ -191,3 +202,68 @@ END SUBROUTINE process_dgemm_stack_c + PURE SUBROUTINE internal_CGEMM(& + TRANSA,TRANSB,M,N,K,ALPHA,A,LDA,B,LDB,BETA,C,LDC) + LOGICAL, INTENT(IN) :: TRANSA, TRANSB + INTEGER, INTENT(IN) :: M, N, K + INTEGER, INTENT(IN) :: LDC + COMPLEX(kind=real_4), INTENT(INOUT) :: C(LDC,*) + COMPLEX(kind=real_4), INTENT(IN) :: BETA + INTEGER, INTENT(IN) :: LDB + COMPLEX(kind=real_4), INTENT(IN) :: B(LDB,*) + INTEGER, INTENT(IN) :: LDA + COMPLEX(kind=real_4), INTENT(IN) :: A(LDA,*), ALPHA + IF (.NOT. transa .AND. .NOT. transb) THEN + CALL internal_CGEMM_nn(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (.NOT. transa .AND. transb) THEN + CALL internal_CGEMM_nt(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (transa .AND. .NOT. transb) THEN + CALL internal_CGEMM_tn(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (transa .AND. transb) THEN + CALL internal_CGEMM_tt(M,N,K,ALPHA,A,B,BETA,C) + ENDIF + END SUBROUTINE internal_CGEMM + + PURE SUBROUTINE internal_CGEMM_nn(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + COMPLEX(kind=real_4), INTENT(INOUT) :: C(M,N) + COMPLEX(kind=real_4), INTENT(IN) :: BETA + COMPLEX(kind=real_4), INTENT(IN) :: B(K,N) + COMPLEX(kind=real_4), INTENT(IN) :: A(M,K), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (A, B) + END SUBROUTINE internal_CGEMM_nn + + PURE SUBROUTINE internal_CGEMM_nt(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + COMPLEX(kind=real_4), INTENT(INOUT) :: C(M,N) + COMPLEX(kind=real_4), INTENT(IN) :: BETA + COMPLEX(kind=real_4), INTENT(IN) :: B(N,K) + COMPLEX(kind=real_4), INTENT(IN) :: A(M,K), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (A, TRANSPOSE(B)) + END SUBROUTINE internal_CGEMM_nt + + PURE SUBROUTINE internal_CGEMM_tn(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + COMPLEX(kind=real_4), INTENT(INOUT) :: C(M,N) + COMPLEX(kind=real_4), INTENT(IN) :: BETA + COMPLEX(kind=real_4), INTENT(IN) :: B(K,N) + COMPLEX(kind=real_4), INTENT(IN) :: A(K,M), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (TRANSPOSE(A), B) + END SUBROUTINE internal_CGEMM_tn + + PURE SUBROUTINE internal_CGEMM_tt(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + COMPLEX(kind=real_4), INTENT(INOUT) :: C(M,N) + COMPLEX(kind=real_4), INTENT(IN) :: BETA + COMPLEX(kind=real_4), INTENT(IN) :: B(N,K) + COMPLEX(kind=real_4), INTENT(IN) :: A(K,M), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (TRANSPOSE(A), TRANSPOSE(B)) + END SUBROUTINE internal_CGEMM_tt diff --git a/src/lib/dbcsr_internal_operations_d.F b/src/lib/dbcsr_internal_operations_d.F index 32c744cc32..15e211fae4 100644 --- a/src/lib/dbcsr_internal_operations_d.F +++ b/src/lib/dbcsr_internal_operations_d.F @@ -12,7 +12,7 @@ SUBROUTINE process_dgemm_stack_d(params,& params_a, params_b, params_c,& stack_size, stack_size_a, stack_size_b, stack_size_c,& - left_data_a, right_data_a, product_data_area, use_plasma, lflop, ltime, error) + left_data_a, right_data_a, product_data_a, use_plasma, lflop, ltime, error) INTEGER, INTENT(IN) :: stack_size, stack_size_a,& stack_size_b, stack_size_c TYPE(dgemm_join), DIMENSION(1:stack_size), & @@ -22,7 +22,7 @@ TYPE(block_parameters), DIMENSION(1:stack_size_c), INTENT(IN) :: params_c REAL(kind=real_8), DIMENSION(:), INTENT(IN) :: left_data_a, & right_data_a - TYPE(dbcsr_data_obj), INTENT(INOUT) :: product_data_area + REAL(kind=real_8), DIMENSION(:), INTENT(INOUT) :: product_data_a LOGICAL, INTENT(IN), OPTIONAL :: use_plasma INTEGER(KIND=int_8), INTENT(OUT), & OPTIONAL :: lflop @@ -33,10 +33,10 @@ routineP = moduleN//':'//routineN REAL, PARAMETER :: resize_factor = 1.618034 - INTEGER :: c, plasma_info, lda, ldb, ldc, maxs, r, sp,& + INTEGER :: c, plasma_info, lda, ldb, ldc, r, sp,& poff INTEGER :: sp_c, sp_a, sp_b - LOGICAL :: do_resize, my_use_plasma + LOGICAL :: my_use_plasma #ifdef __PLASMA INCLUDE 'plasmaf.h' @@ -49,23 +49,6 @@ IF (PRESENT(use_plasma)) my_use_plasma = use_plasma IF (PRESENT (lflop)) lflop = INT(0, int_8) IF (PRESENT (ltime)) ltime = 0.0_real_8 - ! Increase product data area size if necessary. - do_resize = .FALSE. - maxs = dbcsr_data_get_size(product_data_area) - DO sp = 1, stack_size - sp_c = params(sp)%p_c - IF (params_c(sp_c)%offset+params_c(sp_c)%nze-1 .GT. maxs) THEN - maxs = params_c(sp_c)%offset+params_c(sp_c)%nze-1 - do_resize = .TRUE. - ENDIF - ENDDO - IF (maxs .GT. dbcsr_data_get_size_referenced (product_data_area)) THEN - CALL dbcsr_data_set_size_referenced (product_data_area, maxs) - ENDIF - IF (do_resize) THEN - CALL dbcsr_data_ensure_size (product_data_area,& - maxs, factor=resize_factor, error=error) - ENDIF ! Setup encapsulated data area DO sp = 1, stack_size sp_c = params(sp)%p_c @@ -84,19 +67,25 @@ IF (params(sp)%beta%r_dp .EQ. 0.0_real_8) THEN IF (.NOT. params_c(sp_c)%tr & .OR. params(sp)%last_n .EQ. params_c(sp_c)%logical_cols) THEN - CALL dbcsr_data_clear (product_data_area,& - lb=poff,& - ub=poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) + product_data_a(poff : poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) = 0.0_real_8 + !CALL dbcsr_data_clear (product_data_area,& + ! lb=poff,& + ! ub=poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) ELSE ! It must be transposed. + FORALL (r = 1 : params_c(sp_c)%logical_rows,& + c = 1 : params(sp)%last_n) + product_data_a(poff-1+(r-1)*params_c(sp_c)%logical_cols+c) =& + 0.0_real_8 + END FORALL !### optimize out the inner loop (turn into a range) - DO r = 1, params_c(sp_c)%logical_rows - DO c = 1, params(sp)%last_n - CALL dbcsr_data_clear(product_data_area,& - lb=poff-1+(r-1)*params_c(sp_c)%logical_cols+c,& - ub=poff-1+(r-1)*params_c(sp_c)%logical_cols+c) - ENDDO - ENDDO + !DO r = 1, params_c(sp_c)%logical_rows + ! DO c = 1, params(sp)%last_n + ! CALL dbcsr_data_clear(product_data_area,& + ! lb=poff-1+(r-1)*params_c(sp_c)%logical_cols+c,& + ! ub=poff-1+(r-1)*params_c(sp_c)%logical_cols+c) + ! ENDDO + !ENDDO !FORALL (r = 1:params(sp)%lrows_c, c = 1:params(sp)%last_n) ! product_data_d((r-1)*params(sp)%lcols_c+c) = 0.0_dp !END FORALL @@ -123,14 +112,25 @@ left_data_a(params_a(sp_a)%offset), lda,& right_data_a(params_b(sp_b)%offset), ldb,& params(sp)%beta%r_dp,& - product_data_area%d%r_dp(poff), ldc,& + product_data_a(poff), ldc,& plasma_info) - CALL dbcsr_assert( plasma_info==0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& + CALL dbcsr_assert( plasma_info, "EQ", 0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma_gemm failed", __LINE__, error=error) #else CALL dbcsr_assert( .FALSE., dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma badly set", __LINE__, error=error) #endif + ELSEIF (internal_gemm) THEN + CALL internal_DGEMM(& + params_a(sp_a)%tr,& + params_b(sp_b)%tr,& + params_c(sp_c)%logical_rows, params(sp)%last_n,& !m, n + params(sp)%last_k,& ! k + params(sp)%alpha%r_dp,& + left_data_a(params_a(sp_a)%offset:), lda,& + right_data_a(params_b(sp_b)%offset:), ldb,& + params(sp)%beta%r_dp,& + product_data_a(poff:), ldc) ELSE CALL DGEMM(& blas_mat_type(params_a(sp_a)%tr),& @@ -141,7 +141,7 @@ left_data_a(params_a(sp_a)%offset), lda,& right_data_a(params_b(sp_b)%offset), ldb,& params(sp)%beta%r_dp,& - product_data_area%d%r_dp(poff), ldc) + product_data_a(poff), ldc) ENDIF IF (PRESENT (ltime)) ltime = m_walltime() - ltime IF (PRESENT (lflop)) & @@ -161,14 +161,25 @@ right_data_a(params_b(sp_b)%offset), ldb,& left_data_a(params_a(sp_a)%offset), lda,& params(sp)%beta%r_dp,& - product_data_area%d%r_dp(poff), ldc,& + product_data_a(poff), ldc,& plasma_info) - CALL dbcsr_assert( plasma_info==0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& + CALL dbcsr_assert( plasma_info, "EQ", 0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma_gemm failed", __LINE__, error=error) #else CALL dbcsr_assert( .FALSE., dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma badly set", __LINE__, error=error) #endif + ELSEIF (internal_gemm) THEN + CALL internal_DGEMM(& + .NOT.params_b(sp_b)%tr,& + .NOT.params_a(sp_a)%tr,& + params_c(sp_c)%logical_cols, params_c(sp_c)%logical_rows,& !m, n (switched) + params(sp)%last_k,& ! k + params(sp)%alpha%r_dp,& + right_data_a(params_b(sp_b)%offset:), ldb,& + left_data_a(params_a(sp_a)%offset:), lda,& + params(sp)%beta%r_dp,& + product_data_a(poff:), ldc) ELSE CALL DGEMM(& blas_mat_type(.NOT.params_b(sp_b)%tr),& @@ -179,7 +190,7 @@ right_data_a(params_b(sp_b)%offset), ldb,& left_data_a(params_a(sp_a)%offset), lda,& params(sp)%beta%r_dp,& - product_data_area%d%r_dp(poff), ldc) + product_data_a(poff), ldc) ENDIF IF (PRESENT (ltime)) ltime = m_walltime() - ltime IF (PRESENT (lflop)) & @@ -191,3 +202,68 @@ END SUBROUTINE process_dgemm_stack_d + PURE SUBROUTINE internal_DGEMM(& + TRANSA,TRANSB,M,N,K,ALPHA,A,LDA,B,LDB,BETA,C,LDC) + LOGICAL, INTENT(IN) :: TRANSA, TRANSB + INTEGER, INTENT(IN) :: M, N, K + INTEGER, INTENT(IN) :: LDC + REAL(kind=real_8), INTENT(INOUT) :: C(LDC,*) + REAL(kind=real_8), INTENT(IN) :: BETA + INTEGER, INTENT(IN) :: LDB + REAL(kind=real_8), INTENT(IN) :: B(LDB,*) + INTEGER, INTENT(IN) :: LDA + REAL(kind=real_8), INTENT(IN) :: A(LDA,*), ALPHA + IF (.NOT. transa .AND. .NOT. transb) THEN + CALL internal_DGEMM_nn(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (.NOT. transa .AND. transb) THEN + CALL internal_DGEMM_nt(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (transa .AND. .NOT. transb) THEN + CALL internal_DGEMM_tn(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (transa .AND. transb) THEN + CALL internal_DGEMM_tt(M,N,K,ALPHA,A,B,BETA,C) + ENDIF + END SUBROUTINE internal_DGEMM + + PURE SUBROUTINE internal_DGEMM_nn(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + REAL(kind=real_8), INTENT(INOUT) :: C(M,N) + REAL(kind=real_8), INTENT(IN) :: BETA + REAL(kind=real_8), INTENT(IN) :: B(K,N) + REAL(kind=real_8), INTENT(IN) :: A(M,K), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (A, B) + END SUBROUTINE internal_DGEMM_nn + + PURE SUBROUTINE internal_DGEMM_nt(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + REAL(kind=real_8), INTENT(INOUT) :: C(M,N) + REAL(kind=real_8), INTENT(IN) :: BETA + REAL(kind=real_8), INTENT(IN) :: B(N,K) + REAL(kind=real_8), INTENT(IN) :: A(M,K), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (A, TRANSPOSE(B)) + END SUBROUTINE internal_DGEMM_nt + + PURE SUBROUTINE internal_DGEMM_tn(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + REAL(kind=real_8), INTENT(INOUT) :: C(M,N) + REAL(kind=real_8), INTENT(IN) :: BETA + REAL(kind=real_8), INTENT(IN) :: B(K,N) + REAL(kind=real_8), INTENT(IN) :: A(K,M), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (TRANSPOSE(A), B) + END SUBROUTINE internal_DGEMM_tn + + PURE SUBROUTINE internal_DGEMM_tt(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + REAL(kind=real_8), INTENT(INOUT) :: C(M,N) + REAL(kind=real_8), INTENT(IN) :: BETA + REAL(kind=real_8), INTENT(IN) :: B(N,K) + REAL(kind=real_8), INTENT(IN) :: A(K,M), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (TRANSPOSE(A), TRANSPOSE(B)) + END SUBROUTINE internal_DGEMM_tt diff --git a/src/lib/dbcsr_internal_operations_s.F b/src/lib/dbcsr_internal_operations_s.F index 5b547e205e..4f7191d444 100644 --- a/src/lib/dbcsr_internal_operations_s.F +++ b/src/lib/dbcsr_internal_operations_s.F @@ -12,7 +12,7 @@ SUBROUTINE process_dgemm_stack_s(params,& params_a, params_b, params_c,& stack_size, stack_size_a, stack_size_b, stack_size_c,& - left_data_a, right_data_a, product_data_area, use_plasma, lflop, ltime, error) + left_data_a, right_data_a, product_data_a, use_plasma, lflop, ltime, error) INTEGER, INTENT(IN) :: stack_size, stack_size_a,& stack_size_b, stack_size_c TYPE(dgemm_join), DIMENSION(1:stack_size), & @@ -22,7 +22,7 @@ TYPE(block_parameters), DIMENSION(1:stack_size_c), INTENT(IN) :: params_c REAL(kind=real_4), DIMENSION(:), INTENT(IN) :: left_data_a, & right_data_a - TYPE(dbcsr_data_obj), INTENT(INOUT) :: product_data_area + REAL(kind=real_4), DIMENSION(:), INTENT(INOUT) :: product_data_a LOGICAL, INTENT(IN), OPTIONAL :: use_plasma INTEGER(KIND=int_8), INTENT(OUT), & OPTIONAL :: lflop @@ -33,10 +33,10 @@ routineP = moduleN//':'//routineN REAL, PARAMETER :: resize_factor = 1.618034 - INTEGER :: c, plasma_info, lda, ldb, ldc, maxs, r, sp,& + INTEGER :: c, plasma_info, lda, ldb, ldc, r, sp,& poff INTEGER :: sp_c, sp_a, sp_b - LOGICAL :: do_resize, my_use_plasma + LOGICAL :: my_use_plasma #ifdef __PLASMA INCLUDE 'plasmaf.h' @@ -49,23 +49,6 @@ IF (PRESENT(use_plasma)) my_use_plasma = use_plasma IF (PRESENT (lflop)) lflop = INT(0, int_8) IF (PRESENT (ltime)) ltime = 0.0_real_8 - ! Increase product data area size if necessary. - do_resize = .FALSE. - maxs = dbcsr_data_get_size(product_data_area) - DO sp = 1, stack_size - sp_c = params(sp)%p_c - IF (params_c(sp_c)%offset+params_c(sp_c)%nze-1 .GT. maxs) THEN - maxs = params_c(sp_c)%offset+params_c(sp_c)%nze-1 - do_resize = .TRUE. - ENDIF - ENDDO - IF (maxs .GT. dbcsr_data_get_size_referenced (product_data_area)) THEN - CALL dbcsr_data_set_size_referenced (product_data_area, maxs) - ENDIF - IF (do_resize) THEN - CALL dbcsr_data_ensure_size (product_data_area,& - maxs, factor=resize_factor, error=error) - ENDIF ! Setup encapsulated data area DO sp = 1, stack_size sp_c = params(sp)%p_c @@ -84,19 +67,25 @@ IF (params(sp)%beta%r_sp .EQ. 0.0_real_4) THEN IF (.NOT. params_c(sp_c)%tr & .OR. params(sp)%last_n .EQ. params_c(sp_c)%logical_cols) THEN - CALL dbcsr_data_clear (product_data_area,& - lb=poff,& - ub=poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) + product_data_a(poff : poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) = 0.0_real_4 + !CALL dbcsr_data_clear (product_data_area,& + ! lb=poff,& + ! ub=poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) ELSE ! It must be transposed. + FORALL (r = 1 : params_c(sp_c)%logical_rows,& + c = 1 : params(sp)%last_n) + product_data_a(poff-1+(r-1)*params_c(sp_c)%logical_cols+c) =& + 0.0_real_4 + END FORALL !### optimize out the inner loop (turn into a range) - DO r = 1, params_c(sp_c)%logical_rows - DO c = 1, params(sp)%last_n - CALL dbcsr_data_clear(product_data_area,& - lb=poff-1+(r-1)*params_c(sp_c)%logical_cols+c,& - ub=poff-1+(r-1)*params_c(sp_c)%logical_cols+c) - ENDDO - ENDDO + !DO r = 1, params_c(sp_c)%logical_rows + ! DO c = 1, params(sp)%last_n + ! CALL dbcsr_data_clear(product_data_area,& + ! lb=poff-1+(r-1)*params_c(sp_c)%logical_cols+c,& + ! ub=poff-1+(r-1)*params_c(sp_c)%logical_cols+c) + ! ENDDO + !ENDDO !FORALL (r = 1:params(sp)%lrows_c, c = 1:params(sp)%last_n) ! product_data_d((r-1)*params(sp)%lcols_c+c) = 0.0_dp !END FORALL @@ -123,14 +112,25 @@ left_data_a(params_a(sp_a)%offset), lda,& right_data_a(params_b(sp_b)%offset), ldb,& params(sp)%beta%r_sp,& - product_data_area%d%r_sp(poff), ldc,& + product_data_a(poff), ldc,& plasma_info) - CALL dbcsr_assert( plasma_info==0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& + CALL dbcsr_assert( plasma_info, "EQ", 0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma_gemm failed", __LINE__, error=error) #else CALL dbcsr_assert( .FALSE., dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma badly set", __LINE__, error=error) #endif + ELSEIF (internal_gemm) THEN + CALL internal_SGEMM(& + params_a(sp_a)%tr,& + params_b(sp_b)%tr,& + params_c(sp_c)%logical_rows, params(sp)%last_n,& !m, n + params(sp)%last_k,& ! k + params(sp)%alpha%r_sp,& + left_data_a(params_a(sp_a)%offset:), lda,& + right_data_a(params_b(sp_b)%offset:), ldb,& + params(sp)%beta%r_sp,& + product_data_a(poff:), ldc) ELSE CALL SGEMM(& blas_mat_type(params_a(sp_a)%tr),& @@ -141,7 +141,7 @@ left_data_a(params_a(sp_a)%offset), lda,& right_data_a(params_b(sp_b)%offset), ldb,& params(sp)%beta%r_sp,& - product_data_area%d%r_sp(poff), ldc) + product_data_a(poff), ldc) ENDIF IF (PRESENT (ltime)) ltime = m_walltime() - ltime IF (PRESENT (lflop)) & @@ -161,14 +161,25 @@ right_data_a(params_b(sp_b)%offset), ldb,& left_data_a(params_a(sp_a)%offset), lda,& params(sp)%beta%r_sp,& - product_data_area%d%r_sp(poff), ldc,& + product_data_a(poff), ldc,& plasma_info) - CALL dbcsr_assert( plasma_info==0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& + CALL dbcsr_assert( plasma_info, "EQ", 0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma_gemm failed", __LINE__, error=error) #else CALL dbcsr_assert( .FALSE., dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma badly set", __LINE__, error=error) #endif + ELSEIF (internal_gemm) THEN + CALL internal_SGEMM(& + .NOT.params_b(sp_b)%tr,& + .NOT.params_a(sp_a)%tr,& + params_c(sp_c)%logical_cols, params_c(sp_c)%logical_rows,& !m, n (switched) + params(sp)%last_k,& ! k + params(sp)%alpha%r_sp,& + right_data_a(params_b(sp_b)%offset:), ldb,& + left_data_a(params_a(sp_a)%offset:), lda,& + params(sp)%beta%r_sp,& + product_data_a(poff:), ldc) ELSE CALL SGEMM(& blas_mat_type(.NOT.params_b(sp_b)%tr),& @@ -179,7 +190,7 @@ right_data_a(params_b(sp_b)%offset), ldb,& left_data_a(params_a(sp_a)%offset), lda,& params(sp)%beta%r_sp,& - product_data_area%d%r_sp(poff), ldc) + product_data_a(poff), ldc) ENDIF IF (PRESENT (ltime)) ltime = m_walltime() - ltime IF (PRESENT (lflop)) & @@ -191,3 +202,68 @@ END SUBROUTINE process_dgemm_stack_s + PURE SUBROUTINE internal_SGEMM(& + TRANSA,TRANSB,M,N,K,ALPHA,A,LDA,B,LDB,BETA,C,LDC) + LOGICAL, INTENT(IN) :: TRANSA, TRANSB + INTEGER, INTENT(IN) :: M, N, K + INTEGER, INTENT(IN) :: LDC + REAL(kind=real_4), INTENT(INOUT) :: C(LDC,*) + REAL(kind=real_4), INTENT(IN) :: BETA + INTEGER, INTENT(IN) :: LDB + REAL(kind=real_4), INTENT(IN) :: B(LDB,*) + INTEGER, INTENT(IN) :: LDA + REAL(kind=real_4), INTENT(IN) :: A(LDA,*), ALPHA + IF (.NOT. transa .AND. .NOT. transb) THEN + CALL internal_SGEMM_nn(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (.NOT. transa .AND. transb) THEN + CALL internal_SGEMM_nt(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (transa .AND. .NOT. transb) THEN + CALL internal_SGEMM_tn(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (transa .AND. transb) THEN + CALL internal_SGEMM_tt(M,N,K,ALPHA,A,B,BETA,C) + ENDIF + END SUBROUTINE internal_SGEMM + + PURE SUBROUTINE internal_SGEMM_nn(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + REAL(kind=real_4), INTENT(INOUT) :: C(M,N) + REAL(kind=real_4), INTENT(IN) :: BETA + REAL(kind=real_4), INTENT(IN) :: B(K,N) + REAL(kind=real_4), INTENT(IN) :: A(M,K), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (A, B) + END SUBROUTINE internal_SGEMM_nn + + PURE SUBROUTINE internal_SGEMM_nt(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + REAL(kind=real_4), INTENT(INOUT) :: C(M,N) + REAL(kind=real_4), INTENT(IN) :: BETA + REAL(kind=real_4), INTENT(IN) :: B(N,K) + REAL(kind=real_4), INTENT(IN) :: A(M,K), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (A, TRANSPOSE(B)) + END SUBROUTINE internal_SGEMM_nt + + PURE SUBROUTINE internal_SGEMM_tn(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + REAL(kind=real_4), INTENT(INOUT) :: C(M,N) + REAL(kind=real_4), INTENT(IN) :: BETA + REAL(kind=real_4), INTENT(IN) :: B(K,N) + REAL(kind=real_4), INTENT(IN) :: A(K,M), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (TRANSPOSE(A), B) + END SUBROUTINE internal_SGEMM_tn + + PURE SUBROUTINE internal_SGEMM_tt(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + REAL(kind=real_4), INTENT(INOUT) :: C(M,N) + REAL(kind=real_4), INTENT(IN) :: BETA + REAL(kind=real_4), INTENT(IN) :: B(N,K) + REAL(kind=real_4), INTENT(IN) :: A(K,M), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (TRANSPOSE(A), TRANSPOSE(B)) + END SUBROUTINE internal_SGEMM_tt diff --git a/src/lib/dbcsr_internal_operations_z.F b/src/lib/dbcsr_internal_operations_z.F index c170153929..72b552fd54 100644 --- a/src/lib/dbcsr_internal_operations_z.F +++ b/src/lib/dbcsr_internal_operations_z.F @@ -12,7 +12,7 @@ SUBROUTINE process_dgemm_stack_z(params,& params_a, params_b, params_c,& stack_size, stack_size_a, stack_size_b, stack_size_c,& - left_data_a, right_data_a, product_data_area, use_plasma, lflop, ltime, error) + left_data_a, right_data_a, product_data_a, use_plasma, lflop, ltime, error) INTEGER, INTENT(IN) :: stack_size, stack_size_a,& stack_size_b, stack_size_c TYPE(dgemm_join), DIMENSION(1:stack_size), & @@ -22,7 +22,7 @@ TYPE(block_parameters), DIMENSION(1:stack_size_c), INTENT(IN) :: params_c COMPLEX(kind=real_8), DIMENSION(:), INTENT(IN) :: left_data_a, & right_data_a - TYPE(dbcsr_data_obj), INTENT(INOUT) :: product_data_area + COMPLEX(kind=real_8), DIMENSION(:), INTENT(INOUT) :: product_data_a LOGICAL, INTENT(IN), OPTIONAL :: use_plasma INTEGER(KIND=int_8), INTENT(OUT), & OPTIONAL :: lflop @@ -33,10 +33,10 @@ routineP = moduleN//':'//routineN REAL, PARAMETER :: resize_factor = 1.618034 - INTEGER :: c, plasma_info, lda, ldb, ldc, maxs, r, sp,& + INTEGER :: c, plasma_info, lda, ldb, ldc, r, sp,& poff INTEGER :: sp_c, sp_a, sp_b - LOGICAL :: do_resize, my_use_plasma + LOGICAL :: my_use_plasma #ifdef __PLASMA INCLUDE 'plasmaf.h' @@ -49,23 +49,6 @@ IF (PRESENT(use_plasma)) my_use_plasma = use_plasma IF (PRESENT (lflop)) lflop = INT(0, int_8) IF (PRESENT (ltime)) ltime = 0.0_real_8 - ! Increase product data area size if necessary. - do_resize = .FALSE. - maxs = dbcsr_data_get_size(product_data_area) - DO sp = 1, stack_size - sp_c = params(sp)%p_c - IF (params_c(sp_c)%offset+params_c(sp_c)%nze-1 .GT. maxs) THEN - maxs = params_c(sp_c)%offset+params_c(sp_c)%nze-1 - do_resize = .TRUE. - ENDIF - ENDDO - IF (maxs .GT. dbcsr_data_get_size_referenced (product_data_area)) THEN - CALL dbcsr_data_set_size_referenced (product_data_area, maxs) - ENDIF - IF (do_resize) THEN - CALL dbcsr_data_ensure_size (product_data_area,& - maxs, factor=resize_factor, error=error) - ENDIF ! Setup encapsulated data area DO sp = 1, stack_size sp_c = params(sp)%p_c @@ -84,19 +67,25 @@ IF (params(sp)%beta%c_dp .EQ. CMPLX(0.0, 0.0, real_8)) THEN IF (.NOT. params_c(sp_c)%tr & .OR. params(sp)%last_n .EQ. params_c(sp_c)%logical_cols) THEN - CALL dbcsr_data_clear (product_data_area,& - lb=poff,& - ub=poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) + product_data_a(poff : poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) = CMPLX(0.0, 0.0, real_8) + !CALL dbcsr_data_clear (product_data_area,& + ! lb=poff,& + ! ub=poff-1+params_c(sp_c)%logical_rows*params(sp)%last_n) ELSE ! It must be transposed. + FORALL (r = 1 : params_c(sp_c)%logical_rows,& + c = 1 : params(sp)%last_n) + product_data_a(poff-1+(r-1)*params_c(sp_c)%logical_cols+c) =& + CMPLX(0.0, 0.0, real_8) + END FORALL !### optimize out the inner loop (turn into a range) - DO r = 1, params_c(sp_c)%logical_rows - DO c = 1, params(sp)%last_n - CALL dbcsr_data_clear(product_data_area,& - lb=poff-1+(r-1)*params_c(sp_c)%logical_cols+c,& - ub=poff-1+(r-1)*params_c(sp_c)%logical_cols+c) - ENDDO - ENDDO + !DO r = 1, params_c(sp_c)%logical_rows + ! DO c = 1, params(sp)%last_n + ! CALL dbcsr_data_clear(product_data_area,& + ! lb=poff-1+(r-1)*params_c(sp_c)%logical_cols+c,& + ! ub=poff-1+(r-1)*params_c(sp_c)%logical_cols+c) + ! ENDDO + !ENDDO !FORALL (r = 1:params(sp)%lrows_c, c = 1:params(sp)%last_n) ! product_data_d((r-1)*params(sp)%lcols_c+c) = 0.0_dp !END FORALL @@ -123,14 +112,25 @@ left_data_a(params_a(sp_a)%offset), lda,& right_data_a(params_b(sp_b)%offset), ldb,& params(sp)%beta%c_dp,& - product_data_area%d%c_dp(poff), ldc,& + product_data_a(poff), ldc,& plasma_info) - CALL dbcsr_assert( plasma_info==0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& + CALL dbcsr_assert( plasma_info, "EQ", 0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma_gemm failed", __LINE__, error=error) #else CALL dbcsr_assert( .FALSE., dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma badly set", __LINE__, error=error) #endif + ELSEIF (internal_gemm) THEN + CALL internal_ZGEMM(& + params_a(sp_a)%tr,& + params_b(sp_b)%tr,& + params_c(sp_c)%logical_rows, params(sp)%last_n,& !m, n + params(sp)%last_k,& ! k + params(sp)%alpha%c_dp,& + left_data_a(params_a(sp_a)%offset:), lda,& + right_data_a(params_b(sp_b)%offset:), ldb,& + params(sp)%beta%c_dp,& + product_data_a(poff:), ldc) ELSE CALL ZGEMM(& blas_mat_type(params_a(sp_a)%tr),& @@ -141,7 +141,7 @@ left_data_a(params_a(sp_a)%offset), lda,& right_data_a(params_b(sp_b)%offset), ldb,& params(sp)%beta%c_dp,& - product_data_area%d%c_dp(poff), ldc) + product_data_a(poff), ldc) ENDIF IF (PRESENT (ltime)) ltime = m_walltime() - ltime IF (PRESENT (lflop)) & @@ -161,14 +161,25 @@ right_data_a(params_b(sp_b)%offset), ldb,& left_data_a(params_a(sp_a)%offset), lda,& params(sp)%beta%c_dp,& - product_data_area%d%c_dp(poff), ldc,& + product_data_a(poff), ldc,& plasma_info) - CALL dbcsr_assert( plasma_info==0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& + CALL dbcsr_assert( plasma_info, "EQ", 0, dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma_gemm failed", __LINE__, error=error) #else CALL dbcsr_assert( .FALSE., dbcsr_fatal_level, dbcsr_internal_error, routineN,& "plasma badly set", __LINE__, error=error) #endif + ELSEIF (internal_gemm) THEN + CALL internal_ZGEMM(& + .NOT.params_b(sp_b)%tr,& + .NOT.params_a(sp_a)%tr,& + params_c(sp_c)%logical_cols, params_c(sp_c)%logical_rows,& !m, n (switched) + params(sp)%last_k,& ! k + params(sp)%alpha%c_dp,& + right_data_a(params_b(sp_b)%offset:), ldb,& + left_data_a(params_a(sp_a)%offset:), lda,& + params(sp)%beta%c_dp,& + product_data_a(poff:), ldc) ELSE CALL ZGEMM(& blas_mat_type(.NOT.params_b(sp_b)%tr),& @@ -179,7 +190,7 @@ right_data_a(params_b(sp_b)%offset), ldb,& left_data_a(params_a(sp_a)%offset), lda,& params(sp)%beta%c_dp,& - product_data_area%d%c_dp(poff), ldc) + product_data_a(poff), ldc) ENDIF IF (PRESENT (ltime)) ltime = m_walltime() - ltime IF (PRESENT (lflop)) & @@ -191,3 +202,68 @@ END SUBROUTINE process_dgemm_stack_z + PURE SUBROUTINE internal_ZGEMM(& + TRANSA,TRANSB,M,N,K,ALPHA,A,LDA,B,LDB,BETA,C,LDC) + LOGICAL, INTENT(IN) :: TRANSA, TRANSB + INTEGER, INTENT(IN) :: M, N, K + INTEGER, INTENT(IN) :: LDC + COMPLEX(kind=real_8), INTENT(INOUT) :: C(LDC,*) + COMPLEX(kind=real_8), INTENT(IN) :: BETA + INTEGER, INTENT(IN) :: LDB + COMPLEX(kind=real_8), INTENT(IN) :: B(LDB,*) + INTEGER, INTENT(IN) :: LDA + COMPLEX(kind=real_8), INTENT(IN) :: A(LDA,*), ALPHA + IF (.NOT. transa .AND. .NOT. transb) THEN + CALL internal_ZGEMM_nn(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (.NOT. transa .AND. transb) THEN + CALL internal_ZGEMM_nt(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (transa .AND. .NOT. transb) THEN + CALL internal_ZGEMM_tn(M,N,K,ALPHA,A,B,BETA,C) + ELSEIF (transa .AND. transb) THEN + CALL internal_ZGEMM_tt(M,N,K,ALPHA,A,B,BETA,C) + ENDIF + END SUBROUTINE internal_ZGEMM + + PURE SUBROUTINE internal_ZGEMM_nn(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + COMPLEX(kind=real_8), INTENT(INOUT) :: C(M,N) + COMPLEX(kind=real_8), INTENT(IN) :: BETA + COMPLEX(kind=real_8), INTENT(IN) :: B(K,N) + COMPLEX(kind=real_8), INTENT(IN) :: A(M,K), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (A, B) + END SUBROUTINE internal_ZGEMM_nn + + PURE SUBROUTINE internal_ZGEMM_nt(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + COMPLEX(kind=real_8), INTENT(INOUT) :: C(M,N) + COMPLEX(kind=real_8), INTENT(IN) :: BETA + COMPLEX(kind=real_8), INTENT(IN) :: B(N,K) + COMPLEX(kind=real_8), INTENT(IN) :: A(M,K), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (A, TRANSPOSE(B)) + END SUBROUTINE internal_ZGEMM_nt + + PURE SUBROUTINE internal_ZGEMM_tn(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + COMPLEX(kind=real_8), INTENT(INOUT) :: C(M,N) + COMPLEX(kind=real_8), INTENT(IN) :: BETA + COMPLEX(kind=real_8), INTENT(IN) :: B(K,N) + COMPLEX(kind=real_8), INTENT(IN) :: A(K,M), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (TRANSPOSE(A), B) + END SUBROUTINE internal_ZGEMM_tn + + PURE SUBROUTINE internal_ZGEMM_tt(& + M,N,K,ALPHA,A,B,BETA,C) + INTEGER, INTENT(IN) :: M, N, K + COMPLEX(kind=real_8), INTENT(INOUT) :: C(M,N) + COMPLEX(kind=real_8), INTENT(IN) :: BETA + COMPLEX(kind=real_8), INTENT(IN) :: B(N,K) + COMPLEX(kind=real_8), INTENT(IN) :: A(K,M), ALPHA + C(:,:) = BETA * C(:,:) & + + ALPHA * MATMUL (TRANSPOSE(A), TRANSPOSE(B)) + END SUBROUTINE internal_ZGEMM_tt