#!/usr/bin/env python """ 知识图谱查询脚本 示例用法:python scripts/query_kg.py --kg output/kg.json --query "城市更新" """ import argparse import json import logging import sys from pathlib import Path # 添加src到路径 sys.path.insert(0, str(Path(__file__).parent.parent)) from src.kg_builder.graph import KnowledgeGraph from src.query.global_search import GlobalSearcher from src.query.local_search import LocalSearcher from src.utils.config import load_config logging.basicConfig( level=logging.INFO, format='%(asctime)s - %(name)s - %(levelname)s - %(message)s' ) logger = logging.getLogger(__name__) def main(): parser = argparse.ArgumentParser(description="查询知识图谱") parser.add_argument( "--kg", type=str, required=True, help="知识图谱文件路径(JSON格式)" ) parser.add_argument( "--query", type=str, required=True, help="查询文本" ) parser.add_argument( "--mode", type=str, choices=["global", "local", "drift"], default="global", help="查询模式" ) parser.add_argument( "--output", type=str, default=None, help="输出结果文件路径" ) args = parser.parse_args() # 加载知识图谱 logger.info(f"加载知识图谱: {args.kg}") kg = KnowledgeGraph() kg.load(str(args.kg), format="json") # 获取社区信息(从元数据) communities = {} if hasattr(kg, 'graph') and kg.graph: metadata = kg.graph.graph.get("metadata", {}) communities = metadata.get("summaries", {}) # 执行查询 logger.info(f"执行{args.mode}查询: {args.query}") kg_graph = kg.graph if hasattr(kg, 'graph') and kg.graph else None if args.mode == "global": searcher = GlobalSearcher(kg_graph, communities) result = searcher.search(args.query) elif args.mode == "local": searcher = LocalSearcher(kg_graph) # 尝试从查询中提取实体 entity = args.query # 简化处理,可以将查询作为实体 result = searcher.search(entity, depth=2) else: # DRIFT search from src.query.drift_search import DriftSearcher searcher = DriftSearcher(kg_graph, communities) result = searcher.search(args.query) # 输出结果 print("\n查询结果:") print(json.dumps(result, ensure_ascii=False, indent=2)) if args.output: with open(args.output, "w", encoding="utf-8") as f: json.dump(result, f, ensure_ascii=False, indent=2) logger.info(f"结果已保存到: {args.output}") if __name__ == "__main__": main()