Merge pull request #1056 from omarkahmed/omarkahmed/ccsdt_offload_batched_gemm-fixup

Fixups for Intel Xe CCSD(T) OpenMP Offload
This commit is contained in:
Edoardo Aprà 2024-12-11 22:30:26 +08:00 committed by GitHub
commit 0da5c01b24
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 15 additions and 86 deletions

View file

@ -969,9 +969,11 @@ c print *,'call trpdrv ',nvpass
!$omp target
!$omp end target
#ifndef USE_BATCHDGEMM_TRPDRV
! Dummy interop object
!$omp interop init(prefer_type("sycl"),targetsync: dummy_obj)
!$omp interop destroy(dummy_obj)
#endif
#endif
tx(2)=tcgtime()

View file

@ -169,6 +169,12 @@
logical got_ak
integer thread_num
! timers
double precision :: tt0, tt1, tc0, tc1
double precision :: t_dgemm0, t_dgemm1, t_dgemm_total
double precision :: t_red0, t_red1, t_red_total
#ifndef USE_BATCHDGEMM_TRPDRV
! OpenMP interop objects
integer(kind = omp_interop_kind) :: obj0 = omp_interop_none
integer(kind = omp_interop_kind) :: obj1 = omp_interop_none
@ -180,34 +186,6 @@
integer(kind = omp_interop_kind) :: obj7 = omp_interop_none
integer(kind = omp_interop_kind) :: obj_lev0 = omp_interop_none
! timers
double precision :: tt0, tt1, tc0, tc1
double precision :: t_dgemm0, t_dgemm1, t_dgemm_total
double precision :: t_red0, t_red1, t_red_total
#if 0
!$omp interop init(targetsync:obj0)
!$omp interop init(targetsync:obj1)
!$omp interop init(targetsync:obj2)
!$omp interop init(targetsync:obj3)
!$omp interop init(targetsync:obj4)
!$omp interop init(targetsync:obj5)
!$omp interop init(targetsync:obj6)
!$omp interop init(targetsync:obj7)
#endif
#if 0
!$omp interop init(prefer_type("level_zero"),targetsync: obj0)
!$omp interop init(prefer_type("level_zero"),targetsync: obj1)
!$omp interop init(prefer_type("level_zero"),targetsync: obj2)
!$omp interop init(prefer_type("level_zero"),targetsync: obj3)
!$omp interop init(prefer_type("level_zero"),targetsync: obj4)
!$omp interop init(prefer_type("level_zero"),targetsync: obj5)
!$omp interop init(prefer_type("level_zero"),targetsync: obj6)
!$omp interop init(prefer_type("level_zero"),targetsync: obj7)
#endif
#if 1
!$omp interop init(prefer_type("sycl"),targetsync: obj0)
!$omp interop init(prefer_type("sycl"),targetsync: obj1)
!$omp interop init(prefer_type("sycl"),targetsync: obj2)
@ -216,7 +194,7 @@
!$omp interop init(prefer_type("sycl"),targetsync: obj5)
!$omp interop init(prefer_type("sycl"),targetsync: obj6)
!$omp interop init(prefer_type("sycl"),targetsync: obj7)
!$omp interop init(prefer_type("level_zero"),targetsync: obj_lev0)
#endif
@ -255,7 +233,6 @@
! & stat=alloc_error)
! if (alloc_error.ne.0) call errquit('f[1234][tn]',8,MA_ERR)
#if 1
!$omp allocate allocator(omp_target_device_mem_alloc)
allocate( f1n(1:nvir,1:nvir) )
!$omp allocate allocator(omp_target_device_mem_alloc)
@ -272,25 +249,6 @@
allocate( f4n(1:nvir,1:nvir) )
!$omp allocate allocator(omp_target_device_mem_alloc)
allocate( f4t(1:nvir,1:nvir) )
#else
!$omp allocate allocator(omp_target_host_mem_alloc)
allocate( f1n(1:nvir,1:nvir) )
!$omp allocate allocator(omp_target_host_mem_alloc)
allocate( f1t(1:nvir,1:nvir) )
!$omp allocate allocator(omp_target_host_mem_alloc)
allocate( f2n(1:nvir,1:nvir) )
!$omp allocate allocator(omp_target_host_mem_alloc)
allocate( f2t(1:nvir,1:nvir) )
!$omp allocate allocator(omp_target_host_mem_alloc)
allocate( f3n(1:nvir,1:nvir) )
!$omp allocate allocator(omp_target_host_mem_alloc)
allocate( f3t(1:nvir,1:nvir) )
!$omp allocate allocator(omp_target_host_mem_alloc)
allocate( f4n(1:nvir,1:nvir) )
!$omp allocate allocator(omp_target_host_mem_alloc)
allocate( f4t(1:nvir,1:nvir) )
#endif
!
! device-only copy of input eorb
!
@ -394,7 +352,10 @@
& map(to:dintc1,dintx1,t1v1,dintc2,dintx2,t1v2)
& map(to:Jia,Tkj,Tia,Kkj,Kia,Tka,Kij)
& map(to:Xka,Jij,Xia,Jkj,Jka,Tij,Kka)
#ifdef USE_BATCHDGEMM_TRPDRV
& map(to:a_array1, b_array1, c_array1)
& map(to:a_array2, b_array2, c_array2)
#endif
do klo = 1, nocc, kchunk
akold=0
khi = min(nocc, klo+kchunk-1)
@ -925,13 +886,7 @@
1 Xka(1+(k-klo)*lnov),nvir,Jij,nocc,1.0d0,
2 f4t,nvir)
!$omp end parallel sections
#endif
#if 0
!$omp interop use(obj0)
#endif
#if 1
! "omp interop use(obj_lev0)" generates a device barrier
! on level0 queue. This behavior is specific to Intel and
! works because of the mapping of level0 queues to
@ -944,7 +899,6 @@
!$omp interop use(obj_lev0) nowait
#endif
t_dgemm1 = util_wallsec()
t_dgemm_total = t_dgemm_total + (t_dgemm1 - t_dgemm0)
@ -1043,35 +997,6 @@
emp5 = emp5 + emp5k
#else
#if 0
!use two separate calls to C implementation
!$omp target update to(emp4i,emp5i)
call ccsd_trpdrv_omp_reduce_01(f1n, f1t, f2n, f2t,
& f3n, f3t, f4n, f4t,
& eorb,
& ncor, nocc, nvir,
! & emp4i, emp5i,
& eaijk,
& dintc1, dintx1, t1v1)
!$omp target update from(emp4i,emp5i)
if (i.ne.k) then
!$omp target update to(emp4k,emp5k)
call ccsd_trpdrv_omp_reduce_02(f1n, f1t, f2n, f2t,
& f3n, f3t, f4n, f4t,
& eorb,
& ncor, nocc, nvir,
! & emp4k, emp5k,
& eaijk,
& dintc2, dintx2, t1v2)
!$omp target update from(emp4k,emp5k)
end if ! (i.ne.k)
emp4 = emp4 + emp4i
emp5 = emp5 + emp5i
emp4 = emp4 + emp4k
emp5 = emp5 + emp5k
#endif
call ccsd_trpdrv_omp_fbody_reduce_new (f1n, f1t, f2n, f2t,
& f3n, f3t, f4n, f4t,
@ -1163,6 +1088,7 @@
! end mapping of all data below
!$omp end target data
#ifndef USE_BATCHDGEMM_TRPDRV
!$omp interop destroy(obj0)
!$omp interop destroy(obj1)
!$omp interop destroy(obj2)
@ -1172,6 +1098,7 @@
!$omp interop destroy(obj6)
!$omp interop destroy(obj7)
!$omp interop destroy(obj_lev0)
#endif
call ga_sync()
next=nxtask(-nodes, 1)