pdfchecking / app.py
felixkky's picture
Update app.py
982b503 verified
Raw
History Blame Contribute Delete
6.87 kB
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()