codekingpro/portable-devtools
114k
1import time2from dataclasses import dataclass, field3from datetime import datetime4from enum import Enum5from typing import TYPE_CHECKING, Optional6 7from huggingface_hub.errors import InferenceEndpointError, InferenceEndpointTimeoutError8 9from .utils import get_session, logging, parse_datetime10 11 12if TYPE_CHECKING:13 from .hf_api import HfApi14 from .inference._client import InferenceClient15 from .inference._generated._async_client import AsyncInferenceClient16 17logger = logging.get_logger(__name__)18 19 20class InferenceEndpointStatus(str, Enum):21 PENDING = "pending"22 INITIALIZING = "initializing"23 UPDATING = "updating"24 UPDATE_FAILED = "updateFailed"25 RUNNING = "running"26 PAUSED = "paused"27 FAILED = "failed"28 SCALED_TO_ZERO = "scaledToZero"29 30 31class InferenceEndpointType(str, Enum):32 PUBlIC = "public"33 PROTECTED = "protected"34 PRIVATE = "private"35 36 37class InferenceEndpointScalingMetric(str, Enum):38 PENDING_REQUESTS = "pendingRequests"39 HARDWARE_USAGE = "hardwareUsage"40 41 42@dataclass43class InferenceEndpoint:44 """45 Contains information about a deployed Inference Endpoint.46 47 Args:48 name (`str`):49 The unique name of the Inference Endpoint.50 namespace (`str`):51 The namespace where the Inference Endpoint is located.52 repository (`str`):53 The name of the model repository deployed on this Inference Endpoint.54 status ([`InferenceEndpointStatus`]):55 The current status of the Inference Endpoint.56 url (`str`, *optional*):57 The URL of the Inference Endpoint, if available. Only a deployed Inference Endpoint will have a URL.58 framework (`str`):59 The machine learning framework used for the model.60 revision (`str`):61 The specific model revision deployed on the Inference Endpoint.62 task (`str`):63 The task associated with the deployed model.64 created_at (`datetime.datetime`):65 The timestamp when the Inference Endpoint was created.66 updated_at (`datetime.datetime`):67 The timestamp of the last update of the Inference Endpoint.68 type ([`InferenceEndpointType`]):69 The type of the Inference Endpoint (public, protected, private).70 raw (`dict`):71 The raw dictionary data returned from the API.72 token (`str` or `bool`, *optional*):73 Authentication token for the Inference Endpoint, if set when requesting the API. Will default to the74 locally saved token if not provided. Pass `token=False` if you don't want to send your token to the server.75 76 Example:77 ```python78 >>> from huggingface_hub import get_inference_endpoint79 >>> endpoint = get_inference_endpoint("my-text-to-image")80 >>> endpoint81 InferenceEndpoint(name='my-text-to-image', ...)82 83 # Get status84 >>> endpoint.status85 'running'86 >>> endpoint.url87 'https://my-text-to-image.region.vendor.endpoints.huggingface.cloud'88 89 # Run inference90 >>> endpoint.client.text_to_image(...)91 92 # Pause endpoint to save $$$93 >>> endpoint.pause()94 95 # ...96 # Resume and wait for deployment97 >>> endpoint.resume()98 >>> endpoint.wait()99 >>> endpoint.client.text_to_image(...)100 ```101 """102 103 # Field in __repr__104 name: str = field(init=False)105 namespace: str106 repository: str = field(init=False)107 status: InferenceEndpointStatus = field(init=False)108 health_route: str = field(init=False)109 url: str | None = field(init=False)110 111 # Other fields112 framework: str = field(repr=False, init=False)113 revision: str = field(repr=False, init=False)114 task: str = field(repr=False, init=False)115 created_at: datetime = field(repr=False, init=False)116 updated_at: datetime = field(repr=False, init=False)117 type: InferenceEndpointType = field(repr=False, init=False)118 119 # Raw dict from the API120 raw: dict = field(repr=False)121 122 # Internal fields123 _token: str | bool | None = field(repr=False, compare=False)124 _api: "HfApi" = field(repr=False, compare=False)125 126 @classmethod127 def from_raw(128 cls, raw: dict, namespace: str, token: str | bool | None = None, api: Optional["HfApi"] = None129 ) -> "InferenceEndpoint":130 """Initialize object from raw dictionary."""131 if api is None:132 from .hf_api import HfApi133 134 api = HfApi()135 if token is None:136 token = api.token137 138 # All other fields are populated in __post_init__139 return cls(raw=raw, namespace=namespace, _token=token, _api=api)140 141 def __post_init__(self) -> None:142 """Populate fields from raw dictionary."""143 self._populate_from_raw()144 145 @property146 def client(self) -> "InferenceClient":147 """Returns a client to make predictions on this Inference Endpoint.148 149 Returns:150 [`InferenceClient`]: an inference client pointing to the deployed endpoint.151 152 Raises:153 [`InferenceEndpointError`]: If the Inference Endpoint is not yet deployed.154 """155 if self.url is None:156 raise InferenceEndpointError(157 "Cannot create a client for this Inference Endpoint as it is not yet deployed. "158 "Please wait for the Inference Endpoint to be deployed using `endpoint.wait()` and try again."159 )160 from .inference._client import InferenceClient161 162 return InferenceClient(163 model=self.url,164 token=self._token, # type: ignore # boolean token shouldn't be possible. In practice it's ok.165 )166 167 @property168 def async_client(self) -> "AsyncInferenceClient":169 """Returns a client to make predictions on this Inference Endpoint.170 171 Returns:172 [`AsyncInferenceClient`]: an asyncio-compatible inference client pointing to the deployed endpoint.173 174 Raises:175 [`InferenceEndpointError`]: If the Inference Endpoint is not yet deployed.176 """177 if self.url is None:178 raise InferenceEndpointError(179 "Cannot create a client for this Inference Endpoint as it is not yet deployed. "180 "Please wait for the Inference Endpoint to be deployed using `endpoint.wait()` and try again."181 )182 from .inference._generated._async_client import AsyncInferenceClient183 184 return AsyncInferenceClient(185 model=self.url,186 token=self._token, # type: ignore # boolean token shouldn't be possible. In practice it's ok.187 )188 189 def wait(self, timeout: int | None = None, refresh_every: int = 5) -> "InferenceEndpoint":190 """Wait for the Inference Endpoint to be deployed.191 192 Information from the server will be fetched every 1s. If the Inference Endpoint is not deployed after `timeout`193 seconds, a [`InferenceEndpointTimeoutError`] will be raised. The [`InferenceEndpoint`] will be mutated in place with the latest194 data.195 196 Args:197 timeout (`int`, *optional*):198 The maximum time to wait for the Inference Endpoint to be deployed, in seconds. If `None`, will wait199 indefinitely.200 refresh_every (`int`, *optional*):201 The time to wait between each fetch of the Inference Endpoint status, in seconds. Defaults to 5s.202 203 Returns:204 [`InferenceEndpoint`]: the same Inference Endpoint, mutated in place with the latest data.205 206 Raises:207 [`InferenceEndpointError`]208 If the Inference Endpoint ended up in a failed state.209 [`InferenceEndpointTimeoutError`]210 If the Inference Endpoint is not deployed after `timeout` seconds.211 """212 if timeout is not None and timeout < 0:213 raise ValueError("`timeout` cannot be negative.")214 if refresh_every <= 0:215 raise ValueError("`refresh_every` must be positive.")216 217 start = time.time()218 while True:219 if self.status == InferenceEndpointStatus.FAILED:220 raise InferenceEndpointError(221 f"Inference Endpoint {self.name} failed to deploy. Please check the logs for more information."222 )223 if self.status == InferenceEndpointStatus.UPDATE_FAILED:224 raise InferenceEndpointError(225 f"Inference Endpoint {self.name} failed to update. Please check the logs for more information."226 )227 if self.status == InferenceEndpointStatus.RUNNING and self.url is not None:228 # Verify the endpoint is actually reachable229 _health_url = f"{self.url.rstrip('/')}/{self.health_route.lstrip('/')}"230 response = get_session().get(_health_url, headers=self._api._build_hf_headers(token=self._token))231 if response.status_code == 200:232 logger.info("Inference Endpoint is ready to be used.")233 return self234 235 if timeout is not None:236 if time.time() - start > timeout:237 raise InferenceEndpointTimeoutError("Timeout while waiting for Inference Endpoint to be deployed.")238 logger.info(f"Inference Endpoint is not deployed yet ({self.status}). Waiting {refresh_every}s...")239 time.sleep(refresh_every)240 self.fetch()241 242 def fetch(self) -> "InferenceEndpoint":243 """Fetch latest information about the Inference Endpoint.244 245 Returns:246 [`InferenceEndpoint`]: the same Inference Endpoint, mutated in place with the latest data.247 """248 obj = self._api.get_inference_endpoint(name=self.name, namespace=self.namespace, token=self._token) # type: ignore [arg-type]249 self.raw = obj.raw250 self._populate_from_raw()251 return self252 253 def update(254 self,255 *,256 # Compute update257 accelerator: str | None = None,258 instance_size: str | None = None,259 instance_type: str | None = None,260 min_replica: int | None = None,261 max_replica: int | None = None,262 scale_to_zero_timeout: int | None = None,263 # Model update264 repository: str | None = None,265 framework: str | None = None,266 revision: str | None = None,267 task: str | None = None,268 custom_image: dict | None = None,269 secrets: dict[str, str] | None = None,270 ) -> "InferenceEndpoint":271 """Update the Inference Endpoint.272 273 This method allows the update of either the compute configuration, the deployed model, or both. All arguments are274 optional but at least one must be provided.275 276 This is an alias for [`HfApi.update_inference_endpoint`]. The current object is mutated in place with the277 latest data from the server.278 279 Args:280 accelerator (`str`, *optional*):281 The hardware accelerator to be used for inference (e.g. `"cpu"`).282 instance_size (`str`, *optional*):283 The size or type of the instance to be used for hosting the model (e.g. `"x4"`).284 instance_type (`str`, *optional*):285 The cloud instance type where the Inference Endpoint will be deployed (e.g. `"intel-icl"`).286 min_replica (`int`, *optional*):287 The minimum number of replicas (instances) to keep running for the Inference Endpoint.288 max_replica (`int`, *optional*):289 The maximum number of replicas (instances) to scale to for the Inference Endpoint.290 scale_to_zero_timeout (`int`, *optional*):291 The duration in minutes before an inactive endpoint is scaled to zero.292 293 repository (`str`, *optional*):294 The name of the model repository associated with the Inference Endpoint (e.g. `"gpt2"`).295 framework (`str`, *optional*):296 The machine learning framework used for the model (e.g. `"custom"`).297 revision (`str`, *optional*):298 The specific model revision to deploy on the Inference Endpoint (e.g. `"6c0e6080953db56375760c0471a8c5f2929baf11"`).299 task (`str`, *optional*):300 The task on which to deploy the model (e.g. `"text-classification"`).301 custom_image (`dict`, *optional*):302 A custom Docker image to use for the Inference Endpoint. This is useful if you want to deploy an303 Inference Endpoint running on the `text-generation-inference` (TGI) framework (see examples).304 secrets (`dict[str, str]`, *optional*):305 Secret values to inject in the container environment.306 Returns:307 [`InferenceEndpoint`]: the same Inference Endpoint, mutated in place with the latest data.308 """309 # Make API call310 obj = self._api.update_inference_endpoint(311 name=self.name,312 namespace=self.namespace,313 accelerator=accelerator,314 instance_size=instance_size,315 instance_type=instance_type,316 min_replica=min_replica,317 max_replica=max_replica,318 scale_to_zero_timeout=scale_to_zero_timeout,319 repository=repository,320 framework=framework,321 revision=revision,322 task=task,323 custom_image=custom_image,324 secrets=secrets,325 token=self._token, # type: ignore [arg-type]326 )327 328 # Mutate current object329 self.raw = obj.raw330 self._populate_from_raw()331 return self332 333 def pause(self) -> "InferenceEndpoint":334 """Pause the Inference Endpoint.335 336 A paused Inference Endpoint will not be charged. It can be resumed at any time using [`InferenceEndpoint.resume`].337 This is different from scaling the Inference Endpoint to zero with [`InferenceEndpoint.scale_to_zero`], which338 would be automatically restarted when a request is made to it.339 340 This is an alias for [`HfApi.pause_inference_endpoint`]. The current object is mutated in place with the341 latest data from the server.342 343 Returns:344 [`InferenceEndpoint`]: the same Inference Endpoint, mutated in place with the latest data.345 """346 obj = self._api.pause_inference_endpoint(name=self.name, namespace=self.namespace, token=self._token) # type: ignore [arg-type]347 self.raw = obj.raw348 self._populate_from_raw()349 return self350 351 def resume(self, running_ok: bool = True) -> "InferenceEndpoint":352 """Resume the Inference Endpoint.353 354 This is an alias for [`HfApi.resume_inference_endpoint`]. The current object is mutated in place with the355 latest data from the server.356 357 Args:358 running_ok (`bool`, *optional*):359 If `True`, the method will not raise an error if the Inference Endpoint is already running. Defaults to360 `True`.361 362 Returns:363 [`InferenceEndpoint`]: the same Inference Endpoint, mutated in place with the latest data.364 """365 obj = self._api.resume_inference_endpoint(366 name=self.name, namespace=self.namespace, running_ok=running_ok, token=self._token367 ) # type: ignore [arg-type]368 self.raw = obj.raw369 self._populate_from_raw()370 return self371 372 def scale_to_zero(self) -> "InferenceEndpoint":373 """Scale Inference Endpoint to zero.374 375 An Inference Endpoint scaled to zero will not be charged. It will be resumed on the next request to it, with a376 cold start delay. This is different from pausing the Inference Endpoint with [`InferenceEndpoint.pause`], which377 would require a manual resume with [`InferenceEndpoint.resume`].378 379 This is an alias for [`HfApi.scale_to_zero_inference_endpoint`]. The current object is mutated in place with the380 latest data from the server.381 382 Returns:383 [`InferenceEndpoint`]: the same Inference Endpoint, mutated in place with the latest data.384 """385 obj = self._api.scale_to_zero_inference_endpoint(name=self.name, namespace=self.namespace, token=self._token) # type: ignore [arg-type]386 self.raw = obj.raw387 self._populate_from_raw()388 return self389 390 def delete(self) -> None:391 """Delete the Inference Endpoint.392 393 This operation is not reversible. If you don't want to be charged for an Inference Endpoint, it is preferable394 to pause it with [`InferenceEndpoint.pause`] or scale it to zero with [`InferenceEndpoint.scale_to_zero`].395 396 This is an alias for [`HfApi.delete_inference_endpoint`].397 """398 self._api.delete_inference_endpoint(name=self.name, namespace=self.namespace, token=self._token) # type: ignore [arg-type]399 400 def _populate_from_raw(self) -> None:401 """Populate fields from raw dictionary.402 403 Called in __post_init__ + each time the Inference Endpoint is updated.404 """405 # Repr fields406 self.name = self.raw["name"]407 self.repository = self.raw["model"]["repository"]408 self.status = self.raw["status"]["state"]409 self.url = self.raw["status"].get("url")410 self.health_route = self.raw["healthRoute"]411 412 # Other fields413 self.framework = self.raw["model"]["framework"]414 self.revision = self.raw["model"]["revision"]415 self.task = self.raw["model"]["task"]416 self.created_at = parse_datetime(self.raw["status"]["createdAt"])417 self.updated_at = parse_datetime(self.raw["status"]["updatedAt"])418 self.type = self.raw["type"]419 