Spaces:
Sleeping
Sleeping
| import gradio as gr | |
| import fitz # PyMuPDF | |
| import json | |
| import numpy as np | |
| from PIL import Image, ImageDraw | |
| from mistralai.client import Mistral | |
| from paddleocr import PaddleOCR | |
| import paddle | |
| # 1. 強制關閉 Paddle 的 oneDNN CPU 加速引擎,徹底避開 Hugging Face 環境的底層 Bug | |
| paddle.set_flags({'FLAGS_use_onednn': False}) | |
| # 2. 初始化 PaddleOCR,拿掉方向分類器(考卷字體都是正的,不用它反而跑得更快、更穩!) | |
| ocr = PaddleOCR(lang="ch") | |
| def process_document(file, api_key, progress=gr.Progress()): | |
| if not api_key: | |
| return None, "錯誤:請提供 Mistral API 金鑰。" | |
| if not file: | |
| return None, "錯誤:請上傳檔案。" | |
| progress(0.05, desc="初始化 Mistral 客戶端...") | |
| client = Mistral(api_key=api_key) | |
| images = [] | |
| # --- 1. 處理檔案上傳 --- | |
| file_name = file.name.lower() | |
| try: | |
| progress(0.1, desc="正在讀取與轉換圖片...") | |
| if file_name.endswith('.pdf'): | |
| doc = fitz.open(file.name) | |
| for i in range(len(doc)): | |
| page = doc.load_page(i) | |
| pix = page.get_pixmap(dpi=150) | |
| img = Image.frombytes("RGB", [pix.width, pix.height], pix.samples) | |
| images.append(img) | |
| else: | |
| img = Image.open(file.name).convert("RGB") | |
| images.append(img) | |
| except Exception as e: | |
| return None, f"檔案讀取錯誤:{str(e)}" | |
| drawn_images = [] | |
| all_json_results = [] | |
| # --- 2. 準備給 Mistral 的語意判斷 Prompt --- | |
| prompt_template = """ | |
| 你是一個考卷批改助手。以下是一份透過 OCR 掃描出來的考卷內容資料,包含了「印刷體的考卷題目」以及「學生手寫的回答」。 | |
| 請透過閱讀上下文的邏輯,幫我過濾掉所有的題目,**只保留「學生的回答」**。 | |
| 請以嚴格的 JSON 格式輸出,回傳一個名為 "student_answer_ids" 的陣列,裡面包含你認為是「學生回答」的 ID。 | |
| 格式範例: | |
| { | |
| "student_answer_ids": [3, 5, 8, 12] | |
| } | |
| 以下是 OCR 提取的原始資料: | |
| """ | |
| # --- 3. 逐頁處理 --- | |
| for page_num, img in enumerate(images): | |
| try: | |
| # === Step A: 專門 OCR 提取所有文字與座標 === | |
| progress(0.3, desc=f"第 {page_num + 1} 頁 [Step 1/2]:PaddleOCR 精準掃描中...") | |
| # 將 PIL Image 轉為 numpy array 供 PaddleOCR 使用 | |
| img_np = np.array(img) | |
| ocr_results = ocr.ocr(img_np) | |
| # 整理 OCR 結果,賦予每個文字框一個 ID | |
| extracted_data = [] | |
| if ocr_results and ocr_results[0]: | |
| for idx, line in enumerate(ocr_results[0]): | |
| box = line[0] # 座標 [[x1,y1], [x2,y2], [x3,y3], [x4,y4]] | |
| text = line[1][0] # 文字內容 | |
| extracted_data.append({ | |
| "id": idx, | |
| "text": text, | |
| "box": box | |
| }) | |
| if not extracted_data: | |
| all_json_results.append({f"Page {page_num + 1}": "未偵測到任何文字"}) | |
| drawn_images.append(img) | |
| continue | |
| # 建立一份只包含 id 和 text 的輕量資料給 LLM(不給座標以節省 Token) | |
| llm_input_data = [{"id": item["id"], "text": item["text"]} for item in extracted_data] | |
| # === Step B: Mistral Large 語意過濾 === | |
| progress(0.7, desc=f"第 {page_num + 1} 頁 [Step 2/2]:Mistral 語意分析過濾中...") | |
| response = client.chat.complete( | |
| model="mistral-large-latest", | |
| messages=[ | |
| { | |
| "role": "user", | |
| "content": prompt_template + json.dumps(llm_input_data, ensure_ascii=False) | |
| } | |
| ], | |
| response_format={"type": "json_object"} | |
| ) | |
| llm_result_text = response.choices[0].message.content | |
| # 清理潛在的 Markdown 標記 | |
| clean_content = llm_result_text.replace("```json", "").replace("```", "").strip() | |
| answer_data = json.loads(clean_content) | |
| student_ids = answer_data.get("student_answer_ids", []) | |
| # === Step C: 整合資料並畫框 === | |
| progress(0.9, desc=f"第 {page_num + 1} 頁:正在繪製學生答案的邊界框...") | |
| draw = ImageDraw.Draw(img) | |
| final_page_result = {"answers": []} | |
| for item in extracted_data: | |
| if item["id"] in student_ids: | |
| # 記錄到最終結果 | |
| final_page_result["answers"].append({ | |
| "text": item["text"], | |
| "box": item["box"] | |
| }) | |
| # 畫多邊形紅框 (PaddleOCR 輸出的是四個角的確切座標) | |
| polygon = [tuple(point) for point in item["box"]] | |
| draw.polygon(polygon, outline="red", width=3) | |
| all_json_results.append({f"Page {page_num + 1}": final_page_result}) | |
| drawn_images.append(img) | |
| except Exception as e: | |
| error_msg = f"處理發生錯誤: {str(e)}" | |
| print(error_msg) | |
| all_json_results.append({f"Page {page_num + 1} Error": error_msg}) | |
| drawn_images.append(img) | |
| progress(1.0, desc="處理完成!") | |
| return drawn_images, json.dumps(all_json_results, ensure_ascii=False, indent=2) | |
| # --- Gradio 介面設計 --- | |
| with gr.Blocks(title="考卷手寫答案偵測 (OCR + LLM)") as demo: | |
| gr.Markdown("## 📝 考卷手寫答案偵測 (PaddleOCR + Mistral-large)") | |
| gr.Markdown("**業界標準架構:** 底層由 `PaddleOCR` 負責精準抓字與計算座標,高層由 `Mistral-large` 判讀上下文語意,自動過濾掉題目,只框出學生的回答!") | |
| with gr.Row(): | |
| with gr.Column(): | |
| api_key_input = gr.Textbox(label="Mistral API Key", type="password") | |
| file_input = gr.File(label="上傳 PDF 或圖片", file_types=[".pdf", ".jpg", ".png", ".jpeg"]) | |
| submit_btn = gr.Button("開始分析與畫框", variant="primary") | |
| with gr.Column(): | |
| output_gallery = gr.Gallery(label="畫框結果預覽", columns=1, height="auto") | |
| output_json = gr.JSON(label="偵測結果與完美座標 (JSON)") | |
| submit_btn.click( | |
| fn=process_document, | |
| inputs=[file_input, api_key_input], | |
| outputs=[output_gallery, output_json] | |
| ) | |
| if __name__ == "__main__": | |
| demo.launch() | |