diff --git a/Text/nodes.py b/Text/nodes.py new file mode 100644 index 0000000..5c1f2f4 --- /dev/null +++ b/Text/nodes.py @@ -0,0 +1,90 @@ +import torch +from ..session import session + +class Translation: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "endpoint": ("STRING", {}), + "text": ("STRING", {"multiline": True}), + } + } + + RETURN_TYPES = ("STRING",) + FUNCTION = "inference" + CATEGORY = "HF_Inference/Text/Translation" + TITLE = "HF Text Translation" + + def inference(self, endpoint, text): + json = { + 'inputs': text, + } + response = session.post(endpoint, json=json) + if response.status_code != 200: + raise Exception(response.text) + result = response.json() + translation = ''.join(x['translation_text'] for x in result) + return translation + +class QuestionAnswering: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "endpoint": ("STRING", {}), + "question": ("STRING", {"multiline": True}), + "context": ("STRING", {"multiline": True}), + } + } + + RETURN_TYPES = ("STRING",) + FUNCTION = "inference" + CATEGORY = "HF_Inference/Text/QuestionAnswering" + TITLE = "HF Text Question Answering" + + def inference(self, endpoint, question, context): + json = { + 'inputs': { + 'question': question, + 'context': context, + }, + } + response = session.post(endpoint, json=json) + if response.status_code != 200: + raise Exception(response.text) + result = response.json() + answer = result['answer'] + return answer + +class FeatureExtraction: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "endpoint": ("STRING", {}), + "text": ("STRING", {"multiline": True}), + } + } + + RETURN_TYPES = ("CONDITIONING",) + FUNCTION = "inference" + CATEGORY = "HF_Inference/Text/FeatureExtraction" + TITLE = "HF Text Feature Extraction" + + def inference(self, endpoint, text): + json = { + 'inputs': text, + } + response = session.post(endpoint, json=json) + if response.status_code != 200: + raise Exception(response.text) + result = response.json() + cond = torch.tensor(result, dtype=torch.float16).to('cuda') + return ([[cond, {}]],) + +NODE_CLASS_MAPPINGS = { + "TextFeatureExtraction": FeatureExtraction, + "QuestionAnswering": QuestionAnswering, + "Translation": Translation, +} \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..422976b --- /dev/null +++ b/__init__.py @@ -0,0 +1,14 @@ +try: + import comfy.utils +except ImportError: + pass +else: + NODE_CLASS_MAPPINGS = {} + + from .Text.nodes import NODE_CLASS_MAPPINGS as Text_Nodes + NODE_CLASS_MAPPINGS.update(Text_Nodes) + + NODE_DISPLAY_NAME_MAPPINGS = {k:v.TITLE for k,v in NODE_CLASS_MAPPINGS.items()} + print(NODE_CLASS_MAPPINGS) + print(NODE_DISPLAY_NAME_MAPPINGS) + __all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/session.py b/session.py new file mode 100644 index 0000000..cdb52fa --- /dev/null +++ b/session.py @@ -0,0 +1,10 @@ +"""Session handler""" +import os +import requests +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