diff --git a/AIDevGallery/Models/GenAIConfig.cs b/AIDevGallery/Models/GenAIConfig.cs index 738e107e..eecf2194 100644 --- a/AIDevGallery/Models/GenAIConfig.cs +++ b/AIDevGallery/Models/GenAIConfig.cs @@ -1,6 +1,10 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. +using System; +using System.Collections.Generic; +using System.Linq; +using System.Text.Json; using System.Text.Json.Serialization; namespace AIDevGallery.Models; @@ -25,6 +29,8 @@ internal class Decoder { [JsonPropertyName("session_options")] public required GenAISessionOptions SessionOptions { get; set; } + [JsonPropertyName("pipeline")] + public PipelineItem[]? Pipeline { get; set; } } internal class GenAISessionOptions @@ -33,18 +39,53 @@ internal class GenAISessionOptions public required ProviderOptions[] ProviderOptions { get; set; } } -internal class ProviderOptions +internal class PipelineItem { - [JsonPropertyName("dml")] - public Dml? Dml { get; set; } - [JsonPropertyName("cuda")] - public Cuda? Cuda { get; set; } + [JsonExtensionData] + public Dictionary? Stages { get; set; } } -internal class Dml +internal class PipelineStage { + [JsonPropertyName("session_options")] + public GenAISessionOptions? SessionOptions { get; set; } } -internal class Cuda +internal class ProviderOptions { + [JsonExtensionData] + public Dictionary? ExtensionData { get; set; } + + public bool HasProvider(string name) + { + if (ExtensionData == null) + { + return false; + } + + return ExtensionData.Keys.Any(k => k.Equals(name, StringComparison.OrdinalIgnoreCase)); + } + + public Dictionary? GetProviderOptions(string name) + { + if (ExtensionData == null) + { + return null; + } + + var key = ExtensionData.Keys.FirstOrDefault(k => k.Equals(name, StringComparison.OrdinalIgnoreCase)); + if (key == null) + { + return null; + } + + try + { + return JsonSerializer.Deserialize(ExtensionData[key].GetRawText(), AIDevGallery.Utils.SourceGenerationContext.Default.DictionaryStringString); + } + catch + { + return new Dictionary(); + } + } } \ No newline at end of file diff --git a/AIDevGallery/Models/ModelCompatibility.cs b/AIDevGallery/Models/ModelCompatibility.cs index 2676212f..c1c4c894 100644 --- a/AIDevGallery/Models/ModelCompatibility.cs +++ b/AIDevGallery/Models/ModelCompatibility.cs @@ -53,11 +53,27 @@ public static ModelCompatibility GetModelCompatibility(ModelDetails modelDetails compatibility = ModelCompatibilityState.NotCompatible; description = "This model is not currently supported on Arm64 devices."; } - else if (modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.CPU) || - (modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.QNN) && DeviceUtils.IsArm64())) + else if (modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.CPU)) { compatibility = ModelCompatibilityState.Compatible; } + else if (modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.QNN) && DeviceUtils.IsArm64()) + { + compatibility = ModelCompatibilityState.Compatible; + } + else if (modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.NPU)) + { + // Check if any NPU is available using ONNX Runtime's EP detection + if (DeviceUtils.HasNPU()) + { + compatibility = ModelCompatibilityState.Compatible; + } + else + { + compatibility = ModelCompatibilityState.NotCompatible; + description = "This model requires an NPU (Neural Processing Unit). No compatible NPU was detected on your device."; + } + } else if ( (modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.DML) || modelDetails.HardwareAccelerators.Contains(HardwareAccelerator.GPU)) && !DeviceUtils.IsArm64()) diff --git a/AIDevGallery/Pages/Scenarios/ScenarioPage.xaml.cs b/AIDevGallery/Pages/Scenarios/ScenarioPage.xaml.cs index 518836b9..73bad359 100644 --- a/AIDevGallery/Pages/Scenarios/ScenarioPage.xaml.cs +++ b/AIDevGallery/Pages/Scenarios/ScenarioPage.xaml.cs @@ -142,38 +142,38 @@ private static async Task> GetSupportedHardwareAccelerators() switch (epName) { - case "VitisAIExecutionProvider": - supportedHardwareAccelerators.Add(new([HardwareAccelerator.VitisAI, HardwareAccelerator.NPU], "VitisAIExecutionProvider", "VitisAI", "NPU")); + case ExecutionProviderNames.VitisAI: + supportedHardwareAccelerators.Add(new([HardwareAccelerator.VitisAI, HardwareAccelerator.NPU], ExecutionProviderNames.VitisAI, "VitisAI", "NPU")); break; - case "OpenVINOExecutionProvider": + case ExecutionProviderNames.OpenVINO: if (epDeviceTypes.Contains("CPU")) { - supportedHardwareAccelerators.Add(new([HardwareAccelerator.OpenVINO, HardwareAccelerator.CPU], "OpenVINOExecutionProvider", "OpenVINO", "CPU")); + supportedHardwareAccelerators.Add(new([HardwareAccelerator.OpenVINO, HardwareAccelerator.CPU], ExecutionProviderNames.OpenVINO, "OpenVINO", "CPU")); } if (epDeviceTypes.Contains("GPU")) { - supportedHardwareAccelerators.Add(new([HardwareAccelerator.OpenVINO, HardwareAccelerator.GPU], "OpenVINOExecutionProvider", "OpenVINO", "GPU")); + supportedHardwareAccelerators.Add(new([HardwareAccelerator.OpenVINO, HardwareAccelerator.GPU], ExecutionProviderNames.OpenVINO, "OpenVINO", "GPU")); } if (epDeviceTypes.Contains("NPU")) { - supportedHardwareAccelerators.Add(new([HardwareAccelerator.OpenVINO, HardwareAccelerator.NPU], "OpenVINOExecutionProvider", "OpenVINO", "NPU")); + supportedHardwareAccelerators.Add(new([HardwareAccelerator.OpenVINO, HardwareAccelerator.NPU], ExecutionProviderNames.OpenVINO, "OpenVINO", "NPU")); } break; - case "QNNExecutionProvider": - supportedHardwareAccelerators.Add(new([HardwareAccelerator.QNN, HardwareAccelerator.NPU], "QNNExecutionProvider", "QNN", "NPU")); + case ExecutionProviderNames.QNN: + supportedHardwareAccelerators.Add(new([HardwareAccelerator.QNN, HardwareAccelerator.NPU], ExecutionProviderNames.QNN, "QNN", "NPU")); break; - case "DmlExecutionProvider": - supportedHardwareAccelerators.Add(new([HardwareAccelerator.DML, HardwareAccelerator.GPU], "DmlExecutionProvider", "DML", "GPU")); + case ExecutionProviderNames.DML: + supportedHardwareAccelerators.Add(new([HardwareAccelerator.DML, HardwareAccelerator.GPU], ExecutionProviderNames.DML, "DML", "GPU")); break; - case "NvTensorRTRTXExecutionProvider": - supportedHardwareAccelerators.Add(new([HardwareAccelerator.NvTensorRT, HardwareAccelerator.GPU], "NvTensorRTRTXExecutionProvider", "NvTensorRT", "GPU")); + case ExecutionProviderNames.NvTensorRTRTX: + supportedHardwareAccelerators.Add(new([HardwareAccelerator.NvTensorRT, HardwareAccelerator.GPU], ExecutionProviderNames.NvTensorRTRTX, "NvTensorRT", "GPU")); break; } } @@ -193,7 +193,13 @@ private async void HandleModelSelectionChanged(List selectedModel VisualStateManager.GoToState(this, "PageLoading", true); modelDetails.Clear(); - selectedModels.ForEach(modelDetails.Add); + foreach (var model in selectedModels) + { + if (model != null) + { + modelDetails.Add(model!); + } + } // temporary fix EP dropdown list for useradded local languagemodel 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) // TODO: don't load sample if model is not cached, but still let code to be seen // this would probably be handled in the SampleContainer - _ = SampleContainer.LoadSampleAsync(sample, [.. modelDetails], App.AppData.WinMLSampleOptions); + _ = SampleContainer.LoadSampleAsync(sample, modelDetails.Where(m => m != null).Select(m => m!).ToList(), App.AppData.WinMLSampleOptions); _ = App.AppData.AddMru( new MostRecentlyUsedItem() { diff --git a/AIDevGallery/Samples/SharedCode/WinMLHelpers.cs b/AIDevGallery/Samples/SharedCode/WinMLHelpers.cs index 2e76c189..200d5a5a 100644 --- a/AIDevGallery/Samples/SharedCode/WinMLHelpers.cs +++ b/AIDevGallery/Samples/SharedCode/WinMLHelpers.cs @@ -1,6 +1,7 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. +using AIDevGallery.Utils; using Microsoft.ML.OnnxRuntime; using System; using System.Collections.Generic; @@ -12,7 +13,7 @@ namespace AIDevGallery.Samples.SharedCode; internal static class WinMLHelpers { - public static bool AppendExecutionProviderFromEpName(this SessionOptions sessionOptions, string epName, string? deviceType, OrtEnv? environment = null) + public static bool AppendExecutionProviderFromEpName(this SessionOptions sessionOptions, string epName, string? deviceType) { if (epName == "CPU") { @@ -20,24 +21,24 @@ public static bool AppendExecutionProviderFromEpName(this SessionOptions session return true; } - environment ??= OrtEnv.Instance(); - var epDeviceMap = GetEpDeviceMap(environment); + var environment = OrtEnv.Instance(); + var epDeviceMap = GetEpDeviceMap(); if (epDeviceMap.TryGetValue(epName, out var devices)) { Dictionary epOptions = new(StringComparer.OrdinalIgnoreCase); switch (epName) { - case "DmlExecutionProvider": + case ExecutionProviderNames.DML: // Configure performance mode for Dml EP // Dml some times have multiple devices which cause exception, we pick the first one here sessionOptions.AppendExecutionProvider(environment, [devices[0]], epOptions); return true; - case "OpenVINOExecutionProvider": + case ExecutionProviderNames.OpenVINO: var device = devices.Where(d => d.HardwareDevice.Type.ToString().Equals(deviceType, StringComparison.Ordinal)).FirstOrDefault(); sessionOptions.AppendExecutionProvider(environment, [device], epOptions); return true; - case "QNNExecutionProvider": + case ExecutionProviderNames.QNN: // Configure performance mode for QNN EP epOptions["htp_performance_mode"] = "high_performance"; break; @@ -105,10 +106,9 @@ public static bool AppendExecutionProviderFromEpName(this SessionOptions session return null; } - public static Dictionary> GetEpDeviceMap(OrtEnv? environment = null) + public static Dictionary> GetEpDeviceMap() { - environment ??= OrtEnv.Instance(); - IReadOnlyList epDevices = environment.GetEpDevices(); + IReadOnlyList epDevices = DeviceUtils.GetEpDevices(); Dictionary> epDeviceMap = new(StringComparer.OrdinalIgnoreCase); foreach (OrtEpDevice device in epDevices) diff --git a/AIDevGallery/Utils/AppUtils.cs b/AIDevGallery/Utils/AppUtils.cs index 78114158..406a82db 100644 --- a/AIDevGallery/Utils/AppUtils.cs +++ b/AIDevGallery/Utils/AppUtils.cs @@ -285,10 +285,10 @@ public static string GetThemeAssetSuffix() { var accessibilitySettings = new AccessibilitySettings(); bool isHighContrast = accessibilitySettings.HighContrast; - if(isHighContrast) + if (isHighContrast) { string hcThemeName = accessibilitySettings.HighContrastScheme; - if(hcThemeName == "High Contrast White") + if (hcThemeName == "High Contrast White") { return ".light"; } diff --git a/AIDevGallery/Utils/DeviceUtils.cs b/AIDevGallery/Utils/DeviceUtils.cs index 6943dd9a..a0e467e5 100644 --- a/AIDevGallery/Utils/DeviceUtils.cs +++ b/AIDevGallery/Utils/DeviceUtils.cs @@ -1,7 +1,9 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. +using Microsoft.ML.OnnxRuntime; using System; +using System.Linq; using Windows.Win32.Foundation; using Windows.Win32.Graphics.Dxgi; @@ -9,6 +11,9 @@ namespace AIDevGallery.Utils; internal static class DeviceUtils { + private static readonly object _epDevicesLock = new(); + private static System.Collections.Generic.IReadOnlyList? _cachedEpDevices; + public static int GetBestDeviceId() { int deviceId = 0; @@ -107,8 +112,70 @@ public static ulong GetVram() return maxDedicatedVideoMemory; } - public static bool IsArm64() + public static bool IsArm64() => + System.Runtime.InteropServices.RuntimeInformation.OSArchitecture == System.Runtime.InteropServices.Architecture.Arm64; + + public static bool HasNPU() => HasExecutionProvider(device => + device.HardwareDevice.Type.ToString().Equals("NPU", StringComparison.OrdinalIgnoreCase)); + + private static bool HasExecutionProvider(Func predicate) { - return System.Runtime.InteropServices.RuntimeInformation.OSArchitecture == System.Runtime.InteropServices.Architecture.Arm64; + try + { + return GetEpDevices().Any(predicate); + } + catch + { + return false; + } + } + + /// + /// Gets the list of available ONNX Runtime Execution Provider devices. + /// This method ensures that certified EPs (like OpenVINO, QNN, DML) are registered before querying, + /// as OrtEnv.GetEpDevices() only returns already-registered providers. + /// Results are cached to avoid repeated registration overhead. + /// + /// A read-only list of available ONNX Runtime Execution Provider devices. + public static System.Collections.Generic.IReadOnlyList GetEpDevices() + { + if (_cachedEpDevices != null) + { + return _cachedEpDevices; + } + + lock (_epDevicesLock) + { + if (_cachedEpDevices != null) + { + return _cachedEpDevices; + } + + try + { + OrtEnv.Instance(); + var catalog = Microsoft.Windows.AI.MachineLearning.ExecutionProviderCatalog.GetDefault(); + + try + { + catalog.EnsureAndRegisterCertifiedAsync().GetAwaiter().GetResult(); + } + catch (Exception ex) + { + // Log but continue + Telemetry.TelemetryFactory.Get().LogException("GetEpDevices_RegistrationFailed", ex); + } + + _cachedEpDevices = OrtEnv.Instance().GetEpDevices(); + } + catch (Exception ex) + { + // Log the failure to get EP devices - this could indicate ONNX Runtime initialization issues + Telemetry.TelemetryFactory.Get().LogException("GetEpDevices_Failed", ex); + _cachedEpDevices = System.Array.Empty(); + } + + return _cachedEpDevices; + } } } \ No newline at end of file diff --git a/AIDevGallery/Utils/ExecutionProviderNames.cs b/AIDevGallery/Utils/ExecutionProviderNames.cs new file mode 100644 index 00000000..c727c949 --- /dev/null +++ b/AIDevGallery/Utils/ExecutionProviderNames.cs @@ -0,0 +1,51 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +namespace AIDevGallery.Utils; + +/// +/// Provides standard execution provider names used by ONNX Runtime. +/// These constants match the EpName values returned by OrtEpDevice. +/// +internal static class ExecutionProviderNames +{ + /// + /// CPU Execution Provider - always available + /// + public const string CPU = "CPUExecutionProvider"; + + /// + /// DirectML Execution Provider - for DirectX 12 GPU acceleration on Windows + /// + public const string DML = "DmlExecutionProvider"; + + /// + /// Qualcomm Neural Network (QNN) Execution Provider - for Qualcomm NPU + /// + public const string QNN = "QNNExecutionProvider"; + + /// + /// OpenVINO Execution Provider - supports CPU, GPU, and NPU on Intel hardware + /// + public const string OpenVINO = "OpenVINOExecutionProvider"; + + /// + /// Vitis AI Execution Provider - for AMD/Xilinx NPU acceleration + /// + public const string VitisAI = "VitisAIExecutionProvider"; + + /// + /// CUDA Execution Provider - for NVIDIA GPU acceleration + /// + public const string CUDA = "CUDAExecutionProvider"; + + /// + /// TensorRT Execution Provider - for optimized inference on NVIDIA GPUs + /// + public const string TensorRT = "TensorrtExecutionProvider"; + + /// + /// NVIDIA TensorRT RTX Execution Provider + /// + public const string NvTensorRTRTX = "NvTensorRTRTXExecutionProvider"; +} \ No newline at end of file diff --git a/AIDevGallery/Utils/SourceGenerationContext.cs b/AIDevGallery/Utils/SourceGenerationContext.cs index 6386d30a..31a8bdb8 100644 --- a/AIDevGallery/Utils/SourceGenerationContext.cs +++ b/AIDevGallery/Utils/SourceGenerationContext.cs @@ -9,6 +9,8 @@ namespace AIDevGallery.Utils; [JsonSerializable(typeof(List))] [JsonSerializable(typeof(GenAIConfig))] +[JsonSerializable(typeof(PipelineStage))] +[JsonSerializable(typeof(Dictionary))] internal partial class SourceGenerationContext : JsonSerializerContext { } \ No newline at end of file diff --git a/AIDevGallery/Utils/UserAddedModelUtil.cs b/AIDevGallery/Utils/UserAddedModelUtil.cs index 743853bd..9d859a85 100644 --- a/AIDevGallery/Utils/UserAddedModelUtil.cs +++ b/AIDevGallery/Utils/UserAddedModelUtil.cs @@ -30,8 +30,8 @@ public static async Task OpenAddLanguageModelFlow(XamlRoot root) if (folder != null) { - var files = Directory.GetFiles(folder.Path); - var config = files.Where(r => Path.GetFileName(r) == "genai_config.json").FirstOrDefault(); + var config = Directory.GetFiles(folder.Path) + .FirstOrDefault(r => Path.GetFileName(r) == "genai_config.json"); if (string.IsNullOrEmpty(config) || App.ModelCache.Models.Any(m => m.Path == folder.Path)) { @@ -52,10 +52,10 @@ public static async Task OpenAddLanguageModelFlow(XamlRoot root) } HardwareAccelerator accelerator = HardwareAccelerator.CPU; + string configContents = string.Empty; try { - string configContents = string.Empty; configContents = await File.ReadAllTextAsync(config); accelerator = GetHardwareAcceleratorFromConfig(configContents); } @@ -73,6 +73,35 @@ public static async Task OpenAddLanguageModelFlow(XamlRoot root) return; } + var (isValid, unavailableProviders) = ValidateExecutionProviders(configContents); + if (!isValid) + { + var warningMessage = "This model requires execution providers that are not available on your device:\n\n" + + string.Join(", ", unavailableProviders) + + "\n\nThe model may fail to load or run. Do you want to add it anyway?"; + + ContentDialog warningDialog = new() + { + Title = "Incompatible Execution Providers", + Content = new TextBlock() + { + Text = warningMessage, + TextWrapping = TextWrapping.Wrap + }, + XamlRoot = root, + CloseButtonText = "Cancel", + PrimaryButtonText = "Add Anyway", + DefaultButton = ContentDialogButton.Close, + Style = Application.Current.Resources["DefaultContentDialogStyle"] as Style + }; + + var warningResult = await warningDialog.ShowAsync(); + if (warningResult != ContentDialogResult.Primary) + { + return; + } + } + var nameTextBox = new TextBox() { Text = Path.GetFileName(folder.Path), @@ -220,17 +249,10 @@ public static async Task OpenAddModelFlow(XamlRoot targetRoot, List GetValidatedModelTypesForUploadedOnnxModel(string modelFilepath, List modelTypes) { - List validatedModelTypes = new(); - - foreach (var (type, models) in ModelDetailsHelper.GetModelDetailsForModelTypes(modelTypes)) - { - if (ValidateUserAddedModelDimensionsForModelTypeModelDetails(models, modelFilepath)) - { - validatedModelTypes.Add(type); - } - } - - return validatedModelTypes; + return ModelDetailsHelper.GetModelDetailsForModelTypes(modelTypes) + .Where(pair => ValidateUserAddedModelDimensionsForModelTypeModelDetails(pair.Value, modelFilepath)) + .Select(pair => pair.Key) + .ToList(); } private static bool ValidateUserAddedModelDimensionsForModelTypeModelDetails(List modelDetailsList, string modelFilepath) @@ -239,21 +261,11 @@ private static bool ValidateUserAddedModelDimensionsForModelTypeModelDetails(Lis sessionOptions.RegisterOrtExtensions(); using InferenceSession inferenceSession = new(modelFilepath, sessionOptions); - List inputDimensions = new(); - List outputDimensions = new(); + var inputDimensions = inferenceSession.InputMetadata.Select(kvp => kvp.Value.Dimensions).ToList(); + var outputDimensions = inferenceSession.OutputMetadata.Select(kvp => kvp.Value.Dimensions).ToList(); - inputDimensions.AddRange(inferenceSession.InputMetadata.Select(kvp => kvp.Value.Dimensions)); - outputDimensions.AddRange(inferenceSession.OutputMetadata.Select(kvp => kvp.Value.Dimensions)); - - foreach (ModelDetails modelDetails in modelDetailsList) - { - if (ValidateUserAddedModelAgainstModelDimensions(inputDimensions, outputDimensions, modelDetails)) - { - return true; - } - } - - return false; + return modelDetailsList.Any(modelDetails => + ValidateUserAddedModelAgainstModelDimensions(inputDimensions, outputDimensions, modelDetails)); } private static bool ValidateUserAddedModelAgainstModelDimensions(List inputDimensions, List outputDimensions, ModelDetails modelDetails) @@ -304,36 +316,195 @@ private static bool CompareDimension(int[] dimensionA, int[] dimensionB) public static bool IsModelsDetailsListUploadCompatible(this IEnumerable modelDetailsList) { - foreach (ModelDetails modelDetails in modelDetailsList) + return modelDetailsList.Any(m => m.InputDimensions != null && m.OutputDimensions != null); + } + + public static HardwareAccelerator GetHardwareAcceleratorFromConfig(string configContents) + { + if (configContents.Contains(""""backend_path": "QnnHtp.dll"""", StringComparison.OrdinalIgnoreCase)) + { + return HardwareAccelerator.QNN; + } + + var config = JsonSerializer.Deserialize(configContents, SourceGenerationContext.Default.GenAIConfig); + if (config == null) { - if (modelDetails.InputDimensions != null && modelDetails.OutputDimensions != null) + throw new InvalidDataException("genai_config.json is not valid"); + } + + // Return based on priority: QNN > DML > NPU > GPU > CPU + bool hasGpu = false; + bool hasNpu = false; + bool hasCpu = false; + + // Check all provider options from decoder-level and pipeline-level + var allProviderOptions = GetAllProviderOptions(config); + foreach (var provider in allProviderOptions) + { + var accelerator = CheckProviderForAccelerator(provider, ref hasGpu, ref hasNpu, ref hasCpu); + if (accelerator.HasValue) { - return true; + return accelerator.Value; } } - return false; + if (hasNpu) + { + return HardwareAccelerator.NPU; + } + + if (hasGpu) + { + return HardwareAccelerator.GPU; + } + + return HardwareAccelerator.CPU; } - public static HardwareAccelerator GetHardwareAcceleratorFromConfig(string configContents) + private static IEnumerable GetAllProviderOptions(GenAIConfig config) { - if (configContents.Contains(""""backend_path": "QnnHtp.dll"""", StringComparison.OrdinalIgnoreCase)) + foreach (var provider in config.Model.Decoder.SessionOptions.ProviderOptions) + { + yield return provider; + } + + if (config.Model.Decoder.Pipeline == null) + { + yield break; + } + + foreach (var pipelineItem in config.Model.Decoder.Pipeline) + { + if (pipelineItem.Stages == null) + { + continue; + } + + foreach (var stageEntry in pipelineItem.Stages) + { + PipelineStage? stage = null; + try + { + stage = JsonSerializer.Deserialize(stageEntry.Value.GetRawText(), SourceGenerationContext.Default.PipelineStage); + } + catch (JsonException) + { + continue; + } + + if (stage?.SessionOptions?.ProviderOptions != null) + { + foreach (var provider in stage.SessionOptions.ProviderOptions) + { + yield return provider; + } + } + } + } + } + + private static HardwareAccelerator? CheckProviderForAccelerator(ProviderOptions provider, ref bool hasGpu, ref bool hasNpu, ref bool hasCpu) + { + if (provider.HasProvider("qnn")) { return HardwareAccelerator.QNN; } + if (provider.HasProvider("dml")) + { + return HardwareAccelerator.DML; + } + + var openvinoOptions = provider.GetProviderOptions("OpenVINO"); + if (openvinoOptions != null && openvinoOptions.TryGetValue("device_type", out var deviceType)) + { + var devType = deviceType.ToLowerInvariant(); + if (devType == "npu") + { + hasNpu = true; + } + else if (devType == "gpu") + { + hasGpu = true; + } + else if (devType == "cpu") + { + hasCpu = true; + } + } + + if (provider.HasProvider("vitisai")) + { + hasNpu = true; + } + + if (provider.HasProvider("cpu")) + { + hasCpu = true; + } + + return null; + } + + /// + /// Validates that the execution providers specified in the genai_config.json are available on this device. + /// + /// A tuple with (isValid, unavailableProviders) + private static (bool IsValid, List UnavailableProviders) ValidateExecutionProviders(string configContents) + { var config = JsonSerializer.Deserialize(configContents, SourceGenerationContext.Default.GenAIConfig); if (config == null) { - throw new InvalidDataException("genai_config.json is not valid"); + return (false, new List { "Invalid genai_config.json" }); } - if (config.Model.Decoder.SessionOptions.ProviderOptions.Any(p => p.Dml != null)) + var availableEPs = DeviceUtils.GetEpDevices() + .Select(device => device.EpName) + .Distinct() + .ToList(); + var unavailableProviders = new List(); + + var providerMapping = new Dictionary(StringComparer.OrdinalIgnoreCase) { - return HardwareAccelerator.DML; + { "qnn", ExecutionProviderNames.QNN }, + { "dml", ExecutionProviderNames.DML }, + { "openvino", ExecutionProviderNames.OpenVINO }, + { "vitisai", ExecutionProviderNames.VitisAI }, + { "cuda", ExecutionProviderNames.CUDA }, + { "tensorrt", ExecutionProviderNames.TensorRT }, + { "cpu", ExecutionProviderNames.CPU } + }; + + var allProviderOptions = GetAllProviderOptions(config); + foreach (var provider in allProviderOptions) + { + if (provider.ExtensionData == null) + { + continue; + } + + foreach (var providerKey in provider.ExtensionData.Keys) + { + // Skip CPU as it's always available + if (providerKey.Equals("cpu", StringComparison.OrdinalIgnoreCase)) + { + continue; + } + + if (providerMapping.TryGetValue(providerKey, out var expectedEP)) + { + if (!availableEPs.Any(ep => ep.Equals(expectedEP, StringComparison.OrdinalIgnoreCase))) + { + if (!unavailableProviders.Contains(providerKey, StringComparer.OrdinalIgnoreCase)) + { + unavailableProviders.Add(providerKey); + } + } + } + } } - return HardwareAccelerator.CPU; + return (unavailableProviders.Count == 0, unavailableProviders); } public static async Task AddModelFromLocalFilePath(string filepath, string name, List modelTypes)