diff --git a/backends/cuda-gen/ceed-cuda-gen-operator.c b/backends/cuda-gen/ceed-cuda-gen-operator.c index f8228fbfbd..249bc14628 100644 --- a/backends/cuda-gen/ceed-cuda-gen-operator.c +++ b/backends/cuda-gen/ceed-cuda-gen-operator.c @@ -7,10 +7,13 @@ #include #include +#include #include #include #include #include +#include +#include #include #include "../cuda/ceed-cuda-common.h" @@ -31,7 +34,20 @@ static int CeedOperatorDestroy_Cuda_gen(CeedOperator op) { if (impl->module_assemble_full) CeedCallCuda(ceed, cuModuleUnload(impl->module_assemble_full)); if (impl->module_assemble_diagonal) CeedCallCuda(ceed, cuModuleUnload(impl->module_assemble_diagonal)); if (impl->module_assemble_qfunction) CeedCallCuda(ceed, cuModuleUnload(impl->module_assemble_qfunction)); - if (impl->points.num_per_elem) CeedCallCuda(ceed, cudaFree((void **)impl->points.num_per_elem)); + if (impl->points.num_per_elem) CeedCallCuda(ceed, cudaFree((void *)impl->points.num_per_elem)); + + if (impl->graph_instance) { + CeedCallCuda(ceed, cudaGraphExecDestroy(impl->graph_instance)); + impl->graph_instance = NULL; + } + if (impl->graph) { + CeedCallCuda(ceed, cudaGraphDestroy(impl->graph)); + impl->graph = NULL; + } + impl->graph_created = false; + impl->captured_input_ptr = NULL; + impl->captured_output_ptr = NULL; + CeedCallBackend(CeedFree(&impl)); CeedCallBackend(CeedDestroy(&ceed)); return CEED_ERROR_SUCCESS; @@ -284,7 +300,10 @@ static int CeedOperatorApplyAdd_Cuda_gen(CeedOperator op, CeedVector input_vec, // Try to run kernel if (input_vec != CEED_VECTOR_NONE) CeedCallBackend(CeedVectorGetArrayRead(input_vec, CEED_MEM_DEVICE, &input_arr)); if (output_vec != CEED_VECTOR_NONE) CeedCallBackend(CeedVectorGetArray(output_vec, CEED_MEM_DEVICE, &output_arr)); - CeedCallBackend(CeedOperatorApplyAddCore_Cuda_gen(op, NULL, input_arr, output_arr, &is_run_good, request)); + enum cudaStreamCaptureStatus capture_status; + cudaStreamIsCapturing(cudaStreamPerThread, &capture_status); + CUstream stream_to_use = (capture_status != cudaStreamCaptureStatusNone) ? cudaStreamPerThread : NULL; + CeedCallBackend(CeedOperatorApplyAddCore_Cuda_gen(op, stream_to_use, input_arr, output_arr, &is_run_good, request)); if (input_vec != CEED_VECTOR_NONE) CeedCallBackend(CeedVectorRestoreArrayRead(input_vec, &input_arr)); if (output_vec != CEED_VECTOR_NONE) CeedCallBackend(CeedVectorRestoreArray(output_vec, &output_arr)); @@ -299,7 +318,58 @@ static int CeedOperatorApplyAdd_Cuda_gen(CeedOperator op, CeedVector input_vec, return CEED_ERROR_SUCCESS; } -static int CeedOperatorApplyAddComposite_Cuda_gen(CeedOperator op, CeedVector input_vec, CeedVector output_vec, CeedRequest *request) { +// Push passive inputs and QFunction context to device before graph replay. +static int CeedCompositeRefreshForReplay_Cuda_gen(CeedOperator *sub_operators, CeedInt num_suboperators) { + for (CeedInt i = 0; i < num_suboperators; i++) { + bool is_at_points; + CeedInt num_input_fields, num_output_fields; + CeedOperatorField *op_input_fields, *op_output_fields; + CeedQFunction qf = NULL; + CeedQFunctionField *qf_input_fields; + void *d_c = NULL; + + CeedCallBackend(CeedOperatorGetFields(sub_operators[i], &num_input_fields, &op_input_fields, &num_output_fields, &op_output_fields)); + CeedCallBackend(CeedOperatorGetQFunction(sub_operators[i], &qf)); + CeedCallBackend(CeedQFunctionGetFields(qf, NULL, &qf_input_fields, NULL, NULL)); + + for (CeedInt j = 0; j < num_input_fields; j++) { + CeedEvalMode eval_mode; + + CeedCallBackend(CeedQFunctionFieldGetEvalMode(qf_input_fields[j], &eval_mode)); + if (eval_mode == CEED_EVAL_WEIGHT) continue; + { + const CeedScalar *arr; + CeedVector vec; + + CeedCallBackend(CeedOperatorFieldGetVector(op_input_fields[j], &vec)); + if (vec != CEED_VECTOR_ACTIVE && vec != CEED_VECTOR_NONE) { + CeedCallBackend(CeedVectorGetArrayRead(vec, CEED_MEM_DEVICE, &arr)); + CeedCallBackend(CeedVectorRestoreArrayRead(vec, &arr)); + } + CeedCallBackend(CeedVectorDestroy(&vec)); + } + } + + CeedCallBackend(CeedOperatorIsAtPoints(sub_operators[i], &is_at_points)); + if (is_at_points) { + const CeedScalar *arr; + CeedVector vec; + + CeedCallBackend(CeedOperatorAtPointsGetPoints(sub_operators[i], NULL, &vec)); + CeedCallBackend(CeedVectorGetArrayRead(vec, CEED_MEM_DEVICE, &arr)); + CeedCallBackend(CeedVectorRestoreArrayRead(vec, &arr)); + CeedCallBackend(CeedVectorDestroy(&vec)); + } + + CeedCallBackend(CeedQFunctionGetInnerContextData(qf, CEED_MEM_DEVICE, &d_c)); + CeedCallBackend(CeedQFunctionRestoreInnerContextData(qf, &d_c)); + CeedCallBackend(CeedQFunctionDestroy(&qf)); + } + return CEED_ERROR_SUCCESS; +} + +// Composite apply without CUDA graphs. +static int CeedOperatorApplyAddComposite_NoGraph_Cuda_gen(CeedOperator op, CeedVector input_vec, CeedVector output_vec, CeedRequest *request) { bool is_run_good[CEED_COMPOSITE_MAX] = {false}, is_sequential; CeedInt num_suboperators; const CeedScalar *input_arr = NULL; @@ -309,16 +379,16 @@ static int CeedOperatorApplyAddComposite_Cuda_gen(CeedOperator op, CeedVector in cudaStream_t stream = NULL; CeedCallBackend(CeedOperatorGetCeed(op, &ceed)); - CeedCall(CeedOperatorCompositeGetNumSub(op, &num_suboperators)); - CeedCall(CeedOperatorCompositeGetSubList(op, &sub_operators)); - CeedCall(CeedOperatorCompositeIsSequential(op, &is_sequential)); + CeedCallBackend(CeedOperatorCompositeGetNumSub(op, &num_suboperators)); + CeedCallBackend(CeedOperatorCompositeGetSubList(op, &sub_operators)); + CeedCallBackend(CeedOperatorCompositeIsSequential(op, &is_sequential)); if (input_vec != CEED_VECTOR_NONE) CeedCallBackend(CeedVectorGetArrayRead(input_vec, CEED_MEM_DEVICE, &input_arr)); if (output_vec != CEED_VECTOR_NONE) CeedCallBackend(CeedVectorGetArray(output_vec, CEED_MEM_DEVICE, &output_arr)); if (is_sequential) CeedCallCuda(ceed, cudaStreamCreate(&stream)); for (CeedInt i = 0; i < num_suboperators; i++) { CeedInt num_elem = 0; - CeedCall(CeedOperatorGetNumElements(sub_operators[i], &num_elem)); + CeedCallBackend(CeedOperatorGetNumElements(sub_operators[i], &num_elem)); if (num_elem > 0) { if (!is_sequential) CeedCallCuda(ceed, cudaStreamCreate(&stream)); CeedCallBackend(CeedOperatorApplyAddCore_Cuda_gen(sub_operators[i], stream, input_arr, output_arr, &is_run_good[i], request)); @@ -330,7 +400,7 @@ static int CeedOperatorApplyAddComposite_Cuda_gen(CeedOperator op, CeedVector in if (output_vec != CEED_VECTOR_NONE) CeedCallBackend(CeedVectorRestoreArray(output_vec, &output_arr)); CeedCallCuda(ceed, cudaDeviceSynchronize()); - // Fallback on unsuccessful run + // Fall back to /gpu/cuda/ref for any sub-operator that couldn't run here for (CeedInt i = 0; i < num_suboperators; i++) { if (!is_run_good[i]) { CeedOperator op_fallback; @@ -344,6 +414,134 @@ static int CeedOperatorApplyAddComposite_Cuda_gen(CeedOperator op, CeedVector in return CEED_ERROR_SUCCESS; } +static int CeedOperatorApplyAddComposite_Cuda_gen(CeedOperator op, CeedVector input_vec, CeedVector output_vec, CeedRequest *request) { + Ceed ceed; + CeedOperator_Cuda_gen *impl; + CeedOperator *sub_operators; + CeedInt num_suboperators; + + ceed = CeedOperatorReturnCeed(op); + CeedCallBackend(CeedOperatorCompositeGetNumSub(op, &num_suboperators)); + CeedCallBackend(CeedOperatorCompositeGetSubList(op, &sub_operators)); + CeedCallBackend(CeedOperatorGetData(op, &impl)); + + if (!impl->use_graph || (input_vec == CEED_VECTOR_NONE && output_vec == CEED_VECTOR_NONE)) { + return CeedOperatorApplyAddComposite_NoGraph_Cuda_gen(op, input_vec, output_vec, request); + } + + if (!impl->warmup_done) { + CeedCallBackend(CeedOperatorApplyAddComposite_NoGraph_Cuda_gen(op, input_vec, output_vec, request)); + impl->warmup_done = true; + return CEED_ERROR_SUCCESS; + } + + bool need_build = !impl->graph_created; + + if (!need_build && input_vec != CEED_VECTOR_NONE) { + const CeedScalar *in_ptr; + + CeedCallBackend(CeedVectorGetArrayRead(input_vec, CEED_MEM_DEVICE, &in_ptr)); + need_build = in_ptr != impl->captured_input_ptr; + CeedCallBackend(CeedVectorRestoreArrayRead(input_vec, &in_ptr)); + } + if (!need_build && output_vec != CEED_VECTOR_NONE) { + CeedScalar *out_ptr; + + CeedCallBackend(CeedVectorGetArray(output_vec, CEED_MEM_DEVICE, &out_ptr)); + need_build = out_ptr != impl->captured_output_ptr; + CeedCallBackend(CeedVectorRestoreArray(output_vec, &out_ptr)); + } + + if (need_build) { + const CeedScalar *input_arr = NULL; + CeedScalar *output_arr = NULL; + cudaStream_t capture_stream = cudaStreamPerThread; + cudaGraph_t graph = NULL; + bool capture_ok = true; + cudaError_t err; + + if (impl->graph_instance) CeedCallCuda(ceed, cudaGraphExecDestroy(impl->graph_instance)); + if (impl->graph) CeedCallCuda(ceed, cudaGraphDestroy(impl->graph)); + impl->graph = NULL; + impl->graph_instance = NULL; + + if (input_vec != CEED_VECTOR_NONE) CeedCallBackend(CeedVectorGetArrayRead(input_vec, CEED_MEM_DEVICE, &input_arr)); + if (output_vec != CEED_VECTOR_NONE) CeedCallBackend(CeedVectorGetArray(output_vec, CEED_MEM_DEVICE, &output_arr)); + impl->captured_input_ptr = input_arr; + impl->captured_output_ptr = output_arr; + + err = cudaStreamBeginCapture(capture_stream, cudaStreamCaptureModeThreadLocal); + if (err != cudaSuccess) capture_ok = false; + if (capture_ok) { + // Still call EndCapture if capture is invalidated mid-way. + for (CeedInt i = 0; i < num_suboperators; i++) { + bool is_run_good = true; + + if (CeedOperatorApplyAddCore_Cuda_gen(sub_operators[i], capture_stream, input_arr, output_arr, &is_run_good, request) || !is_run_good) { + capture_ok = false; + break; + } + } + } + err = cudaStreamEndCapture(capture_stream, &graph); + + if (capture_ok && (err != cudaSuccess || !graph)) capture_ok = false; + if (capture_ok) { + impl->graph = graph; + if (cudaGraphInstantiate(&impl->graph_instance, impl->graph, 0) != cudaSuccess) { + CeedCallCuda(ceed, cudaGraphDestroy(impl->graph)); + impl->graph = NULL; + capture_ok = false; + } + } else if (graph) { + CeedCallCuda(ceed, cudaGraphDestroy(graph)); + } + + if (input_vec != CEED_VECTOR_NONE) CeedCallBackend(CeedVectorRestoreArrayRead(input_vec, &input_arr)); + if (output_vec != CEED_VECTOR_NONE) CeedCallBackend(CeedVectorRestoreArray(output_vec, &output_arr)); + + if (!capture_ok) { + cudaGetLastError(); + CeedCallCuda(ceed, cudaDeviceSynchronize()); + cudaGetLastError(); + impl->graph_created = false; + impl->captured_input_ptr = NULL; + impl->captured_output_ptr = NULL; + CeedCallBackend(CeedOperatorSetEnableCudaGraph(op, false)); + return CeedOperatorApplyAddComposite_NoGraph_Cuda_gen(op, input_vec, output_vec, request); + } + impl->graph_created = true; + } + + if (input_vec != CEED_VECTOR_NONE) { + const CeedScalar *in_arr; + + CeedCallBackend(CeedVectorGetArrayRead(input_vec, CEED_MEM_DEVICE, &in_arr)); + CeedCallBackend(CeedVectorRestoreArrayRead(input_vec, &in_arr)); + } + if (output_vec != CEED_VECTOR_NONE) { + CeedScalar *out_arr; + + CeedCallBackend(CeedVectorGetArray(output_vec, CEED_MEM_DEVICE, &out_arr)); + CeedCallBackend(CeedVectorRestoreArray(output_vec, &out_arr)); + } + CeedCallBackend(CeedCompositeRefreshForReplay_Cuda_gen(sub_operators, num_suboperators)); + + if (cudaGraphLaunch(impl->graph_instance, NULL) != cudaSuccess) { + cudaGetLastError(); + if (impl->graph_instance) CeedCallCuda(ceed, cudaGraphExecDestroy(impl->graph_instance)); + if (impl->graph) CeedCallCuda(ceed, cudaGraphDestroy(impl->graph)); + impl->graph = NULL; + impl->graph_instance = NULL; + impl->graph_created = false; + impl->captured_input_ptr = NULL; + impl->captured_output_ptr = NULL; + CeedCallBackend(CeedOperatorSetEnableCudaGraph(op, false)); + return CeedOperatorApplyAddComposite_NoGraph_Cuda_gen(op, input_vec, output_vec, request); + } + return CEED_ERROR_SUCCESS; +} + //------------------------------------------------------------------------------ // QFunction assembly //------------------------------------------------------------------------------ @@ -465,7 +663,7 @@ static int CeedOperatorLinearAssembleQFunctionCore_Cuda_gen(CeedOperator op, boo // Assemble QFunction void *opargs[] = {(void *)&num_elem, &qf_data->d_c, &data->indices, &data->fields, &data->B, &data->G, &data->W, &data->points, &assembled_array}; - bool is_tensor = false; + bool is_tensor; int max_threads_per_block, min_grid_size, grid; CeedCallBackend(CeedOperatorHasTensorBases(op, &is_tensor)); @@ -874,6 +1072,17 @@ static int CeedOperatorAssembleSingleAtPoints_Cuda_gen(CeedOperator op, CeedInt return CEED_ERROR_SUCCESS; } +//------------------------------------------------------------------------------ +// Set CUDA Graph use +//------------------------------------------------------------------------------ +static int CeedOperatorSetEnableCudaGraph_Cuda_gen(CeedOperator op, bool enable_graph) { + CeedOperator_Cuda_gen *impl; + + CeedCallBackend(CeedOperatorGetData(op, &impl)); + impl->use_graph = enable_graph; + return CEED_ERROR_SUCCESS; +} + //------------------------------------------------------------------------------ // Create operator //------------------------------------------------------------------------------ @@ -885,24 +1094,36 @@ int CeedOperatorCreate_Cuda_gen(CeedOperator op) { CeedCallBackend(CeedOperatorGetCeed(op, &ceed)); CeedCallBackend(CeedCalloc(1, &impl)); CeedCallBackend(CeedOperatorSetData(op, impl)); - CeedCall(CeedOperatorIsComposite(op, &is_composite)); + + CeedCallBackend(CeedOperatorIsComposite(op, &is_composite)); if (is_composite) { CeedCallBackend(CeedSetBackendFunction(ceed, "Operator", op, "ApplyAddComposite", CeedOperatorApplyAddComposite_Cuda_gen)); } else { CeedCallBackend(CeedSetBackendFunction(ceed, "Operator", op, "ApplyAdd", CeedOperatorApplyAdd_Cuda_gen)); } - CeedCall(CeedOperatorIsAtPoints(op, &is_at_points)); + CeedCallBackend(CeedOperatorIsAtPoints(op, &is_at_points)); if (is_at_points) { CeedCallBackend(CeedSetBackendFunction(ceed, "Operator", op, "LinearAssembleAddDiagonal", CeedOperatorLinearAssembleAddDiagonalAtPoints_Cuda_gen)); CeedCallBackend(CeedSetBackendFunction(ceed, "Operator", op, "LinearAssembleSingle", CeedOperatorAssembleSingleAtPoints_Cuda_gen)); } + if (!is_at_points) { CeedCallBackend(CeedSetBackendFunction(ceed, "Operator", op, "LinearAssembleQFunction", CeedOperatorLinearAssembleQFunction_Cuda_gen)); CeedCallBackend(CeedSetBackendFunction(ceed, "Operator", op, "LinearAssembleQFunctionUpdate", CeedOperatorLinearAssembleQFunctionUpdate_Cuda_gen)); } + CeedCallBackend(CeedSetBackendFunction(ceed, "Operator", op, "SetEnableCudaGraph", CeedOperatorSetEnableCudaGraph_Cuda_gen)); CeedCallBackend(CeedSetBackendFunction(ceed, "Operator", op, "Destroy", CeedOperatorDestroy_Cuda_gen)); + + { + const char *env_val = getenv("CEED_ENABLE_CUDA_GRAPH"); + bool enable_graph = true; + + if (env_val) enable_graph = strcmp(env_val, "0") && strcmp(env_val, "false"); + CeedCallBackend(CeedOperatorSetEnableCudaGraph(op, enable_graph)); + } + CeedCallBackend(CeedDestroy(&ceed)); return CEED_ERROR_SUCCESS; } diff --git a/backends/cuda-gen/ceed-cuda-gen.h b/backends/cuda-gen/ceed-cuda-gen.h index 0e04f3c4e4..5ec193095a 100644 --- a/backends/cuda-gen/ceed-cuda-gen.h +++ b/backends/cuda-gen/ceed-cuda-gen.h @@ -10,6 +10,7 @@ #include #include #include +#include typedef struct { bool use_fallback, use_assembly_fallback; @@ -25,6 +26,14 @@ typedef struct { Fields_Cuda G; CeedScalar *W; Points_Cuda points; + + bool use_graph; + bool graph_created; + bool warmup_done; + cudaGraph_t graph; + cudaGraphExec_t graph_instance; + const CeedScalar *captured_input_ptr; + CeedScalar *captured_output_ptr; } CeedOperator_Cuda_gen; typedef struct { diff --git a/backends/cuda-ref/ceed-cuda-ref-qfunctioncontext.c b/backends/cuda-ref/ceed-cuda-ref-qfunctioncontext.c index 491e658338..e608cce9f3 100644 --- a/backends/cuda-ref/ceed-cuda-ref-qfunctioncontext.c +++ b/backends/cuda-ref/ceed-cuda-ref-qfunctioncontext.c @@ -36,7 +36,17 @@ static inline int CeedQFunctionContextSyncH2D_Cuda(const CeedQFunctionContext ct CeedCallCuda(ceed, cudaMalloc((void **)&impl->d_data_owned, ctx_size)); impl->d_data = impl->d_data_owned; } - CeedCallCuda(ceed, cudaMemcpy(impl->d_data, impl->h_data, ctx_size, cudaMemcpyHostToDevice)); + + // Use async memcpy during CUDA Graph capture for compatibility + enum cudaStreamCaptureStatus capture_status; + + cudaStreamIsCapturing(cudaStreamPerThread, &capture_status); + if (capture_status != cudaStreamCaptureStatusNone) { + CeedCallCuda(ceed, cudaMemcpyAsync(impl->d_data, impl->h_data, ctx_size, cudaMemcpyHostToDevice, cudaStreamPerThread)); + } else { + CeedCallCuda(ceed, cudaMemcpy(impl->d_data, impl->h_data, ctx_size, cudaMemcpyHostToDevice)); + } + CeedCallBackend(CeedDestroy(&ceed)); return CEED_ERROR_SUCCESS; } diff --git a/backends/cuda-ref/ceed-cuda-ref-vector.c b/backends/cuda-ref/ceed-cuda-ref-vector.c index 980d2f0583..3ede631308 100644 --- a/backends/cuda-ref/ceed-cuda-ref-vector.c +++ b/backends/cuda-ref/ceed-cuda-ref-vector.c @@ -326,7 +326,17 @@ static int CeedVectorSetValue_Cuda(CeedVector vec, CeedScalar val) { } if (impl->d_array) { if (val == 0) { - CeedCallCuda(CeedVectorReturnCeed(vec), cudaMemset(impl->d_array, 0, length * sizeof(CeedScalar))); + // Check if we're in CUDA Graph capture mode + enum cudaStreamCaptureStatus capture_status; + + cudaStreamIsCapturing(cudaStreamPerThread, &capture_status); + if (capture_status != cudaStreamCaptureStatusNone) { + // During capture, use async memset with cudaStreamPerThread + CeedCallCuda(CeedVectorReturnCeed(vec), cudaMemsetAsync(impl->d_array, 0, length * sizeof(CeedScalar), cudaStreamPerThread)); + } else { + // Normal execution, use blocking memset + CeedCallCuda(CeedVectorReturnCeed(vec), cudaMemset(impl->d_array, 0, length * sizeof(CeedScalar))); + } } else { CeedCallBackend(CeedDeviceSetValue_Cuda(impl->d_array, length, val)); } diff --git a/include/ceed-impl.h b/include/ceed-impl.h index 549dde8308..3135d571ae 100644 --- a/include/ceed-impl.h +++ b/include/ceed-impl.h @@ -368,6 +368,7 @@ struct CeedOperator_private { int (*ApplyAdd)(CeedOperator, CeedVector, CeedVector, CeedRequest *); int (*ApplyAddComposite)(CeedOperator, CeedVector, CeedVector, CeedRequest *); int (*ApplyJacobian)(CeedOperator, CeedVector, CeedVector, CeedVector, CeedVector, CeedRequest *); + int (*SetEnableCudaGraph)(CeedOperator, bool); int (*Destroy)(CeedOperator); CeedOperatorField *input_fields; CeedOperatorField *output_fields; diff --git a/include/ceed/cuda.h b/include/ceed/cuda.h index eb9ac3e9cb..7ef3a9abfe 100644 --- a/include/ceed/cuda.h +++ b/include/ceed/cuda.h @@ -13,3 +13,4 @@ #include CEED_EXTERN int CeedQFunctionSetCUDAUserFunction(CeedQFunction qf, CUfunction f); +CEED_EXTERN int CeedOperatorSetEnableCudaGraph(CeedOperator op, bool enable_graph); diff --git a/interface/ceed-cuda.c b/interface/ceed-cuda.c index ea15d46735..d3bda703cf 100644 --- a/interface/ceed-cuda.c +++ b/interface/ceed-cuda.c @@ -12,7 +12,10 @@ #include /** - @brief Set CUDA function pointer to evaluate action at quadrature points + @brief Set CUDA function pointer to evaluate action at quadrature points. + + If the backend does not support `CUfunction` pointers for QFunctions, then the call succeeds without effect. + When unsupported, a message is emitted via `CeedDebug`. @param[in,out] qf `CeedQFunction` to set device pointer @param[in] f Device function pointer to evaluate action at quadrature points @@ -29,3 +32,25 @@ int CeedQFunctionSetCUDAUserFunction(CeedQFunction qf, CUfunction f) { } return CEED_ERROR_SUCCESS; } + +/** + @brief Enable or disable CUDA Graph capture/replay for a `CeedOperator`. + + If the backend does not support CUDA Graphs for operators, then the call succeeds without effect. + When unsupported, a message is emitted via `CeedDebug`. + + @param[in,out] op `CeedOperator` + @param[in] enable_graph Boolean flag to enable CUDA Graph use + + @return An error code: 0 - success, otherwise - failure + + @ref User +**/ +int CeedOperatorSetEnableCudaGraph(CeedOperator op, bool enable_graph) { + if (!op->SetEnableCudaGraph) { + CeedDebug(CeedOperatorReturnCeed(op), "Backend does not support CUDA Graphs for operators."); + } else { + CeedCall(op->SetEnableCudaGraph(op, enable_graph)); + } + return CEED_ERROR_SUCCESS; +} diff --git a/interface/ceed.c b/interface/ceed.c index abe987d18f..8094c85c58 100644 --- a/interface/ceed.c +++ b/interface/ceed.c @@ -1366,6 +1366,7 @@ int CeedInit(const char *resource, Ceed *ceed) { CEED_FTABLE_ENTRY(CeedOperator, ApplyAdd), CEED_FTABLE_ENTRY(CeedOperator, ApplyAddComposite), CEED_FTABLE_ENTRY(CeedOperator, ApplyJacobian), + CEED_FTABLE_ENTRY(CeedOperator, SetEnableCudaGraph), CEED_FTABLE_ENTRY(CeedOperator, Destroy), {NULL, 0} // End of lookup table - used in SetBackendFunction loop };