Add image inference, clean up categories, and update readme
This commit is contained in:
+39
-12
@@ -16,7 +16,7 @@ class TextToImage:
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "inference"
|
||||
CATEGORY = "HF_Inference/Image/TextToImage"
|
||||
CATEGORY = "HF_Inference/Image"
|
||||
TITLE = "HF Image TextToImage"
|
||||
|
||||
def inference(self, endpoint, prompt):
|
||||
@@ -42,13 +42,13 @@ class Classification:
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "inference"
|
||||
CATEGORY = "HF_Inference/Image/Classification"
|
||||
CATEGORY = "HF_Inference/Image"
|
||||
TITLE = "HF Image Classification"
|
||||
|
||||
def inference(self, endpoint, image):
|
||||
response = post(endpoint, data=image)
|
||||
result = response.json()
|
||||
return {"ui": {"text": result}}
|
||||
return result
|
||||
|
||||
class ObjectDetection:
|
||||
@classmethod
|
||||
@@ -62,33 +62,60 @@ class ObjectDetection:
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "inference"
|
||||
CATEGORY = "HF_Inference/Image/ObjectDetection"
|
||||
CATEGORY = "HF_Inference/Image"
|
||||
TITLE = "HF Image Object Detection"
|
||||
|
||||
def inference(self, endpoint, image):
|
||||
response = post(endpoint, data=image)
|
||||
result = response.json()
|
||||
return {"ui": {"text": result}}
|
||||
|
||||
from base64 import b64decode
|
||||
|
||||
class Segmentation:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"endpoint": ("STRING", {}),
|
||||
"image": ("IMAGE",),
|
||||
"images": ("IMAGE",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_TYPES = ()
|
||||
FUNCTION = "inference"
|
||||
CATEGORY = "HF_Inference/Image/Segmentation"
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "HF_Inference/Image"
|
||||
TITLE = "HF Image Segmentation"
|
||||
|
||||
def inference(self, endpoint, image):
|
||||
response = post(endpoint, data=image)
|
||||
result = response.json()
|
||||
return {"ui": {"text": result}}
|
||||
def inference(self, endpoint, images):
|
||||
for image in images:
|
||||
image_bytes = BytesIO()
|
||||
pil_img = Image.fromarray(
|
||||
np.clip(
|
||||
image.cpu().numpy() * 255.0,
|
||||
0,
|
||||
255,
|
||||
).astype(np.uint8)
|
||||
)
|
||||
pil_img.save(image_bytes, format='png')
|
||||
|
||||
image_bytes.seek(0)
|
||||
image_bytes = image_bytes.read()
|
||||
print(len(image_bytes))
|
||||
|
||||
response = post(endpoint, data=image_bytes)
|
||||
result = response.json()
|
||||
|
||||
for item in result:
|
||||
label = item['label']
|
||||
mask_data = item['mask']
|
||||
|
||||
mask_img = Image.open(BytesIO(b64decode(mask_data)))
|
||||
|
||||
pil_img.paste(mask_img, (0, 0), mask_img)
|
||||
pil_img = np.array(pil_img).astype(np.float32) / 255.0
|
||||
pil_img = torch.from_numpy(pil_img)[None,]
|
||||
return {"ui": {"images": [pil_img, ]}}
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"Classification": Classification,
|
||||
|
||||
@@ -20,17 +20,29 @@ Export HF_AUTH_TOKEN with one of your [Hugging Face tokens](https://huggingface.
|
||||
### Run ComfyUI
|
||||
`HF_AUTH_TOKEN=hf_1111111111111111111111111111111111 python main.py`
|
||||
|
||||
## Usage
|
||||
## Nodes
|
||||
|
||||
> [!WARNING]
|
||||
> Inference API (serverless) requires a model 10GB or below and fails for random reasons on different models.
|
||||
|
||||
### Text
|
||||
|
||||
* Feature Extraction (ie: T5 encoder embeddings, BERT, etc.)
|
||||
* Question Answering (ie: roberta-base for QA)
|
||||
* Translation (ie: T5 Small)
|
||||
* Generation (ie: zephyr-7b-beta)
|
||||
* Feature Extraction
|
||||
- [facebook/bart-base](https://huggingface.co/facebook/bart-base)
|
||||
* Question Answering
|
||||
- [deepset/roberta-base-squad2](https://huggingface.co/deepset/roberta-base-squad2)
|
||||
* Translation
|
||||
- [google-t5/t5-base](https://huggingface.co/google-t5/t5-base)
|
||||
* Generation
|
||||
- [HuggingFaceH4/zephyr-7b-beta](https://huggingface.co/HuggingFaceH4/zephyr-7b-beta)
|
||||
|
||||
### Image
|
||||
|
||||
* Classification (ie: vit-base-patch16-224)
|
||||
* Object Detection (ie: detr-resnet-50)
|
||||
* Segmentation (ie: detr-resnet-50)
|
||||
* Classification
|
||||
- [google/vit-base-patch16-224](https://huggingface.co/google/vit-base-patch16-224)
|
||||
* Object Detection
|
||||
- [facebook/detr-resnet-50](https://huggingface.co/facebook/detr-resnet-50)
|
||||
* Segmentation
|
||||
- [facebook/detr-resnet-50-panoptic](https://huggingface.co/facebook/detr-resnet-50-panoptic)
|
||||
* TextToImage
|
||||
- [sd-community/sdxl-flash](https://huggingface.co/sd-community/sdxl-flash)
|
||||
+5
-5
@@ -13,7 +13,7 @@ class Generation:
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "inference"
|
||||
CATEGORY = "HF_Inference/Text/Generation"
|
||||
CATEGORY = "HF_Inference/Text"
|
||||
TITLE = "HF Text Generation"
|
||||
|
||||
def inference(self, endpoint, text):
|
||||
@@ -23,7 +23,7 @@ class Generation:
|
||||
response = post(endpoint, json=json)
|
||||
result = response.json()
|
||||
generated = ''.join(x['generated_text'] for x in result)
|
||||
return generated
|
||||
return generated
|
||||
|
||||
class Translation:
|
||||
@classmethod
|
||||
@@ -37,7 +37,7 @@ class Translation:
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "inference"
|
||||
CATEGORY = "HF_Inference/Text/Translation"
|
||||
CATEGORY = "HF_Inference/Text"
|
||||
TITLE = "HF Text Translation"
|
||||
|
||||
def inference(self, endpoint, text):
|
||||
@@ -62,7 +62,7 @@ class QuestionAnswering:
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
FUNCTION = "inference"
|
||||
CATEGORY = "HF_Inference/Text/QuestionAnswering"
|
||||
CATEGORY = "HF_Inference/Text"
|
||||
TITLE = "HF Text Question Answering"
|
||||
|
||||
def inference(self, endpoint, question, context):
|
||||
@@ -89,7 +89,7 @@ class FeatureExtraction:
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
FUNCTION = "inference"
|
||||
CATEGORY = "HF_Inference/Text/FeatureExtraction"
|
||||
CATEGORY = "HF_Inference/Text"
|
||||
TITLE = "HF Text Feature Extraction"
|
||||
|
||||
def inference(self, endpoint, text):
|
||||
|
||||
Reference in New Issue
Block a user