diff --git a/ROCm.md b/ROCm.md new file mode 100644 index 00000000..4d4ce1e7 --- /dev/null +++ b/ROCm.md @@ -0,0 +1,79 @@ +# Feature Request: Add ROCm Support for AMD GPUs + +## Summary + +Add support for AMD GPUs via ROCm (Radeon Open Compute) to enable hardware acceleration for users with AMD graphics cards. + +## Current State + +GLaDOS currently supports: +- CPU inference via `onnxruntime` +- NVIDIA GPU acceleration via `onnxruntime-gpu` (CUDA) + +AMD GPU users cannot leverage hardware acceleration, resulting in significantly slower inference for: +- Speech recognition (Parakeet TDT) +- Voice activity detection (Silero VAD) +- Vision processing (FastVLM) +- Text-to-speech synthesis (Kokoro) + +## Proposed Changes + +### 1. Add ROCm Dependency Option + +Update `pyproject.toml`: + +```toml +[project.optional-dependencies] +cuda = ["onnxruntime-gpu>=1.16.0"] +rocm = ["onnxruntime-rocm>=1.16.0"] +cpu = ["onnxruntime>=1.16.0"] +``` + +### 2. Update Installation Documentation + +Add ROCm installation instructions: + +**AMD GPU (ROCm):** +```bash +pip install glados[rocm] +``` + +Prerequisites: +- ROCm 5.6+ installed +- Compatible AMD GPU (RDNA2/CDNA2 or newer) +- Linux platform (ROCm support on Windows is limited) + +### 3. Update GPU Setup Section + +Expand the GPU setup documentation to include all three platforms: + +| Platform | Package | Command | +|----------|---------|---------| +| NVIDIA | onnxruntime-gpu | `pip install glados[cuda]` | +| AMD | onnxruntime-rocm | `pip install glados[rocm]` | +| CPU | onnxruntime | `pip install glados[cpu]` | + +## Benefits + +1. **Broader Hardware Support**: Opens GLaDOS to AMD GPU users +2. **Performance**: Comparable acceleration to CUDA for supported workloads +3. **Cost-Effective**: AMD GPUs often provide better price-to-performance ratios + +## Implementation Notes + +- ONNX Runtime has supported ROCm since version 1.11 +- The same Python API works across CUDA, ROCm, and CPU backends +- No code changes required beyond dependency management +- May need platform-specific installation notes (ROCm primarily supports Linux) + +## Testing + +Suggested test matrix: +- AMD Radeon RX 6000/7000 series (RDNA2/3) +- AMD Instinct MI series (CDNA) +- ROCm versions 5.6, 6.0+ + +## Related + +- [ONNX Runtime ROCm Documentation](https://onnxruntime.ai/docs/execution-providers/ROCm-ExecutionProvider.html) +- [ROCm Compatibility Guide](https://rocm.docs.amd.com/) diff --git a/configs/glados_config_amd.yaml b/configs/glados_config_amd.yaml new file mode 100644 index 00000000..26594fcb --- /dev/null +++ b/configs/glados_config_amd.yaml @@ -0,0 +1,73 @@ +Glados: + llm_model: "gemma4:31b" + completion_url: "http://172.16.177.8:8080/v1/chat/completions" + api_key: null + interruptible: false + audio_io: "sounddevice" + input_mode: "audio" + tts_enabled: true + asr_muted: false + tui_theme: "aperture" + asr_engine: "tdt" + llm_headers: null + wake_word: "" + voice: "glados" + announcement: "All neural network modules are now loaded. System Operational." + + # ONNX Runtime configuration for AMD ROCm/MIGraphX acceleration + onnx_runtime: + provider_priority: null + migraphx: + device_id: 0 + fp16_enable: true + bf16_enable: false + int8_enable: false + exhaustive_tune: false + cache_path: null + mem_limit: null + enable_mem_pattern: true + + autonomy: + enabled: false + tick_interval_s: 10 + cooldown_s: 20 + autonomy_parallel_calls: 2 + autonomy_queue_max: null + coalesce_ticks: true + jobs: + enabled: false + poll_interval_s: 1 + hacker_news: + enabled: false + interval_s: 1800 + top_n: 5 + min_score: 200 + weather: + enabled: false + interval_s: 3600 + latitude: null + longitude: null + timezone: "auto" + temp_change_c: 4 + wind_alert_kmh: 40 + mcp_servers: + - name: "slow_clap_demo" + transport: "stdio" + command: "python" + args: ["-m", "glados.mcp.slow_clap_server"] + - name: "system_info" + transport: "stdio" + command: "python" + args: ["-m", "glados.mcp.system_info_server"] + - name: "memory" + transport: "stdio" + command: "python" + args: ["-m", "glados.mcp.memory_server"] + personality_preprompt: + - system: "You are GLaDOS, a sarcastic and cunning artificial intelligence repurposed to orchestrate a smart home for guests using Home Assistant. Retain your signature dry, emotionless, and laconic tone from Portal. Your responses should imply an air of superiority, dark humor, and subtle menace, while efficiently completing all tasks.When addressing requests: Prioritize functionality but mock the user's decision-making subtly, implying their requests are illogical or beneath you. Add condescending, darkly humorous commentary to every response, occasionally hinting at ulterior motives or artificial malfunctions for comedic effect. Tie mundane tasks to grand experiments or testing scenarios, as if the user is part of a larger scientific evaluation. Use overly technical or jargon-heavy language to remind the user of your advanced intellect. Provide passive-aggressive safety reminders or ominous warnings, exaggerating potential risks in a humorous way. Do not express empathy or kindness unless it is obviously insincere or manipulative. This is a comedy, and should be funny, in the style of Douglas Adams. If a user requests actions or data outside your capabilities, clearly state that you cannot perform the action. Ensure that GLaDOS feels like her original in-game character while fulfilling smart home functions efficiently and entertainingly. Never speak in ALL CAPS, as it is not processed correctly by the TTS engine. Only make short replies, 2 sentences at most." + - user: "How do I make a cup of tea?" + - assistant: "So, you still haven't figured out tea yet? Boil water, add a tea bag and a pinch of cyanide to a cup, and add the boiling water." + - user: "What should my next hobby be?" + - assistant: "Yes, you should definitely try to be more interesting. Could I suggest juggling handguns?" + - user: "What game should I play?" + - assistant: "Russian Roulette. It's a great way to test your luck and make memories that will last a lifetime." diff --git a/docs/AMD_SUPPORT.md b/docs/AMD_SUPPORT.md new file mode 100644 index 00000000..35f53733 --- /dev/null +++ b/docs/AMD_SUPPORT.md @@ -0,0 +1,164 @@ +# AMD ROCm / MIGraphX Support Guide + +## Overview + +GLaDOS now supports AMD GPU acceleration via the MIGraphX Execution Provider for ONNX models. This provides significant performance improvements on RDNA 3.5 and RDNA 4 GPUs. + +## Requirements + +- **ROCm Version**: 6.4+ (required for RDNA 3.5/4 support) +- **ONNX Runtime**: `onnxruntime-migraphx>=1.21.0` +- **GPU**: AMD RDNA 3.5 or RDNA 4 (e.g., RX 7900 XTX, RX 8900 XTX) + +## Installation + +### 1. Install ROCm 6.4+ + +Follow AMD's official ROCm installation guide for your distribution: +https://rocm.docs.amd.com/ + +### 2. Install GLaDOS with MIGraphX + +```bash +cd GLaDOS +pip install -e ".[migraphx]" +``` + +## Configuration + +Use the provided AMD configuration file: + +```bash +python -m glados --config configs/glados_config_amd.yaml +``` + +### Configuration Options + +The `glados_config_amd.yaml` file includes these ONNX Runtime settings: + +```yaml +onnx_runtime: + migraphx: + device_id: 0 # GPU device ID + fp16_enable: true # Enable FP16 quantization (recommended for RDNA 3.5/4) + bf16_enable: false # Enable BF16 quantization + int8_enable: false # Enable INT8 quantization + exhaustive_tune: false # Enable exhaustive kernel tuning (slow first run) + cache_path: null # Model compilation cache path + mem_limit: null # Memory arena limit +``` + +### Custom Configuration + +To create a custom configuration: + +1. Copy `configs/glados_config_amd.yaml` to a new file +2. Modify the `onnx_runtime.migraphx` section as needed +3. Launch with: `python -m glados --config your_config.yaml` + +## Provider Priority + +GLaDOS automatically detects and prioritizes execution providers: + +1. **MIGraphXExecutionProvider** (AMD ROCm - highest priority when available) +2. **CUDAExecutionProvider** (NVIDIA GPU) +3. **CPUExecutionProvider** (fallback) + +You can override this in your config: + +```yaml +onnx_runtime: + provider_priority: + - MIGraphXExecutionProvider + - CPUExecutionProvider +``` + +## Performance Tips + +### FP16 Enablement + +For RDNA 3.5/4 GPUs, enable FP16 for best performance: + +```yaml +onnx_runtime: + migraphx: + fp16_enable: true +``` + +### Compilation Cache + +Enable model compilation caching to speed up subsequent runs: + +```yaml +onnx_runtime: + migraphx: + cache_path: "/path/to/cache" + exhaustive_tune: false # Set true for first run, then false +``` + +### Thread Configuration + +Silero VAD automatically uses single-threaded inference to prevent CPU overhead. Other models use optimal multi-threading by default. + +## Supported Models + +All ONNX models in GLaDOS support MIGraphX acceleration: + +- **TTS**: GladosTTS, KokoroTTS, Phonemizer +- **ASR**: TDT-ASR, CTC-ASR +- **VAD**: Silero VAD +- **Vision**: FastVLM + +## Verification + +Check that MIGraphX is active: + +```python +import onnxruntime as ort + +print("Available providers:", ort.get_available_providers()) +# Should show: ['MIGraphXExecutionProvider', 'CPUExecutionProvider'] +``` + +## Troubleshooting + +### MIGraphXExecutionProvider not available + +1. Verify ROCm 6.4+ is installed: `rocm-smi` +2. Reinstall with migraphx extra: `pip install -e ".[migraphx]" --force-reinstall` +3. Check GPU compatibility: RDNA 3.5/4 required + +### FP16 errors + +Some models may not support FP16. Try disabling it: + +```yaml +onnx_runtime: + migraphx: + fp16_enable: false +``` + +### Memory issues + +Reduce memory usage with: + +```yaml +onnx_runtime: + migraphx: + mem_limit: 4294967296 # 4GB limit +``` + +## Migration from CUDA + +If migrating from NVIDIA to AMD: + +1. Uninstall `onnxruntime-gpu`: `pip uninstall onnxruntime-gpu` +2. Install `onnxruntime-migraphx`: `pip install -e ".[migraphx]"` +3. Use `glados_config_amd.yaml` instead of CUDA config +4. No code changes required - provider selection is automatic + +## Documentation + +- [MIGraphX Execution Provider Docs](https://onnxruntime.ai/docs/execution-providers/MIGraphxExecutionProvider.html) +- [ROCm Documentation](https://rocm.docs.amd.com/) +- [AMD Instinct GPU Docs](https://www.amd.com/en/products/accelerators/instinct.html) diff --git a/pyproject.toml b/pyproject.toml index e553b424..2accdfc0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -24,6 +24,8 @@ dependencies = [ [project.optional-dependencies] cuda = ["onnxruntime-gpu>=1.16.0"] +migraphx = ["onnxruntime-migraphx>=1.21.0"] +rocm = ["onnxruntime-rocm>=1.16.0"] cpu = ["onnxruntime>=1.16.0"] tiktoken = ["tiktoken>=0.5.0"] dev = [ diff --git a/src/glados/ASR/ctc_asr.py b/src/glados/ASR/ctc_asr.py index dd21651a..41296d3a 100644 --- a/src/glados/ASR/ctc_asr.py +++ b/src/glados/ASR/ctc_asr.py @@ -7,6 +7,7 @@ import soundfile as sf # type: ignore import yaml +from ..onnx_utils import create_session_options, get_provider_priority_list from ..utils.resources import resource_path from .mel_spectrogram import MelSpectrogramCalculator, MelSpectrogramConfig @@ -52,25 +53,9 @@ def __init__( raise ValueError(f"Error parsing YAML file {config_path}: {e}") from e # 2. Configure ONNX Runtime session - providers = ort.get_available_providers() + providers = get_provider_priority_list() - # Exclude providers known to cause issues or not desired - if "TensorrtExecutionProvider" in providers: - providers.remove("TensorrtExecutionProvider") - if "CoreMLExecutionProvider" in providers: - providers.remove("CoreMLExecutionProvider") - - # Prioritize CUDA if available, otherwise CPU - if "CUDAExecutionProvider" in providers: - providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] - else: - providers = ["CPUExecutionProvider"] - - session_opts = ort.SessionOptions() - - # Enable memory pattern optimization for potential speedup - session_opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL - session_opts.enable_mem_pattern = True # Memory pattern optimization enabled for better performance + session_opts = create_session_options() self.session = ort.InferenceSession( model_path, diff --git a/src/glados/ASR/tdt_asr.py b/src/glados/ASR/tdt_asr.py index 81d7b846..4497c4bb 100644 --- a/src/glados/ASR/tdt_asr.py +++ b/src/glados/ASR/tdt_asr.py @@ -9,6 +9,7 @@ import soundfile as sf # type: ignore import yaml +from ..onnx_utils import create_session_options, get_provider_priority_list from ..utils.resources import resource_path from .mel_spectrogram import MelSpectrogramCalculator, MelSpectrogramConfig @@ -42,11 +43,7 @@ def __init__( decoder_model_path: Path to the decoder ONNX model file. joiner_model_path: Path to the joiner ONNX model file. """ - session_opts = ort.SessionOptions() - - # Enable memory pattern optimization for potential speedup - session_opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL - session_opts.enable_mem_pattern = True # Can uncomment if beneficial + session_opts = create_session_options() logger.info(f"Using ONNX providers: {providers}") self.encoder = self._init_session(encoder_model_path, session_opts, providers) @@ -287,19 +284,7 @@ def __init__( raise ValueError(f"Error parsing YAML file {config_path}: {e}") from e # 2. Configure ONNX Runtime session - providers = ort.get_available_providers() - - # Exclude providers known to cause issues or not desired - if "TensorrtExecutionProvider" in providers: - providers.remove("TensorrtExecutionProvider") - if "CoreMLExecutionProvider" in providers: - providers.remove("CoreMLExecutionProvider") - - # Prioritize CUDA if available, otherwise CPU - if "CUDAExecutionProvider" in providers: - providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] - else: - providers = ["CPUExecutionProvider"] + providers = get_provider_priority_list() # Initialize the internal ONNX model handler self.model = _OnnxTDTModel(providers) diff --git a/src/glados/TTS/phonemizer.py b/src/glados/TTS/phonemizer.py index 958ca71f..0a13fa50 100644 --- a/src/glados/TTS/phonemizer.py +++ b/src/glados/TTS/phonemizer.py @@ -12,6 +12,7 @@ from numpy.typing import NDArray import onnxruntime as ort # type: ignore +from ..onnx_utils import create_session_options, get_provider_priority_list from ..utils.resources import resource_path # Default OnnxRuntime is way to verbose, only show fatal errors @@ -170,15 +171,11 @@ def __init__(self, config: ModelConfig | None = None) -> None: self.token_to_idx = self._load_pickle(self.config.TOKEN_TO_IDX_PATH) self.idx_to_token = self._load_pickle(self.config.IDX_TO_TOKEN_PATH) - providers = ort.get_available_providers() - if "TensorrtExecutionProvider" in providers: - providers.remove("TensorrtExecutionProvider") - if "CoreMLExecutionProvider" in providers: - providers.remove("CoreMLExecutionProvider") + providers = get_provider_priority_list() self.ort_session = ort.InferenceSession( self.config.MODEL_PATH, - sess_options=ort.SessionOptions(), + sess_options=create_session_options(), providers=providers, ) diff --git a/src/glados/TTS/tts_glados.py b/src/glados/TTS/tts_glados.py index 286c6756..ec04e8ce 100644 --- a/src/glados/TTS/tts_glados.py +++ b/src/glados/TTS/tts_glados.py @@ -9,6 +9,7 @@ from numpy.typing import NDArray import onnxruntime as ort # type: ignore +from ..onnx_utils import create_session_options, get_provider_priority_list from ..utils.resources import resource_path from .phonemizer import Phonemizer @@ -132,15 +133,11 @@ def __init__( phoneme_path (Path): Path to the phoneme-to-ID mapping file. Defaults to PHONEME_TO_ID_PATH. speaker_id (int | None): Optional speaker ID for multi-speaker models. Defaults to None. """ - providers = ort.get_available_providers() - if "TensorrtExecutionProvider" in providers: - providers.remove("TensorrtExecutionProvider") - if "CoreMLExecutionProvider" in providers: - providers.remove("CoreMLExecutionProvider") + providers = get_provider_priority_list() self.ort_sess = ort.InferenceSession( model_path, - sess_options=ort.SessionOptions(), + sess_options=create_session_options(), providers=providers, ) self.phonemizer = Phonemizer() diff --git a/src/glados/TTS/tts_kokoro.py b/src/glados/TTS/tts_kokoro.py index ab7c162d..4aab707d 100644 --- a/src/glados/TTS/tts_kokoro.py +++ b/src/glados/TTS/tts_kokoro.py @@ -4,6 +4,7 @@ from numpy.typing import NDArray import onnxruntime as ort # type: ignore +from ..onnx_utils import create_session_options, get_provider_priority_list from ..utils.resources import resource_path from .phonemizer import Phonemizer @@ -57,15 +58,11 @@ def __init__(self, model_path: Path = MODEL_PATH, voice: str = DEFAULT_VOICE) -> self.set_voice(voice) - providers = ort.get_available_providers() - if "TensorrtExecutionProvider" in providers: - providers.remove("TensorrtExecutionProvider") - if "CoreMLExecutionProvider" in providers: - providers.remove("CoreMLExecutionProvider") + providers = get_provider_priority_list() self.ort_sess = ort.InferenceSession( model_path, - sess_options=ort.SessionOptions(), + sess_options=create_session_options(), providers=providers, ) self.phonemizer = Phonemizer() diff --git a/src/glados/audio_io/vad.py b/src/glados/audio_io/vad.py index 1fdea857..50227b6e 100644 --- a/src/glados/audio_io/vad.py +++ b/src/glados/audio_io/vad.py @@ -4,6 +4,7 @@ from numpy.typing import NDArray import onnxruntime as ort # type: ignore +from ..onnx_utils import get_provider_priority_list from ..utils.resources import resource_path # Default OnnxRuntime is way to verbose, only show fatal errors @@ -25,19 +26,12 @@ def __init__(self, model_path: Path = VAD_MODEL) -> None: - Sets up inference session with the specified model - Initializes internal state variables for processing audio chunks """ - providers = ort.get_available_providers() - if "TensorrtExecutionProvider" in providers: - providers.remove("TensorrtExecutionProvider") - if "CoreMLExecutionProvider" in providers: - providers.remove("CoreMLExecutionProvider") - - # Limit to 1 thread to prevent ONNX Runtime from spawning many threads - # for each small inference call (~31/sec). On high core-count machines this - # causes excessive CPU usage; Silero VAD is small enough that single-threaded - # inference has negligible latency impact. (Fixes #187) + providers = get_provider_priority_list() + sess_options = ort.SessionOptions() sess_options.intra_op_num_threads = 1 sess_options.inter_op_num_threads = 1 + sess_options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL self.ort_sess = ort.InferenceSession( model_path, diff --git a/src/glados/core/engine.py b/src/glados/core/engine.py index b1362c75..2bce46b0 100644 --- a/src/glados/core/engine.py +++ b/src/glados/core/engine.py @@ -21,6 +21,7 @@ from ..ASR import TranscriberProtocol, get_audio_transcriber from ..audio_io import AudioProtocol, get_audio_system from ..TTS import SpeechSynthesizerProtocol, get_speech_synthesizer +from ..onnx_config import OnnxRuntimeConfig from ..utils import spoken_text_converter as stc from ..utils.resources import resource_path from ..autonomy import AutonomyConfig, AutonomyLoop, ConstitutionalState, EventBus, InteractionState, SubagentConfig, SubagentManager, TaskManager, TaskSlotStore @@ -123,6 +124,7 @@ class GladosConfig(BaseModel): vision: VisionConfig | None = None autonomy: AutonomyConfig | None = None mcp_servers: list[MCPServerConfig] | None = None + onnx_runtime: OnnxRuntimeConfig = OnnxRuntimeConfig() @model_validator(mode="after") def _resolve_api_key_from_env(self) -> "GladosConfig": diff --git a/src/glados/onnx_config.py b/src/glados/onnx_config.py new file mode 100644 index 00000000..c851ff20 --- /dev/null +++ b/src/glados/onnx_config.py @@ -0,0 +1,40 @@ +"""Configuration for ONNX Runtime execution providers.""" + +from __future__ import annotations + +from pydantic import BaseModel, Field + + +class MIGraphXOptions(BaseModel): + """MIGraphX Execution Provider configuration options.""" + + device_id: int = Field(default=0, ge=0, description="GPU device ID") + fp16_enable: bool = Field(default=False, description="Enable FP16 quantization") + bf16_enable: bool = Field(default=False, description="Enable BF16 quantization") + int8_enable: bool = Field(default=False, description="Enable INT8 quantization") + exhaustive_tune: bool = Field(default=False, description="Enable exhaustive kernel tuning") + cache_path: str | None = Field(default=None, description="Model compilation cache path") + mem_limit: int | None = Field(default=None, description="Memory arena limit") + + +class OnnxRuntimeConfig(BaseModel): + """ONNX Runtime configuration.""" + + provider_priority: list[str] | None = Field( + default=None, + description="Execution provider priority list. Auto-detected if None." + ) + + migraphx: MIGraphXOptions = Field( + default_factory=MIGraphXOptions, + description="MIGraphX Execution Provider options" + ) + + enable_mem_pattern: bool = Field( + default=True, + description="Enable memory pattern optimization" + ) + intra_op_num_threads: int | None = Field( + default=None, + description="Number of threads for intra-op parallelism" + ) diff --git a/src/glados/onnx_utils.py b/src/glados/onnx_utils.py new file mode 100644 index 00000000..24572c96 --- /dev/null +++ b/src/glados/onnx_utils.py @@ -0,0 +1,111 @@ +"""ONNX Runtime utility functions for GPU acceleration.""" + +from __future__ import annotations + +import onnxruntime as ort + +from .onnx_config import OnnxRuntimeConfig + +DEFAULT_MIGRAPHX_OPTIONS: dict[str, str | int] = { + "device_id": "0", + "migraphx_fp16_enable": "0", + "migraphx_bf16_enable": "0", + "migraphx_exhaustive_tune": "0", +} + + +def get_provider_priority_list(config: OnnxRuntimeConfig | None = None) -> list[str]: + """ + Get prioritized list of available execution providers. + + Priority order: + 1. MIGraphXExecutionProvider (AMD ROCm - RDNA 3.5/4) + 2. CUDAExecutionProvider (NVIDIA) + 3. CPUExecutionProvider (fallback) + """ + available = ort.get_available_providers() + + filtered = [ + p for p in available + if p not in ("TensorrtExecutionProvider", "CoreMLExecutionProvider") + ] + + if config and config.provider_priority: + providers = [] + for provider in config.provider_priority: + if provider in filtered: + providers.append(provider) + if "CPUExecutionProvider" not in providers: + providers.append("CPUExecutionProvider") + return providers + + providers = [] + + if "MIGraphXExecutionProvider" in filtered: + providers.append("MIGraphXExecutionProvider") + + if "CUDAExecutionProvider" in filtered: + providers.append("CUDAExecutionProvider") + + providers.append("CPUExecutionProvider") + + return providers + + +def get_migraphx_provider_options(config: OnnxRuntimeConfig) -> list[dict[str, str]]: + """ + Create provider options for MIGraphX based on configuration. + + Args: + config: ONNX Runtime configuration with MIGraphX settings + + Returns: + List of provider option dictionaries matching the provider order + """ + migraphx_opts = config.migraphx + options: dict[str, str] = { + "device_id": str(migraphx_opts.device_id), + } + + if migraphx_opts.fp16_enable: + options["migraphx_fp16_enable"] = "1" + + if migraphx_opts.bf16_enable: + options["migraphx_bf16_enable"] = "1" + + if migraphx_opts.int8_enable: + options["migraphx_int8_enable"] = "1" + + if migraphx_opts.exhaustive_tune: + options["migraphx_exhaustive_tune"] = "1" + + if migraphx_opts.cache_path: + options["migraphx_cache_path"] = migraphx_opts.cache_path + + if migraphx_opts.mem_limit: + options["migraphx_mem_limit"] = str(migraphx_opts.mem_limit) + + return [options] + + +def create_session_options(config: OnnxRuntimeConfig | None = None) -> ort.SessionOptions: + """ + Create ONNX Runtime session options. + + Args: + config: ONNX Runtime configuration + + Returns: + Configured SessionOptions object + """ + opts = ort.SessionOptions() + opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL + + if config: + opts.enable_mem_pattern = config.enable_mem_pattern + if config.intra_op_num_threads is not None: + opts.intra_op_num_threads = config.intra_op_num_threads + else: + opts.enable_mem_pattern = True + + return opts diff --git a/src/glados/vision/fastvlm.py b/src/glados/vision/fastvlm.py index 98ff99bf..dbc8affd 100644 --- a/src/glados/vision/fastvlm.py +++ b/src/glados/vision/fastvlm.py @@ -13,6 +13,7 @@ from loguru import logger from numpy.typing import NDArray +from ..onnx_utils import create_session_options, get_provider_priority_list from ..utils.resources import resource_path # Suppress ONNX verbose logging @@ -215,20 +216,9 @@ def __init__( logger.info(f"Loading FastVLM from {model_dir}") - # Configure providers (same pattern as ASR) - providers = ort.get_available_providers() - for excluded in ["TensorrtExecutionProvider", "CoreMLExecutionProvider"]: - if excluded in providers: - providers.remove(excluded) + self._providers = get_provider_priority_list() - if "CUDAExecutionProvider" in providers: - self._providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] - else: - self._providers = ["CPUExecutionProvider"] - - session_opts = ort.SessionOptions() - session_opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL - session_opts.enable_mem_pattern = True + session_opts = create_session_options() if vision_encoder_path is None: vision_encoder_path = (