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