diff --git a/src/models/onnxruntime_inline.h b/src/models/onnxruntime_inline.h index 979e1c57fd..a25e618a5d 100644 --- a/src/models/onnxruntime_inline.h +++ b/src/models/onnxruntime_inline.h @@ -1531,7 +1531,7 @@ inline std::unique_ptr OrtLoraAdapter::Create(const ORTCHAR_T* a return std::unique_ptr{p}; } -#if ORT_API_VERSION >= 28 && ORT_GENAI_HAS_EXPERIMENTAL_C_API +#if ORT_GENAI_HAS_MODEL_PACKAGE namespace Ort { @@ -1541,30 +1541,31 @@ inline const ModelPackageApi& GetModelPackageApi() { if (api == nullptr) { return f; } - f.CreateModelPackageOptionsFromSessionOptions = - Experimental::Get_OrtModelPackageApi_CreateModelPackageOptionsFromSessionOptions_SinceV28_Fn(api); - f.ReleaseModelPackageOptions = - Experimental::Get_OrtModelPackageApi_ReleaseModelPackageOptions_SinceV28_Fn(api); - f.CreateModelPackageContext = - Experimental::Get_OrtModelPackageApi_CreateModelPackageContext_SinceV28_Fn(api); - f.ReleaseModelPackageContext = - Experimental::Get_OrtModelPackageApi_ReleaseModelPackageContext_SinceV28_Fn(api); - f.ModelPackage_GetComponentCount = - Experimental::Get_OrtModelPackageApi_ModelPackage_GetComponentCount_SinceV28_Fn(api); - f.ModelPackage_GetComponentNames = - Experimental::Get_OrtModelPackageApi_ModelPackage_GetComponentNames_SinceV28_Fn(api); - f.ModelPackage_GetVariantCount = - Experimental::Get_OrtModelPackageApi_ModelPackage_GetVariantCount_SinceV28_Fn(api); - f.ModelPackage_GetVariantNames = - Experimental::Get_OrtModelPackageApi_ModelPackage_GetVariantNames_SinceV28_Fn(api); - f.ModelPackage_GetVariantEpName = - Experimental::Get_OrtModelPackageApi_ModelPackage_GetVariantEpName_SinceV28_Fn(api); - f.SelectComponent = - Experimental::Get_OrtModelPackageApi_SelectComponent_SinceV28_Fn(api); - f.ReleaseModelPackageComponentContext = - Experimental::Get_OrtModelPackageApi_ReleaseModelPackageComponentContext_SinceV28_Fn(api); + // Resolve OrtModelPackageApi entries via the C API to avoid including + // onnxruntime_experimental_cxx_api.h. That header defines the + // Ort::Experimental::Get_OrtModelPackageApi_*_SinceV28_Fn accessors but transitively + // pulls in onnxruntime_cxx_api.h, which redefines the Ort:: types that genai has + // vendored in onnxruntime_api.h. +#define GENAI_MP_V28_FN(NAME) \ + reinterpret_cast( \ + api->GetExperimentalFunction( \ + kOrtExperimental_OrtModelPackageApi_##NAME##_SinceV28_FnName)) + + f.CreateModelPackageOptionsFromSessionOptions = GENAI_MP_V28_FN(CreateModelPackageOptionsFromSessionOptions); + f.ReleaseModelPackageOptions = GENAI_MP_V28_FN(ReleaseModelPackageOptions); + f.CreateModelPackageContext = GENAI_MP_V28_FN(CreateModelPackageContext); + f.ReleaseModelPackageContext = GENAI_MP_V28_FN(ReleaseModelPackageContext); + f.ModelPackage_GetComponentCount = GENAI_MP_V28_FN(ModelPackage_GetComponentCount); + f.ModelPackage_GetComponentNames = GENAI_MP_V28_FN(ModelPackage_GetComponentNames); + f.ModelPackage_GetVariantCount = GENAI_MP_V28_FN(ModelPackage_GetVariantCount); + f.ModelPackage_GetVariantNames = GENAI_MP_V28_FN(ModelPackage_GetVariantNames); + f.ModelPackage_GetVariantEpName = GENAI_MP_V28_FN(ModelPackage_GetVariantEpName); + f.SelectComponent = GENAI_MP_V28_FN(SelectComponent); + f.ReleaseModelPackageComponentContext = GENAI_MP_V28_FN(ReleaseModelPackageComponentContext); f.ModelPackageComponent_GetSelectedVariantFolderPath = - Experimental::Get_OrtModelPackageApi_ModelPackageComponent_GetSelectedVariantFolderPath_SinceV28_Fn(api); + GENAI_MP_V28_FN(ModelPackageComponent_GetSelectedVariantFolderPath); + +#undef GENAI_MP_V28_FN return f; }(); if (fns.CreateModelPackageContext == nullptr) { @@ -1634,4 +1635,4 @@ inline std::basic_string OrtModelPackageComponentContext::GetSelected return path == nullptr ? std::basic_string{} : std::basic_string{path}; } -#endif // ORT_API_VERSION >= 28 && ORT_GENAI_HAS_EXPERIMENTAL_C_API +#endif // ORT_GENAI_HAS_MODEL_PACKAGE diff --git a/src/models/recurrent_state.cpp b/src/models/recurrent_state.cpp index 1583efe5fe..cb9c9421dc 100644 --- a/src/models/recurrent_state.cpp +++ b/src/models/recurrent_state.cpp @@ -88,20 +88,15 @@ RecurrentState::RecurrentState(State& state) const int num_layers = static_cast(layer_indices_.size()); - if (!state_.params_->IsPastPresentShareBufferEnabled(model_.config_->model.type)) { + share_buffers_ = state_.params_->IsPastPresentShareBufferEnabled(model_.config_->model.type); + + if (state_.params_->use_graph_capture && !share_buffers_) { throw std::runtime_error( - "RecurrentState requires past_present_share_buffer=true. " - "Set past_present_share_buffer to true in genai_config.json."); + "Graph capture requires past/present buffer sharing for models with recurrent state. " + "Ensure past_present_share_buffer=true in genai_config.json and num_beams=1 " + "(beam search disables buffer sharing)."); } - // WebGPU prohibits binding the same buffer as both read-only (input) and - // read-write (output) storage in the same compute pass, so it must use - // separate past/present buffers with swap. All other EPs share buffers - // for stable addresses (required by TRT-RTX graph replay, beneficial elsewhere). - // TODO: Remove WebGPU special case once the ORT WebGPU EP adds a - // LinearAttention kernel with native past/present buffer sharing support. - share_buffers_ = model_.p_device_kvcache_->GetType() != DeviceType::WEBGPU; - if (!share_buffers_) { pasts_.resize(num_layers * 2); } @@ -132,8 +127,8 @@ void RecurrentState::Add() { const int num_layers = static_cast(layer_indices_.size()); for (int i = 0; i < num_layers * 2; ++i) { - // Shared: alias input=output for stable addresses. - // WebGPU: separate past/present buffers to avoid aliasing violation. + // Shared buffers: alias input=output for stable addresses (required for graph capture). + // Separate buffers: use distinct past/present allocations with per-step pointer swap. state_.inputs_.push_back(share_buffers_ ? presents_[i].get() : pasts_[i].get()); state_.input_names_.push_back(input_name_strings_[i].c_str()); state_.outputs_.push_back(presents_[i].get()); diff --git a/src/models/recurrent_state.h b/src/models/recurrent_state.h index a6787020bc..eca95439ec 100644 --- a/src/models/recurrent_state.h +++ b/src/models/recurrent_state.h @@ -30,7 +30,8 @@ struct RecurrentState { std::vector> pasts_; std::vector> presents_; - // WebGPU cannot alias input/output buffers, so it uses separate past/present\n // with swap. All other EPs share buffers for stable addresses. + // Mirrors past_present_share_buffer config: true means inputs alias outputs (same allocation, + // stable handles for graph capture). False uses separate past/present buffers with per-step swap. bool share_buffers_{false}; size_t input_index_{~0U}; size_t output_index_{~0U};