Actual source code: vecseqcupm_impl.hpp

  1: #pragma once

  3: #include "vecseqcupm.hpp"

  5: #include <petsc/private/randomimpl.h>

  7: #include "../src/sys/objects/device/impls/cupm/cupmthrustutility.hpp"
  8: #include "../src/sys/objects/device/impls/cupm/kernels.hpp"

 10: #if PetscDefined(USE_COMPLEX)
 11:   #include <thrust/transform_reduce.h>
 12: #endif
 13: #include <thrust/transform.h>
 14: #include <thrust/reduce.h>
 15: #include <thrust/functional.h>
 16: #include <thrust/tuple.h>
 17: #include <thrust/device_ptr.h>
 18: #include <thrust/iterator/zip_iterator.h>
 19: #include <thrust/iterator/counting_iterator.h>
 20: #include <thrust/iterator/constant_iterator.h>
 21: #include <thrust/inner_product.h>
 22: #if CCCL_VERSION >= 3004000
 23:   #include <cuda/iterator>
 24: #endif

 26: namespace Petsc
 27: {

 29: namespace vec
 30: {

 32: namespace cupm
 33: {

 35: namespace impl
 36: {

 38: // ==========================================================================================
 39: // VecSeq_CUPM - Private API
 40: // ==========================================================================================

 42: template <device::cupm::DeviceType T>
 43: inline Vec_Seq *VecSeq_CUPM<T>::VecIMPLCast_(Vec v) noexcept
 44: {
 45:   return static_cast<Vec_Seq *>(v->data);
 46: }

 48: template <device::cupm::DeviceType T>
 49: inline constexpr VecType VecSeq_CUPM<T>::VECIMPLCUPM_() noexcept
 50: {
 51:   return VECSEQCUPM();
 52: }

 54: template <device::cupm::DeviceType T>
 55: inline constexpr VecType VecSeq_CUPM<T>::VECIMPL_() noexcept
 56: {
 57:   return VECSEQ;
 58: }

 60: template <device::cupm::DeviceType T>
 61: inline PetscErrorCode VecSeq_CUPM<T>::ClearAsyncFunctions(Vec v) noexcept
 62: {
 63:   PetscFunctionBegin;
 64:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Abs), nullptr));
 65:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(AXPBY), nullptr));
 66:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(AXPBYPCZ), nullptr));
 67:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(AXPY), nullptr));
 68:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(AYPX), nullptr));
 69:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Conjugate), nullptr));
 70:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Copy), nullptr));
 71:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Exp), nullptr));
 72:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Log), nullptr));
 73:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(MAXPY), nullptr));
 74:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(PointwiseDivide), nullptr));
 75:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(PointwiseMax), nullptr));
 76:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(PointwiseMaxAbs), nullptr));
 77:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(PointwiseMin), nullptr));
 78:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(PointwiseMult), nullptr));
 79:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(PointwiseSign), nullptr));
 80:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Reciprocal), nullptr));
 81:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Scale), nullptr));
 82:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Set), nullptr));
 83:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Shift), nullptr));
 84:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(SqrtAbs), nullptr));
 85:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Swap), nullptr));
 86:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(WAXPY), nullptr));
 87:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(SetStdBasis), nullptr));
 88:   PetscFunctionReturn(PETSC_SUCCESS);
 89: }

 91: template <device::cupm::DeviceType T>
 92: inline PetscErrorCode VecSeq_CUPM<T>::InitializeAsyncFunctions(Vec v) noexcept
 93: {
 94:   PetscFunctionBegin;
 95:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Abs), VecSeq_CUPM<T>::AbsAsync));
 96:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(AXPBY), VecSeq_CUPM<T>::AXPBYAsync));
 97:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(AXPBYPCZ), VecSeq_CUPM<T>::AXPBYPCZAsync));
 98:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(AXPY), VecSeq_CUPM<T>::AXPYAsync));
 99:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(AYPX), VecSeq_CUPM<T>::AYPXAsync));
100:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Conjugate), VecSeq_CUPM<T>::ConjugateAsync));
101:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Copy), VecSeq_CUPM<T>::CopyAsync));
102:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Exp), VecSeq_CUPM<T>::ExpAsync));
103:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Log), VecSeq_CUPM<T>::LogAsync));
104:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(MAXPY), VecSeq_CUPM<T>::MAXPYAsync));
105:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(PointwiseDivide), VecSeq_CUPM<T>::PointwiseDivideAsync));
106:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(PointwiseMax), VecSeq_CUPM<T>::PointwiseMaxAsync));
107:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(PointwiseMaxAbs), VecSeq_CUPM<T>::PointwiseMaxAbsAsync));
108:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(PointwiseMin), VecSeq_CUPM<T>::PointwiseMinAsync));
109:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(PointwiseMult), VecSeq_CUPM<T>::PointwiseMultAsync));
110:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(PointwiseSign), VecSeq_CUPM<T>::PointwiseSignAsync));
111:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Reciprocal), VecSeq_CUPM<T>::ReciprocalAsync));
112:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Scale), VecSeq_CUPM<T>::ScaleAsync));
113:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Set), VecSeq_CUPM<T>::SetAsync));
114:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Shift), VecSeq_CUPM<T>::ShiftAsync));
115:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(SqrtAbs), VecSeq_CUPM<T>::SqrtAbsAsync));
116:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(Swap), VecSeq_CUPM<T>::SwapAsync));
117:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(WAXPY), VecSeq_CUPM<T>::WAXPYAsync));
118:   PetscCall(PetscObjectComposeFunction(PetscObjectCast(v), VecAsyncFnName(SetStdBasis), VecSeq_CUPM<T>::SetStdBasisAsync));
119:   PetscFunctionReturn(PETSC_SUCCESS);
120: }

122: template <device::cupm::DeviceType T>
123: inline PetscErrorCode VecSeq_CUPM<T>::VecDestroy_IMPL_(Vec v) noexcept
124: {
125:   PetscFunctionBegin;
126:   PetscCall(ClearAsyncFunctions(v));
127:   PetscCall(VecDestroy_Seq(v));
128:   PetscFunctionReturn(PETSC_SUCCESS);
129: }

131: template <device::cupm::DeviceType T>
132: inline PetscErrorCode VecSeq_CUPM<T>::VecResetArray_IMPL_(Vec v) noexcept
133: {
134:   return VecResetArray_Seq(v);
135: }

137: template <device::cupm::DeviceType T>
138: inline PetscErrorCode VecSeq_CUPM<T>::VecPlaceArray_IMPL_(Vec v, const PetscScalar *a) noexcept
139: {
140:   return VecPlaceArray_Seq(v, a);
141: }

143: template <device::cupm::DeviceType T>
144: inline PetscErrorCode VecSeq_CUPM<T>::VecCreate_IMPL_Private_(Vec v, PetscBool *alloc_missing, PetscInt, PetscScalar *host_array) noexcept
145: {
146:   PetscMPIInt size;

148:   PetscFunctionBegin;
149:   if (alloc_missing) *alloc_missing = PETSC_FALSE;
150:   PetscCallMPI(MPI_Comm_size(PetscObjectComm(PetscObjectCast(v)), &size));
151:   PetscCheck(size <= 1, PETSC_COMM_SELF, PETSC_ERR_ARG_WRONG, "Must create VecSeq on communicator of size 1, have size %d", size);
152:   PetscCall(VecCreate_Seq_Private(v, host_array));
153:   PetscCall(InitializeAsyncFunctions(v));
154:   PetscFunctionReturn(PETSC_SUCCESS);
155: }

157: // for functions with an early return based one vec size we still need to artificially bump the
158: // object state. This is to prevent the following:
159: //
160: // 0. Suppose you have a Vec {
161: //   rank 0: [0],
162: //   rank 1: []
163: // }
164: // 1. both ranks have Vec with PetscObjectState = 0, stashed norm of 0
165: // 2. Vec enters e.g. VecSet(10)
166: // 3. rank 1 has local size 0 and bails immediately
167: // 4. rank 0 has local size 1 and enters function, eventually calls DeviceArrayWrite()
168: // 5. DeviceArrayWrite() calls PetscObjectStateIncrease(), now state = 1
169: // 6. Vec enters VecNorm(), and calls VecNormAvailable()
170: // 7. rank 1 has object state = 0, equal to stash and returns early with norm = 0
171: // 8. rank 0 has object state = 1, not equal to stash, continues to impl function
172: // 9. rank 0 deadlocks on MPI_Allreduce() because rank 1 bailed early
173: template <device::cupm::DeviceType T>
174: inline PetscErrorCode VecSeq_CUPM<T>::MaybeIncrementEmptyLocalVec(Vec v) noexcept
175: {
176:   PetscFunctionBegin;
177:   if (PetscUnlikely((v->map->n == 0) && (v->map->N != 0))) PetscCall(PetscObjectStateIncrease(PetscObjectCast(v)));
178:   PetscFunctionReturn(PETSC_SUCCESS);
179: }

181: template <device::cupm::DeviceType T>
182: inline PetscErrorCode VecSeq_CUPM<T>::CreateSeqCUPM_(Vec v, PetscDeviceContext dctx, PetscScalar *host_array, PetscScalar *device_array) noexcept
183: {
184:   PetscFunctionBegin;
185:   PetscCall(base_type::VecCreate_IMPL_Private(v, nullptr, 0, host_array));
186:   PetscCall(Initialize_CUPMBase(v, PETSC_FALSE, host_array, device_array, dctx));
187:   PetscFunctionReturn(PETSC_SUCCESS);
188: }

190: template <device::cupm::DeviceType T>
191: template <typename BinaryFuncT>
192: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseBinary_(BinaryFuncT &&binary, Vec xin, Vec yin, Vec zout, PetscDeviceContext dctx) noexcept
193: {
194:   PetscFunctionBegin;
195:   if (const auto n = zout->map->n) {
196:     cupmStream_t stream;

198:     PetscCall(PetscDeviceContextGetOptionalNullContext_Internal(&dctx));
199:     PetscCall(GetHandlesFrom_(dctx, &stream));
200:     // clang-format off
201:     PetscCallThrust(
202:       const auto dxptr = thrust::device_pointer_cast(DeviceArrayRead(dctx, xin).data());

204:       THRUST_CALL(
205:         thrust::transform,
206:         stream,
207:         dxptr, dxptr + n,
208:         thrust::device_pointer_cast(DeviceArrayRead(dctx, yin).data()),
209:         thrust::device_pointer_cast(DeviceArrayWrite(dctx, zout).data()),
210:         std::forward<BinaryFuncT>(binary)
211:       )
212:     );
213:     // clang-format on
214:     PetscCall(PetscLogGpuFlops(n));
215:     PetscCall(PetscDeviceContextSynchronizeIfWithBarrier_Internal(dctx));
216:   } else {
217:     PetscCall(MaybeIncrementEmptyLocalVec(zout));
218:   }
219:   PetscFunctionReturn(PETSC_SUCCESS);
220: }

222: template <device::cupm::DeviceType T>
223: template <typename BinaryFuncT>
224: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseBinaryDispatch_(PetscErrorCode (*VecSeqFunction)(Vec, Vec, Vec), BinaryFuncT &&binary, Vec wout, Vec xin, Vec yin, PetscDeviceContext dctx) noexcept
225: {
226:   PetscFunctionBegin;
227:   if (xin->boundtocpu || yin->boundtocpu) PetscCall((*VecSeqFunction)(wout, xin, yin));
228:   else PetscCall(PointwiseBinary_(std::forward<BinaryFuncT>(binary), xin, yin, wout, dctx)); // note order of arguments! xin and yin are read, wout is written!
229:   PetscFunctionReturn(PETSC_SUCCESS);
230: }

232: template <device::cupm::DeviceType T>
233: template <typename UnaryFuncT>
234: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseUnary_(UnaryFuncT &&unary, Vec xinout, Vec yout, PetscDeviceContext dctx) noexcept
235: {
236:   const auto inplace = !yout || (xinout == yout);

238:   PetscFunctionBegin;
239:   if (const auto n = xinout->map->n) {
240:     cupmStream_t stream;
241:     const auto   apply = [&](PetscScalar *xinout, PetscScalar *yout = nullptr) {
242:       PetscFunctionBegin;
243:       // clang-format off
244:       PetscCallThrust(
245:         const auto xptr = thrust::device_pointer_cast(xinout);

247:         THRUST_CALL(
248:           thrust::transform,
249:           stream,
250:           xptr, xptr + n,
251:           (yout && (yout != xinout)) ? thrust::device_pointer_cast(yout) : xptr,
252:           std::forward<UnaryFuncT>(unary)
253:         )
254:       );
255:       // clang-format on
256:       PetscFunctionReturn(PETSC_SUCCESS);
257:     };

259:     PetscCall(PetscDeviceContextGetOptionalNullContext_Internal(&dctx));
260:     PetscCall(GetHandlesFrom_(dctx, &stream));
261:     if (inplace) {
262:       PetscCall(apply(DeviceArrayReadWrite(dctx, xinout).data()));
263:     } else {
264:       PetscCall(apply(DeviceArrayRead(dctx, xinout).data(), DeviceArrayWrite(dctx, yout).data()));
265:     }
266:     PetscCall(PetscLogGpuFlops(n));
267:     PetscCall(PetscDeviceContextSynchronizeIfWithBarrier_Internal(dctx));
268:   } else {
269:     if (inplace) {
270:       PetscCall(MaybeIncrementEmptyLocalVec(xinout));
271:     } else {
272:       PetscCall(MaybeIncrementEmptyLocalVec(yout));
273:     }
274:   }
275:   PetscFunctionReturn(PETSC_SUCCESS);
276: }

278: // ==========================================================================================
279: // VecSeq_CUPM - Public API - Constructors
280: // ==========================================================================================

282: // VecCreateSeqCUPM()
283: template <device::cupm::DeviceType T>
284: inline PetscErrorCode VecSeq_CUPM<T>::CreateSeqCUPM(MPI_Comm comm, PetscInt bs, PetscInt n, Vec *v, PetscBool call_set_type) noexcept
285: {
286:   PetscFunctionBegin;
287:   PetscCall(Create_CUPMBase(comm, bs, n, n, v, call_set_type));
288:   PetscFunctionReturn(PETSC_SUCCESS);
289: }

291: // VecCreateSeqCUPMWithArrays()
292: template <device::cupm::DeviceType T>
293: inline PetscErrorCode VecSeq_CUPM<T>::CreateSeqCUPMWithBothArrays(MPI_Comm comm, PetscInt bs, PetscInt n, const PetscScalar host_array[], const PetscScalar device_array[], Vec *v) noexcept
294: {
295:   PetscDeviceContext dctx;

297:   PetscFunctionBegin;
298:   PetscCall(GetHandles_(&dctx));
299:   // do NOT call VecSetType(), otherwise ops->create() -> create() ->
300:   // CreateSeqCUPM_() is called!
301:   PetscCall(CreateSeqCUPM(comm, bs, n, v, PETSC_FALSE));
302:   PetscCall(CreateSeqCUPM_(*v, dctx, PetscRemoveConstCast(host_array), PetscRemoveConstCast(device_array)));
303:   PetscFunctionReturn(PETSC_SUCCESS);
304: }

306: // v->ops->duplicate
307: template <device::cupm::DeviceType T>
308: inline PetscErrorCode VecSeq_CUPM<T>::Duplicate(Vec v, Vec *y) noexcept
309: {
310:   PetscDeviceContext dctx;

312:   PetscFunctionBegin;
313:   PetscCall(GetHandles_(&dctx));
314:   PetscCall(Duplicate_CUPMBase(v, y, dctx));
315:   PetscFunctionReturn(PETSC_SUCCESS);
316: }

318: // ==========================================================================================
319: // VecSeq_CUPM - Public API - Utility
320: // ==========================================================================================

322: // v->ops->bindtocpu
323: template <device::cupm::DeviceType T>
324: inline PetscErrorCode VecSeq_CUPM<T>::BindToCPU(Vec v, PetscBool usehost) noexcept
325: {
326:   PetscDeviceContext dctx;

328:   PetscFunctionBegin;
329:   PetscCall(GetHandles_(&dctx));
330:   PetscCall(BindToCPU_CUPMBase(v, usehost, dctx));

332:   // REVIEW ME: this absolutely should be some sort of bulk mempcy rather than this mess
333:   VecSetOp_CUPM(dot, VecDot_Seq, Dot);
334:   VecSetOp_CUPM(norm, VecNorm_Seq, Norm);
335:   VecSetOp_CUPM(tdot, VecTDot_Seq, TDot);
336:   VecSetOp_CUPM(mdot, VecMDot_Seq, MDot);
337:   VecSetOp_CUPM(resetarray, VecResetArray_Seq, base_type::template ResetArray<PETSC_MEMTYPE_HOST>);
338:   VecSetOp_CUPM(placearray, VecPlaceArray_Seq, base_type::template PlaceArray<PETSC_MEMTYPE_HOST>);
339:   v->ops->mtdot = v->ops->mtdot_local = VecMTDot_Seq;
340:   VecSetOp_CUPM(max, VecMax_Seq, Max);
341:   VecSetOp_CUPM(min, VecMin_Seq, Min);
342:   VecSetOp_CUPM(setpreallocationcoo, VecSetPreallocationCOO_Seq, SetPreallocationCOO);
343:   VecSetOp_CUPM(setvaluescoo, VecSetValuesCOO_Seq, SetValuesCOO);
344:   PetscFunctionReturn(PETSC_SUCCESS);
345: }

347: // ==========================================================================================
348: // VecSeq_CUPM - Public API - Mutators
349: // ==========================================================================================

351: // v->ops->getlocalvector or v->ops->getlocalvectorread
352: template <device::cupm::DeviceType T>
353: template <PetscMemoryAccessMode access>
354: inline PetscErrorCode VecSeq_CUPM<T>::GetLocalVector(Vec v, Vec w) noexcept
355: {
356:   PetscBool wisseqcupm;

358:   PetscFunctionBegin;
359:   PetscCheckTypeNames(v, VECSEQCUPM(), VECMPICUPM());
360:   PetscCall(PetscObjectTypeCompare(PetscObjectCast(w), VECSEQCUPM(), &wisseqcupm));
361:   if (wisseqcupm) {
362:     if (const auto wseq = VecIMPLCast(w)) {
363:       if (auto &alloced = wseq->array_allocated) {
364:         const auto useit = UseCUPMHostAlloc(util::exchange(w->pinned_memory, PETSC_FALSE));

366:         PetscCall(PetscFree(alloced));
367:       }
368:       wseq->array         = nullptr;
369:       wseq->unplacedarray = nullptr;
370:     }
371:     if (const auto wcu = VecCUPMCast(w)) {
372:       if (auto &device_array = wcu->array_d) {
373:         cupmStream_t stream;

375:         PetscCall(GetHandles_(&stream));
376:         PetscCallCUPM(cupmFreeAsync(device_array, stream));
377:       }
378:       PetscCall(PetscFree(w->spptr /* wcu */));
379:     }
380:   }
381:   if (v->petscnative && wisseqcupm) {
382:     PetscCall(PetscFree(w->data));
383:     w->data          = v->data;
384:     w->offloadmask   = v->offloadmask;
385:     w->pinned_memory = v->pinned_memory;
386:     w->spptr         = v->spptr;
387:     PetscCall(PetscObjectStateIncrease(PetscObjectCast(w)));
388:   } else {
389:     const auto array = &VecIMPLCast(w)->array;

391:     if (access == PETSC_MEMORY_ACCESS_READ) {
392:       PetscCall(VecGetArrayRead(v, const_cast<const PetscScalar **>(array)));
393:     } else {
394:       PetscCall(VecGetArray(v, array));
395:     }
396:     w->offloadmask = PETSC_OFFLOAD_CPU;
397:     if (wisseqcupm) {
398:       PetscDeviceContext dctx;

400:       PetscCall(GetHandles_(&dctx));
401:       PetscCall(DeviceAllocateCheck_(dctx, w));
402:     }
403:   }
404:   PetscFunctionReturn(PETSC_SUCCESS);
405: }

407: // v->ops->restorelocalvector or v->ops->restorelocalvectorread
408: template <device::cupm::DeviceType T>
409: template <PetscMemoryAccessMode access>
410: inline PetscErrorCode VecSeq_CUPM<T>::RestoreLocalVector(Vec v, Vec w) noexcept
411: {
412:   PetscBool wisseqcupm;

414:   PetscFunctionBegin;
415:   PetscCheckTypeNames(v, VECSEQCUPM(), VECMPICUPM());
416:   PetscCall(PetscObjectTypeCompare(PetscObjectCast(w), VECSEQCUPM(), &wisseqcupm));
417:   if (v->petscnative && wisseqcupm) {
418:     // the assignments to nullptr are __critical__, as w may persist after this call returns
419:     // and shouldn't share data with v!
420:     v->pinned_memory = w->pinned_memory;
421:     v->offloadmask   = util::exchange(w->offloadmask, PETSC_OFFLOAD_UNALLOCATED);
422:     v->data          = util::exchange(w->data, nullptr);
423:     v->spptr         = util::exchange(w->spptr, nullptr);
424:   } else {
425:     const auto array = &VecIMPLCast(w)->array;

427:     if (access == PETSC_MEMORY_ACCESS_READ) {
428:       PetscCall(VecRestoreArrayRead(v, const_cast<const PetscScalar **>(array)));
429:     } else {
430:       PetscCall(VecRestoreArray(v, array));
431:     }
432:     if (w->spptr && wisseqcupm) {
433:       cupmStream_t stream;

435:       PetscCall(GetHandles_(&stream));
436:       PetscCallCUPM(cupmFreeAsync(VecCUPMCast(w)->array_d, stream));
437:       PetscCall(PetscFree(w->spptr));
438:     }
439:   }
440:   PetscFunctionReturn(PETSC_SUCCESS);
441: }

443: // ==========================================================================================
444: // VecSeq_CUPM - Public API - Compute Methods
445: // ==========================================================================================

447: // VecAYPXAsync_Private
448: template <device::cupm::DeviceType T>
449: inline PetscErrorCode VecSeq_CUPM<T>::AYPXAsync(Vec yin, PetscScalar alpha, Vec xin, PetscDeviceContext dctx) noexcept
450: {
451:   const auto n = static_cast<cupmBlasInt_t>(yin->map->n);
452:   PetscBool  xiscupm;

454:   PetscFunctionBegin;
455:   PetscCall(PetscObjectTypeCompareAny(PetscObjectCast(xin), &xiscupm, VECSEQCUPM(), VECMPICUPM(), ""));
456:   if (!xiscupm) {
457:     PetscCall(VecAYPX_Seq(yin, alpha, xin));
458:     PetscFunctionReturn(PETSC_SUCCESS);
459:   }
460:   PetscCall(PetscDeviceContextGetOptionalNullContext_Internal(&dctx));
461:   if (alpha == PetscScalar(0.0)) {
462:     cupmStream_t stream;

464:     PetscCall(GetHandlesFrom_(dctx, &stream));
465:     PetscCall(PetscLogGpuTimeBegin());
466:     PetscCall(PetscCUPMMemcpyAsync(DeviceArrayWrite(dctx, yin).data(), DeviceArrayRead(dctx, xin).data(), n, cupmMemcpyDeviceToDevice, stream));
467:     PetscCall(PetscLogGpuTimeEnd());
468:   } else if (n) {
469:     const auto       alphaIsOne = alpha == PetscScalar(1.0);
470:     const auto       calpha     = cupmScalarPtrCast(&alpha);
471:     cupmBlasHandle_t cupmBlasHandle;

473:     PetscCall(GetHandlesFrom_(dctx, &cupmBlasHandle));
474:     {
475:       const auto yptr = DeviceArrayReadWrite(dctx, yin);
476:       const auto xptr = DeviceArrayRead(dctx, xin);

478:       PetscCall(PetscLogGpuTimeBegin());
479:       if (alphaIsOne) {
480:         PetscCallCUPMBLAS(cupmBlasXaxpy(cupmBlasHandle, n, calpha, xptr.cupmdata(), 1, yptr.cupmdata(), 1));
481:       } else {
482:         const auto one = cupmScalarCast(1.0);

484:         PetscCallCUPMBLAS(cupmBlasXscal(cupmBlasHandle, n, calpha, yptr.cupmdata(), 1));
485:         PetscCallCUPMBLAS(cupmBlasXaxpy(cupmBlasHandle, n, &one, xptr.cupmdata(), 1, yptr.cupmdata(), 1));
486:       }
487:       PetscCall(PetscLogGpuTimeEnd());
488:     }
489:     PetscCall(PetscLogGpuFlops((alphaIsOne ? 1 : 2) * n));
490:   }
491:   if (n > 0) PetscCall(PetscDeviceContextSynchronizeIfWithBarrier_Internal(dctx));
492:   PetscFunctionReturn(PETSC_SUCCESS);
493: }

495: // v->ops->aypx
496: template <device::cupm::DeviceType T>
497: inline PetscErrorCode VecSeq_CUPM<T>::AYPX(Vec yin, PetscScalar alpha, Vec xin) noexcept
498: {
499:   PetscFunctionBegin;
500:   PetscCall(AYPXAsync(yin, alpha, xin, nullptr));
501:   PetscFunctionReturn(PETSC_SUCCESS);
502: }

504: // VecAXPYAsync_Private
505: template <device::cupm::DeviceType T>
506: inline PetscErrorCode VecSeq_CUPM<T>::AXPYAsync(Vec yin, PetscScalar alpha, Vec xin, PetscDeviceContext dctx) noexcept
507: {
508:   PetscBool xiscupm;

510:   PetscFunctionBegin;
511:   PetscCall(PetscObjectTypeCompareAny(PetscObjectCast(xin), &xiscupm, VECSEQCUPM(), VECMPICUPM(), ""));
512:   if (xiscupm) {
513:     const auto       n = static_cast<cupmBlasInt_t>(yin->map->n);
514:     cupmBlasHandle_t cupmBlasHandle;

516:     PetscCall(PetscDeviceContextGetOptionalNullContext_Internal(&dctx));
517:     PetscCall(GetHandlesFrom_(dctx, &cupmBlasHandle));
518:     PetscCall(PetscLogGpuTimeBegin());
519:     PetscCallCUPMBLAS(cupmBlasXaxpy(cupmBlasHandle, n, cupmScalarPtrCast(&alpha), DeviceArrayRead(dctx, xin), 1, DeviceArrayReadWrite(dctx, yin), 1));
520:     PetscCall(PetscLogGpuTimeEnd());
521:     PetscCall(PetscLogGpuFlops(2 * n));
522:     if (n > 0) PetscCall(PetscDeviceContextSynchronizeIfWithBarrier_Internal(dctx));
523:   } else {
524:     PetscCall(VecAXPY_Seq(yin, alpha, xin));
525:   }
526:   PetscFunctionReturn(PETSC_SUCCESS);
527: }

529: // v->ops->axpy
530: template <device::cupm::DeviceType T>
531: inline PetscErrorCode VecSeq_CUPM<T>::AXPY(Vec yin, PetscScalar alpha, Vec xin) noexcept
532: {
533:   PetscFunctionBegin;
534:   PetscCall(AXPYAsync(yin, alpha, xin, nullptr));
535:   PetscFunctionReturn(PETSC_SUCCESS);
536: }

538: namespace detail
539: {

541: struct divides {
542:   PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &lhs, const PetscScalar &rhs) const noexcept { return rhs == PetscScalar{0.0} ? (lhs == PetscScalar{0.0} ? PetscScalar{1.0} : rhs) : lhs / rhs; }
543: };

545: } // namespace detail

547: // VecPointwiseDivideAsync_Private
548: template <device::cupm::DeviceType T>
549: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseDivideAsync(Vec wout, Vec xin, Vec yin, PetscDeviceContext dctx) noexcept
550: {
551:   PetscFunctionBegin;
552:   PetscCall(PointwiseBinaryDispatch_(VecPointwiseDivide_Seq, detail::divides{}, wout, xin, yin, dctx));
553:   PetscFunctionReturn(PETSC_SUCCESS);
554: }

556: // v->ops->pointwisedivide
557: template <device::cupm::DeviceType T>
558: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseDivide(Vec wout, Vec xin, Vec yin) noexcept
559: {
560:   PetscFunctionBegin;
561:   PetscCall(PointwiseDivideAsync(wout, xin, yin, nullptr));
562:   PetscFunctionReturn(PETSC_SUCCESS);
563: }

565: namespace detail
566: {

568: struct multiplies {
569:   PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &lhs, const PetscScalar &rhs) const noexcept { return lhs * rhs; }
570: };

572: } // namespace detail

574: // VecPointwiseMultAsync_Private
575: template <device::cupm::DeviceType T>
576: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseMultAsync(Vec wout, Vec xin, Vec yin, PetscDeviceContext dctx) noexcept
577: {
578:   PetscFunctionBegin;
579:   PetscCall(PointwiseBinaryDispatch_(VecPointwiseMult_Seq, detail::multiplies{}, wout, xin, yin, dctx));
580:   PetscFunctionReturn(PETSC_SUCCESS);
581: }

583: // v->ops->pointwisemult
584: template <device::cupm::DeviceType T>
585: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseMult(Vec wout, Vec xin, Vec yin) noexcept
586: {
587:   PetscFunctionBegin;
588:   PetscCall(PointwiseMultAsync(wout, xin, yin, nullptr));
589:   PetscFunctionReturn(PETSC_SUCCESS);
590: }

592: namespace detail
593: {

595: struct MaximumRealPart {
596:   PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &lhs, const PetscScalar &rhs) const noexcept { return thrust::maximum<PetscReal>{}(PetscRealPart(lhs), PetscRealPart(rhs)); }
597: };

599: } // namespace detail

601: // VecPointwiseMaxAsync_Private
602: template <device::cupm::DeviceType T>
603: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseMaxAsync(Vec wout, Vec xin, Vec yin, PetscDeviceContext dctx) noexcept
604: {
605:   PetscFunctionBegin;
606:   PetscCall(PointwiseBinaryDispatch_(VecPointwiseMax_Seq, detail::MaximumRealPart{}, wout, xin, yin, dctx));
607:   PetscFunctionReturn(PETSC_SUCCESS);
608: }

610: // v->ops->pointwisemax
611: template <device::cupm::DeviceType T>
612: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseMax(Vec wout, Vec xin, Vec yin) noexcept
613: {
614:   PetscFunctionBegin;
615:   PetscCall(PointwiseMaxAsync(wout, xin, yin, nullptr));
616:   PetscFunctionReturn(PETSC_SUCCESS);
617: }

619: namespace detail
620: {

622: struct MaximumAbsoluteValue {
623:   PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &lhs, const PetscScalar &rhs) const noexcept { return thrust::maximum<PetscReal>{}(PetscAbsScalar(lhs), PetscAbsScalar(rhs)); }
624: };

626: } // namespace detail

628: // VecPointwiseMaxAbsAsync_Private
629: template <device::cupm::DeviceType T>
630: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseMaxAbsAsync(Vec wout, Vec xin, Vec yin, PetscDeviceContext dctx) noexcept
631: {
632:   PetscFunctionBegin;
633:   PetscCall(PointwiseBinaryDispatch_(VecPointwiseMaxAbs_Seq, detail::MaximumAbsoluteValue{}, wout, xin, yin, dctx));
634:   PetscFunctionReturn(PETSC_SUCCESS);
635: }

637: // v->ops->pointwisemaxabs
638: template <device::cupm::DeviceType T>
639: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseMaxAbs(Vec wout, Vec xin, Vec yin) noexcept
640: {
641:   PetscFunctionBegin;
642:   PetscCall(PointwiseMaxAbsAsync(wout, xin, yin, nullptr));
643:   PetscFunctionReturn(PETSC_SUCCESS);
644: }

646: namespace detail
647: {

649: struct MinimumRealPart {
650:   PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &lhs, const PetscScalar &rhs) const noexcept { return thrust::minimum<PetscReal>{}(PetscRealPart(lhs), PetscRealPart(rhs)); }
651: };

653: } // namespace detail

655: // VecPointwiseMinAsync_Private
656: template <device::cupm::DeviceType T>
657: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseMinAsync(Vec wout, Vec xin, Vec yin, PetscDeviceContext dctx) noexcept
658: {
659:   PetscFunctionBegin;
660:   PetscCall(PointwiseBinaryDispatch_(VecPointwiseMin_Seq, detail::MinimumRealPart{}, wout, xin, yin, dctx));
661:   PetscFunctionReturn(PETSC_SUCCESS);
662: }

664: // v->ops->pointwisemin
665: template <device::cupm::DeviceType T>
666: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseMin(Vec wout, Vec xin, Vec yin) noexcept
667: {
668:   PetscFunctionBegin;
669:   PetscCall(PointwiseMinAsync(wout, xin, yin, nullptr));
670:   PetscFunctionReturn(PETSC_SUCCESS);
671: }

673: namespace detail
674: {

676: struct Reciprocal {
677:   PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &s) const noexcept
678:   {
679:     // yes all of this verbosity is needed because sometimes PetscScalar is a thrust::complex
680:     // and then it matters whether we do s ? true : false vs s == 0, as well as whether we wrap
681:     // everything in PetscScalar...
682:     return s == PetscScalar{0.0} ? s : PetscScalar{1.0} / s;
683:   }
684: };

686: } // namespace detail

688: // VecReciprocalAsync_Private
689: template <device::cupm::DeviceType T>
690: inline PetscErrorCode VecSeq_CUPM<T>::ReciprocalAsync(Vec xin, PetscDeviceContext dctx) noexcept
691: {
692:   PetscFunctionBegin;
693:   PetscCall(PointwiseUnary_(detail::Reciprocal{}, xin, nullptr, dctx));
694:   PetscFunctionReturn(PETSC_SUCCESS);
695: }

697: // v->ops->reciprocal
698: template <device::cupm::DeviceType T>
699: inline PetscErrorCode VecSeq_CUPM<T>::Reciprocal(Vec xin) noexcept
700: {
701:   PetscFunctionBegin;
702:   PetscCall(ReciprocalAsync(xin, nullptr));
703:   PetscFunctionReturn(PETSC_SUCCESS);
704: }

706: namespace detail
707: {

709: struct AbsoluteValue {
710:   PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &s) const noexcept { return PetscAbsScalar(s); }
711: };

713: } // namespace detail

715: // VecAbsAsync_Private
716: template <device::cupm::DeviceType T>
717: inline PetscErrorCode VecSeq_CUPM<T>::AbsAsync(Vec xin, PetscDeviceContext dctx) noexcept
718: {
719:   PetscFunctionBegin;
720:   PetscCall(PointwiseUnary_(detail::AbsoluteValue{}, xin, nullptr, dctx));
721:   PetscFunctionReturn(PETSC_SUCCESS);
722: }

724: // v->ops->abs
725: template <device::cupm::DeviceType T>
726: inline PetscErrorCode VecSeq_CUPM<T>::Abs(Vec xin) noexcept
727: {
728:   PetscFunctionBegin;
729:   PetscCall(AbsAsync(xin, nullptr));
730:   PetscFunctionReturn(PETSC_SUCCESS);
731: }

733: namespace detail
734: {

736: struct SignZeroToSignedUnit {
737:   PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &s) const noexcept { return VecSignZeroToSignedUnit_Private(PetscRealPart(s)); }
738: };

740: struct SignZeroToZero {
741:   PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &s) const noexcept { return VecSignZeroToZero_Private(PetscRealPart(s)); }
742: };

744: struct SignZeroToSignedZero {
745:   PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &s) const noexcept { return VecSignZeroToSignedZero_Private(PetscRealPart(s)); }
746: };

748: } // namespace detail

750: // VecPointwiseSignAsync_Private
751: template <device::cupm::DeviceType T>
752: inline PetscErrorCode VecSeq_CUPM<T>::PointwiseSignAsync(Vec yout, Vec xin, VecSignMode sign_type, PetscDeviceContext dctx) noexcept
753: {
754:   PetscFunctionBegin;
755:   switch (sign_type) {
756:   case VEC_SIGN_ZERO_TO_ZERO:
757:     PetscCall(PointwiseUnary_(detail::SignZeroToZero{}, xin, yout, dctx));
758:     break;
759:   case VEC_SIGN_ZERO_TO_SIGNED_ZERO:
760:     PetscCall(PointwiseUnary_(detail::SignZeroToSignedZero{}, xin, yout, dctx));
761:     break;
762:   case VEC_SIGN_ZERO_TO_SIGNED_UNIT:
763:     PetscCall(PointwiseUnary_(detail::SignZeroToSignedUnit{}, xin, yout, dctx));
764:     break;
765:   }
766:   PetscFunctionReturn(PETSC_SUCCESS);
767: }

769: namespace detail
770: {

772: struct SquareRootAbsoluteValue {
773:   PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &s) const noexcept { return PetscSqrtReal(PetscAbsScalar(s)); }
774: };

776: } // namespace detail

778: // VecSqrtAbsAsync_Private
779: template <device::cupm::DeviceType T>
780: inline PetscErrorCode VecSeq_CUPM<T>::SqrtAbsAsync(Vec xin, PetscDeviceContext dctx) noexcept
781: {
782:   PetscFunctionBegin;
783:   PetscCall(PointwiseUnary_(detail::SquareRootAbsoluteValue{}, xin, nullptr, dctx));
784:   PetscFunctionReturn(PETSC_SUCCESS);
785: }

787: // v->ops->sqrt
788: template <device::cupm::DeviceType T>
789: inline PetscErrorCode VecSeq_CUPM<T>::SqrtAbs(Vec xin) noexcept
790: {
791:   PetscFunctionBegin;
792:   PetscCall(SqrtAbsAsync(xin, nullptr));
793:   PetscFunctionReturn(PETSC_SUCCESS);
794: }

796: namespace detail
797: {

799: struct Exponent {
800:   PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &s) const noexcept { return PetscExpScalar(s); }
801: };

803: } // namespace detail

805: // VecExpAsync_Private
806: template <device::cupm::DeviceType T>
807: inline PetscErrorCode VecSeq_CUPM<T>::ExpAsync(Vec xin, PetscDeviceContext dctx) noexcept
808: {
809:   PetscFunctionBegin;
810:   PetscCall(PointwiseUnary_(detail::Exponent{}, xin, nullptr, dctx));
811:   PetscFunctionReturn(PETSC_SUCCESS);
812: }

814: // v->ops->exp
815: template <device::cupm::DeviceType T>
816: inline PetscErrorCode VecSeq_CUPM<T>::Exp(Vec xin) noexcept
817: {
818:   PetscFunctionBegin;
819:   PetscCall(ExpAsync(xin, nullptr));
820:   PetscFunctionReturn(PETSC_SUCCESS);
821: }

823: namespace detail
824: {

826: struct Logarithm {
827:   PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &s) const noexcept { return PetscLogScalar(s); }
828: };

830: } // namespace detail

832: // VecLogAsync_Private
833: template <device::cupm::DeviceType T>
834: inline PetscErrorCode VecSeq_CUPM<T>::LogAsync(Vec xin, PetscDeviceContext dctx) noexcept
835: {
836:   PetscFunctionBegin;
837:   PetscCall(PointwiseUnary_(detail::Logarithm{}, xin, nullptr, dctx));
838:   PetscFunctionReturn(PETSC_SUCCESS);
839: }

841: // v->ops->log
842: template <device::cupm::DeviceType T>
843: inline PetscErrorCode VecSeq_CUPM<T>::Log(Vec xin) noexcept
844: {
845:   PetscFunctionBegin;
846:   PetscCall(LogAsync(xin, nullptr));
847:   PetscFunctionReturn(PETSC_SUCCESS);
848: }

850: // v->ops->waxpy
851: template <device::cupm::DeviceType T>
852: inline PetscErrorCode VecSeq_CUPM<T>::WAXPYAsync(Vec win, PetscScalar alpha, Vec xin, Vec yin, PetscDeviceContext dctx) noexcept
853: {
854:   PetscBool xiscupm, yiscupm;

856:   PetscFunctionBegin;
857:   PetscCall(PetscObjectTypeCompareAny(PetscObjectCast(xin), &xiscupm, VECSEQCUPM(), VECMPICUPM(), ""));
858:   PetscCall(PetscObjectTypeCompareAny(PetscObjectCast(yin), &yiscupm, VECSEQCUPM(), VECMPICUPM(), ""));
859:   if (!xiscupm || !yiscupm) {
860:     PetscCall(VecWAXPY_Seq(win, alpha, xin, yin));
861:     PetscFunctionReturn(PETSC_SUCCESS);
862:   }
863:   PetscCall(PetscDeviceContextGetOptionalNullContext_Internal(&dctx));
864:   if (alpha == PetscScalar(0.0)) {
865:     PetscCall(CopyAsync(yin, win, dctx));
866:   } else if (const auto n = static_cast<cupmBlasInt_t>(win->map->n)) {
867:     cupmBlasHandle_t cupmBlasHandle;
868:     cupmStream_t     stream;
869:     PetscBool        xiscupm, yiscupm;

871:     PetscCall(PetscObjectTypeCompareAny(PetscObjectCast(xin), &xiscupm, VECSEQCUPM(), VECMPICUPM(), ""));
872:     PetscCall(PetscObjectTypeCompareAny(PetscObjectCast(yin), &yiscupm, VECSEQCUPM(), VECMPICUPM(), ""));
873:     if (!xiscupm || !yiscupm) {
874:       PetscCall(VecWAXPY_Seq(win, alpha, xin, yin));
875:       PetscFunctionReturn(PETSC_SUCCESS);
876:     }
877:     PetscCall(GetHandlesFrom_(dctx, &cupmBlasHandle, NULL, &stream));
878:     {
879:       const auto wptr = DeviceArrayWrite(dctx, win);

881:       PetscCall(PetscLogGpuTimeBegin());
882:       PetscCall(PetscCUPMMemcpyAsync(wptr.data(), DeviceArrayRead(dctx, yin).data(), n, cupmMemcpyDeviceToDevice, stream, true));
883:       PetscCallCUPMBLAS(cupmBlasXaxpy(cupmBlasHandle, n, cupmScalarPtrCast(&alpha), DeviceArrayRead(dctx, xin), 1, wptr.cupmdata(), 1));
884:       PetscCall(PetscLogGpuTimeEnd());
885:     }
886:     PetscCall(PetscLogGpuFlops(2 * n));
887:     PetscCall(PetscDeviceContextSynchronizeIfWithBarrier_Internal(dctx));
888:   }
889:   PetscFunctionReturn(PETSC_SUCCESS);
890: }

892: // v->ops->waxpy
893: template <device::cupm::DeviceType T>
894: inline PetscErrorCode VecSeq_CUPM<T>::WAXPY(Vec win, PetscScalar alpha, Vec xin, Vec yin) noexcept
895: {
896:   PetscFunctionBegin;
897:   PetscCall(WAXPYAsync(win, alpha, xin, yin, nullptr));
898:   PetscFunctionReturn(PETSC_SUCCESS);
899: }

901: namespace detail
902: {
903: struct stdbasis_functor {
904:   PetscInt _i;

906:   stdbasis_functor(PetscInt i) : _i(i) { }

908:   PETSC_NODISCARD PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &, const PetscScalar &i) const noexcept { return i == _i ? 1.0 : 0.0; }
909: };
910: } // namespace detail

912: // v->ops->setstdbasis
913: template <device::cupm::DeviceType T>
914: inline PetscErrorCode VecSeq_CUPM<T>::SetStdBasisAsync(Vec s, PetscInt i, PetscDeviceContext dctx) noexcept
915: {
916:   cupmStream_t stream;

918:   PetscFunctionBegin;
919:   PetscCall(PetscDeviceContextGetOptionalNullContext_Internal(&dctx));
920:   PetscCall(GetHandlesFrom_(dctx, &stream));
921:   {
922:     const auto nl   = s->map->n;
923:     const auto st   = s->map->rstart;
924:     auto       nit  = thrust::make_counting_iterator(PetscInt{st});
925:     const auto sptr = thrust::device_pointer_cast(DeviceArrayWrite(dctx, s).data());
926:     PetscCall(PetscLogGpuTimeBegin());
927:     // clang-format off
928:     PetscCallThrust(
929:       THRUST_CALL(
930:         thrust::transform,
931:         stream,
932:         sptr,
933:         sptr + nl,
934:         nit,
935:         sptr,
936:         detail::stdbasis_functor(i)
937:       )
938:     );
939:     // clang-format on
940:     PetscCall(PetscLogGpuFlops(nl));
941:     PetscCall(PetscDeviceContextSynchronizeIfWithBarrier_Internal(dctx));
942:   }
943:   PetscFunctionReturn(PETSC_SUCCESS);
944: }

946: // v->ops->setstdbasis
947: template <device::cupm::DeviceType T>
948: inline PetscErrorCode VecSeq_CUPM<T>::SetStdBasis(Vec s, PetscInt i) noexcept
949: {
950:   PetscFunctionBegin;
951:   PetscCall(SetStdBasisAsync(s, i, nullptr));
952:   PetscFunctionReturn(PETSC_SUCCESS);
953: }

955: namespace kernels
956: {

958: template <typename... Args>
959: PETSC_KERNEL_DECL static void MAXPY_kernel(const PetscInt size, PetscScalar *PETSC_RESTRICT xptr, const PetscScalar *PETSC_RESTRICT aptr, Args... yptr)
960: {
961:   constexpr int      N        = sizeof...(Args);
962:   const auto         tx       = threadIdx.x;
963:   const PetscScalar *yptr_p[] = {yptr...};

965:   PETSC_SHAREDMEM_DECL PetscScalar aptr_shmem[N];

967:   // load a to shared memory
968:   if (tx < N) aptr_shmem[tx] = aptr[tx];
969:   __syncthreads();

971:   ::Petsc::device::cupm::kernels::util::grid_stride_1D(size, [&](PetscInt i) {
972:   // these may look the same but give different results!
973: #if 0
974:     PetscScalar sum = 0.0;

976:   #pragma unroll
977:     for (auto j = 0; j < N; ++j) sum += aptr_shmem[j]*yptr_p[j][i];
978:     xptr[i] += sum;
979: #else
980:     auto sum = xptr[i];

982:   #pragma unroll
983:     for (auto j = 0; j < N; ++j) sum += aptr_shmem[j] * yptr_p[j][i];
984:     xptr[i] = sum;
985: #endif
986:   });
987:   return;
988: }

990: } // namespace kernels

992: namespace detail
993: {

995: // a helper-struct to gobble the size_t input, it is used with template parameter pack
996: // expansion such that
997: // typename repeat_type...
998: // expands to
999: // MyType, MyType, MyType, ... [repeated sizeof...(IdxParamPack) times]
1000: template <typename T, std::size_t>
1001: struct repeat_type {
1002:   using type = T;
1003: };

1005: } // namespace detail

1007: template <device::cupm::DeviceType T>
1008: template <std::size_t... Idx>
1009: inline PetscErrorCode VecSeq_CUPM<T>::MAXPY_kernel_dispatch_(PetscDeviceContext dctx, cupmStream_t stream, PetscScalar *xptr, const PetscScalar *aptr, const Vec *yin, PetscInt size, util::index_sequence<Idx...>) noexcept
1010: {
1011:   PetscFunctionBegin;
1012:   // clang-format off
1013:   PetscCall(
1014:     PetscCUPMLaunchKernel1D(
1015:       size, 0, stream,
1016:       kernels::MAXPY_kernel<typename detail::repeat_type<const PetscScalar *, Idx>::type...>,
1017:       size, xptr, aptr, DeviceArrayRead(dctx, yin[Idx]).data()...
1018:     )
1019:   );
1020:   // clang-format on
1021:   PetscFunctionReturn(PETSC_SUCCESS);
1022: }

1024: template <device::cupm::DeviceType T>
1025: template <int N>
1026: inline PetscErrorCode VecSeq_CUPM<T>::MAXPY_kernel_dispatch_(PetscDeviceContext dctx, cupmStream_t stream, PetscScalar *xptr, const PetscScalar *aptr, const Vec *yin, PetscInt size, PetscInt &yidx) noexcept
1027: {
1028:   PetscFunctionBegin;
1029:   PetscCall(MAXPY_kernel_dispatch_(dctx, stream, xptr, aptr + yidx, yin + yidx, size, util::make_index_sequence<N>{}));
1030:   yidx += N;
1031:   PetscFunctionReturn(PETSC_SUCCESS);
1032: }

1034: // VecMAXPYAsync_Private
1035: template <device::cupm::DeviceType T>
1036: inline PetscErrorCode VecSeq_CUPM<T>::MAXPYAsync(Vec xin, PetscInt nv, const PetscScalar *alpha, Vec *yin, PetscDeviceContext dctx) noexcept
1037: {
1038:   const auto   n = xin->map->n;
1039:   cupmStream_t stream;
1040:   PetscBool    yiscupm = PETSC_TRUE;

1042:   PetscFunctionBegin;
1043:   for (PetscInt i = 0; i < nv && yiscupm; i++) PetscCall(PetscObjectTypeCompareAny(PetscObjectCast(yin[i]), &yiscupm, VECSEQCUPM(), VECMPICUPM(), ""));
1044:   if (!yiscupm) {
1045:     PetscCall(VecMAXPY_Seq(xin, nv, alpha, yin));
1046:     PetscFunctionReturn(PETSC_SUCCESS);
1047:   }
1048:   PetscCall(PetscDeviceContextGetOptionalNullContext_Internal(&dctx));
1049:   PetscCall(GetHandlesFrom_(dctx, &stream));
1050:   {
1051:     const auto   xptr    = DeviceArrayReadWrite(dctx, xin);
1052:     PetscScalar *d_alpha = nullptr;
1053:     PetscInt     yidx    = 0;

1055:     // placement of early-return is deliberate, we would like to capture the
1056:     // DeviceArrayReadWrite() call (which calls PetscObjectStateIncreate()) before we bail
1057:     if (!n || !nv) PetscFunctionReturn(PETSC_SUCCESS);
1058:     PetscCall(PetscDeviceMalloc(dctx, PETSC_MEMTYPE_CUPM(), nv, &d_alpha));
1059:     PetscCall(PetscCUPMMemcpyAsync(d_alpha, alpha, nv, cupmMemcpyHostToDevice, stream));
1060:     PetscCall(PetscLogGpuTimeBegin());
1061:     do {
1062:       switch (nv - yidx) {
1063:       case 7:
1064:         PetscCall(MAXPY_kernel_dispatch_<7>(dctx, stream, xptr.data(), d_alpha, yin, n, yidx));
1065:         break;
1066:       case 6:
1067:         PetscCall(MAXPY_kernel_dispatch_<6>(dctx, stream, xptr.data(), d_alpha, yin, n, yidx));
1068:         break;
1069:       case 5:
1070:         PetscCall(MAXPY_kernel_dispatch_<5>(dctx, stream, xptr.data(), d_alpha, yin, n, yidx));
1071:         break;
1072:       case 4:
1073:         PetscCall(MAXPY_kernel_dispatch_<4>(dctx, stream, xptr.data(), d_alpha, yin, n, yidx));
1074:         break;
1075:       case 3:
1076:         PetscCall(MAXPY_kernel_dispatch_<3>(dctx, stream, xptr.data(), d_alpha, yin, n, yidx));
1077:         break;
1078:       case 2:
1079:         PetscCall(MAXPY_kernel_dispatch_<2>(dctx, stream, xptr.data(), d_alpha, yin, n, yidx));
1080:         break;
1081:       case 1:
1082:         PetscCall(MAXPY_kernel_dispatch_<1>(dctx, stream, xptr.data(), d_alpha, yin, n, yidx));
1083:         break;
1084:       default: // 8 or more
1085:         PetscCall(MAXPY_kernel_dispatch_<8>(dctx, stream, xptr.data(), d_alpha, yin, n, yidx));
1086:         break;
1087:       }
1088:     } while (yidx < nv);
1089:     PetscCall(PetscLogGpuTimeEnd());
1090:     PetscCall(PetscDeviceFree(dctx, d_alpha));
1091:     PetscCall(PetscDeviceContextSynchronizeIfWithBarrier_Internal(dctx));
1092:   }
1093:   PetscCall(PetscLogGpuFlops(nv * 2 * n));
1094:   PetscFunctionReturn(PETSC_SUCCESS);
1095: }

1097: // v->ops->maxpy
1098: template <device::cupm::DeviceType T>
1099: inline PetscErrorCode VecSeq_CUPM<T>::MAXPY(Vec xin, PetscInt nv, const PetscScalar *alpha, Vec *yin) noexcept
1100: {
1101:   PetscFunctionBegin;
1102:   PetscCall(MAXPYAsync(xin, nv, alpha, yin, nullptr));
1103:   PetscFunctionReturn(PETSC_SUCCESS);
1104: }

1106: template <device::cupm::DeviceType T>
1107: inline PetscErrorCode VecSeq_CUPM<T>::Dot(Vec xin, Vec yin, PetscScalar *z) noexcept
1108: {
1109:   PetscBool yiscupm;

1111:   PetscFunctionBegin;
1112:   PetscCall(PetscObjectTypeCompareAny(PetscObjectCast(yin), &yiscupm, VECSEQCUPM(), VECMPICUPM(), ""));
1113:   if (!yiscupm) {
1114:     PetscCall(VecDot_Seq(xin, yin, z));
1115:     PetscFunctionReturn(PETSC_SUCCESS);
1116:   }
1117:   if (const auto n = static_cast<cupmBlasInt_t>(xin->map->n)) {
1118:     PetscDeviceContext dctx;
1119:     cupmBlasHandle_t   cupmBlasHandle;

1121:     PetscCall(GetHandles_(&dctx, &cupmBlasHandle));
1122:     // arguments y, x are reversed because BLAS complex conjugates the first argument, PETSc the
1123:     // second
1124:     PetscCall(PetscLogGpuTimeBegin());
1125:     PetscCallCUPMBLAS(cupmBlasXdot(cupmBlasHandle, n, DeviceArrayRead(dctx, yin), 1, DeviceArrayRead(dctx, xin), 1, cupmScalarPtrCast(z)));
1126:     PetscCall(PetscLogGpuTimeEnd());
1127:     PetscCall(PetscLogGpuFlops(2 * n - 1));
1128:   } else {
1129:     *z = 0.0;
1130:   }
1131:   PetscFunctionReturn(PETSC_SUCCESS);
1132: }

1134: #define MDOT_WORKGROUP_NUM  128
1135: #define MDOT_WORKGROUP_SIZE MDOT_WORKGROUP_NUM

1137: namespace kernels
1138: {

1140: PETSC_DEVICE_INLINE_DECL static PetscInt EntriesPerGroup(const PetscInt size) noexcept
1141: {
1142:   const auto group_entries = (size - 1) / gridDim.x + 1;
1143:   // for very small vectors, a group should still do some work
1144:   return group_entries ? group_entries : 1;
1145: }

1147: template <typename... ConstPetscScalarPointer>
1148: PETSC_KERNEL_DECL static void MDot_kernel(const PetscScalar *PETSC_RESTRICT x, const PetscInt size, PetscScalar *PETSC_RESTRICT results, ConstPetscScalarPointer... y)
1149: {
1150:   constexpr int      N        = sizeof...(ConstPetscScalarPointer);
1151:   const PetscScalar *ylocal[] = {y...};
1152:   PetscScalar        sumlocal[N];

1154:   PETSC_SHAREDMEM_DECL PetscScalar shmem[N * MDOT_WORKGROUP_SIZE];

1156:   // HIP -- for whatever reason -- has threadIdx, blockIdx, blockDim, and gridDim as separate
1157:   // types, so each of these go on separate lines...
1158:   const auto tx       = threadIdx.x;
1159:   const auto bx       = blockIdx.x;
1160:   const auto bdx      = blockDim.x;
1161:   const auto gdx      = gridDim.x;
1162:   const auto worksize = EntriesPerGroup(size);
1163:   const auto begin    = tx + bx * worksize;
1164:   const auto end      = min((bx + 1) * worksize, size);

1166: #pragma unroll
1167:   for (auto i = 0; i < N; ++i) sumlocal[i] = 0;

1169:   for (auto i = begin; i < end; i += bdx) {
1170:     const auto xi = x[i]; // load only once from global memory!

1172: #pragma unroll
1173:     for (auto j = 0; j < N; ++j) sumlocal[j] += ylocal[j][i] * xi;
1174:   }

1176: #pragma unroll
1177:   for (auto i = 0; i < N; ++i) shmem[tx + i * MDOT_WORKGROUP_SIZE] = sumlocal[i];

1179:   // parallel reduction
1180:   for (auto stride = bdx / 2; stride > 0; stride /= 2) {
1181:     __syncthreads();
1182:     if (tx < stride) {
1183: #pragma unroll
1184:       for (auto i = 0; i < N; ++i) shmem[tx + i * MDOT_WORKGROUP_SIZE] += shmem[tx + stride + i * MDOT_WORKGROUP_SIZE];
1185:     }
1186:   }
1187:   // bottom N threads per block write to global memory
1188:   // REVIEW ME: I am ~pretty~ sure we don't need another __syncthreads() here since each thread
1189:   // writes to the same sections in the above loop that it is about to read from below, but
1190:   // running this under the racecheck tool of compute-sanitizer reports a write-after-write hazard.
1191:   __syncthreads();
1192:   if (tx < N) results[bx + tx * gdx] = shmem[tx * MDOT_WORKGROUP_SIZE];
1193:   return;
1194: }

1196: namespace
1197: {

1199: PETSC_KERNEL_DECL void sum_kernel(const PetscInt size, PetscScalar *PETSC_RESTRICT results)
1200: {
1201:   int         local_i = 0;
1202:   PetscScalar local_results[8];

1204:   // each thread sums up MDOT_WORKGROUP_NUM entries of the result, storing it in a local buffer
1205:   //
1206:   // *-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*
1207:   // | | | | | | | | | | | | | | | | | | | | | | | | | | | | | | ...
1208:   // *-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*-*
1209:   //  |  ______________________________________________________/
1210:   //  | /            <- MDOT_WORKGROUP_NUM ->
1211:   //  |/
1212:   //  +
1213:   //  v
1214:   // *-*-*
1215:   // | | | ...
1216:   // *-*-*
1217:   //
1218:   ::Petsc::device::cupm::kernels::util::grid_stride_1D(size, [&](PetscInt i) {
1219:     PetscScalar z_sum = 0;

1221:     for (auto j = i * MDOT_WORKGROUP_SIZE; j < (i + 1) * MDOT_WORKGROUP_SIZE; ++j) z_sum += results[j];
1222:     local_results[local_i++] = z_sum;
1223:   });
1224:   // if we needed more than 1 workgroup to handle the vector we should sync since other threads
1225:   // may currently be reading from results
1226:   if (size >= MDOT_WORKGROUP_SIZE) __syncthreads();
1227:   // Local buffer is now written to global memory
1228:   ::Petsc::device::cupm::kernels::util::grid_stride_1D(size, [&](PetscInt i) {
1229:     const auto j = --local_i;

1231:     if (j >= 0) results[i] = local_results[j];
1232:   });
1233:   return;
1234: }

1236: } // namespace

1238: #if PetscDefined(USING_HCC)
1239: namespace do_not_use
1240: {

1242: inline void silence_warning_function_sum_kernel_is_not_needed_and_will_not_be_emitted()
1243: {
1244:   (void)sum_kernel;
1245: }

1247: } // namespace do_not_use
1248: #endif

1250: } // namespace kernels

1252: template <device::cupm::DeviceType T>
1253: template <std::size_t... Idx>
1254: inline PetscErrorCode VecSeq_CUPM<T>::MDot_kernel_dispatch_(PetscDeviceContext dctx, cupmStream_t stream, const PetscScalar *xarr, const Vec yin[], PetscInt size, PetscScalar *results, util::index_sequence<Idx...>) noexcept
1255: {
1256:   PetscFunctionBegin;
1257:   // REVIEW ME: convert this kernel launch to PetscCUPMLaunchKernel1D(), it currently launches
1258:   // 128 blocks of 128 threads every time which may be wasteful
1259:   // clang-format off
1260:   PetscCallCUPM(
1261:     cupmLaunchKernel(
1262:       kernels::MDot_kernel<typename detail::repeat_type<const PetscScalar *, Idx>::type...>,
1263:       MDOT_WORKGROUP_NUM, MDOT_WORKGROUP_SIZE, 0, stream,
1264:       xarr, size, results, DeviceArrayRead(dctx, yin[Idx]).data()...
1265:     )
1266:   );
1267:   // clang-format on
1268:   PetscFunctionReturn(PETSC_SUCCESS);
1269: }

1271: template <device::cupm::DeviceType T>
1272: template <int N>
1273: inline PetscErrorCode VecSeq_CUPM<T>::MDot_kernel_dispatch_(PetscDeviceContext dctx, cupmStream_t stream, const PetscScalar *xarr, const Vec yin[], PetscInt size, PetscScalar *results, PetscInt &yidx) noexcept
1274: {
1275:   PetscFunctionBegin;
1276:   PetscCall(MDot_kernel_dispatch_(dctx, stream, xarr, yin + yidx, size, results + yidx * MDOT_WORKGROUP_NUM, util::make_index_sequence<N>{}));
1277:   yidx += N;
1278:   PetscFunctionReturn(PETSC_SUCCESS);
1279: }

1281: template <device::cupm::DeviceType T>
1282: inline PetscErrorCode VecSeq_CUPM<T>::MDot_(std::false_type, Vec xin, PetscInt nv, const Vec yin[], PetscScalar *z, PetscDeviceContext dctx) noexcept
1283: {
1284:   // the largest possible size of a batch
1285:   constexpr PetscInt batchsize = 8;
1286:   // how many sub streams to create, if nv <= batchsize we can do this without looping, so we
1287:   // do not create substreams. Note we don't create more than 8 streams, in practice we could
1288:   // not get more parallelism with higher numbers.
1289:   const auto   num_sub_streams = nv > batchsize ? std::min((nv + batchsize) / batchsize, batchsize) : 0;
1290:   const auto   n               = xin->map->n;
1291:   const auto   nwork           = nv * MDOT_WORKGROUP_NUM;
1292:   PetscScalar *d_results;
1293:   cupmStream_t stream;

1295:   PetscFunctionBegin;
1296:   PetscCall(GetHandlesFrom_(dctx, &stream));
1297:   // allocate scratchpad memory for the results of individual work groups
1298:   PetscCall(PetscDeviceMalloc(dctx, PETSC_MEMTYPE_CUPM(), nwork, &d_results));
1299:   {
1300:     const auto          xptr       = DeviceArrayRead(dctx, xin);
1301:     PetscInt            yidx       = 0;
1302:     auto                subidx     = 0;
1303:     auto                cur_stream = stream;
1304:     auto                cur_ctx    = dctx;
1305:     PetscDeviceContext *sub        = nullptr;
1306:     PetscStreamType     stype;

1308:     // REVIEW ME: maybe PetscDeviceContextFork() should insert dctx into the first entry of
1309:     // sub. Ideally the parent context should also join in on the fork, but it is extremely
1310:     // fiddly to do so presently
1311:     PetscCall(PetscDeviceContextGetStreamType(dctx, &stype));
1312:     if (stype == PETSC_STREAM_DEFAULT || stype == PETSC_STREAM_DEFAULT_WITH_BARRIER) stype = PETSC_STREAM_NONBLOCKING;
1313:     // If we have a default stream create nonblocking streams instead (as we can
1314:     // locally exploit the parallelism). Otherwise use the prescribed stream type.
1315:     PetscCall(PetscDeviceContextForkWithStreamType(dctx, stype, num_sub_streams, &sub));
1316:     PetscCall(PetscLogGpuTimeBegin());
1317:     do {
1318:       if (num_sub_streams) {
1319:         cur_ctx = sub[subidx++ % num_sub_streams];
1320:         PetscCall(GetHandlesFrom_(cur_ctx, &cur_stream));
1321:       }
1322:       // REVIEW ME: Should probably try and load-balance these. Consider the case where nv = 9;
1323:       // it is very likely better to do 4+5 rather than 8+1
1324:       switch (nv - yidx) {
1325:       case 7:
1326:         PetscCall(MDot_kernel_dispatch_<7>(cur_ctx, cur_stream, xptr.data(), yin, n, d_results, yidx));
1327:         break;
1328:       case 6:
1329:         PetscCall(MDot_kernel_dispatch_<6>(cur_ctx, cur_stream, xptr.data(), yin, n, d_results, yidx));
1330:         break;
1331:       case 5:
1332:         PetscCall(MDot_kernel_dispatch_<5>(cur_ctx, cur_stream, xptr.data(), yin, n, d_results, yidx));
1333:         break;
1334:       case 4:
1335:         PetscCall(MDot_kernel_dispatch_<4>(cur_ctx, cur_stream, xptr.data(), yin, n, d_results, yidx));
1336:         break;
1337:       case 3:
1338:         PetscCall(MDot_kernel_dispatch_<3>(cur_ctx, cur_stream, xptr.data(), yin, n, d_results, yidx));
1339:         break;
1340:       case 2:
1341:         PetscCall(MDot_kernel_dispatch_<2>(cur_ctx, cur_stream, xptr.data(), yin, n, d_results, yidx));
1342:         break;
1343:       case 1:
1344:         PetscCall(MDot_kernel_dispatch_<1>(cur_ctx, cur_stream, xptr.data(), yin, n, d_results, yidx));
1345:         break;
1346:       default: // 8 or more
1347:         PetscCall(MDot_kernel_dispatch_<8>(cur_ctx, cur_stream, xptr.data(), yin, n, d_results, yidx));
1348:         break;
1349:       }
1350:     } while (yidx < nv);
1351:     PetscCall(PetscLogGpuTimeEnd());
1352:     PetscCall(PetscDeviceContextJoin(dctx, num_sub_streams, PETSC_DEVICE_CONTEXT_JOIN_DESTROY, &sub));
1353:   }

1355:   PetscCall(PetscCUPMLaunchKernel1D(nv, 0, stream, kernels::sum_kernel, nv, d_results));
1356:   // copy result of device reduction to host
1357:   PetscCall(PetscCUPMMemcpyAsync(z, d_results, nv, cupmMemcpyDeviceToHost, stream));
1358:   // do these now while final reduction is in flight
1359:   PetscCall(PetscLogGpuFlops(nwork));
1360:   PetscCall(PetscDeviceFree(dctx, d_results));
1361:   PetscFunctionReturn(PETSC_SUCCESS);
1362: }

1364: #undef MDOT_WORKGROUP_NUM
1365: #undef MDOT_WORKGROUP_SIZE

1367: template <device::cupm::DeviceType T>
1368: inline PetscErrorCode VecSeq_CUPM<T>::MDot_(std::true_type, Vec xin, PetscInt nv, const Vec yin[], PetscScalar *z, PetscDeviceContext dctx) noexcept
1369: {
1370:   // probably not worth it to run more than 8 of these at a time?
1371:   const auto          n_sub = PetscMin(nv, 8);
1372:   const auto          n     = static_cast<cupmBlasInt_t>(xin->map->n);
1373:   const auto          xptr  = DeviceArrayRead(dctx, xin);
1374:   PetscScalar        *d_z;
1375:   PetscDeviceContext *subctx;
1376:   cupmStream_t        stream;

1378:   PetscFunctionBegin;
1379:   PetscCall(GetHandlesFrom_(dctx, &stream));
1380:   PetscCall(PetscDeviceMalloc(dctx, PETSC_MEMTYPE_CUPM(), nv, &d_z));
1381:   PetscCall(PetscDeviceContextFork(dctx, n_sub, &subctx));
1382:   PetscCall(PetscLogGpuTimeBegin());
1383:   for (PetscInt i = 0; i < nv; ++i) {
1384:     const auto            sub = subctx[i % n_sub];
1385:     cupmBlasHandle_t      handle;
1386:     cupmBlasPointerMode_t old_mode;

1388:     PetscCall(GetHandlesFrom_(sub, &handle));
1389:     PetscCallCUPMBLAS(cupmBlasGetPointerMode(handle, &old_mode));
1390:     if (old_mode != CUPMBLAS_POINTER_MODE_DEVICE) PetscCallCUPMBLAS(cupmBlasSetPointerMode(handle, CUPMBLAS_POINTER_MODE_DEVICE));
1391:     PetscCallCUPMBLAS(cupmBlasXdot(handle, n, DeviceArrayRead(sub, yin[i]), 1, xptr.cupmdata(), 1, cupmScalarPtrCast(d_z + i)));
1392:     if (old_mode != CUPMBLAS_POINTER_MODE_DEVICE) PetscCallCUPMBLAS(cupmBlasSetPointerMode(handle, old_mode));
1393:   }
1394:   PetscCall(PetscLogGpuTimeEnd());
1395:   PetscCall(PetscDeviceContextJoin(dctx, n_sub, PETSC_DEVICE_CONTEXT_JOIN_DESTROY, &subctx));
1396:   PetscCall(PetscCUPMMemcpyAsync(z, d_z, nv, cupmMemcpyDeviceToHost, stream));
1397:   PetscCall(PetscDeviceFree(dctx, d_z));
1398:   // REVIEW ME: flops?????
1399:   PetscFunctionReturn(PETSC_SUCCESS);
1400: }

1402: // v->ops->mdot
1403: template <device::cupm::DeviceType T>
1404: inline PetscErrorCode VecSeq_CUPM<T>::MDot(Vec xin, PetscInt nv, const Vec yin[], PetscScalar *z) noexcept
1405: {
1406:   PetscFunctionBegin;
1407:   if (PetscUnlikely(nv == 1)) {
1408:     // dot handles nv = 0 correctly
1409:     PetscCall(Dot(xin, const_cast<Vec>(yin[0]), z));
1410:   } else if (const auto n = xin->map->n) {
1411:     PetscDeviceContext dctx;

1413:     PetscCheck(nv > 0, PETSC_COMM_SELF, PETSC_ERR_LIB, "Number of vectors provided to %s %" PetscInt_FMT " not positive", PETSC_FUNCTION_NAME, nv);
1414:     PetscCall(GetHandles_(&dctx));
1415:     PetscCall(MDot_(std::integral_constant<bool, PetscDefined(USE_COMPLEX)>{}, xin, nv, yin, z, dctx));
1416:     // REVIEW ME: double count of flops??
1417:     PetscCall(PetscLogGpuFlops(nv * (2 * n - 1)));
1418:     PetscCall(PetscDeviceContextSynchronize(dctx));
1419:   } else {
1420:     PetscCall(PetscArrayzero(z, nv));
1421:   }
1422:   PetscFunctionReturn(PETSC_SUCCESS);
1423: }

1425: // VecSetAsync_Private
1426: template <device::cupm::DeviceType T>
1427: inline PetscErrorCode VecSeq_CUPM<T>::SetAsync(Vec xin, PetscScalar alpha, PetscDeviceContext dctx) noexcept
1428: {
1429:   const auto   n = xin->map->n;
1430:   cupmStream_t stream;

1432:   PetscFunctionBegin;
1433:   PetscCall(PetscDeviceContextGetOptionalNullContext_Internal(&dctx));
1434:   PetscCall(GetHandlesFrom_(dctx, &stream));
1435:   {
1436:     const auto xptr = DeviceArrayWrite(dctx, xin);

1438:     if (alpha == PetscScalar(0.0)) {
1439:       PetscCall(PetscCUPMMemsetAsync(xptr.data(), 0, n, stream));
1440:     } else {
1441:       const auto dptr = thrust::device_pointer_cast(xptr.data());

1443:       PetscCallThrust(THRUST_CALL(thrust::fill, stream, dptr, dptr + n, alpha));
1444:     }
1445:   }
1446:   if (n > 0) PetscCall(PetscDeviceContextSynchronizeIfWithBarrier_Internal(dctx));
1447:   PetscFunctionReturn(PETSC_SUCCESS);
1448: }

1450: // v->ops->set
1451: template <device::cupm::DeviceType T>
1452: inline PetscErrorCode VecSeq_CUPM<T>::Set(Vec xin, PetscScalar alpha) noexcept
1453: {
1454:   PetscFunctionBegin;
1455:   PetscCall(SetAsync(xin, alpha, nullptr));
1456:   PetscFunctionReturn(PETSC_SUCCESS);
1457: }

1459: // VecScaleAsync_Private
1460: template <device::cupm::DeviceType T>
1461: inline PetscErrorCode VecSeq_CUPM<T>::ScaleAsync(Vec xin, PetscScalar alpha, PetscDeviceContext dctx) noexcept
1462: {
1463:   PetscFunctionBegin;
1464:   if (PetscUnlikely(alpha == PetscScalar(1.0))) PetscFunctionReturn(PETSC_SUCCESS);
1465:   PetscCall(PetscDeviceContextGetOptionalNullContext_Internal(&dctx));
1466:   if (PetscUnlikely(alpha == PetscScalar(0.0))) {
1467:     PetscCall(SetAsync(xin, alpha, dctx));
1468:   } else if (const auto n = static_cast<cupmBlasInt_t>(xin->map->n)) {
1469:     cupmBlasHandle_t cupmBlasHandle;

1471:     PetscCall(GetHandlesFrom_(dctx, &cupmBlasHandle));
1472:     PetscCall(PetscLogGpuTimeBegin());
1473:     PetscCallCUPMBLAS(cupmBlasXscal(cupmBlasHandle, n, cupmScalarPtrCast(&alpha), DeviceArrayReadWrite(dctx, xin), 1));
1474:     PetscCall(PetscLogGpuTimeEnd());
1475:     PetscCall(PetscLogGpuFlops(n));
1476:     PetscCall(PetscDeviceContextSynchronizeIfWithBarrier_Internal(dctx));
1477:   } else {
1478:     PetscCall(MaybeIncrementEmptyLocalVec(xin));
1479:   }
1480:   PetscFunctionReturn(PETSC_SUCCESS);
1481: }

1483: // v->ops->scale
1484: template <device::cupm::DeviceType T>
1485: inline PetscErrorCode VecSeq_CUPM<T>::Scale(Vec xin, PetscScalar alpha) noexcept
1486: {
1487:   PetscFunctionBegin;
1488:   PetscCall(ScaleAsync(xin, alpha, nullptr));
1489:   PetscFunctionReturn(PETSC_SUCCESS);
1490: }

1492: // v->ops->tdot
1493: template <device::cupm::DeviceType T>
1494: inline PetscErrorCode VecSeq_CUPM<T>::TDot(Vec xin, Vec yin, PetscScalar *z) noexcept
1495: {
1496:   PetscBool yiscupm;

1498:   PetscFunctionBegin;
1499:   PetscCall(PetscObjectTypeCompareAny(PetscObjectCast(yin), &yiscupm, VECSEQCUPM(), VECMPICUPM(), ""));
1500:   if (!yiscupm) {
1501:     PetscCall(VecTDot_Seq(xin, yin, z));
1502:     PetscFunctionReturn(PETSC_SUCCESS);
1503:   }
1504:   if (const auto n = static_cast<cupmBlasInt_t>(xin->map->n)) {
1505:     PetscDeviceContext dctx;
1506:     cupmBlasHandle_t   cupmBlasHandle;

1508:     PetscCall(GetHandles_(&dctx, &cupmBlasHandle));
1509:     PetscCall(PetscLogGpuTimeBegin());
1510:     PetscCallCUPMBLAS(cupmBlasXdotu(cupmBlasHandle, n, DeviceArrayRead(dctx, xin), 1, DeviceArrayRead(dctx, yin), 1, cupmScalarPtrCast(z)));
1511:     PetscCall(PetscLogGpuTimeEnd());
1512:     PetscCall(PetscLogGpuFlops(2 * n - 1));
1513:   } else {
1514:     *z = 0.0;
1515:   }
1516:   PetscFunctionReturn(PETSC_SUCCESS);
1517: }

1519: // VecCopyAsync_Private
1520: template <device::cupm::DeviceType T>
1521: inline PetscErrorCode VecSeq_CUPM<T>::CopyAsync(Vec xin, Vec yout, PetscDeviceContext dctx) noexcept
1522: {
1523:   PetscFunctionBegin;
1524:   if (xin == yout) PetscFunctionReturn(PETSC_SUCCESS);
1525:   if (const auto n = xin->map->n) {
1526:     const auto xmask = xin->offloadmask;
1527:     // silence buggy gcc warning: mode may be used uninitialized in this function
1528:     auto         mode = cupmMemcpyDeviceToDevice;
1529:     cupmStream_t stream;

1531:     // translate from PetscOffloadMask to cupmMemcpyKind
1532:     PetscCall(PetscDeviceContextGetOptionalNullContext_Internal(&dctx));
1533:     switch (const auto ymask = yout->offloadmask) {
1534:     case PETSC_OFFLOAD_CPU:
1535:     case PETSC_OFFLOAD_UNALLOCATED: {
1536:       PetscBool yiscupm;

1538:       PetscCall(PetscObjectTypeCompareAny(PetscObjectCast(yout), &yiscupm, VECSEQCUPM(), VECMPICUPM(), ""));
1539:       if (yiscupm && !yout->boundtocpu) {
1540:         /* If GPU vector, also ensure output is on GPU unless explicitly bound to CPU */
1541:         mode = PetscOffloadDevice(xmask) ? cupmMemcpyDeviceToDevice : cupmMemcpyHostToDevice;
1542:       } else {
1543:         mode = PetscOffloadDevice(xmask) ? cupmMemcpyDeviceToHost : cupmMemcpyHostToHost;
1544:       }
1545:       break;
1546:     }
1547:     case PETSC_OFFLOAD_BOTH:
1548:     case PETSC_OFFLOAD_GPU:
1549:       mode = PetscOffloadDevice(xmask) ? cupmMemcpyDeviceToDevice : cupmMemcpyHostToDevice;
1550:       break;
1551:     default:
1552:       SETERRQ(PETSC_COMM_SELF, PETSC_ERR_ARG_INCOMP, "Incompatible offload mask %s", PetscOffloadMaskToString(ymask));
1553:     }

1555:     PetscCall(GetHandlesFrom_(dctx, &stream));
1556:     switch (mode) {
1557:     case cupmMemcpyDeviceToDevice: // the best case
1558:     case cupmMemcpyHostToDevice: { // not terrible
1559:       const auto yptr = DeviceArrayWrite(dctx, yout);
1560:       const auto xptr = mode == cupmMemcpyDeviceToDevice ? DeviceArrayRead(dctx, xin).data() : HostArrayRead(dctx, xin).data();

1562:       PetscCall(PetscLogGpuTimeBegin());
1563:       PetscCall(PetscCUPMMemcpyAsync(yptr.data(), xptr, n, mode, stream));
1564:       PetscCall(PetscLogGpuTimeEnd());
1565:     } break;
1566:     case cupmMemcpyDeviceToHost: // not great
1567:     case cupmMemcpyHostToHost: { // worst case
1568:       const auto   xptr = mode == cupmMemcpyDeviceToHost ? DeviceArrayRead(dctx, xin).data() : HostArrayRead(dctx, xin).data();
1569:       PetscScalar *yptr;

1571:       PetscCall(VecGetArrayWrite(yout, &yptr));
1572:       if (mode == cupmMemcpyDeviceToHost) PetscCall(PetscLogGpuTimeBegin());
1573:       PetscCall(PetscCUPMMemcpyAsync(yptr, xptr, n, mode, stream, /* force async */ true));
1574:       if (mode == cupmMemcpyDeviceToHost) PetscCall(PetscLogGpuTimeEnd());
1575:       PetscCall(VecRestoreArrayWrite(yout, &yptr));
1576:     } break;
1577:     default:
1578:       SETERRQ(PETSC_COMM_SELF, PETSC_ERR_GPU, "Unknown cupmMemcpyKind %d", static_cast<int>(mode));
1579:     }
1580:     PetscCall(PetscDeviceContextSynchronizeIfWithBarrier_Internal(dctx));
1581:   } else {
1582:     PetscCall(MaybeIncrementEmptyLocalVec(yout));
1583:   }
1584:   PetscFunctionReturn(PETSC_SUCCESS);
1585: }

1587: // v->ops->copy
1588: template <device::cupm::DeviceType T>
1589: inline PetscErrorCode VecSeq_CUPM<T>::Copy(Vec xin, Vec yout) noexcept
1590: {
1591:   PetscFunctionBegin;
1592:   PetscCall(CopyAsync(xin, yout, nullptr));
1593:   PetscFunctionReturn(PETSC_SUCCESS);
1594: }

1596: // VecSwapAsync_Private
1597: template <device::cupm::DeviceType T>
1598: inline PetscErrorCode VecSeq_CUPM<T>::SwapAsync(Vec xin, Vec yin, PetscDeviceContext dctx) noexcept
1599: {
1600:   PetscBool yiscupm;

1602:   PetscFunctionBegin;
1603:   if (xin == yin) PetscFunctionReturn(PETSC_SUCCESS);
1604:   PetscCall(PetscObjectTypeCompareAny(PetscObjectCast(yin), &yiscupm, VECSEQCUPM(), VECMPICUPM(), ""));
1605:   PetscCheck(yiscupm, PetscObjectComm(PetscObjectCast(yin)), PETSC_ERR_SUP, "Cannot swap with Y of type %s", PetscObjectCast(yin)->type_name);
1606:   PetscCall(PetscDeviceContextGetOptionalNullContext_Internal(&dctx));
1607:   if (const auto n = static_cast<cupmBlasInt_t>(xin->map->n)) {
1608:     cupmBlasHandle_t cupmBlasHandle;

1610:     PetscCall(GetHandlesFrom_(dctx, &cupmBlasHandle));
1611:     PetscCall(PetscLogGpuTimeBegin());
1612:     PetscCallCUPMBLAS(cupmBlasXswap(cupmBlasHandle, n, DeviceArrayReadWrite(dctx, xin), 1, DeviceArrayReadWrite(dctx, yin), 1));
1613:     PetscCall(PetscLogGpuTimeEnd());
1614:     PetscCall(PetscDeviceContextSynchronizeIfWithBarrier_Internal(dctx));
1615:   } else {
1616:     PetscCall(MaybeIncrementEmptyLocalVec(xin));
1617:     PetscCall(MaybeIncrementEmptyLocalVec(yin));
1618:   }
1619:   PetscFunctionReturn(PETSC_SUCCESS);
1620: }

1622: // v->ops->swap
1623: template <device::cupm::DeviceType T>
1624: inline PetscErrorCode VecSeq_CUPM<T>::Swap(Vec xin, Vec yin) noexcept
1625: {
1626:   PetscFunctionBegin;
1627:   PetscCall(SwapAsync(xin, yin, nullptr));
1628:   PetscFunctionReturn(PETSC_SUCCESS);
1629: }

1631: // VecAXPYBYAsync_Private
1632: template <device::cupm::DeviceType T>
1633: inline PetscErrorCode VecSeq_CUPM<T>::AXPBYAsync(Vec yin, PetscScalar alpha, PetscScalar beta, Vec xin, PetscDeviceContext dctx) noexcept
1634: {
1635:   PetscBool xiscupm;

1637:   PetscFunctionBegin;
1638:   PetscCall(PetscObjectTypeCompareAny(PetscObjectCast(xin), &xiscupm, VECSEQCUPM(), VECMPICUPM(), ""));
1639:   if (!xiscupm) {
1640:     PetscCall(VecAXPBY_Seq(yin, alpha, beta, xin));
1641:     PetscFunctionReturn(PETSC_SUCCESS);
1642:   }
1643:   PetscCall(PetscDeviceContextGetOptionalNullContext_Internal(&dctx));
1644:   if (alpha == PetscScalar(0.0)) {
1645:     PetscCall(ScaleAsync(yin, beta, dctx));
1646:   } else if (beta == PetscScalar(1.0)) {
1647:     PetscCall(AXPYAsync(yin, alpha, xin, dctx));
1648:   } else if (alpha == PetscScalar(1.0)) {
1649:     PetscCall(AYPXAsync(yin, beta, xin, dctx));
1650:   } else if (const auto n = static_cast<cupmBlasInt_t>(yin->map->n)) {
1651:     PetscBool xiscupm;

1653:     PetscCall(PetscObjectTypeCompareAny(PetscObjectCast(xin), &xiscupm, VECSEQCUPM(), VECMPICUPM(), ""));
1654:     if (!xiscupm) {
1655:       PetscCall(VecAXPBY_Seq(yin, alpha, beta, xin));
1656:       PetscFunctionReturn(PETSC_SUCCESS);
1657:     }

1659:     const auto       betaIsZero = beta == PetscScalar(0.0);
1660:     const auto       aptr       = cupmScalarPtrCast(&alpha);
1661:     cupmBlasHandle_t cupmBlasHandle;

1663:     PetscCall(GetHandlesFrom_(dctx, &cupmBlasHandle));
1664:     {
1665:       const auto xptr = DeviceArrayRead(dctx, xin);

1667:       if (betaIsZero /* beta = 0 */) {
1668:         // here we can get away with purely write-only as we memcpy into it first
1669:         const auto   yptr = DeviceArrayWrite(dctx, yin);
1670:         cupmStream_t stream;

1672:         PetscCall(GetHandlesFrom_(dctx, &stream));
1673:         PetscCall(PetscLogGpuTimeBegin());
1674:         PetscCall(PetscCUPMMemcpyAsync(yptr.data(), xptr.data(), n, cupmMemcpyDeviceToDevice, stream));
1675:         PetscCallCUPMBLAS(cupmBlasXscal(cupmBlasHandle, n, aptr, yptr.cupmdata(), 1));
1676:       } else {
1677:         const auto yptr = DeviceArrayReadWrite(dctx, yin);

1679:         PetscCall(PetscLogGpuTimeBegin());
1680:         PetscCallCUPMBLAS(cupmBlasXscal(cupmBlasHandle, n, cupmScalarPtrCast(&beta), yptr.cupmdata(), 1));
1681:         PetscCallCUPMBLAS(cupmBlasXaxpy(cupmBlasHandle, n, aptr, xptr.cupmdata(), 1, yptr.cupmdata(), 1));
1682:       }
1683:     }
1684:     PetscCall(PetscLogGpuTimeEnd());
1685:     PetscCall(PetscLogGpuFlops((betaIsZero ? 1 : 3) * n));
1686:     PetscCall(PetscDeviceContextSynchronizeIfWithBarrier_Internal(dctx));
1687:   } else {
1688:     PetscCall(MaybeIncrementEmptyLocalVec(yin));
1689:   }
1690:   PetscFunctionReturn(PETSC_SUCCESS);
1691: }

1693: // v->ops->axpby
1694: template <device::cupm::DeviceType T>
1695: inline PetscErrorCode VecSeq_CUPM<T>::AXPBY(Vec yin, PetscScalar alpha, PetscScalar beta, Vec xin) noexcept
1696: {
1697:   PetscFunctionBegin;
1698:   PetscCall(AXPBYAsync(yin, alpha, beta, xin, nullptr));
1699:   PetscFunctionReturn(PETSC_SUCCESS);
1700: }

1702: // VecAXPBYPCZAsync_Private
1703: template <device::cupm::DeviceType T>
1704: inline PetscErrorCode VecSeq_CUPM<T>::AXPBYPCZAsync(Vec zin, PetscScalar alpha, PetscScalar beta, PetscScalar gamma, Vec xin, Vec yin, PetscDeviceContext dctx) noexcept
1705: {
1706:   PetscFunctionBegin;
1707:   PetscCall(PetscDeviceContextGetOptionalNullContext_Internal(&dctx));
1708:   if (gamma != PetscScalar(1.0)) PetscCall(ScaleAsync(zin, gamma, dctx));
1709:   PetscCall(AXPYAsync(zin, alpha, xin, dctx));
1710:   PetscCall(AXPYAsync(zin, beta, yin, dctx));
1711:   PetscFunctionReturn(PETSC_SUCCESS);
1712: }

1714: // v->ops->axpbypcz
1715: template <device::cupm::DeviceType T>
1716: inline PetscErrorCode VecSeq_CUPM<T>::AXPBYPCZ(Vec zin, PetscScalar alpha, PetscScalar beta, PetscScalar gamma, Vec xin, Vec yin) noexcept
1717: {
1718:   PetscFunctionBegin;
1719:   PetscCall(AXPBYPCZAsync(zin, alpha, beta, gamma, xin, yin, nullptr));
1720:   PetscFunctionReturn(PETSC_SUCCESS);
1721: }

1723: // v->ops->norm
1724: template <device::cupm::DeviceType T>
1725: inline PetscErrorCode VecSeq_CUPM<T>::Norm(Vec xin, NormType type, PetscReal *z) noexcept
1726: {
1727:   PetscDeviceContext dctx;
1728:   cupmBlasHandle_t   cupmBlasHandle;

1730:   PetscFunctionBegin;
1731:   PetscCall(GetHandles_(&dctx, &cupmBlasHandle));
1732:   if (const auto n = static_cast<cupmBlasInt_t>(xin->map->n)) {
1733:     const auto xptr      = DeviceArrayRead(dctx, xin);
1734:     PetscInt   flopCount = 0;

1736:     PetscCall(PetscLogGpuTimeBegin());
1737:     switch (type) {
1738:     case NORM_1_AND_2:
1739:     case NORM_1:
1740:       PetscCallCUPMBLAS(cupmBlasXasum(cupmBlasHandle, n, xptr.cupmdata(), 1, cupmRealPtrCast(z)));
1741:       flopCount = std::max(n - 1, 0);
1742:       if (type == NORM_1) break;
1743:       ++z; // fall-through
1744: #if PETSC_CPP_VERSION >= 17
1745:       [[fallthrough]];
1746: #endif
1747:     case NORM_2:
1748:     case NORM_FROBENIUS:
1749:       PetscCallCUPMBLAS(cupmBlasXnrm2(cupmBlasHandle, n, xptr.cupmdata(), 1, cupmRealPtrCast(z)));
1750:       flopCount += std::max(2 * n - 1, 0); // += in case we've fallen through from NORM_1_AND_2
1751:       break;
1752:     case NORM_INFINITY: {
1753:       cupmBlasInt_t max_loc = 0;
1754:       PetscScalar   xv      = 0.;
1755:       cupmStream_t  stream;

1757:       PetscCall(GetHandlesFrom_(dctx, &stream));
1758:       PetscCallCUPMBLAS(cupmBlasXamax(cupmBlasHandle, n, xptr.cupmdata(), 1, &max_loc));
1759:       PetscCall(PetscCUPMMemcpyAsync(&xv, xptr.data() + max_loc - 1, 1, cupmMemcpyDeviceToHost, stream));
1760:       *z = PetscAbsScalar(xv);
1761:       // REVIEW ME: flopCount = ???
1762:     } break;
1763:     }
1764:     PetscCall(PetscLogGpuTimeEnd());
1765:     PetscCall(PetscLogGpuFlops(flopCount));
1766:   } else {
1767:     z[0]                    = 0.0;
1768:     z[type == NORM_1_AND_2] = 0.0;
1769:   }
1770:   PetscFunctionReturn(PETSC_SUCCESS);
1771: }

1773: namespace detail
1774: {

1776: template <NormType wnormtype>
1777: class ErrorWNormTransformBase {
1778: public:
1779:   using result_type = thrust::tuple<PetscReal, PetscReal, PetscReal, PetscInt, PetscInt, PetscInt>;

1781:   constexpr explicit ErrorWNormTransformBase(PetscReal v) noexcept : ignore_max_{v} { }

1783: protected:
1784:   struct NormTuple {
1785:     PetscReal norm;
1786:     PetscInt  loc;
1787:   };

1789:   PETSC_NODISCARD PETSC_HOSTDEVICE_INLINE_DECL static NormTuple compute_norm_(PetscReal err, PetscReal tol) noexcept
1790:   {
1791:     if (tol > 0.) {
1792:       const auto val = err / tol;

1794:       return {wnormtype == NORM_INFINITY ? val : PetscSqr(val), 1};
1795:     } else {
1796:       return {0.0, 0};
1797:     }
1798:   }

1800:   PetscReal ignore_max_;
1801: };

1803: template <NormType wnormtype>
1804: struct ErrorWNormTransform : ErrorWNormTransformBase<wnormtype> {
1805:   using base_type     = ErrorWNormTransformBase<wnormtype>;
1806:   using result_type   = typename base_type::result_type;
1807:   using argument_type = thrust::tuple<PetscScalar, PetscScalar, PetscScalar, PetscScalar>;

1809:   using base_type::base_type;

1811:   PETSC_NODISCARD PETSC_HOSTDEVICE_INLINE_DECL result_type operator()(const argument_type &x) const noexcept
1812:   {
1813:     const auto u     = thrust::get<0>(x); // with x.get<0>(), cuda-12.4.0 gives error: class "cuda::std::__4::tuple" has no member "get"
1814:     const auto y     = thrust::get<1>(x);
1815:     const auto au    = PetscAbsScalar(u);
1816:     const auto ay    = PetscAbsScalar(y);
1817:     const auto skip  = au < this->ignore_max_ || ay < this->ignore_max_;
1818:     const auto tola  = skip ? 0.0 : PetscRealPart(thrust::get<2>(x));
1819:     const auto tolr  = skip ? 0.0 : PetscRealPart(thrust::get<3>(x)) * PetscMax(au, ay);
1820:     const auto tol   = tola + tolr;
1821:     const auto err   = PetscAbsScalar(u - y);
1822:     const auto tup_a = this->compute_norm_(err, tola);
1823:     const auto tup_r = this->compute_norm_(err, tolr);
1824:     const auto tup_n = this->compute_norm_(err, tol);

1826:     return {tup_n.norm, tup_a.norm, tup_r.norm, tup_n.loc, tup_a.loc, tup_r.loc};
1827:   }
1828: };

1830: template <NormType wnormtype>
1831: struct ErrorWNormETransform : ErrorWNormTransformBase<wnormtype> {
1832:   using base_type     = ErrorWNormTransformBase<wnormtype>;
1833:   using result_type   = typename base_type::result_type;
1834:   using argument_type = thrust::tuple<PetscScalar, PetscScalar, PetscScalar, PetscScalar, PetscScalar>;

1836:   using base_type::base_type;

1838:   PETSC_NODISCARD PETSC_HOSTDEVICE_INLINE_DECL result_type operator()(const argument_type &x) const noexcept
1839:   {
1840:     const auto au    = PetscAbsScalar(thrust::get<0>(x));
1841:     const auto ay    = PetscAbsScalar(thrust::get<1>(x));
1842:     const auto skip  = au < this->ignore_max_ || ay < this->ignore_max_;
1843:     const auto tola  = skip ? 0.0 : PetscRealPart(thrust::get<3>(x));
1844:     const auto tolr  = skip ? 0.0 : PetscRealPart(thrust::get<4>(x)) * PetscMax(au, ay);
1845:     const auto tol   = tola + tolr;
1846:     const auto err   = PetscAbsScalar(thrust::get<2>(x));
1847:     const auto tup_a = this->compute_norm_(err, tola);
1848:     const auto tup_r = this->compute_norm_(err, tolr);
1849:     const auto tup_n = this->compute_norm_(err, tol);

1851:     return {tup_n.norm, tup_a.norm, tup_r.norm, tup_n.loc, tup_a.loc, tup_r.loc};
1852:   }
1853: };

1855: template <NormType wnormtype>
1856: struct ErrorWNormReduce {
1857:   using value_type = typename ErrorWNormTransformBase<wnormtype>::result_type;

1859:   PETSC_NODISCARD PETSC_HOSTDEVICE_INLINE_DECL value_type operator()(const value_type &lhs, const value_type &rhs) const noexcept
1860:   {
1861:     // cannot use lhs.get<0>() etc since the using decl above ambiguates the fact that
1862:     // result_type is a template, so in order to fix this we would need to write:
1863:     //
1864:     // lhs.template get<0>()
1865:     //
1866:     // which is unseemly.
1867:     if (wnormtype == NORM_INFINITY) {
1868:       // clang-format off
1869:       return {
1870:         PetscMax(thrust::get<0>(lhs), thrust::get<0>(rhs)),
1871:         PetscMax(thrust::get<1>(lhs), thrust::get<1>(rhs)),
1872:         PetscMax(thrust::get<2>(lhs), thrust::get<2>(rhs)),
1873:         thrust::get<3>(lhs) + thrust::get<3>(rhs),
1874:         thrust::get<4>(lhs) + thrust::get<4>(rhs),
1875:         thrust::get<5>(lhs) + thrust::get<5>(rhs)
1876:       };
1877:       // clang-format on
1878:     } else {
1879:       // clang-format off
1880:       return {
1881:         thrust::get<0>(lhs) + thrust::get<0>(rhs),
1882:         thrust::get<1>(lhs) + thrust::get<1>(rhs),
1883:         thrust::get<2>(lhs) + thrust::get<2>(rhs),
1884:         thrust::get<3>(lhs) + thrust::get<3>(rhs),
1885:         thrust::get<4>(lhs) + thrust::get<4>(rhs),
1886:         thrust::get<5>(lhs) + thrust::get<5>(rhs)
1887:       };
1888:       // clang-format on
1889:     }
1890:   }
1891: };

1893: template <template <NormType> class WNormTransformType, typename Tuple, typename cupmStream_t>
1894: inline PetscErrorCode ExecuteWNorm(Tuple &&first, Tuple &&last, NormType wnormtype, cupmStream_t stream, PetscReal ignore_max, PetscReal *norm, PetscInt *norm_loc, PetscReal *norma, PetscInt *norma_loc, PetscReal *normr, PetscInt *normr_loc) noexcept
1895: {
1896:   auto      begin = thrust::make_zip_iterator(std::forward<Tuple>(first));
1897:   auto      end   = thrust::make_zip_iterator(std::forward<Tuple>(last));
1898:   PetscReal n = 0, na = 0, nr = 0;
1899:   PetscInt  n_loc = 0, na_loc = 0, nr_loc = 0;

1901:   PetscFunctionBegin;
1902:   // clang-format off
1903:   if (wnormtype == NORM_INFINITY) {
1904:     PetscCallThrust(
1905:       thrust::tie(*norm, *norma, *normr, *norm_loc, *norma_loc, *normr_loc) = THRUST_CALL(
1906:         thrust::transform_reduce,
1907:         stream,
1908:         std::move(begin),
1909:         std::move(end),
1910:         WNormTransformType<NORM_INFINITY>{ignore_max},
1911:         thrust::make_tuple(n, na, nr, n_loc, na_loc, nr_loc),
1912:         ErrorWNormReduce<NORM_INFINITY>{}
1913:       )
1914:     );
1915:   } else {
1916:     PetscCallThrust(
1917:       thrust::tie(*norm, *norma, *normr, *norm_loc, *norma_loc, *normr_loc) = THRUST_CALL(
1918:         thrust::transform_reduce,
1919:         stream,
1920:         std::move(begin),
1921:         std::move(end),
1922:         WNormTransformType<NORM_2>{ignore_max},
1923:         thrust::make_tuple(n, na, nr, n_loc, na_loc, nr_loc),
1924:         ErrorWNormReduce<NORM_2>{}
1925:       )
1926:     );
1927:   }
1928:   // clang-format on
1929:   if (wnormtype == NORM_2) {
1930:     *norm  = PetscSqrtReal(*norm);
1931:     *norma = PetscSqrtReal(*norma);
1932:     *normr = PetscSqrtReal(*normr);
1933:   }
1934:   PetscFunctionReturn(PETSC_SUCCESS);
1935: }

1937: } // namespace detail

1939: // v->ops->errorwnorm
1940: template <device::cupm::DeviceType T>
1941: inline PetscErrorCode VecSeq_CUPM<T>::ErrorWnorm(Vec U, Vec Y, Vec E, NormType wnormtype, PetscReal atol, Vec vatol, PetscReal rtol, Vec vrtol, PetscReal ignore_max, PetscReal *norm, PetscInt *norm_loc, PetscReal *norma, PetscInt *norma_loc, PetscReal *normr, PetscInt *normr_loc) noexcept
1942: {
1943:   const auto nl = U->map->n;
1944: #if CCCL_VERSION >= 3004000
1945:   auto ait = cuda::make_constant_iterator(static_cast<PetscScalar>(atol));
1946:   auto rit = cuda::make_constant_iterator(static_cast<PetscScalar>(rtol));
1947: #else
1948:   auto ait = thrust::make_constant_iterator(static_cast<PetscScalar>(atol));
1949:   auto rit = thrust::make_constant_iterator(static_cast<PetscScalar>(rtol));
1950: #endif
1951:   PetscDeviceContext dctx;
1952:   cupmStream_t       stream;

1954:   PetscFunctionBegin;
1955:   PetscCall(GetHandles_(&dctx, &stream));
1956:   {
1957:     const auto ConditionalDeviceArrayRead = [&](Vec v) {
1958:       if (v) {
1959:         return thrust::device_pointer_cast(DeviceArrayRead(dctx, v).data());
1960:       } else {
1961:         return thrust::device_ptr<PetscScalar>{nullptr};
1962:       }
1963:     };

1965:     const auto uarr = DeviceArrayRead(dctx, U);
1966:     const auto yarr = DeviceArrayRead(dctx, Y);
1967:     const auto uptr = thrust::device_pointer_cast(uarr.data());
1968:     const auto yptr = thrust::device_pointer_cast(yarr.data());
1969:     const auto eptr = ConditionalDeviceArrayRead(E);
1970:     const auto rptr = ConditionalDeviceArrayRead(vrtol);
1971:     const auto aptr = ConditionalDeviceArrayRead(vatol);

1973:     if (!vatol && !vrtol) {
1974:       if (E) {
1975:         // clang-format off
1976:         PetscCall(
1977:           detail::ExecuteWNorm<detail::ErrorWNormETransform>(
1978:             thrust::make_tuple(uptr, yptr, eptr, ait, rit),
1979:             thrust::make_tuple(uptr + nl, yptr + nl, eptr + nl, ait, rit),
1980:             wnormtype, stream, ignore_max, norm, norm_loc, norma, norma_loc, normr, normr_loc
1981:           )
1982:         );
1983:         // clang-format on
1984:       } else {
1985:         // clang-format off
1986:         PetscCall(
1987:           detail::ExecuteWNorm<detail::ErrorWNormTransform>(
1988:             thrust::make_tuple(uptr, yptr, ait, rit),
1989:             thrust::make_tuple(uptr + nl, yptr + nl, ait, rit),
1990:             wnormtype, stream, ignore_max, norm, norm_loc, norma, norma_loc, normr, normr_loc
1991:           )
1992:         );
1993:         // clang-format on
1994:       }
1995:     } else if (!vatol) {
1996:       if (E) {
1997:         // clang-format off
1998:         PetscCall(
1999:           detail::ExecuteWNorm<detail::ErrorWNormETransform>(
2000:             thrust::make_tuple(uptr, yptr, eptr, ait, rptr),
2001:             thrust::make_tuple(uptr + nl, yptr + nl, eptr + nl, ait, rptr + nl),
2002:             wnormtype, stream, ignore_max, norm, norm_loc, norma, norma_loc, normr, normr_loc
2003:           )
2004:         );
2005:         // clang-format on
2006:       } else {
2007:         // clang-format off
2008:         PetscCall(
2009:           detail::ExecuteWNorm<detail::ErrorWNormTransform>(
2010:             thrust::make_tuple(uptr, yptr, ait, rptr),
2011:             thrust::make_tuple(uptr + nl, yptr + nl, ait, rptr + nl),
2012:             wnormtype, stream, ignore_max, norm, norm_loc, norma, norma_loc, normr, normr_loc
2013:           )
2014:         );
2015:         // clang-format on
2016:       }
2017:     } else if (!vrtol) {
2018:       if (E) {
2019:         // clang-format off
2020:           PetscCall(
2021:             detail::ExecuteWNorm<detail::ErrorWNormETransform>(
2022:               thrust::make_tuple(uptr, yptr, eptr, aptr, rit),
2023:               thrust::make_tuple(uptr + nl, yptr + nl, eptr + nl, aptr + nl, rit),
2024:               wnormtype, stream, ignore_max, norm, norm_loc, norma, norma_loc, normr, normr_loc
2025:             )
2026:           );
2027:         // clang-format on
2028:       } else {
2029:         // clang-format off
2030:           PetscCall(
2031:             detail::ExecuteWNorm<detail::ErrorWNormTransform>(
2032:               thrust::make_tuple(uptr, yptr, aptr, rit),
2033:               thrust::make_tuple(uptr + nl, yptr + nl, aptr + nl, rit),
2034:               wnormtype, stream, ignore_max, norm, norm_loc, norma, norma_loc, normr, normr_loc
2035:             )
2036:           );
2037:         // clang-format on
2038:       }
2039:     } else {
2040:       if (E) {
2041:         // clang-format off
2042:           PetscCall(
2043:             detail::ExecuteWNorm<detail::ErrorWNormETransform>(
2044:               thrust::make_tuple(uptr, yptr, eptr, aptr, rptr),
2045:               thrust::make_tuple(uptr + nl, yptr + nl, eptr + nl, aptr + nl, rptr + nl),
2046:               wnormtype, stream, ignore_max, norm, norm_loc, norma, norma_loc, normr, normr_loc
2047:             )
2048:           );
2049:         // clang-format on
2050:       } else {
2051:         // clang-format off
2052:           PetscCall(
2053:             detail::ExecuteWNorm<detail::ErrorWNormTransform>(
2054:               thrust::make_tuple(uptr, yptr, aptr, rptr),
2055:               thrust::make_tuple(uptr + nl, yptr + nl, aptr + nl, rptr + nl),
2056:               wnormtype, stream, ignore_max, norm, norm_loc, norma, norma_loc, normr, normr_loc
2057:             )
2058:           );
2059:         // clang-format on
2060:       }
2061:     }
2062:   }
2063:   PetscFunctionReturn(PETSC_SUCCESS);
2064: }

2066: namespace detail
2067: {
2068: struct dotnorm2_mult {
2069:   PETSC_NODISCARD PETSC_HOSTDEVICE_INLINE_DECL thrust::tuple<PetscScalar, PetscScalar> operator()(const PetscScalar &s, const PetscScalar &t) const noexcept
2070:   {
2071:     const auto conjt = PetscConj(t);

2073:     return {s * conjt, t * conjt};
2074:   }
2075: };

2077: // it is positively __bananas__ that thrust does not define default operator+ for tuples... I
2078: // would do it myself but now I am worried that they do so on purpose...
2079: struct dotnorm2_tuple_plus {
2080:   using value_type = thrust::tuple<PetscScalar, PetscScalar>;

2082:   PETSC_NODISCARD PETSC_HOSTDEVICE_INLINE_DECL value_type operator()(const value_type &lhs, const value_type &rhs) const noexcept { return {thrust::get<0>(lhs) + thrust::get<0>(rhs), thrust::get<1>(lhs) + thrust::get<1>(rhs)}; }
2083: };

2085: } // namespace detail

2087: // v->ops->dotnorm2
2088: template <device::cupm::DeviceType T>
2089: inline PetscErrorCode VecSeq_CUPM<T>::DotNorm2(Vec s, Vec t, PetscScalar *dp, PetscScalar *nm) noexcept
2090: {
2091:   PetscDeviceContext dctx;
2092:   cupmStream_t       stream;

2094:   PetscFunctionBegin;
2095:   PetscCall(GetHandles_(&dctx, &stream));
2096:   {
2097:     PetscScalar dpt = 0.0, nmt = 0.0;
2098:     const auto  sdptr = thrust::device_pointer_cast(DeviceArrayRead(dctx, s).data());

2100:     // clang-format off
2101:     PetscCallThrust(
2102:       thrust::tie(*dp, *nm) = THRUST_CALL(
2103:         thrust::inner_product,
2104:         stream,
2105:         sdptr, sdptr+s->map->n, thrust::device_pointer_cast(DeviceArrayRead(dctx, t).data()),
2106:         thrust::make_tuple(dpt, nmt),
2107:         detail::dotnorm2_tuple_plus{}, detail::dotnorm2_mult{}
2108:       );
2109:     );
2110:     // clang-format on
2111:   }
2112:   PetscFunctionReturn(PETSC_SUCCESS);
2113: }

2115: namespace detail
2116: {
2117: struct conjugate {
2118:   PETSC_NODISCARD PETSC_HOSTDEVICE_INLINE_DECL PetscScalar operator()(const PetscScalar &x) const noexcept { return PetscConj(x); }
2119: };

2121: } // namespace detail

2123: // v->ops->conjugate
2124: template <device::cupm::DeviceType T>
2125: inline PetscErrorCode VecSeq_CUPM<T>::ConjugateAsync(Vec xin, PetscDeviceContext dctx) noexcept
2126: {
2127:   PetscFunctionBegin;
2128:   if (PetscDefined(USE_COMPLEX)) PetscCall(PointwiseUnary_(detail::conjugate{}, xin, nullptr, dctx));
2129:   PetscFunctionReturn(PETSC_SUCCESS);
2130: }

2132: // v->ops->conjugate
2133: template <device::cupm::DeviceType T>
2134: inline PetscErrorCode VecSeq_CUPM<T>::Conjugate(Vec xin) noexcept
2135: {
2136:   PetscFunctionBegin;
2137:   PetscCall(ConjugateAsync(xin, nullptr));
2138:   PetscFunctionReturn(PETSC_SUCCESS);
2139: }

2141: namespace detail
2142: {

2144: struct real_part {
2145:   PETSC_NODISCARD PETSC_HOSTDEVICE_INLINE_DECL thrust::tuple<PetscReal, PetscInt> operator()(const thrust::tuple<PetscScalar, PetscInt> &x) const noexcept { return {PetscRealPart(thrust::get<0>(x)), thrust::get<1>(x)}; }

2147:   PETSC_NODISCARD PETSC_HOSTDEVICE_INLINE_DECL PetscReal operator()(const PetscScalar &x) const noexcept { return PetscRealPart(x); }
2148: };

2150: // deriving from Operator allows us to "store" an instance of the operator in the class but
2151: // also take advantage of empty base class optimization if the operator is stateless
2152: template <typename Operator>
2153: class tuple_compare : Operator {
2154: public:
2155:   using tuple_type    = thrust::tuple<PetscReal, PetscInt>;
2156:   using operator_type = Operator;

2158:   PETSC_NODISCARD PETSC_HOSTDEVICE_INLINE_DECL tuple_type operator()(const tuple_type &x, const tuple_type &y) const noexcept
2159:   {
2160:     if (op_()(thrust::get<0>(y), thrust::get<0>(x))) {
2161:       // if y is strictly greater/less than x, return y
2162:       return y;
2163:     } else if (thrust::get<0>(y) == thrust::get<0>(x)) {
2164:       // if equal, prefer lower index
2165:       return thrust::get<1>(y) < thrust::get<1>(x) ? y : x;
2166:     }
2167:     // otherwise return x
2168:     return x;
2169:   }

2171: private:
2172:   PETSC_NODISCARD PETSC_HOSTDEVICE_INLINE_DECL const operator_type &op_() const noexcept { return *this; }
2173: };

2175: } // namespace detail

2177: template <device::cupm::DeviceType T>
2178: template <typename TupleFuncT, typename UnaryFuncT>
2179: inline PetscErrorCode VecSeq_CUPM<T>::MinMax_(TupleFuncT &&tuple_ftr, UnaryFuncT &&unary_ftr, Vec v, PetscInt *p, PetscReal *m) noexcept
2180: {
2181:   PetscFunctionBegin;
2182:   PetscCheckTypeNames(v, VECSEQCUPM(), VECMPICUPM());
2183:   if (p) *p = -1;
2184:   if (const auto n = v->map->n) {
2185:     PetscDeviceContext dctx;
2186:     cupmStream_t       stream;

2188:     PetscCall(GetHandles_(&dctx, &stream));
2189:     // needed to:
2190:     // 1. switch between transform_reduce and reduce
2191:     // 2. strip the real_part functor from the arguments
2192: #if PetscDefined(USE_COMPLEX)
2193:   #define THRUST_MINMAX_REDUCE(...) THRUST_CALL(thrust::transform_reduce, __VA_ARGS__)
2194: #else
2195:   #define THRUST_MINMAX_REDUCE(s, b, e, real_part__, ...) THRUST_CALL(thrust::reduce, s, b, e, __VA_ARGS__)
2196: #endif
2197:     {
2198:       const auto vptr = thrust::device_pointer_cast(DeviceArrayRead(dctx, v).data());

2200:       if (p) {
2201:         // clang-format off
2202:         const auto zip = thrust::make_zip_iterator(
2203:           thrust::make_tuple(std::move(vptr), thrust::make_counting_iterator(PetscInt{0}))
2204:         );
2205:         // clang-format on
2206:         // need to use preprocessor conditionals since otherwise thrust complains about not being
2207:         // able to convert a thrust::device_reference to a PetscReal on complex
2208:         // builds...
2209:         // clang-format off
2210:         PetscCallThrust(
2211:           thrust::tie(*m, *p) = THRUST_MINMAX_REDUCE(
2212:             stream, zip, zip + n, detail::real_part{},
2213:             thrust::make_tuple(*m, *p), std::forward<TupleFuncT>(tuple_ftr)
2214:           );
2215:         );
2216:         // clang-format on
2217:       } else {
2218:         // clang-format off
2219:         PetscCallThrust(
2220:           *m = THRUST_MINMAX_REDUCE(
2221:             stream, vptr, vptr + n, detail::real_part{},
2222:             *m, std::forward<UnaryFuncT>(unary_ftr)
2223:           );
2224:         );
2225:         // clang-format on
2226:       }
2227:     }
2228: #undef THRUST_MINMAX_REDUCE
2229:   }
2230:   // REVIEW ME: flops?
2231:   PetscFunctionReturn(PETSC_SUCCESS);
2232: }

2234: // v->ops->max
2235: template <device::cupm::DeviceType T>
2236: inline PetscErrorCode VecSeq_CUPM<T>::Max(Vec v, PetscInt *p, PetscReal *m) noexcept
2237: {
2238: #if CCCL_VERSION >= 3001000
2239:   using tuple_functor = detail::tuple_compare<cuda::std::greater<PetscReal>>;
2240:   using unary_functor = cuda::maximum<PetscReal>;
2241: #else
2242:   using tuple_functor = detail::tuple_compare<thrust::greater<PetscReal>>;
2243:   using unary_functor = thrust::maximum<PetscReal>;
2244: #endif

2246:   PetscFunctionBegin;
2247:   *m = PETSC_MIN_REAL;
2248:   // use {} constructor syntax otherwise most vexing parse
2249:   PetscCall(MinMax_(tuple_functor{}, unary_functor{}, v, p, m));
2250:   PetscFunctionReturn(PETSC_SUCCESS);
2251: }

2253: // v->ops->min
2254: template <device::cupm::DeviceType T>
2255: inline PetscErrorCode VecSeq_CUPM<T>::Min(Vec v, PetscInt *p, PetscReal *m) noexcept
2256: {
2257: #if CCCL_VERSION >= 3001000
2258:   using tuple_functor = detail::tuple_compare<cuda::std::less<PetscReal>>;
2259:   using unary_functor = cuda::minimum<PetscReal>;
2260: #else
2261:   using tuple_functor = detail::tuple_compare<thrust::less<PetscReal>>;
2262:   using unary_functor = thrust::minimum<PetscReal>;
2263: #endif

2265:   PetscFunctionBegin;
2266:   *m = PETSC_MAX_REAL;
2267:   // use {} constructor syntax otherwise most vexing parse
2268:   PetscCall(MinMax_(tuple_functor{}, unary_functor{}, v, p, m));
2269:   PetscFunctionReturn(PETSC_SUCCESS);
2270: }

2272: // v->ops->sum
2273: template <device::cupm::DeviceType T>
2274: inline PetscErrorCode VecSeq_CUPM<T>::Sum(Vec v, PetscScalar *sum) noexcept
2275: {
2276:   PetscFunctionBegin;
2277:   if (const auto n = v->map->n) {
2278:     PetscDeviceContext dctx;
2279:     cupmStream_t       stream;

2281:     PetscCall(GetHandles_(&dctx, &stream));
2282:     const auto dptr = thrust::device_pointer_cast(DeviceArrayRead(dctx, v).data());
2283:     // REVIEW ME: why not cupmBlasXasum()?
2284:     PetscCallThrust(*sum = THRUST_CALL(thrust::reduce, stream, dptr, dptr + n, PetscScalar{0.0}););
2285:     // REVIEW ME: must be at least n additions
2286:     PetscCall(PetscLogGpuFlops(n));
2287:   } else {
2288:     *sum = 0.0;
2289:   }
2290:   PetscFunctionReturn(PETSC_SUCCESS);
2291: }

2293: template <device::cupm::DeviceType T>
2294: inline PetscErrorCode VecSeq_CUPM<T>::ShiftAsync(Vec v, PetscScalar shift, PetscDeviceContext dctx) noexcept
2295: {
2296:   PetscFunctionBegin;
2297:   PetscCall(PointwiseUnary_(device::cupm::functors::make_plus_equals(shift), v, nullptr, dctx));
2298:   PetscFunctionReturn(PETSC_SUCCESS);
2299: }

2301: template <device::cupm::DeviceType T>
2302: inline PetscErrorCode VecSeq_CUPM<T>::Shift(Vec v, PetscScalar shift) noexcept
2303: {
2304:   PetscFunctionBegin;
2305:   PetscCall(ShiftAsync(v, shift, nullptr));
2306:   PetscFunctionReturn(PETSC_SUCCESS);
2307: }

2309: template <device::cupm::DeviceType T>
2310: inline PetscErrorCode VecSeq_CUPM<T>::SetRandom(Vec v, PetscRandom rand) noexcept
2311: {
2312:   PetscFunctionBegin;
2313:   if (const auto n = v->map->n) {
2314:     PetscBool          iscurand;
2315:     PetscDeviceContext dctx;

2317:     PetscCall(GetHandles_(&dctx));
2318:     PetscCall(PetscObjectTypeCompare(PetscObjectCast(rand), PETSCCURAND, &iscurand));
2319:     if (iscurand) PetscCall(PetscRandomGetValues(rand, n, DeviceArrayWrite(dctx, v)));
2320:     else PetscCall(PetscRandomGetValues(rand, n, HostArrayWrite(dctx, v)));
2321:   } else {
2322:     PetscCall(MaybeIncrementEmptyLocalVec(v));
2323:   }
2324:   // REVIEW ME: flops????
2325:   // REVIEW ME: Timing???
2326:   PetscFunctionReturn(PETSC_SUCCESS);
2327: }

2329: // v->ops->setpreallocation
2330: template <device::cupm::DeviceType T>
2331: inline PetscErrorCode VecSeq_CUPM<T>::SetPreallocationCOO(Vec v, PetscCount ncoo, const PetscInt coo_i[]) noexcept
2332: {
2333:   PetscDeviceContext dctx;

2335:   PetscFunctionBegin;
2336:   PetscCall(GetHandles_(&dctx));
2337:   PetscCall(VecSetPreallocationCOO_Seq(v, ncoo, coo_i));
2338:   PetscCall(SetPreallocationCOO_CUPMBase(v, ncoo, coo_i, dctx));
2339:   PetscFunctionReturn(PETSC_SUCCESS);
2340: }

2342: // v->ops->setvaluescoo
2343: template <device::cupm::DeviceType T>
2344: inline PetscErrorCode VecSeq_CUPM<T>::SetValuesCOO(Vec x, const PetscScalar v[], InsertMode imode) noexcept
2345: {
2346:   auto               vv = const_cast<PetscScalar *>(v);
2347:   PetscMemType       memtype;
2348:   PetscDeviceContext dctx;
2349:   cupmStream_t       stream;

2351:   PetscFunctionBegin;
2352:   PetscCall(GetHandles_(&dctx, &stream));
2353:   PetscCall(PetscGetMemType(v, &memtype));
2354:   if (PetscMemTypeHost(memtype)) {
2355:     const auto size = VecIMPLCast(x)->coo_n;

2357:     // If user gave v[] in host, we might need to copy it to device if any
2358:     PetscCall(PetscDeviceMalloc(dctx, PETSC_MEMTYPE_CUPM(), size, &vv));
2359:     PetscCall(PetscCUPMMemcpyAsync(vv, v, size, cupmMemcpyHostToDevice, stream));
2360:     PetscCall(PetscLogCpuToGpu(size * sizeof(PetscScalar)));
2361:   }

2363:   if (const auto n = x->map->n) {
2364:     const auto vcu = VecCUPMCast(x);

2366:     PetscCall(PetscCUPMLaunchKernel1D(n, 0, stream, kernels::add_coo_values, vv, n, vcu->jmap1_d, vcu->perm1_d, imode, imode == INSERT_VALUES ? DeviceArrayWrite(dctx, x).data() : DeviceArrayReadWrite(dctx, x).data()));
2367:   } else {
2368:     PetscCall(MaybeIncrementEmptyLocalVec(x));
2369:   }

2371:   if (PetscMemTypeHost(memtype)) PetscCall(PetscDeviceFree(dctx, vv));
2372:   PetscCall(PetscDeviceContextSynchronize(dctx));
2373:   PetscFunctionReturn(PETSC_SUCCESS);
2374: }

2376: } // namespace impl

2378: } // namespace cupm

2380: } // namespace vec

2382: } // namespace Petsc