Automatically selects the most suitable model based on available GPU memory. Attributes: model_requirements (dict): Model names mapped to their minimum GPU requirements (MB) sorted_models (list): Models sorted by memory requirements (descending) and HF priority extra
| 19 | |
| 20 | |
| 21 | class AutoModelSwitcher: |
| 22 | """Automatically selects the most suitable model based on available GPU memory. |
| 23 | |
| 24 | Attributes: |
| 25 | model_requirements (dict): Model names mapped to their minimum GPU requirements (MB) |
| 26 | sorted_models (list): Models sorted by memory requirements (descending) and HF priority |
| 27 | extra_memory (int): Additional memory buffer to reserve for other processes |
| 28 | get_memory (callable): Function to check available GPU memory |
| 29 | available_mb (float): Current available GPU memory in MB |
| 30 | """ |
| 31 | |
| 32 | def __init__(self, model_requirements, get_memory_func=None, extra_memory=0): |
| 33 | """Initialize the model switcher. |
| 34 | |
| 35 | Args: |
| 36 | model_requirements (dict): {model_name: min_required_memory(MB)} |
| 37 | get_memory_func (callable, optional): Custom function to check available GPU memory |
| 38 | extra_memory (int, optional): Additional memory buffer to reserve (MB) |
| 39 | """ |
| 40 | self.model_requirements = model_requirements |
| 41 | # Sort models by: 1. Memory requirements (descending) 2. HF models first |
| 42 | self.sorted_models = sorted( |
| 43 | model_requirements.items(), |
| 44 | key=lambda x: (-x[1], '-HF' not in x[0]) |
| 45 | ) |
| 46 | |
| 47 | self.extra_memory = extra_memory |
| 48 | |
| 49 | # Initialize memory checking method |
| 50 | if get_memory_func is None: |
| 51 | self.get_memory = self._default_memory_check |
| 52 | else: |
| 53 | self.get_memory = get_memory_func |
| 54 | |
| 55 | self.available_mb = self._default_memory_check() |
| 56 | |
| 57 | def _default_memory_check(self, gpu_id=0): |
| 58 | """Check available GPU memory using GPUtil. |
| 59 | |
| 60 | Args: |
| 61 | gpu_id (int, optional): Target GPU device ID |
| 62 | |
| 63 | Returns: |
| 64 | float: Available memory in MB |
| 65 | |
| 66 | Raises: |
| 67 | RuntimeError: If no GPUs are found |
| 68 | IndexError: If specified GPU ID is invalid |
| 69 | """ |
| 70 | gpus = GPUtil.getGPUs() |
| 71 | |
| 72 | if not gpus: |
| 73 | raise RuntimeError("No available GPUs detected") |
| 74 | |
| 75 | if gpu_id >= len(gpus): |
| 76 | raise IndexError(f"Invalid GPU ID {gpu_id}. Only {len(gpus)} GPUs available") |
| 77 | |
| 78 | gpu = gpus[gpu_id] |
no outgoing calls
no test coverage detected