DocTron/DocTron-Formula
1
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()