| 79 | |
| 80 | |
| 81 | class ModelReference: |
| 82 | def __init__(self, model_url, model_port, device_id, max_n_images, model_name=None, dtype='auto', enforce_eager=False, use_external_endpoint=False): |
| 83 | self.model_url_to_start = model_url |
| 84 | self.model_port = model_port |
| 85 | self.model_url_to_log = model_url |
| 86 | self.model_to_log = None |
| 87 | self.model_prefix = None |
| 88 | self.device_id = device_id |
| 89 | self.model_name = model_name |
| 90 | self.max_n_images = max_n_images |
| 91 | self.dtype = dtype |
| 92 | self.enforce_eager = enforce_eager |
| 93 | self.use_external_endpoint = use_external_endpoint |
| 94 | |
| 95 | if self.use_external_endpoint: |
| 96 | # If using external endpoint, model_url is expected to be a config dict or path to config file |
| 97 | self.model_url_to_log = model_url |
| 98 | self.model_to_log = model_url |
| 99 | self.model_prefix = "external_model" |
| 100 | return |
| 101 | elif model_url is None: |
| 102 | if self.model_name: |
| 103 | logging.info(f'Using provided model name: {self.model_name}') |
| 104 | self.model_prefix = self.model_name.replace('/', '_').replace(':', '_') |
| 105 | else: |
| 106 | response = requests.get(f'http://localhost:{model_port}/model') |
| 107 | if response.status_code == 200: |
| 108 | model_name_from_response = response.json()['model'] |
| 109 | self.model_url_to_log = response.json()['model_url'] |
| 110 | self.model_to_log = self.model_url_to_log |
| 111 | self.model_prefix = model_name_from_response |
| 112 | else: |
| 113 | raise Exception(f"Failed to get model info from VLLM server, status code: {response.status_code}") |
| 114 | |
| 115 | else: |
| 116 | if _is_azure_blob_url(self.model_url_to_log): |
| 117 | raise NotImplementedError("Logging Azure Blob URLs is not implemented in this version.") |
| 118 | else: |
| 119 | # It's a local directory |
| 120 | self.model_to_log = self.model_url_to_log |
| 121 | self.model_prefix = Path(self.model_url_to_log).name.replace('/', '_').replace(':', '_') |
| 122 | |
| 123 | def log_2_mlflow(self): |
| 124 | mlflow.log_param('model', self.model_to_log) |
| 125 | mlflow.log_param('model_url', self.model_url_to_log) |
| 126 | |
| 127 | class Callback: |
| 128 | def __init__(self, callbacks = None): |