This commit is contained in:
kijai
2024-06-26 11:26:25 +03:00
2 changed files with 57 additions and 13 deletions
+23 -7
View File
@@ -5,21 +5,37 @@ Florence-2 can interpret simple text prompts to perform tasks like captioning, o
It leverages our FLD-5B dataset, containing 5.4 billion annotations across 126 million images, to master multi-task learning.
The model's sequence-to-sequence architecture enables it to excel in both zero-shot and fine-tuned settings, proving to be a competitive vision foundation model.
## New Feature: Document Visual Question Answering (DocVQA)
This fork includes support for Document Visual Question Answering (DocVQA) using the Florence2 model. DocVQA allows you to ask questions about the content of document images, and the model will provide answers based on the visual and textual information in the document. This feature is particularly useful for extracting information from scanned documents, forms, receipts, and other text-heavy images.
## Installation:
- clone this repository to 'ComfyUI/custom_nodes` -folder.
Only real dependency is new enough transformers version.
- Clone this repository to 'ComfyUI/custom_nodes` folder.
- The main dependency is a new enough transformers version.
![image](https://github.com/kijai/ComfyUI-Florence2/assets/40791699/4d537ac7-5490-470f-92f5-3007da7b9cc7)
![image](https://github.com/kijai/ComfyUI-Florence2/assets/40791699/512357b7-39ee-43ee-bb63-7347b0a8d07d)
Supports the following models, they are automatically downloaded to `ComfyUI/LLM`:
Supports the following models, which are automatically downloaded to `ComfyUI/LLM`:
https://huggingface.co/microsoft/Florence-2-base
https://huggingface.co/microsoft/Florence-2-base-ft
https://huggingface.co/microsoft/Florence-2-large
https://huggingface.co/microsoft/Florence-2-large-ft
https://huggingface.co/HuggingFaceM4/Florence-2-DocVQA
## Using DocVQA
To use the DocVQA feature:
1. Load a document image into ComfyUI.
2. Connect the image to the Florence2 DocVQA node.
3. Input your question about the document.
4. The node will output the answer based on the document's content.
Example questions:
- "What is the total amount on this receipt?"
- "What is the date mentioned in this form?"
- "Who is the sender of this letter?"
Note: The accuracy of answers depends on the quality of the input image and the complexity of the question.
+34 -6
View File
@@ -42,6 +42,7 @@ class DownloadAndLoadFlorence2Model:
'microsoft/Florence-2-base-ft',
'microsoft/Florence-2-large',
'microsoft/Florence-2-large-ft',
'HuggingFaceM4/Florence-2-DocVQA'
],
{
"default": 'microsoft/Florence-2-base'
@@ -111,8 +112,8 @@ class Florence2Run:
'caption_to_phrase_grounding',
'referring_expression_segmentation',
'ocr',
'ocr_with_region'
'ocr_with_region',
'docvqa'
],
),
"fill_mask": ("BOOLEAN", {"default": True}),
@@ -154,12 +155,13 @@ class Florence2Run:
'caption_to_phrase_grounding': '<CAPTION_TO_PHRASE_GROUNDING>',
'referring_expression_segmentation': '<REFERRING_EXPRESSION_SEGMENTATION>',
'ocr': '<OCR>',
'ocr_with_region': '<OCR_WITH_REGION>'
'ocr_with_region': '<OCR_WITH_REGION>',
'docvqa': '<DocVQA>'
}
task_prompt = prompts.get(task, '<OD>')
if (task!= 'referring_expression_segmentation' and task!= 'caption_to_phrase_grounding') and text_input:
raise ValueError("Text input (prompt) is only supported for 'referring_expression_segmentation' and 'caption_to_phrase_grounding'")
if (task not in ['referring_expression_segmentation', 'caption_to_phrase_grounding', 'docvqa']) and text_input:
raise ValueError("Text input (prompt) is only supported for 'referring_expression_segmentation', 'caption_to_phrase_grounding', and 'docvqa'")
if text_input != "":
prompt = task_prompt + " " + text_input
@@ -361,6 +363,32 @@ class Florence2Run:
image_tensor = image_tensor[:3, :, :].unsqueeze(0).permute(0, 2, 3, 1).cpu().float()
out.append(image_tensor)
elif task == 'docvqa':
if text_input == "":
raise ValueError("Text input (prompt) is required for 'docvqa'")
prompt = "<DocVQA> " + text_input
inputs = processor(text=prompt, images=image_pil, return_tensors="pt", do_rescale=False).to(dtype).to(device)
generated_ids = model.generate(
input_ids=inputs["input_ids"],
pixel_values=inputs["pixel_values"],
max_new_tokens=max_new_tokens,
do_sample=do_sample,
num_beams=num_beams,
)
results = processor.batch_decode(generated_ids, skip_special_tokens=False)[0]
clean_results = results.replace('</s>', '').replace('<s>', '')
if len(image) == 1:
out_results = clean_results
else:
out_results.append(clean_results)
out.append(F.to_tensor(image_pil).unsqueeze(0).permute(0, 2, 3, 1).cpu().float())
pbar.update(1)
if len(out) > 0:
out_tensor = torch.cat(out, dim=0)
else:
@@ -384,4 +412,4 @@ NODE_CLASS_MAPPINGS = {
NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadFlorence2Model": "DownloadAndLoadFlorence2Model",
"Florence2Run": "Florence2Run",
}
}