CoolFace
Apppublic

DocTron/DocTron-Formula

sourceHugging Faceupdated 1y agoView on Hugging Face
1likes
app.py221 linesDownload Raw Back to root
1import argparse2import json3import os4 5import torch6from PIL import Image7from qwen_vl_utils import process_vision_info8from transformers import AutoProcessor, Qwen2_5_VLForConditionalGeneration9import gradio as gr10 11user_prompt = "Analyze the image. Extract and output only the LaTeX formulas present in the image, in LaTeX code format. Ignore inline formulas, all other text, and do not include any explanations."12 13 14def read_input_file(input_file):15    with open(input_file, 'r') as file:16        data = json.load(file)17 18    image_path = data[0]['images'][0]19    gt_latex_code = data[0]['messages'][1]['content']20 21    return image_path, gt_latex_code22 23 24class ImageProcessor:25    def __init__(self, args):26        self.args = args27        self.model, self.vis_processor = self.load_model_and_processor()28        self.generate_kwargs = dict(29            max_new_tokens=2048,30            top_p=0.001,31            top_k=1,32            temperature=0.01,33            repetition_penalty=1.0,34        )35 36    def load_model_and_processor(self):37        # Load model38        checkpoint = self.args.ckpt39        vis_processor = AutoProcessor.from_pretrained(checkpoint)40 41        model = Qwen2_5_VLForConditionalGeneration.from_pretrained(checkpoint, torch_dtype="auto", device_map="auto")42        model.eval()43 44        return model, vis_processor45 46    def process_single_image(self, image_path):47        question = user_prompt48 49        try:50            image_local_path = "file://" + image_path51 52            messages = []53            messages.append(54                {"role": "user", "content": [55                        {"type": "image", "image": image_local_path, "min_pixels": 32 * 32, "max_pixels": 512 * 512},56                        {"type": "text", "text": question},57                    ]58                }59            )60 61            text = self.vis_processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)62            images, videos = process_vision_info([messages])63 64            inputs = self.vis_processor(text=text, images=images, videos=videos, padding=True, return_tensors='pt')65            inputs = inputs.to(self.model.device)66 67            with torch.no_grad():68                generated_ids = self.model.generate(69                    **inputs,70                    **self.generate_kwargs,71                )72            generated_ids = [73                output_ids[len(input_ids):] for input_ids, output_ids in zip(inputs.input_ids, generated_ids)74            ]75            out = self.vis_processor.tokenizer.batch_decode(76                generated_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False77            )78            model_answer = out[0]79        except Exception as e:80            print(e, flush=True)81            model_answer = "None"82 83        return model_answer84 85def save_image_with_auto_naming(image, save_dir="./tmp"):    86    # 确保目录存在87    os.makedirs(save_dir, exist_ok=True)88    89    # 获取目录中现有的文件名90    existing_files = [f for f in os.listdir(save_dir) if f.endswith('.png') and f.split('.')[0].isdigit()]91    92    # 找到最大的数字93    next_num = 094    if existing_files:95        next_num = max([int(f.split('.')[0]) for f in existing_files]) + 196    97    # 生成新文件名98    temp_path = os.path.join(save_dir, f"{next_num}.png")99    100    # 保存图片101    image.save(temp_path)102    103    return temp_path104 105# {{ edit_1 }}106def process_image_for_gradio(image):107    """处理上传的图片并返回LaTeX结果"""108    if image is None:109        return ""110 111    # 保存上传的图片到指定目录,并自动命名112    temp_path = save_image_with_auto_naming(image)113 114    # 处理图片115    pred_latex_code = processor.process_single_image(temp_path)116 117    # 清理临时文件118    if os.path.exists(temp_path):119        os.remove(temp_path)120 121    return pred_latex_code122 123def load_example(example_name):124    """加载示例图片"""125    input_file = os.path.join('./asset/test_jsons', f"{example_name}.json")126    image_path, gt_latex_code = read_input_file(input_file)127    return Image.open(image_path), example_name128 129# {{ edit_2 }}130def create_gradio_interface(processor):131    """创建Gradio界面"""132    with gr.Blocks(title="DocTron-Formula") as demo:133        gr.Markdown("# DocTron-Formula LaTeX公式识别")134        gr.Markdown("上传图片或选择示例来识别LaTeX公式")135 136        with gr.Row():137            with gr.Column():138                # 左侧列139                image_input = gr.Image(type="pil", label="上传图片")140 141                with gr.Row():142                    clear_btn = gr.Button("Clear")143                    submit_btn = gr.Button("Submit", variant="primary")144 145                gr.Markdown("### 示例图片")146                with gr.Row():147                    line_btn = gr.Button("Line-level")148                    paragraph_btn = gr.Button("Paragraph-level")149                    page_btn = gr.Button("Page-level")150 151                # 存储示例名称152                example_name = gr.State()153 154            with gr.Column():155                # 右侧列 - 显示结果156                latex_output = gr.Textbox(label="预测的LaTeX公式", lines=10, interactive=False)157 158        # 按钮事件绑定159        submit_btn.click(160            fn=process_image_for_gradio,161            inputs=[image_input],162            outputs=[latex_output]163        )164 165        clear_btn.click(166            fn=lambda: (None, ""),167            inputs=[],168            outputs=[image_input, latex_output]169        )170 171        # 示例按钮事件172        line_btn.click(173            fn=load_example,174            inputs=gr.Textbox(value="line-level", visible=False),175            outputs=[image_input, example_name]176        ).then(177            fn=lambda img: process_image_for_gradio(img),178            inputs=[image_input],179            outputs=[latex_output]180        )181 182        paragraph_btn.click(183            fn=load_example,184            inputs=gr.Textbox(value="paragraph-level", visible=False),185            outputs=[image_input, example_name]186        ).then(187            fn=lambda img: process_image_for_gradio(img),188            inputs=[image_input],189            outputs=[latex_output]190        )191 192        page_btn.click(193            fn=load_example,194            inputs=gr.Textbox(value="page-level", visible=False),195            outputs=[image_input, example_name]196        ).then(197            fn=lambda img: process_image_for_gradio(img),198            inputs=[image_input],199            outputs=[latex_output]200        )201 202    return demo203 204if __name__ == "__main__":205    parser = argparse.ArgumentParser()206    parser.add_argument("--ckpt", type=str, default="DocTron/DocTron-Formula")207    parser.add_argument("--input_file", type=str, default="line-level")208    args = parser.parse_args()209 210    # Init model211    processor = ImageProcessor(args)212 213    # {{ edit_3 }}214    # 创建并启动Gradio界面215    demo = create_gradio_interface(processor)216    # demo.launch(217    #     server_name="10.238.36.208",218    #     server_port=8000,219    #     share=False220    # )221    demo.launch()