Actual source code: cupmcontext.hpp
1: #pragma once
3: #include <petsc/private/deviceimpl.h>
4: #include <petsc/private/cupmsolverinterface.hpp>
5: #include <petsc/private/logimpl.h>
7: #include <petsc/private/cpp/array.hpp>
9: #include "../segmentedmempool.hpp"
10: #include "cupmallocator.hpp"
11: #include "cupmstream.hpp"
12: #include "cupmevent.hpp"
14: namespace Petsc
15: {
17: namespace device
18: {
20: namespace cupm
21: {
23: namespace impl
24: {
26: template <DeviceType T>
27: class PETSC_SINGLE_LIBRARY_VISIBILITY_INTERNAL DeviceContext : SolverInterface<T> {
28: public:
29: PETSC_CUPMSOLVER_INHERIT_INTERFACE_TYPEDEFS_USING(T);
31: private:
32: template <typename H, std::size_t>
33: struct HandleTag {
34: using type = H;
35: };
37: using stream_tag = HandleTag<cupmStream_t, 0>;
38: using blas_tag = HandleTag<cupmBlasHandle_t, 1>;
39: using solver_tag = HandleTag<cupmSolverHandle_t, 2>;
41: using stream_type = CUPMStream<T>;
42: using event_type = CUPMEvent<T>;
44: public:
45: // This is the canonical PETSc "impls" struct that normally resides in a standalone impls
46: // header, but since we are using the power of templates it must be declared part of
47: // this class to have easy access the same typedefs. Technically one can make a
48: // templated struct outside the class but it's more code for the same result.
49: struct PetscDeviceContext_IMPLS {
50: stream_type stream{};
51: cupmEvent_t event{};
52: cupmEvent_t begin{}; // timer-only
53: cupmEvent_t end{}; // timer-only
54: #if PetscDefined(USE_DEBUG)
55: PetscBool timerInUse{};
56: PetscBool EnergyMeterInUse{};
57: #endif
58: cupmBlasHandle_t blas{};
59: cupmSolverHandle_t solver{};
60: #if PetscDefined(HAVE_NVML)
61: nvmlDevice_t nvmlHandle{};
62: unsigned long long energymeterbegin{};
63: unsigned long long energymeterend{};
64: #endif
66: constexpr PetscDeviceContext_IMPLS() noexcept = default;
68: PETSC_NODISCARD const cupmStream_t &get(stream_tag) const noexcept { return this->stream.get_stream(); }
70: PETSC_NODISCARD const cupmBlasHandle_t &get(blas_tag) const noexcept { return this->blas; }
72: PETSC_NODISCARD const cupmSolverHandle_t &get(solver_tag) const noexcept { return this->solver; }
73: };
75: private:
76: static bool initialized_;
78: static std::array<cupmBlasHandle_t, PETSC_DEVICE_MAX_DEVICES> blashandles_;
79: static std::array<cupmSolverHandle_t, PETSC_DEVICE_MAX_DEVICES> solverhandles_;
81: PETSC_NODISCARD static constexpr PetscDeviceContext_IMPLS *impls_cast_(PetscDeviceContext ptr) noexcept { return static_cast<PetscDeviceContext_IMPLS *>(ptr->data); }
83: PETSC_NODISCARD static constexpr CUPMEvent<T> *event_cast_(PetscEvent event) noexcept { return static_cast<CUPMEvent<T> *>(event->data); }
85: PETSC_NODISCARD static PetscLogEvent CUPMBLAS_HANDLE_CREATE() noexcept { return T == DeviceType::CUDA ? CUBLAS_HANDLE_CREATE : HIPBLAS_HANDLE_CREATE; }
87: PETSC_NODISCARD static PetscLogEvent CUPMSOLVER_HANDLE_CREATE() noexcept { return T == DeviceType::CUDA ? CUSOLVER_HANDLE_CREATE : HIPSOLVER_HANDLE_CREATE; }
89: // this exists purely to satisfy the compiler so the tag-based dispatch works for the other
90: // handles
91: static PetscErrorCode initialize_handle_(stream_tag, PetscDeviceContext) noexcept { return PETSC_SUCCESS; }
93: static PetscErrorCode initialize_handle_(blas_tag, PetscDeviceContext dctx) noexcept
94: {
95: const auto dci = impls_cast_(dctx);
96: auto &handle = blashandles_[dctx->device->deviceId];
98: PetscFunctionBegin;
99: if (!handle) {
100: PetscCall(PetscLogEventsPause());
101: PetscCall(PetscLogEventBegin(CUPMBLAS_HANDLE_CREATE(), 0, 0, 0, 0));
102: for (auto i = 0; i < 3; ++i) {
103: const auto cberr = cupmBlasCreate(handle.ptr_to());
104: if (PetscLikely(cberr == CUPMBLAS_STATUS_SUCCESS)) break;
105: if (PetscUnlikely(cberr != CUPMBLAS_STATUS_ALLOC_FAILED) && (cberr != CUPMBLAS_STATUS_NOT_INITIALIZED)) PetscCallCUPMBLAS(cberr);
106: if (i != 2) {
107: PetscCall(PetscSleep(3));
108: continue;
109: }
110: PetscCheck(cberr == CUPMBLAS_STATUS_SUCCESS, PETSC_COMM_SELF, PETSC_ERR_GPU_RESOURCE, "Unable to initialize %s", cupmBlasName());
111: }
112: PetscCall(PetscLogEventEnd(CUPMBLAS_HANDLE_CREATE(), 0, 0, 0, 0));
113: PetscCall(PetscLogEventsResume());
114: }
115: PetscCallCUPMBLAS(cupmBlasSetStream(handle, dci->stream.get_stream()));
116: dci->blas = handle;
117: PetscFunctionReturn(PETSC_SUCCESS);
118: }
120: static PetscErrorCode initialize_handle_(solver_tag, PetscDeviceContext dctx) noexcept
121: {
122: const auto dci = impls_cast_(dctx);
123: auto &handle = solverhandles_[dctx->device->deviceId];
125: PetscFunctionBegin;
126: if (!handle) {
127: PetscCall(PetscLogEventsPause());
128: PetscCall(PetscLogEventBegin(CUPMSOLVER_HANDLE_CREATE(), 0, 0, 0, 0));
129: for (auto i = 0; i < 3; ++i) {
130: const auto cerr = cupmSolverCreate(&handle);
131: if (PetscLikely(cerr == CUPMSOLVER_STATUS_SUCCESS)) break;
132: if (cerr != CUPMSOLVER_STATUS_NOT_INITIALIZED && cerr != CUPMSOLVER_STATUS_ALLOC_FAILED) PetscCallCUPMSOLVER(cerr);
133: if (i < 2) {
134: PetscCall(PetscSleep(3));
135: continue;
136: }
137: PetscCheck(cerr == CUPMSOLVER_STATUS_SUCCESS, PETSC_COMM_SELF, PETSC_ERR_GPU_RESOURCE, "Unable to initialize %s", cupmSolverName());
138: }
139: PetscCall(PetscLogEventEnd(CUPMSOLVER_HANDLE_CREATE(), 0, 0, 0, 0));
140: PetscCall(PetscLogEventsResume());
141: }
142: PetscCallCUPMSOLVER(cupmSolverSetStream(handle, dci->stream.get_stream()));
143: dci->solver = handle;
144: PetscFunctionReturn(PETSC_SUCCESS);
145: }
147: static PetscErrorCode check_current_device_(PetscDeviceContext dctxl, PetscDeviceContext dctxr) noexcept
148: {
149: const auto devidl = dctxl->device->deviceId, devidr = dctxr->device->deviceId;
151: PetscFunctionBegin;
152: PetscCheck(devidl == devidr, PETSC_COMM_SELF, PETSC_ERR_GPU, "Device contexts must be on the same device; dctx A (id %" PetscInt64_FMT " device id %" PetscInt_FMT ") dctx B (id %" PetscInt64_FMT " device id %" PetscInt_FMT ")",
153: PetscObjectCast(dctxl)->id, devidl, PetscObjectCast(dctxr)->id, devidr);
154: PetscCall(PetscDeviceCheckDeviceCount_Internal(devidl));
155: PetscCall(PetscDeviceCheckDeviceCount_Internal(devidr));
156: PetscCallCUPM(cupmSetDevice(static_cast<int>(devidl)));
157: PetscFunctionReturn(PETSC_SUCCESS);
158: }
160: static PetscErrorCode check_current_device_(PetscDeviceContext dctx) noexcept { return check_current_device_(dctx, dctx); }
162: static PetscErrorCode finalize_() noexcept
163: {
164: PetscFunctionBegin;
165: for (auto &&handle : blashandles_) {
166: if (handle) {
167: PetscCallCUPMBLAS(cupmBlasDestroy(handle));
168: handle = nullptr;
169: }
170: }
171: for (auto &&handle : solverhandles_) {
172: if (handle) {
173: PetscCallCUPMSOLVER(cupmSolverDestroy(handle));
174: handle = nullptr;
175: }
176: }
177: initialized_ = false;
178: PetscFunctionReturn(PETSC_SUCCESS);
179: }
181: template <typename Allocator, typename PoolType = ::Petsc::memory::SegmentedMemoryPool<typename Allocator::value_type, stream_type, Allocator, 256 * sizeof(PetscScalar)>>
182: PETSC_NODISCARD static PoolType &default_pool_() noexcept
183: {
184: static PoolType pool;
185: return pool;
186: }
188: static PetscErrorCode check_memtype_(PetscMemType mtype, const char mess[]) noexcept
189: {
190: PetscFunctionBegin;
191: PetscCheck(PetscMemTypeHost(mtype) || (mtype == PETSC_MEMTYPE_DEVICE) || (mtype == PETSC_MEMTYPE_CUPM()), PETSC_COMM_SELF, PETSC_ERR_SUP, "%s device context can only handle %s (pinned) host or device memory", cupmName(), mess);
192: PetscFunctionReturn(PETSC_SUCCESS);
193: }
195: public:
196: // All of these functions MUST be static in order to be callable from C, otherwise they
197: // get the implicit 'this' pointer tacked on
198: static PetscErrorCode destroy(PetscDeviceContext) noexcept;
199: static PetscErrorCode changeStreamType(PetscDeviceContext, PetscStreamType) noexcept;
200: static PetscErrorCode setUp(PetscDeviceContext) noexcept;
201: static PetscErrorCode query(PetscDeviceContext, PetscBool *) noexcept;
202: static PetscErrorCode waitForContext(PetscDeviceContext, PetscDeviceContext) noexcept;
203: static PetscErrorCode synchronize(PetscDeviceContext) noexcept;
204: template <typename Handle_t>
205: static PetscErrorCode getHandle(PetscDeviceContext, void *) noexcept;
206: template <typename Handle_t>
207: static PetscErrorCode getHandlePtr(PetscDeviceContext, void **) noexcept;
208: static PetscErrorCode beginTimer(PetscDeviceContext) noexcept;
209: static PetscErrorCode endTimer(PetscDeviceContext, PetscLogDouble *) noexcept;
210: static PetscErrorCode getPower(PetscDeviceContext, PetscLogDouble *) noexcept;
211: static PetscErrorCode beginEnergyMeter(PetscDeviceContext) noexcept;
212: static PetscErrorCode endEnergyMeter(PetscDeviceContext, PetscLogDouble *) noexcept;
213: static PetscErrorCode memAlloc(PetscDeviceContext, PetscBool, PetscMemType, std::size_t, std::size_t, void **) noexcept;
214: static PetscErrorCode memFree(PetscDeviceContext, PetscMemType, void **) noexcept;
215: static PetscErrorCode memCopy(PetscDeviceContext, void *PETSC_RESTRICT, const void *PETSC_RESTRICT, std::size_t, PetscDeviceCopyMode) noexcept;
216: static PetscErrorCode memSet(PetscDeviceContext, PetscMemType, void *, PetscInt, std::size_t) noexcept;
217: static PetscErrorCode createEvent(PetscDeviceContext, PetscEvent) noexcept;
218: static PetscErrorCode recordEvent(PetscDeviceContext, PetscEvent) noexcept;
219: static PetscErrorCode waitForEvent(PetscDeviceContext, PetscEvent) noexcept;
221: // not a PetscDeviceContext method, this registers the class
222: static PetscErrorCode initialize(PetscDevice) noexcept;
224: // clang-format off
225: static constexpr _DeviceContextOps ops = {
226: PetscDesignatedInitializer(destroy, destroy),
227: PetscDesignatedInitializer(changestreamtype, changeStreamType),
228: PetscDesignatedInitializer(setup, setUp),
229: PetscDesignatedInitializer(query, query),
230: PetscDesignatedInitializer(waitforcontext, waitForContext),
231: PetscDesignatedInitializer(synchronize, synchronize),
232: PetscDesignatedInitializer(getblashandle, getHandle<blas_tag>),
233: PetscDesignatedInitializer(getsolverhandle, getHandle<solver_tag>),
234: PetscDesignatedInitializer(getstreamhandle, getHandlePtr<stream_tag>),
235: PetscDesignatedInitializer(begintimer, beginTimer),
236: PetscDesignatedInitializer(endtimer, endTimer),
237: #if PetscDefined(HAVE_NVML)
238: PetscDesignatedInitializer(getpower, getPower),
239: PetscDesignatedInitializer(beginenergymeter, beginEnergyMeter),
240: PetscDesignatedInitializer(endenergymeter, endEnergyMeter),
241: #else
242: PetscDesignatedInitializer(getpower, nullptr),
243: PetscDesignatedInitializer(beginenergymeter, nullptr),
244: PetscDesignatedInitializer(endenergymeter, nullptr),
245: #endif
246: PetscDesignatedInitializer(memalloc, memAlloc),
247: PetscDesignatedInitializer(memfree, memFree),
248: PetscDesignatedInitializer(memcopy, memCopy),
249: PetscDesignatedInitializer(memset, memSet),
250: PetscDesignatedInitializer(createevent, createEvent),
251: PetscDesignatedInitializer(recordevent, recordEvent),
252: PetscDesignatedInitializer(waitforevent, waitForEvent)
253: };
254: // clang-format on
255: };
257: // not a PetscDeviceContext method, this initializes the CLASS
258: template <DeviceType T>
259: inline PetscErrorCode DeviceContext<T>::initialize(PetscDevice device) noexcept
260: {
261: PetscFunctionBegin;
262: if (PetscUnlikely(!initialized_)) {
263: uint64_t threshold = UINT64_MAX;
264: cupmMemPool_t mempool;
266: initialized_ = true;
267: PetscCallCUPM(cupmDeviceGetMemPool(&mempool, static_cast<int>(device->deviceId)));
268: PetscCallCUPM(cupmMemPoolSetAttribute(mempool, cupmMemPoolAttrReleaseThreshold, &threshold));
269: blashandles_.fill(nullptr);
270: solverhandles_.fill(nullptr);
271: PetscCall(PetscRegisterFinalize(finalize_));
272: }
273: PetscFunctionReturn(PETSC_SUCCESS);
274: }
276: template <DeviceType T>
277: inline PetscErrorCode DeviceContext<T>::destroy(PetscDeviceContext dctx) noexcept
278: {
279: PetscFunctionBegin;
280: if (const auto dci = impls_cast_(dctx)) {
281: PetscCall(dci->stream.destroy());
282: if (dci->event) PetscCall(cupm_fast_event_pool<T>().deallocate(&dci->event));
283: if (dci->begin) PetscCallCUPM(cupmEventDestroy(dci->begin));
284: if (dci->end) PetscCallCUPM(cupmEventDestroy(dci->end));
285: delete dci;
286: dctx->data = nullptr;
287: }
288: PetscFunctionReturn(PETSC_SUCCESS);
289: }
291: template <DeviceType T>
292: inline PetscErrorCode DeviceContext<T>::changeStreamType(PetscDeviceContext dctx, PETSC_UNUSED PetscStreamType stype) noexcept
293: {
294: const auto dci = impls_cast_(dctx);
296: PetscFunctionBegin;
297: PetscCall(dci->stream.destroy());
298: // set these to null so they aren't usable until setup is called again
299: dci->blas = nullptr;
300: dci->solver = nullptr;
301: PetscFunctionReturn(PETSC_SUCCESS);
302: }
304: template <DeviceType T>
305: inline PetscErrorCode DeviceContext<T>::setUp(PetscDeviceContext dctx) noexcept
306: {
307: const auto dci = impls_cast_(dctx);
308: auto &event = dci->event;
310: PetscFunctionBegin;
311: PetscCall(check_current_device_(dctx));
312: PetscCall(dci->stream.change_type(dctx->streamType));
313: if (!event) PetscCall(cupm_fast_event_pool<T>().allocate(&event));
314: #if PetscDefined(USE_DEBUG)
315: dci->timerInUse = PETSC_FALSE;
316: #endif
317: PetscFunctionReturn(PETSC_SUCCESS);
318: }
320: template <DeviceType T>
321: inline PetscErrorCode DeviceContext<T>::query(PetscDeviceContext dctx, PetscBool *idle) noexcept
322: {
323: PetscFunctionBegin;
324: PetscCall(check_current_device_(dctx));
325: switch (auto cerr = cupmStreamQuery(impls_cast_(dctx)->stream.get_stream())) {
326: case cupmSuccess:
327: *idle = PETSC_TRUE;
328: break;
329: case cupmErrorNotReady:
330: *idle = PETSC_FALSE;
331: // reset the error
332: cerr = cupmGetLastError();
333: static_cast<void>(cerr);
334: break;
335: default:
336: PetscCallCUPM(cerr);
337: PetscUnreachable();
338: }
339: PetscFunctionReturn(PETSC_SUCCESS);
340: }
342: template <DeviceType T>
343: inline PetscErrorCode DeviceContext<T>::waitForContext(PetscDeviceContext dctxa, PetscDeviceContext dctxb) noexcept
344: {
345: const auto dcib = impls_cast_(dctxb);
346: const auto event = dcib->event;
348: PetscFunctionBegin;
349: PetscCall(check_current_device_(dctxa, dctxb));
350: PetscCallCUPM(cupmEventRecord(event, dcib->stream.get_stream()));
351: PetscCallCUPM(cupmStreamWaitEvent(impls_cast_(dctxa)->stream.get_stream(), event, 0));
352: PetscFunctionReturn(PETSC_SUCCESS);
353: }
355: template <DeviceType T>
356: inline PetscErrorCode DeviceContext<T>::synchronize(PetscDeviceContext dctx) noexcept
357: {
358: auto idle = PETSC_TRUE;
360: PetscFunctionBegin;
361: PetscCall(query(dctx, &idle));
362: if (!idle) PetscCallCUPM(cupmStreamSynchronize(impls_cast_(dctx)->stream.get_stream()));
363: PetscFunctionReturn(PETSC_SUCCESS);
364: }
366: template <DeviceType T>
367: template <typename handle_t>
368: inline PetscErrorCode DeviceContext<T>::getHandle(PetscDeviceContext dctx, void *handle) noexcept
369: {
370: PetscFunctionBegin;
371: PetscCall(initialize_handle_(handle_t{}, dctx));
372: *static_cast<typename handle_t::type *>(handle) = impls_cast_(dctx)->get(handle_t{});
373: PetscFunctionReturn(PETSC_SUCCESS);
374: }
376: template <DeviceType T>
377: template <typename handle_t>
378: inline PetscErrorCode DeviceContext<T>::getHandlePtr(PetscDeviceContext dctx, void **handle) noexcept
379: {
380: using handle_type = typename handle_t::type;
382: PetscFunctionBegin;
383: PetscCall(initialize_handle_(handle_t{}, dctx));
384: *reinterpret_cast<handle_type **>(handle) = const_cast<handle_type *>(std::addressof(impls_cast_(dctx)->get(handle_t{})));
385: PetscFunctionReturn(PETSC_SUCCESS);
386: }
388: template <DeviceType T>
389: inline PetscErrorCode DeviceContext<T>::beginTimer(PetscDeviceContext dctx) noexcept
390: {
391: const auto dci = impls_cast_(dctx);
393: PetscFunctionBegin;
394: PetscCall(check_current_device_(dctx));
395: #if PetscDefined(USE_DEBUG)
396: PetscCheck(!dci->timerInUse, PETSC_COMM_SELF, PETSC_ERR_PLIB, "Forgot to call PetscLogGpuTimeEnd()?");
397: dci->timerInUse = PETSC_TRUE;
398: #endif
399: if (!dci->begin) {
400: PetscAssert(!dci->end, PETSC_COMM_SELF, PETSC_ERR_PLIB, "Don't have a 'begin' event, but somehow have an end event");
401: PetscCallCUPM(cupmEventCreate(&dci->begin));
402: PetscCallCUPM(cupmEventCreate(&dci->end));
403: }
404: PetscCallCUPM(cupmEventRecord(dci->begin, dci->stream.get_stream()));
405: PetscFunctionReturn(PETSC_SUCCESS);
406: }
408: template <DeviceType T>
409: inline PetscErrorCode DeviceContext<T>::endTimer(PetscDeviceContext dctx, PetscLogDouble *elapsed) noexcept
410: {
411: float gtime;
412: const auto dci = impls_cast_(dctx);
413: const auto end = dci->end;
415: PetscFunctionBegin;
416: PetscCall(check_current_device_(dctx));
417: #if PetscDefined(USE_DEBUG)
418: PetscCheck(dci->timerInUse, PETSC_COMM_SELF, PETSC_ERR_PLIB, "Forgot to call PetscLogGpuTimeBegin()?");
419: dci->timerInUse = PETSC_FALSE;
420: #endif
421: PetscCallCUPM(cupmEventRecord(end, dci->stream.get_stream()));
422: PetscCallCUPM(cupmEventSynchronize(end));
423: PetscCallCUPM(cupmEventElapsedTime(>ime, dci->begin, end));
424: *elapsed = static_cast<util::remove_pointer_t<decltype(elapsed)>>(gtime);
425: PetscFunctionReturn(PETSC_SUCCESS);
426: }
428: #if PetscDefined(HAVE_NVML)
429: template <DeviceType T>
430: inline PetscErrorCode DeviceContext<T>::getPower(PetscDeviceContext dctx, PetscLogDouble *power) noexcept
431: {
432: const auto dci = impls_cast_(dctx);
433: nvmlFieldValue_t values[1];
435: PetscFunctionBegin;
436: PetscCall(check_current_device_(dctx));
437: PetscCallCUPM(cupmStreamSynchronize(dci->stream.get_stream()));
438: values[0].fieldId = NVML_FI_DEV_POWER_INSTANT;
439: if (!dci->nvmlHandle) PetscCallNVML(nvmlDeviceGetHandleByIndex(dctx->device->deviceId, &dci->nvmlHandle));
440: PetscCallNVML(nvmlDeviceGetFieldValues(dci->nvmlHandle, 1, values));
441: *power = static_cast<util::remove_pointer_t<decltype(power)>>(values[0].value.uiVal);
442: PetscFunctionReturn(PETSC_SUCCESS);
443: }
445: template <DeviceType T>
446: inline PetscErrorCode DeviceContext<T>::beginEnergyMeter(PetscDeviceContext dctx) noexcept
447: {
448: const auto dci = impls_cast_(dctx);
450: PetscFunctionBegin;
451: PetscCall(check_current_device_(dctx));
452: #if PetscDefined(USE_DEBUG)
453: PetscCheck(!dci->EnergyMeterInUse, PETSC_COMM_SELF, PETSC_ERR_PLIB, "Forgot to call PetscLogGpuEnergyMeterEnd()?");
454: dci->EnergyMeterInUse = PETSC_TRUE;
455: #endif
456: if (!dci->nvmlHandle) PetscCallNVML(nvmlDeviceGetHandleByIndex(dctx->device->deviceId, &dci->nvmlHandle));
457: PetscCallNVML(nvmlDeviceGetTotalEnergyConsumption(dci->nvmlHandle, &dci->energymeterbegin));
458: PetscFunctionReturn(PETSC_SUCCESS);
459: }
461: template <DeviceType T>
462: inline PetscErrorCode DeviceContext<T>::endEnergyMeter(PetscDeviceContext dctx, PetscLogDouble *energy) noexcept
463: {
464: const auto dci = impls_cast_(dctx);
466: PetscFunctionBegin;
467: PetscCall(check_current_device_(dctx));
468: #if PetscDefined(USE_DEBUG)
469: PetscCheck(dci->EnergyMeterInUse, PETSC_COMM_SELF, PETSC_ERR_PLIB, "Forgot to call PetscLogGpuEnergyMeterBegin()?");
470: dci->EnergyMeterInUse = PETSC_FALSE;
471: #endif
472: PetscCallCUPM(cupmStreamSynchronize(dci->stream.get_stream()));
473: PetscCallNVML(nvmlDeviceGetTotalEnergyConsumption(dci->nvmlHandle, &dci->energymeterend));
474: *energy = static_cast<util::remove_pointer_t<decltype(energy)>>(dci->energymeterend - dci->energymeterbegin) / 1000; // convert to Joule
475: PetscFunctionReturn(PETSC_SUCCESS);
476: }
477: #endif
479: template <DeviceType T>
480: inline PetscErrorCode DeviceContext<T>::memAlloc(PetscDeviceContext dctx, PetscBool clear, PetscMemType mtype, std::size_t n, std::size_t alignment, void **dest) noexcept
481: {
482: const auto &stream = impls_cast_(dctx)->stream;
484: PetscFunctionBegin;
485: PetscCall(check_current_device_(dctx));
486: PetscCall(check_memtype_(mtype, "allocating"));
487: if (PetscMemTypeHost(mtype)) {
488: PetscCall(default_pool_<HostAllocator<T>>().allocate(n, reinterpret_cast<char **>(dest), &stream, alignment));
489: } else {
490: PetscCall(default_pool_<DeviceAllocator<T>>().allocate(n, reinterpret_cast<char **>(dest), &stream, alignment));
491: }
492: if (clear) PetscCallCUPM(cupmMemsetAsync(*dest, 0, n, stream.get_stream()));
493: PetscFunctionReturn(PETSC_SUCCESS);
494: }
496: template <DeviceType T>
497: inline PetscErrorCode DeviceContext<T>::memFree(PetscDeviceContext dctx, PetscMemType mtype, void **ptr) noexcept
498: {
499: const auto &stream = impls_cast_(dctx)->stream;
501: PetscFunctionBegin;
502: PetscCall(check_current_device_(dctx));
503: PetscCall(check_memtype_(mtype, "freeing"));
504: if (!*ptr) PetscFunctionReturn(PETSC_SUCCESS);
505: if (PetscMemTypeHost(mtype)) {
506: PetscCall(default_pool_<HostAllocator<T>>().deallocate(reinterpret_cast<char **>(ptr), &stream));
507: // if ptr exists still exists the pool didn't own it
508: if (*ptr) {
509: auto registered = PETSC_FALSE, managed = PETSC_FALSE;
511: PetscCall(PetscCUPMGetMemType(*ptr, nullptr, ®istered, &managed));
512: if (registered) {
513: PetscCallCUPM(cupmFreeHost(*ptr));
514: } else if (managed) {
515: PetscCallCUPM(cupmFreeAsync(*ptr, stream.get_stream()));
516: }
517: }
518: } else {
519: PetscCall(default_pool_<DeviceAllocator<T>>().deallocate(reinterpret_cast<char **>(ptr), &stream));
520: // if ptr still exists the pool didn't own it
521: if (*ptr) PetscCallCUPM(cupmFreeAsync(*ptr, stream.get_stream()));
522: }
523: PetscFunctionReturn(PETSC_SUCCESS);
524: }
526: template <DeviceType T>
527: inline PetscErrorCode DeviceContext<T>::memCopy(PetscDeviceContext dctx, void *PETSC_RESTRICT dest, const void *PETSC_RESTRICT src, std::size_t n, PetscDeviceCopyMode mode) noexcept
528: {
529: const auto stream = impls_cast_(dctx)->stream.get_stream();
531: PetscFunctionBegin;
532: // can't use PetscCUPMMemcpyAsync here since we don't know sizeof(*src)...
533: if (mode == PETSC_DEVICE_COPY_HTOH) {
534: const auto cerr = cupmStreamQuery(stream);
536: // yes this is faster
537: if (cerr == cupmSuccess) {
538: PetscCall(PetscMemcpy(dest, src, n));
539: PetscFunctionReturn(PETSC_SUCCESS);
540: } else if (cerr == cupmErrorNotReady) {
541: auto PETSC_UNUSED unused = cupmGetLastError();
543: static_cast<void>(unused);
544: } else {
545: PetscCallCUPM(cerr);
546: }
547: }
548: PetscCallCUPM(cupmMemcpyAsync(dest, src, n, PetscDeviceCopyModeToCUPMMemcpyKind(mode), stream));
549: PetscFunctionReturn(PETSC_SUCCESS);
550: }
552: template <DeviceType T>
553: inline PetscErrorCode DeviceContext<T>::memSet(PetscDeviceContext dctx, PetscMemType mtype, void *ptr, PetscInt v, std::size_t n) noexcept
554: {
555: PetscFunctionBegin;
556: PetscCall(check_current_device_(dctx));
557: PetscCall(check_memtype_(mtype, "zeroing"));
558: PetscCallCUPM(cupmMemsetAsync(ptr, static_cast<int>(v), n, impls_cast_(dctx)->stream.get_stream()));
559: PetscFunctionReturn(PETSC_SUCCESS);
560: }
562: template <DeviceType T>
563: inline PetscErrorCode DeviceContext<T>::createEvent(PetscDeviceContext, PetscEvent event) noexcept
564: {
565: PetscFunctionBegin;
566: PetscCallCXX(event->data = new event_type{});
567: event->destroy = [](PetscEvent event) {
568: PetscFunctionBegin;
569: delete event_cast_(event);
570: event->data = nullptr;
571: PetscFunctionReturn(PETSC_SUCCESS);
572: };
573: PetscFunctionReturn(PETSC_SUCCESS);
574: }
576: template <DeviceType T>
577: inline PetscErrorCode DeviceContext<T>::recordEvent(PetscDeviceContext dctx, PetscEvent event) noexcept
578: {
579: PetscFunctionBegin;
580: PetscCall(impls_cast_(dctx)->stream.record_event(*event_cast_(event)));
581: PetscFunctionReturn(PETSC_SUCCESS);
582: }
584: template <DeviceType T>
585: inline PetscErrorCode DeviceContext<T>::waitForEvent(PetscDeviceContext dctx, PetscEvent event) noexcept
586: {
587: PetscFunctionBegin;
588: PetscCall(impls_cast_(dctx)->stream.wait_for_event(*event_cast_(event)));
589: PetscFunctionReturn(PETSC_SUCCESS);
590: }
592: // initialize the static member variables
593: template <DeviceType T>
594: bool DeviceContext<T>::initialized_ = false;
596: template <DeviceType T>
597: std::array<typename DeviceContext<T>::cupmBlasHandle_t, PETSC_DEVICE_MAX_DEVICES> DeviceContext<T>::blashandles_ = {};
599: template <DeviceType T>
600: std::array<typename DeviceContext<T>::cupmSolverHandle_t, PETSC_DEVICE_MAX_DEVICES> DeviceContext<T>::solverhandles_ = {};
602: template <DeviceType T>
603: constexpr _DeviceContextOps DeviceContext<T>::ops;
605: } // namespace impl
607: // shorten this one up a bit (and instantiate the templates)
608: using CUPMContextCuda = impl::DeviceContext<DeviceType::CUDA>;
609: using CUPMContextHip = impl::DeviceContext<DeviceType::HIP>;
611: // shorthand for what is an EXTREMELY long name
612: #define PetscDeviceContext_(IMPLS) ::Petsc::device::cupm::impl::DeviceContext<::Petsc::device::cupm::DeviceType::IMPLS>::PetscDeviceContext_IMPLS
614: } // namespace cupm
616: } // namespace device
618: } // namespace Petsc