Skip to content
51 changes: 26 additions & 25 deletions src/models/onnxruntime_inline.h
Original file line number Diff line number Diff line change
Expand Up @@ -1531,7 +1531,7 @@ inline std::unique_ptr<OrtLoraAdapter> OrtLoraAdapter::Create(const ORTCHAR_T* a
return std::unique_ptr<OrtLoraAdapter>{p};
}

#if ORT_API_VERSION >= 28 && ORT_GENAI_HAS_EXPERIMENTAL_C_API
#if ORT_GENAI_HAS_MODEL_PACKAGE

namespace Ort {

Expand All @@ -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) \
Comment thread
qjia7 marked this conversation as resolved.
reinterpret_cast<OrtExperimental_OrtModelPackageApi_##NAME##_SinceV28_Fn>( \
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) {
Expand Down Expand Up @@ -1634,4 +1635,4 @@ inline std::basic_string<ORTCHAR_T> OrtModelPackageComponentContext::GetSelected
return path == nullptr ? std::basic_string<ORTCHAR_T>{} : std::basic_string<ORTCHAR_T>{path};
}

#endif // ORT_API_VERSION >= 28 && ORT_GENAI_HAS_EXPERIMENTAL_C_API
#endif // ORT_GENAI_HAS_MODEL_PACKAGE
21 changes: 8 additions & 13 deletions src/models/recurrent_state.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -88,20 +88,15 @@ RecurrentState::RecurrentState(State& state)

const int num_layers = static_cast<int>(layer_indices_.size());

if (!state_.params_->IsPastPresentShareBufferEnabled(model_.config_->model.type)) {
share_buffers_ = state_.params_->IsPastPresentShareBufferEnabled(model_.config_->model.type);

Comment thread
qjia7 marked this conversation as resolved.
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).");
}
Comment thread
qjia7 marked this conversation as resolved.

// 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);
}
Expand Down Expand Up @@ -132,8 +127,8 @@ void RecurrentState::Add() {

const int num_layers = static_cast<int>(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());
Expand Down
3 changes: 2 additions & 1 deletion src/models/recurrent_state.h
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,8 @@ struct RecurrentState {
std::vector<std::unique_ptr<OrtValue>> pasts_;
std::vector<std::unique_ptr<OrtValue>> 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};
Expand Down
Loading