Skip to content

Commit 46ea2a7

Browse files
weiyuanyueMillyWeiCopilot
authored
[Fix]Incorrect CPU Tagging for AITK NPU Models by Enhancing EP Detection (#524)
Co-authored-by: MillyWei <yuanwei@microsoft.com> Co-authored-by: Copilot <198982749+Copilot@users.noreply.github.com>
1 parent 97ae2f9 commit 46ea2a7

9 files changed

Lines changed: 428 additions & 74 deletions

File tree

AIDevGallery/Models/GenAIConfig.cs

Lines changed: 48 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,10 @@
11
// Copyright (c) Microsoft Corporation. All rights reserved.
22
// Licensed under the MIT License.
33

4+
using System;
5+
using System.Collections.Generic;
6+
using System.Linq;
7+
using System.Text.Json;
48
using System.Text.Json.Serialization;
59

610
namespace AIDevGallery.Models;
@@ -25,6 +29,8 @@ internal class Decoder
2529
{
2630
[JsonPropertyName("session_options")]
2731
public required GenAISessionOptions SessionOptions { get; set; }
32+
[JsonPropertyName("pipeline")]
33+
public PipelineItem[]? Pipeline { get; set; }
2834
}
2935

3036
internal class GenAISessionOptions
@@ -33,18 +39,53 @@ internal class GenAISessionOptions
3339
public required ProviderOptions[] ProviderOptions { get; set; }
3440
}
3541

36-
internal class ProviderOptions
42+
internal class PipelineItem
3743
{
38-
[JsonPropertyName("dml")]
39-
public Dml? Dml { get; set; }
40-
[JsonPropertyName("cuda")]
41-
public Cuda? Cuda { get; set; }
44+
[JsonExtensionData]
45+
public Dictionary<string, JsonElement>? Stages { get; set; }
4246
}
4347

44-
internal class Dml
48+
internal class PipelineStage
4549
{
50+
[JsonPropertyName("session_options")]
51+
public GenAISessionOptions? SessionOptions { get; set; }
4652
}
4753

48-
internal class Cuda
54+
internal class ProviderOptions
4955
{
56+
[JsonExtensionData]
57+
public Dictionary<string, JsonElement>? ExtensionData { get; set; }
58+
59+
public bool HasProvider(string name)
60+
{
61+
if (ExtensionData == null)
62+
{
63+
return false;
64+
}
65+
66+
return ExtensionData.Keys.Any(k => k.Equals(name, StringComparison.OrdinalIgnoreCase));
67+
}
68+
69+
public Dictionary<string, string>? GetProviderOptions(string name)
70+
{
71+
if (ExtensionData == null)
72+
{
73+
return null;
74+
}
75+
76+
var key = ExtensionData.Keys.FirstOrDefault(k => k.Equals(name, StringComparison.OrdinalIgnoreCase));
77+
if (key == null)
78+
{
79+
return null;
80+
}
81+
82+
try
83+
{
84+
return JsonSerializer.Deserialize(ExtensionData[key].GetRawText(), AIDevGallery.Utils.SourceGenerationContext.Default.DictionaryStringString);
85+
}
86+
catch
87+
{
88+
return new Dictionary<string, string>();
89+
}
90+
}
5091
}

AIDevGallery/Models/ModelCompatibility.cs

Lines changed: 18 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -53,11 +53,27 @@ public static ModelCompatibility GetModelCompatibility(ModelDetails modelDetails
5353
compatibility = ModelCompatibilityState.NotCompatible;
5454
description = "This model is not currently supported on Arm64 devices.";
5555
}
56-
else if (modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.CPU) ||
57-
(modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.QNN) && DeviceUtils.IsArm64()))
56+
else if (modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.CPU))
5857
{
5958
compatibility = ModelCompatibilityState.Compatible;
6059
}
60+
else if (modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.QNN) && DeviceUtils.IsArm64())
61+
{
62+
compatibility = ModelCompatibilityState.Compatible;
63+
}
64+
else if (modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.NPU))
65+
{
66+
// Check if any NPU is available using ONNX Runtime's EP detection
67+
if (DeviceUtils.HasNPU())
68+
{
69+
compatibility = ModelCompatibilityState.Compatible;
70+
}
71+
else
72+
{
73+
compatibility = ModelCompatibilityState.NotCompatible;
74+
description = "This model requires an NPU (Neural Processing Unit). No compatible NPU was detected on your device.";
75+
}
76+
}
6177
else if (
6278
(modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.DML) || modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.GPU))
6379
&& !DeviceUtils.IsArm64())

AIDevGallery/Pages/Scenarios/ScenarioPage.xaml.cs

Lines changed: 20 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -142,38 +142,38 @@ private static async Task<List<WinMlEp>> GetSupportedHardwareAccelerators()
142142

143143
switch (epName)
144144
{
145-
case "VitisAIExecutionProvider":
146-
supportedHardwareAccelerators.Add(new([HardwareAccelerator.VitisAI, HardwareAccelerator.NPU], "VitisAIExecutionProvider", "VitisAI", "NPU"));
145+
case ExecutionProviderNames.VitisAI:
146+
supportedHardwareAccelerators.Add(new([HardwareAccelerator.VitisAI, HardwareAccelerator.NPU], ExecutionProviderNames.VitisAI, "VitisAI", "NPU"));
147147
break;
148148

149-
case "OpenVINOExecutionProvider":
149+
case ExecutionProviderNames.OpenVINO:
150150
if (epDeviceTypes.Contains("CPU"))
151151
{
152-
supportedHardwareAccelerators.Add(new([HardwareAccelerator.OpenVINO, HardwareAccelerator.CPU], "OpenVINOExecutionProvider", "OpenVINO", "CPU"));
152+
supportedHardwareAccelerators.Add(new([HardwareAccelerator.OpenVINO, HardwareAccelerator.CPU], ExecutionProviderNames.OpenVINO, "OpenVINO", "CPU"));
153153
}
154154

155155
if (epDeviceTypes.Contains("GPU"))
156156
{
157-
supportedHardwareAccelerators.Add(new([HardwareAccelerator.OpenVINO, HardwareAccelerator.GPU], "OpenVINOExecutionProvider", "OpenVINO", "GPU"));
157+
supportedHardwareAccelerators.Add(new([HardwareAccelerator.OpenVINO, HardwareAccelerator.GPU], ExecutionProviderNames.OpenVINO, "OpenVINO", "GPU"));
158158
}
159159

160160
if (epDeviceTypes.Contains("NPU"))
161161
{
162-
supportedHardwareAccelerators.Add(new([HardwareAccelerator.OpenVINO, HardwareAccelerator.NPU], "OpenVINOExecutionProvider", "OpenVINO", "NPU"));
162+
supportedHardwareAccelerators.Add(new([HardwareAccelerator.OpenVINO, HardwareAccelerator.NPU], ExecutionProviderNames.OpenVINO, "OpenVINO", "NPU"));
163163
}
164164

165165
break;
166166

167-
case "QNNExecutionProvider":
168-
supportedHardwareAccelerators.Add(new([HardwareAccelerator.QNN, HardwareAccelerator.NPU], "QNNExecutionProvider", "QNN", "NPU"));
167+
case ExecutionProviderNames.QNN:
168+
supportedHardwareAccelerators.Add(new([HardwareAccelerator.QNN, HardwareAccelerator.NPU], ExecutionProviderNames.QNN, "QNN", "NPU"));
169169
break;
170170

171-
case "DmlExecutionProvider":
172-
supportedHardwareAccelerators.Add(new([HardwareAccelerator.DML, HardwareAccelerator.GPU], "DmlExecutionProvider", "DML", "GPU"));
171+
case ExecutionProviderNames.DML:
172+
supportedHardwareAccelerators.Add(new([HardwareAccelerator.DML, HardwareAccelerator.GPU], ExecutionProviderNames.DML, "DML", "GPU"));
173173
break;
174174

175-
case "NvTensorRTRTXExecutionProvider":
176-
supportedHardwareAccelerators.Add(new([HardwareAccelerator.NvTensorRT, HardwareAccelerator.GPU], "NvTensorRTRTXExecutionProvider", "NvTensorRT", "GPU"));
175+
case ExecutionProviderNames.NvTensorRTRTX:
176+
supportedHardwareAccelerators.Add(new([HardwareAccelerator.NvTensorRT, HardwareAccelerator.GPU], ExecutionProviderNames.NvTensorRTRTX, "NvTensorRT", "GPU"));
177177
break;
178178
}
179179
}
@@ -193,7 +193,13 @@ private async void HandleModelSelectionChanged(List<ModelDetails?> selectedModel
193193
VisualStateManager.GoToState(this, "PageLoading", true);
194194

195195
modelDetails.Clear();
196-
selectedModels.ForEach(modelDetails.Add);
196+
foreach (var model in selectedModels)
197+
{
198+
if (model != null)
199+
{
200+
modelDetails.Add(model!);
201+
}
202+
}
197203

198204
// temporary fix EP dropdown list for useradded local languagemodel
199205
if (selectedModels.Any(m => m != null && m.IsOnnxModel() && string.IsNullOrEmpty(m.ParameterSize) && m.Id.StartsWith("useradded-local-languagemodel", System.StringComparison.InvariantCultureIgnoreCase) == false))
@@ -316,7 +322,7 @@ private void LoadSample(Sample? sampleToLoad)
316322

317323
// TODO: don't load sample if model is not cached, but still let code to be seen
318324
// this would probably be handled in the SampleContainer
319-
_ = SampleContainer.LoadSampleAsync(sample, [.. modelDetails], App.AppData.WinMLSampleOptions);
325+
_ = SampleContainer.LoadSampleAsync(sample, modelDetails.Where(m => m != null).Select(m => m!).ToList(), App.AppData.WinMLSampleOptions);
320326
_ = App.AppData.AddMru(
321327
new MostRecentlyUsedItem()
322328
{

AIDevGallery/Samples/SharedCode/WinMLHelpers.cs

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
// Copyright (c) Microsoft Corporation. All rights reserved.
22
// Licensed under the MIT License.
33

4+
using AIDevGallery.Utils;
45
using Microsoft.ML.OnnxRuntime;
56
using System;
67
using System.Collections.Generic;
@@ -12,32 +13,32 @@ namespace AIDevGallery.Samples.SharedCode;
1213

1314
internal static class WinMLHelpers
1415
{
15-
public static bool AppendExecutionProviderFromEpName(this SessionOptions sessionOptions, string epName, string? deviceType, OrtEnv? environment = null)
16+
public static bool AppendExecutionProviderFromEpName(this SessionOptions sessionOptions, string epName, string? deviceType)
1617
{
1718
if (epName == "CPU")
1819
{
1920
// No need to append CPU execution provider
2021
return true;
2122
}
2223

23-
environment ??= OrtEnv.Instance();
24-
var epDeviceMap = GetEpDeviceMap(environment);
24+
var environment = OrtEnv.Instance();
25+
var epDeviceMap = GetEpDeviceMap();
2526

2627
if (epDeviceMap.TryGetValue(epName, out var devices))
2728
{
2829
Dictionary<string, string> epOptions = new(StringComparer.OrdinalIgnoreCase);
2930
switch (epName)
3031
{
31-
case "DmlExecutionProvider":
32+
case ExecutionProviderNames.DML:
3233
// Configure performance mode for Dml EP
3334
// Dml some times have multiple devices which cause exception, we pick the first one here
3435
sessionOptions.AppendExecutionProvider(environment, [devices[0]], epOptions);
3536
return true;
36-
case "OpenVINOExecutionProvider":
37+
case ExecutionProviderNames.OpenVINO:
3738
var device = devices.Where(d => d.HardwareDevice.Type.ToString().Equals(deviceType, StringComparison.Ordinal)).FirstOrDefault();
3839
sessionOptions.AppendExecutionProvider(environment, [device], epOptions);
3940
return true;
40-
case "QNNExecutionProvider":
41+
case ExecutionProviderNames.QNN:
4142
// Configure performance mode for QNN EP
4243
epOptions["htp_performance_mode"] = "high_performance";
4344
break;
@@ -105,10 +106,9 @@ public static bool AppendExecutionProviderFromEpName(this SessionOptions session
105106
return null;
106107
}
107108

108-
public static Dictionary<string, List<OrtEpDevice>> GetEpDeviceMap(OrtEnv? environment = null)
109+
public static Dictionary<string, List<OrtEpDevice>> GetEpDeviceMap()
109110
{
110-
environment ??= OrtEnv.Instance();
111-
IReadOnlyList<OrtEpDevice> epDevices = environment.GetEpDevices();
111+
IReadOnlyList<OrtEpDevice> epDevices = DeviceUtils.GetEpDevices();
112112
Dictionary<string, List<OrtEpDevice>> epDeviceMap = new(StringComparer.OrdinalIgnoreCase);
113113

114114
foreach (OrtEpDevice device in epDevices)

AIDevGallery/Utils/AppUtils.cs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -285,10 +285,10 @@ public static string GetThemeAssetSuffix()
285285
{
286286
var accessibilitySettings = new AccessibilitySettings();
287287
bool isHighContrast = accessibilitySettings.HighContrast;
288-
if(isHighContrast)
288+
if (isHighContrast)
289289
{
290290
string hcThemeName = accessibilitySettings.HighContrastScheme;
291-
if(hcThemeName == "High Contrast White")
291+
if (hcThemeName == "High Contrast White")
292292
{
293293
return ".light";
294294
}

AIDevGallery/Utils/DeviceUtils.cs

Lines changed: 69 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,14 +1,19 @@
11
// Copyright (c) Microsoft Corporation. All rights reserved.
22
// Licensed under the MIT License.
33

4+
using Microsoft.ML.OnnxRuntime;
45
using System;
6+
using System.Linq;
57
using Windows.Win32.Foundation;
68
using Windows.Win32.Graphics.Dxgi;
79

810
namespace AIDevGallery.Utils;
911

1012
internal static class DeviceUtils
1113
{
14+
private static readonly object _epDevicesLock = new();
15+
private static System.Collections.Generic.IReadOnlyList<OrtEpDevice>? _cachedEpDevices;
16+
1217
public static int GetBestDeviceId()
1318
{
1419
int deviceId = 0;
@@ -107,8 +112,70 @@ public static ulong GetVram()
107112
return maxDedicatedVideoMemory;
108113
}
109114

110-
public static bool IsArm64()
115+
public static bool IsArm64() =>
116+
System.Runtime.InteropServices.RuntimeInformation.OSArchitecture == System.Runtime.InteropServices.Architecture.Arm64;
117+
118+
public static bool HasNPU() => HasExecutionProvider(device =>
119+
device.HardwareDevice.Type.ToString().Equals("NPU", StringComparison.OrdinalIgnoreCase));
120+
121+
private static bool HasExecutionProvider(Func<OrtEpDevice, bool> predicate)
111122
{
112-
return System.Runtime.InteropServices.RuntimeInformation.OSArchitecture == System.Runtime.InteropServices.Architecture.Arm64;
123+
try
124+
{
125+
return GetEpDevices().Any(predicate);
126+
}
127+
catch
128+
{
129+
return false;
130+
}
131+
}
132+
133+
/// <summary>
134+
/// Gets the list of available ONNX Runtime Execution Provider devices.
135+
/// This method ensures that certified EPs (like OpenVINO, QNN, DML) are registered before querying,
136+
/// as OrtEnv.GetEpDevices() only returns already-registered providers.
137+
/// Results are cached to avoid repeated registration overhead.
138+
/// </summary>
139+
/// <returns>A read-only list of available ONNX Runtime Execution Provider devices.</returns>
140+
public static System.Collections.Generic.IReadOnlyList<OrtEpDevice> GetEpDevices()
141+
{
142+
if (_cachedEpDevices != null)
143+
{
144+
return _cachedEpDevices;
145+
}
146+
147+
lock (_epDevicesLock)
148+
{
149+
if (_cachedEpDevices != null)
150+
{
151+
return _cachedEpDevices;
152+
}
153+
154+
try
155+
{
156+
OrtEnv.Instance();
157+
var catalog = Microsoft.Windows.AI.MachineLearning.ExecutionProviderCatalog.GetDefault();
158+
159+
try
160+
{
161+
catalog.EnsureAndRegisterCertifiedAsync().GetAwaiter().GetResult();
162+
}
163+
catch (Exception ex)
164+
{
165+
// Log but continue
166+
Telemetry.TelemetryFactory.Get<Telemetry.ITelemetry>().LogException("GetEpDevices_RegistrationFailed", ex);
167+
}
168+
169+
_cachedEpDevices = OrtEnv.Instance().GetEpDevices();
170+
}
171+
catch (Exception ex)
172+
{
173+
// Log the failure to get EP devices - this could indicate ONNX Runtime initialization issues
174+
Telemetry.TelemetryFactory.Get<Telemetry.ITelemetry>().LogException("GetEpDevices_Failed", ex);
175+
_cachedEpDevices = System.Array.Empty<OrtEpDevice>();
176+
}
177+
178+
return _cachedEpDevices;
179+
}
113180
}
114181
}

0 commit comments

Comments
 (0)