From 087ccaab8ace728c1c237bd78baeef4f37615498 Mon Sep 17 00:00:00 2001 From: damus88 <98550229+damus88@users.noreply.github.com> Date: Tue, 25 Jun 2024 12:25:15 -0500 Subject: [PATCH] added DocVQA model inference --- README.md | 30 +++++++++++++++++++++++------- nodes.py | 40 ++++++++++++++++++++++++++++++++++------ 2 files changed, 57 insertions(+), 13 deletions(-) diff --git a/README.md b/README.md index 52365a4..283cd7b 100644 --- a/README.md +++ b/README.md @@ -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. diff --git a/nodes.py b/nodes.py index fee7003..4fe9622 100644 --- a/nodes.py +++ b/nodes.py @@ -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': '', 'referring_expression_segmentation': '', 'ocr': '', - 'ocr_with_region': '' + 'ocr_with_region': '', + 'docvqa': '' } task_prompt = prompts.get(task, '') - 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 @@ -358,6 +360,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 = " " + 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('', '').replace('', '') + + 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: @@ -381,4 +409,4 @@ NODE_CLASS_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = { "DownloadAndLoadFlorence2Model": "DownloadAndLoadFlorence2Model", "Florence2Run": "Florence2Run", -} \ No newline at end of file +}