"""测试 Heron 版面检测 + 裁剪 + VLM 图片描述流程""" import argparse import asyncio import logging import sys import time from pathlib import Path sys.path.append(str(Path(__file__).parent.parent)) from PIL import Image from src.rag.document_loaders import PDFImageExtractor logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s") logger = logging.getLogger(__name__) PROJECT_ROOT = Path(__file__).parent.parent.parent # course-agent-od/ TEST_PDF = ( PROJECT_ROOT / "data" / "uploads" / "testuser" / "knowledge_bases" / "植物知识图谱" / "ab810d4a-b030-45c0-94d2-bd6e2a23053b.pdf" ) OUTPUT_DIR = PROJECT_ROOT / "data" / "images" / "_test_pdf_extract" def test_extract(max_pages: int | None = None): """阶段1:Heron 版面检测 + 裁剪""" if not TEST_PDF.exists(): logger.error(f"测试PDF不存在: {TEST_PDF}") sys.exit(1) # 清理旧输出 if OUTPUT_DIR.exists(): for f in OUTPUT_DIR.iterdir(): f.unlink() print(f"\n{'='*60}") print(f"[阶段1] Heron 版面检测 + 裁剪测试") print(f" PDF: {TEST_PDF.name}") print(f" 输出: {OUTPUT_DIR}") print(f"{'='*60}\n") t0 = time.time() images = PDFImageExtractor.extract_images( str(TEST_PDF), str(OUTPUT_DIR), ) elapsed = time.time() - t0 if not images: print(" ❌ 未提取到任何图片") return [] if max_pages: images = [img for img in images if img["page"] <= max_pages] # 按页分组统计 from collections import defaultdict by_page = defaultdict(list) for img in images: by_page[img["page"]].append(img) total_size = 0 fig_count = 0 full_count = 0 for page_num in sorted(by_page.keys()): page_imgs = by_page[page_num] for img in page_imgs: path = Path(img["path"]) with Image.open(path) as pil: w, h = pil.size size_kb = img["size"] / 1024 total_size += img["size"] is_full = "_full." in img["filename"] tag = "FULL" if is_full else "FIG" if is_full: full_count += 1 else: fig_count += 1 print( f" page {img['page']:>3d} [{tag}] {img['filename']:<25s} " f"{w}x{h} {size_kb:>7.1f} KB" ) print( f"\n 汇总: {len(images)} 张 ({fig_count} 裁剪 + {full_count} 整页兜底), " f"{total_size/1024/1024:.1f} MB, 耗时 {elapsed:.1f}s\n" ) return images async def test_vlm(images: list[dict], max_images: int = 3): """阶段2:VLM图片描述""" from src.llm.siliconflow import get_llm_client # 优先选裁剪图 fig_images = [img for img in images if "_full." not in img["filename"]] targets = (fig_images or images)[:max_images] print(f"\n{'='*60}") print(f"[阶段2] VLM图片描述测试 ({len(targets)} 张)") print(f"{'='*60}\n") client = get_llm_client() sem = asyncio.Semaphore(3) async def describe_one(idx: int, img: dict): async with sem: t0 = time.time() desc = await client.describe_image(img["path"], img.get("context_text", "")) elapsed = time.time() - t0 return idx, desc, elapsed tasks = [describe_one(i, img) for i, img in enumerate(targets)] results = await asyncio.gather(*tasks, return_exceptions=True) for r in results: if isinstance(r, Exception): print(f" ❌ 失败: {r}\n") continue idx, desc, elapsed = r img = targets[idx] print(f" --- page {img['page']} ({img['filename']}) {elapsed:.1f}s ---") print(f" {desc}\n") def main(): parser = argparse.ArgumentParser(description="测试PDF图片提取") parser.add_argument("--vlm", action="store_true", help="启用VLM图片描述") parser.add_argument("--max-pages", type=int, default=None, help="限制提取页数") parser.add_argument("--max-vlm", type=int, default=3, help="VLM描述最大图片数") args = parser.parse_args() images = test_extract(max_pages=args.max_pages) if args.vlm and images: asyncio.run(test_vlm(images, max_images=args.max_vlm)) if __name__ == "__main__": main()