diff --git a/external/ggml/src/ggml-cuda/ggml-cuda.cu b/external/ggml/src/ggml-cuda/ggml-cuda.cu index 144f484e3..8dc80a82d 100644 --- a/external/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/external/ggml/src/ggml-cuda/ggml-cuda.cu @@ -5897,6 +5897,9 @@ static void * ggml_backend_cuda_reg_get_proc_address(ggml_backend_reg_t reg, con if (strcmp(name, "ggml_backend_cuda_set_stream_priority") == 0) { return (void *)ggml_backend_cuda_set_stream_priority; } + if (strcmp(name, "ggml_backend_cuda_get_stream") == 0) { + return (void *)ggml_backend_cuda_get_stream; + } return nullptr; } diff --git a/include/engine/framework/core/backend.h b/include/engine/framework/core/backend.h index 273c3b571..1bcfa2342 100644 --- a/include/engine/framework/core/backend.h +++ b/include/engine/framework/core/backend.h @@ -59,6 +59,7 @@ void trim_backend_pools(ggml_backend_t backend); // higher priority, 0 = default). No-op on non-CUDA backends and on builds // whose CUDA backend does not export the hook. void set_backend_stream_priority(ggml_backend_t backend, int priority); +void * backend_cuda_stream(ggml_backend_t backend); // evict_cuda_graph_cache=false (the default) is the historical no-op; // true drops the backend's cached compiled-graph state (CUDA/HIP graph // cache) for this cgraph at destruction — opt in per family. diff --git a/src/framework/core/backend.cpp b/src/framework/core/backend.cpp index b21d448ff..ac2be4ac5 100644 --- a/src/framework/core/backend.cpp +++ b/src/framework/core/backend.cpp @@ -346,6 +346,23 @@ void set_backend_stream_priority(ggml_backend_t backend, int priority) { if (fn != nullptr) fn(backend, priority); } +void * backend_cuda_stream(ggml_backend_t backend) { + if (backend == nullptr) return nullptr; + if (!is_cuda_backend_handle(backend) && !is_hip_backend_handle(backend)) return nullptr; + ggml_backend_dev_t device = ggml_backend_get_device(backend); + if (device == nullptr) { + throw std::runtime_error("CUDA backend stream lookup failed: backend has no device"); + } + auto fn = (void * (*)(ggml_backend_t)) + ggml_backend_reg_get_proc_address( + ggml_backend_dev_backend_reg(device), + "ggml_backend_cuda_get_stream"); + if (fn == nullptr) { + throw std::runtime_error("CUDA backend stream lookup failed: backend does not export ggml_backend_cuda_get_stream"); + } + return fn(backend); +} + // evict_cuda_graph_cache defaults to false, preserving historical behavior // for existing call sites: before the CUDA backend exported // ggml_backend_cuda_clear_graph the lookup resolved nothing, and families diff --git a/src/framework/sampling/torch_random.cpp b/src/framework/sampling/torch_random.cpp index 09456b4e3..346d0eac7 100644 --- a/src/framework/sampling/torch_random.cpp +++ b/src/framework/sampling/torch_random.cpp @@ -1,14 +1,10 @@ #include "engine/framework/sampling/torch_random.h" +#include "engine/framework/core/backend.h" #include "engine/framework/debug/trace.h" #include "engine/framework/io/dynamic_library.h" #ifdef ENGINE_HAS_CUDA_TORCH_RANDOM #include "torch_random_cuda_runtime.h" - -#ifdef ENGINE_HAS_CUDA_TORCH_RANDOM -#include "ggml-backend.h" -#include "ggml-cuda.h" -#endif #endif #include @@ -488,7 +484,7 @@ void torch_cuda_sample_topk_exponential_pairs( void * torch_cuda_backend_stream(void * ggml_backend) { #ifdef ENGINE_HAS_CUDA_TORCH_RANDOM - return ggml_backend_cuda_get_stream(static_cast(ggml_backend)); + return core::backend_cuda_stream(static_cast(ggml_backend)); #else (void) ggml_backend; return nullptr;