Add image inference, clean up categories, and update readme

This commit is contained in:
bitaffinity
2024-06-10 18:46:05 -04:00
parent 4cdc73c498
commit 4f35d84928
3 changed files with 64 additions and 25 deletions
+39 -12
View File
@@ -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 -8
View File
@@ -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
View File
@@ -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):