From 98edfb34a10e8c74c61b185362a4d613372f1cd8 Mon Sep 17 00:00:00 2001 From: bitaffinity <168344121+bitaffinity@users.noreply.github.com> Date: Sat, 1 Jun 2024 08:17:31 -0400 Subject: [PATCH] Refactor request code to reduce duplication --- Image/nodes.py | 18 +++++------------- Text/nodes.py | 18 +++++------------- session.py | 20 +++++++++++++++++++- 3 files changed, 29 insertions(+), 27 deletions(-) diff --git a/Image/nodes.py b/Image/nodes.py index 92afc63..2464a3c 100644 --- a/Image/nodes.py +++ b/Image/nodes.py @@ -1,4 +1,4 @@ -from ..session import session +from ..session import post from io import BytesIO from PIL import Image import numpy as np @@ -21,9 +21,7 @@ class TextToImage: def inference(self, endpoint, prompt): payload = {"inputs": prompt} - response = session.post(endpoint, json=payload) - if response.status_code != 200: - raise Exception(response.text) + response = post(endpoint, json=payload) result = BytesIO(response.content) image = Image.open(result) @@ -48,9 +46,7 @@ class Classification: TITLE = "HF Image Classification" def inference(self, endpoint, image): - response = session.post(endpoint, data=image) - if response.status_code != 200: - raise Exception(response.text) + response = post(endpoint, data=image) result = response.json() return {"ui": {"text": result}} @@ -70,9 +66,7 @@ class ObjectDetection: TITLE = "HF Image Object Detection" def inference(self, endpoint, image): - response = session.post(endpoint, data=image) - if response.status_code != 200: - raise Exception(response.text) + response = post(endpoint, data=image) result = response.json() return {"ui": {"text": result}} @@ -92,9 +86,7 @@ class Segmentation: TITLE = "HF Image Segmentation" def inference(self, endpoint, image): - response = session.post(endpoint, data=image) - if response.status_code != 200: - raise Exception(response.text) + response = post(endpoint, data=image) result = response.json() return {"ui": {"text": result}} diff --git a/Text/nodes.py b/Text/nodes.py index 0f7938c..5e755c6 100644 --- a/Text/nodes.py +++ b/Text/nodes.py @@ -1,5 +1,5 @@ import torch -from ..session import session +from ..session import post class Generation: @classmethod @@ -20,9 +20,7 @@ class Generation: json = { 'inputs': text, } - response = session.post(endpoint, json=json) - if response.status_code != 200: - raise Exception(response.text) + response = post(endpoint, json=json) result = response.json() generated = ''.join(x['generated_text'] for x in result) return generated @@ -46,9 +44,7 @@ class Translation: json = { 'inputs': text, } - response = session.post(endpoint, json=json) - if response.status_code != 200: - raise Exception(response.text) + response = post(endpoint, json=json) result = response.json() translation = ''.join(x['translation_text'] for x in result) return translation @@ -76,9 +72,7 @@ class QuestionAnswering: 'context': context, }, } - response = session.post(endpoint, json=json) - if response.status_code != 200: - raise Exception(response.text) + response = post(endpoint, json=json) result = response.json() answer = result['answer'] return answer @@ -102,9 +96,7 @@ class FeatureExtraction: json = { 'inputs': text, } - response = session.post(endpoint, json=json) - if response.status_code != 200: - raise Exception(response.text) + response = post(endpoint, json=json) result = response.json() cond = torch.tensor(result, dtype=torch.float16).to('cuda') return ([[cond, {}]],) diff --git a/session.py b/session.py index cdb52fa..0c8d748 100644 --- a/session.py +++ b/session.py @@ -1,10 +1,28 @@ """Session handler""" import os import requests +from time import sleep session = requests.Session() if "HF_AUTH_TOKEN" in os.environ: session.headers.update({ "Authorization": f"Bearer {os.environ['HF_AUTH_TOKEN']}", }) else: - print("No 'HF_AUTH_TOKEN' set.") \ No newline at end of file + print("No 'HF_AUTH_TOKEN' set.") + +def post(url, **kwargs): + + if not url.startswith('http'): + url = f'https://api-inference.huggingface.co/models/{url}' + + response = session.post(url, **kwargs) + + if response.status_code != 200: + if 'estimated_time' in response.text: + estimated_time = response.json()['estimated_time'] + model_path = '/'.join(url.split('/')[-2:]) + print('Waiting for ', estimated_time, ' to load ', model_path) + sleep(estimated_time) + return post(url, **kwargs) + raise Exception(response.text) + return response \ No newline at end of file