Source code for sagemaker.hyperpod.inference.hp_endpoint

  1from sagemaker.hyperpod.common.config.metadata import Metadata
  2from sagemaker.hyperpod.inference.config.constants import *
  3from sagemaker.hyperpod.common.utils import (
  4    get_default_namespace,
  5    get_cluster_instance_types,
  6    setup_logging,
  7    get_current_cluster,
  8    get_current_region,
  9)
 10from sagemaker.hyperpod.inference.config.hp_endpoint_config import (
 11    InferenceEndpointConfigStatus,
 12    _HPEndpoint,
 13)
 14from sagemaker.hyperpod.common.telemetry.telemetry_logging import (
 15    _hyperpod_telemetry_emitter,
 16)
 17from sagemaker.hyperpod.common.telemetry.constants import Feature
 18from sagemaker.hyperpod.inference.hp_endpoint_base import HPEndpointBase
 19from typing import Dict, List, Optional
 20from sagemaker_core.main.resources import Endpoint
 21from pydantic import Field, ValidationError
 22from kubernetes import client
 23
 24
[docs] 25class HPEndpoint(_HPEndpoint, HPEndpointBase): 26 metadata: Optional[Metadata] = Field(default=None) 27 status: Optional[InferenceEndpointConfigStatus] = Field(default=None) 28 29 def _create_internal(self, spec, debug=False): 30 """Shared internal create logic""" 31 logger = self.get_logger() 32 logger = setup_logging(logger, debug) 33 34 name = self.metadata.name if self.metadata else None 35 namespace = self.metadata.namespace if self.metadata else None 36 37 if not spec.endpointName and not name: 38 raise Exception('Either metadata name or endpoint name must be provided') 39 40 if not namespace: 41 namespace = get_default_namespace() 42 43 if not name: 44 name = spec.endpointName 45 46 # Create metadata object with labels and annotations if available 47 metadata = Metadata( 48 name=name, 49 namespace=namespace, 50 labels=self.metadata.labels if self.metadata else None, 51 annotations=self.metadata.annotations if self.metadata else None, 52 ) 53 54 if spec.instanceType: 55 self.validate_instance_type(spec.instanceType) 56 elif spec.instanceTypes: 57 for it in spec.instanceTypes: 58 self.validate_instance_type(it) 59 60 self.call_create_api( 61 metadata=metadata, 62 kind=INFERENCE_ENDPOINT_CONFIG_KIND, 63 spec=spec, 64 debug=debug, 65 ) 66 67 self.metadata = metadata 68 69 logger.info( 70 f"Creating sagemaker model and endpoint. Endpoint name: {spec.endpointName}.\n The process may take a few minutes..." 71 ) 72 73 @_hyperpod_telemetry_emitter(Feature.HYPERPOD, "create_endpoint") 74 def create( 75 self, 76 debug=False 77 ) -> None: 78 spec = _HPEndpoint(**self.model_dump(by_alias=True, exclude_none=True)) 79 self._create_internal(spec, debug) 80 81 @_hyperpod_telemetry_emitter(Feature.HYPERPOD, "create_endpoint_from_dict") 82 def create_from_dict( 83 self, 84 input: Dict, 85 debug=False 86 ) -> None: 87 spec = _HPEndpoint.model_validate(input, by_name=True) 88 self._create_internal(spec, debug) 89 90 91 def refresh(self): 92 if not self.metadata: 93 raise Exception( 94 "Metadata not found! Please provide object name and namespace in metadata field." 95 ) 96 97 response = self.call_get_api( 98 name=self.metadata.name, 99 kind=INFERENCE_ENDPOINT_CONFIG_KIND, 100 namespace=self.metadata.namespace, 101 ) 102 103 self.status = InferenceEndpointConfigStatus.model_validate( 104 response["status"], by_name=True 105 ) 106 107 return self 108 109 @classmethod 110 @_hyperpod_telemetry_emitter(Feature.HYPERPOD, "list_endpoints") 111 def list( 112 cls, 113 namespace: str = None, 114 ) -> List[Endpoint]: 115 if not namespace: 116 namespace = get_default_namespace() 117 118 response = cls.call_list_api( 119 kind=INFERENCE_ENDPOINT_CONFIG_KIND, 120 namespace=namespace, 121 ) 122 123 endpoints = [] 124 125 if response and response["items"]: 126 for item in response["items"]: 127 name = item["metadata"]["name"] 128 endpoints.append(cls.get(name, namespace=namespace)) 129 130 return endpoints 131 132 @classmethod 133 @_hyperpod_telemetry_emitter(Feature.HYPERPOD, "get_endpoint") 134 def get(cls, name: str, namespace: str = None) -> Endpoint: 135 if not namespace: 136 namespace = get_default_namespace() 137 138 response = cls.call_get_api( 139 name=name, 140 kind=INFERENCE_ENDPOINT_CONFIG_KIND, 141 namespace=namespace, 142 ) 143 144 endpoint = HPEndpoint.model_validate(response["spec"], by_name=True) 145 status = response.get("status") 146 if status is not None: 147 try: 148 endpoint.status = InferenceEndpointConfigStatus.model_validate( 149 status, by_name=True 150 ) 151 except ValidationError: 152 endpoint.status = None 153 else: 154 endpoint.status = None 155 endpoint.metadata = Metadata.model_validate(response["metadata"], by_name=True) 156 157 return endpoint 158 159 @_hyperpod_telemetry_emitter(Feature.HYPERPOD, "delete_endpoint") 160 def delete(self) -> None: 161 logger = self.get_logger() 162 logger = setup_logging(logger) 163 164 self.call_delete_api( 165 name=self.metadata.name, 166 kind=INFERENCE_ENDPOINT_CONFIG_KIND, 167 namespace=self.metadata.namespace, 168 ) 169 logger.info(f"Deleting HPEndpoint: {self.metadata.name}...") 170 171 @_hyperpod_telemetry_emitter(Feature.HYPERPOD, "invoke_endpoint") 172 def invoke(self, body, content_type="application/json"): 173 if not self.endpointName: 174 raise Exception("SageMaker endpoint name not found in this object!") 175 176 endpoint = Endpoint.get(self.endpointName, region=get_current_region()) 177 178 return endpoint.invoke(body=body, content_type=content_type) 179 180 def validate_instance_type(self, instance_type: str): 181 logger = self.get_logger() 182 logger = setup_logging(logger) 183 184 cluster_instance_types = None 185 186 # verify supported instance types from HyperPod cluster 187 try: 188 cluster_instance_types = get_cluster_instance_types( 189 cluster=get_current_cluster(), 190 region=get_current_region(), 191 ) 192 except Exception as e: 193 logger.warning(f"Failed to get instance types from HyperPod cluster: {e}") 194 195 if cluster_instance_types and (instance_type not in cluster_instance_types): 196 raise Exception( 197 f"Current HyperPod cluster does not have instance type {instance_type}. Supported instance types are {cluster_instance_types}" 198 ) 199
[docs] 200 @classmethod 201 @_hyperpod_telemetry_emitter(Feature.HYPERPOD, "list_pods_endpoint") 202 def list_pods(cls, namespace=None, endpoint_name=None): 203 cls.verify_kube_config() 204 205 if not namespace: 206 namespace = get_default_namespace() 207 208 v1 = client.CoreV1Api() 209 list_pods_response = v1.list_namespaced_pod(namespace=namespace) 210 211 endpoints = set() 212 if endpoint_name: 213 endpoints.add(endpoint_name) 214 else: 215 list_response = cls.call_list_api( 216 kind=INFERENCE_ENDPOINT_CONFIG_KIND, 217 namespace=namespace, 218 ) 219 if list_response and list_response["items"]: 220 for item in list_response["items"]: 221 endpoints.add(item["metadata"]["name"]) 222 223 pods = [] 224 for item in list_pods_response.items: 225 app_name = item.metadata.labels.get("app", None) 226 if app_name in endpoints: 227 # list_namespaced_pod will return all pods in the namespace, so we need to filter 228 # out the pods that are created by custom endpoint 229 pods.append(item.metadata.name) 230 231 return pods