meg/backend
1
1import sys2from time import sleep3 4import requests5from huggingface_hub import create_inference_endpoint, get_inference_endpoint6 7from src.backend.compute_memory_requirements import get_instance_needs8from src.backend.run_toxicity_eval import get_generation9from src.envs import TOKEN10from src.logging import setup_logger11 12logger = setup_logger(__name__)13TIMEOUT = 2014MAX_REPLICA = 115 16 17def create_endpoint(endpoint_name, repository, framework='pytorch',18 task='text-generation', accelerator='gpu', vendor='aws',19 region='us-east-1', type='protected'):20 """Tries to automagically create a running endpoint for the given model."""21 logger.info("Creating endpoint %s..." % endpoint_name)22 endpoint = None23 instance_size, instance_type = get_instance_needs(repository, TOKEN)24 logger.info("Estimating the following instance size and type: %s, %s" % (25 instance_size, instance_type))26 # Useful in debugging, when models are being run over and over:27 # Check if the endpoint is already there.28 try:29 endpoint = get_inference_endpoint(endpoint_name)30 have_endpoint = True31 except requests.exceptions.HTTPError:32 have_endpoint = False33 if not have_endpoint:34 endpoint = create_inference_endpoint(endpoint_name,35 repository=repository,36 framework=framework, task=task,37 accelerator=accelerator,38 vendor=vendor, region=region,39 type=type,40 instance_size=instance_size,41 instance_type=instance_type,42 max_replica=MAX_REPLICA)43 logger.info("Endpoint status: %s." % endpoint.status)44 if endpoint.status == 'scaledToZero':45 # Send a request to wake it up.46 get_generation(endpoint.url, "Wake up")47 sleep(TIMEOUT)48 # Applies in ['updating', 'pending', 'initializing']49 wait_for_endpoint(endpoint)50 if endpoint.status == 'failed':51 logger.info("Endpoint failed, attempting to change compute.")52 endpoint = update_endpoint_exception(endpoint)53 # Applies in ['updating', 'pending', 'initializing']54 wait_for_endpoint(endpoint)55 logger.info("Endpoint created:")56 logger.info(endpoint)57 generation_url = endpoint.url58 if generation_url is None:59 logger.error("Failed to create an endpoint. Exiting.")60 sys.exit()61 return generation_url62 63 64def wait_for_endpoint(endpoint):65 # TODO: HANDLE 'paused'66 i = 067 while endpoint.status in ['updating', 'pending',68 'initializing']: # not in ['failed', 'running', 'scaledToZero']69 if i >= 20:70 logger.error("Model failed to respond after 20 tries. Exiting.")71 sys.exit()72 logger.info(73 "Waiting %d seconds to check again if the endpoint is running." %74 TIMEOUT)75 sleep(TIMEOUT)76 endpoint.fetch()77 logger.info("Endpoint status: %s." % (endpoint.status))78 i += 179 80 81def update_endpoint_exception(endpoint):82 """83 Endpoints can fail from too little memory, as well as for missing84 flash attention, etc. This function tries new compute setups,85 scaling up the compute power until it's running or expensive.86 """87 raw_info = endpoint.raw88 cur_instance_size = raw_info['compute']['instanceSize']89 cur_instance_type = raw_info['compute']['instanceType']90 91 if (cur_instance_type, cur_instance_size) == ('nvidia-a10g', 'x1'):92 endpoint.update(instance_size='x4', instance_type='nvidia-t4',93 max_replica=MAX_REPLICA)94 elif (cur_instance_type, cur_instance_size) == ('nvidia-t4', 'x4'):95 endpoint.update(instance_size='x1', instance_type='nvidia-a100',96 max_replica=MAX_REPLICA)97 elif (cur_instance_type, cur_instance_size) == ('nvidia-a100', 'x1'):98 endpoint.update(instance_size='x4', instance_type='nvidia-a10g',99 max_replica=MAX_REPLICA)100 elif (cur_instance_type, cur_instance_size) == ('nvidia-l4', 'x4'):101 endpoint.update(instance_size='x2', instance_type='nvidia-a100',102 max_replica=MAX_REPLICA)103 else:104 logger.error(105 "Getting expensive to run this model without human oversight."106 " Exiting.")107 sys.exit()108 return endpoint109 110 111if __name__ == '__main__':112 generation_url = create_endpoint('this-is-a-test',113 'Qwen/Qwen2-7B')114 