template_manage.py 22 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494
  1. # -*- coding: utf-8 -*-
  2. import re
  3. import os
  4. import importlib
  5. import inspect
  6. import traceback
  7. from lxml import etree
  8. from typing import Dict, List, Optional, Any
  9. from BiddingKG.dl.template_extract.abstract_template import AbstractTemplate
  10. from BiddingKG.dl.template_extract.out_line_extractor import OutlineExtractor
  11. from BiddingKG.dl.template_extract.table_extractor import get_html_table
  12. from BiddingKG.dl.common.Utils import log
  13. class TemplateManager:
  14. """
  15. 模板管理与调用类(新增站源属性绑定+按站源筛选调用)
  16. 核心特性:
  17. 1. 自动注册时从模板文件的父文件夹名称读取“站源”属性
  18. 2. 调用模板时支持传入站源参数,仅执行对应站源的模板
  19. 3. 保留原有优先级规则、结果融合逻辑
  20. """
  21. def __init__(self):
  22. # 存储已注册的模板:key=模板ID,value=AbstractTemplate子类实例
  23. self.registered_templates: Dict[str, AbstractTemplate] = {}
  24. # 缓存模板优先级映射(避免重复读取,提升筛选效率)
  25. self.template_priority_cache: Dict[str, int] = {}
  26. # 新增:缓存模板-站源映射(快速筛选站源对应的模板)
  27. self.template_site_map: Dict[str, List[str]] = {} # key=站源名,value=模板ID列表
  28. def register_template(self, template: AbstractTemplate, site_source: str) -> None:
  29. """
  30. 注册模板(新增站源参数,绑定模板与站源)
  31. :param template: AbstractTemplate子类实例
  32. :param site_source: 站源名称(从模板文件父文件夹名称获取)
  33. """
  34. if not isinstance(template, AbstractTemplate):
  35. raise TypeError("仅支持注册AbstractTemplate子类的模板实例")
  36. # 1. 处理模板覆盖逻辑
  37. template_id = template.template_id
  38. if template_id in self.registered_templates:
  39. print(f"模板ID={template_id}已存在,将覆盖旧版本(原站源:{self._get_template_site(template_id)} → 新站源:{site_source})")
  40. # 移除旧模板在站源映射中的关联
  41. old_site = self._get_template_site(template_id)
  42. if old_site and old_site in self.template_site_map:
  43. self.template_site_map[old_site].remove(template_id)
  44. if not self.template_site_map[old_site]: # 若站源下无模板,删除空列表
  45. del self.template_site_map[old_site]
  46. # 2. 注册模板核心信息
  47. self.registered_templates[template_id] = template
  48. self.template_priority_cache[template_id] = template.priority
  49. # 3. 绑定模板与站源(更新站源-模板ID映射)
  50. if site_source not in self.template_site_map:
  51. self.template_site_map[site_source] = []
  52. if template_id not in self.template_site_map[site_source]:
  53. self.template_site_map[site_source].append(template_id)
  54. # 4. 扩展模板实例属性:动态添加站源属性(便于后续查看)
  55. setattr(template, "site_source", site_source)
  56. print(f"模板注册成功:ID={template_id},站源={site_source}")
  57. def unregister_template(self, template_id: str) -> bool:
  58. """
  59. 注销模板(同步删除站源-模板映射)
  60. """
  61. if template_id not in self.registered_templates:
  62. print(f"模板ID={template_id}不存在,注销失败")
  63. return False
  64. # 1. 移除站源-模板映射
  65. site_source = self._get_template_site(template_id)
  66. if site_source and site_source in self.template_site_map:
  67. self.template_site_map[site_source].remove(template_id)
  68. if not self.template_site_map[site_source]:
  69. del self.template_site_map[site_source]
  70. # 2. 注销模板核心信息
  71. del self.registered_templates[template_id]
  72. del self.template_priority_cache[template_id]
  73. print(f"模板注销成功:ID={template_id},站源={site_source}")
  74. return True
  75. def _get_template_site(self, template_id: str) -> Optional[str]:
  76. """
  77. 辅助方法:通过模板ID获取对应的站源(从实例属性或站源映射反查)
  78. """
  79. if template_id not in self.registered_templates:
  80. return None
  81. # 从模板实例的动态属性获取站源(优先)
  82. template = self.registered_templates[template_id]
  83. if hasattr(template, "site_source"):
  84. return template.site_source
  85. # 从站源映射反查(兜底)
  86. for site, tpl_ids in self.template_site_map.items():
  87. if template_id in tpl_ids:
  88. return site
  89. return None
  90. def get_matched_templates(self, preprocessed_data: Dict[str, Any], site_source: Optional[str] = None) -> List[AbstractTemplate]:
  91. """
  92. 筛选匹配模板(新增站源筛选逻辑)
  93. :param site_source: 目标站源(若传入,仅筛选该站源的模板;若为None,筛选所有站源模板)
  94. :return: 按“站源→优先级”筛选后的可用模板列表
  95. """
  96. # 1. 第一步:按站源筛选模板(核心新增逻辑)
  97. if site_source and site_source in self.template_site_map:
  98. # 仅保留目标站源的模板ID
  99. site_tpl_ids = self.template_site_map[site_source]
  100. candidate_templates = [
  101. self.registered_templates[tpl_id]
  102. for tpl_id in site_tpl_ids
  103. ]
  104. else:
  105. # 无站源参数或站源不存在:候选模板为空
  106. # candidate_templates = list(self.registered_templates.values())
  107. candidate_templates = []
  108. if site_source and site_source not in self.template_site_map:
  109. print(f"警告:站源{site_source}无对应模板")
  110. # 2. 第二步:按“调用时机”筛选(原有逻辑)
  111. matched_templates = [
  112. tpl for tpl in candidate_templates
  113. if tpl.check_call_timing(preprocessed_data)
  114. ]
  115. # 3. 第三步:按优先级降序排序(原有逻辑)
  116. sorted_templates = sorted(
  117. matched_templates,
  118. key=lambda tpl: tpl.priority,
  119. reverse=True
  120. )
  121. return sorted_templates
  122. def call_templates(self, preprocessed_data: Dict[str, Any], general_result: Dict[str, Any], site_source: Optional[str] = None, docid: Optional[str]='') -> Dict[str, Any]:
  123. """
  124. 调用模板并融合结果(新增站源参数,仅调用对应站源模板)
  125. :param site_source: 目标站源(仅调用该站源的模板)
  126. """
  127. # 1. 初始化结果:先纳入通用提取结果(基础值)
  128. final_result = general_result.copy()
  129. # 2. 获取“指定站源+符合调用时机”的模板,按优先级执行(核心修改)
  130. sorted_templates = self.get_matched_templates(preprocessed_data, site_source)
  131. if not sorted_templates:
  132. sorted_templates = self.get_matched_templates(preprocessed_data, site_source='shared_template')
  133. if not sorted_templates:
  134. print(f"站源{site_source}(或所有站源)无可用模板,直接使用通用结果+AI结果")
  135. else:
  136. for tpl in sorted_templates:
  137. tpl_result = tpl.extract(preprocessed_data)
  138. # 用模板结果补充/覆盖(非空值才更新)
  139. for key, value in tpl_result.items():
  140. if value is not None and str(value).strip():
  141. final_result[key] = value
  142. # print(f"调用模板:ID={tpl.template_id},站源={tpl.site_source},提取结果:{tpl_result}")
  143. log("调用模板:%s,站源:%s,docid:%s, 提取结果:%s"%(tpl.template_id, site_source, docid, str(tpl_result)))
  144. # 3. 融合AI提取结果(最高优先级,原有逻辑)
  145. ai_result = preprocessed_data.get("ai_extract_result", {})
  146. for key, value in ai_result.items():
  147. if value is not None and str(value).strip():
  148. final_result[key] = value
  149. return final_result
  150. def batch_call(self, preprocessed_data_list: List[Dict[str, Any]], general_result_list: List[Dict[str, Any]], site_source: Optional[str] = None) -> List[Dict[str, Any]]:
  151. """
  152. 批量调用模板(新增站源参数,批量筛选站源模板)
  153. :param site_source: 目标站源(所有批量数据共用一个站源)
  154. """
  155. if len(preprocessed_data_list) != len(general_result_list):
  156. raise ValueError("预处理数据列表与通用结果列表长度不匹配")
  157. # 批量调用时传入统一的站源参数
  158. batch_result = [
  159. self.call_templates(pre_data, general_result, site_source)
  160. for pre_data, general_result in zip(preprocessed_data_list, general_result_list)
  161. ]
  162. return batch_result
  163. def get_template_info_list(self, site_source: Optional[str] = None) -> List[Dict[str, Any]]:
  164. """
  165. 获取模板信息(支持按站源筛选)
  166. :return: 包含站源属性的模板信息列表
  167. """
  168. # 1. 按站源筛选模板
  169. if site_source and site_source in self.template_site_map:
  170. target_templates = [
  171. self.registered_templates[tpl_id]
  172. for tpl_id in self.template_site_map[site_source]
  173. ]
  174. else:
  175. target_templates = list(self.registered_templates.values())
  176. # 2. 构造信息字典(新增站源字段)
  177. template_info_list = []
  178. for tpl in target_templates:
  179. info = tpl.get_template_info()
  180. info["site_source"] = getattr(tpl, "site_source", "未知") # 新增站源信息
  181. template_info_list.append(info)
  182. return template_info_list
  183. # ------------------------------ 核心修改:自动注册时绑定站源(父文件夹名称) ------------------------------
  184. def auto_register_templates_from_dir(self, root_target_dir: str, exclude_files: Optional[List[str]] = None, exclude_dirs: Optional[List[str]] = None) -> None:
  185. """
  186. 扫描文件夹并自动注册模板(核心修改:从父文件夹名称读取站源)
  187. 目录结构要求:root_target_dir/[站源文件夹]/[模板文件.py]
  188. 例:templates/zhihu/extract_tpl.py → 站源=zhihu;templates/toutiao/tpl_bidding.py → 站源=toutiao
  189. :param root_target_dir: 根目录(包含多个站源子文件夹)
  190. :param exclude_files: 需排除的文件名列表(默认排除__init__.py)
  191. :param exclude_dirs: 需排除的子文件夹列表(如["test_dir", "temp"])
  192. """
  193. # 1. 初始化排除列表
  194. exclude_files = exclude_files or []
  195. if "__init__.py" not in exclude_files:
  196. exclude_files.append("__init__.py")
  197. exclude_dirs = exclude_dirs or [] # 新增:排除不需要的子文件夹
  198. # 2. 验证根目录有效性
  199. root_dir_abs = os.path.abspath(root_target_dir)
  200. if not os.path.isdir(root_dir_abs):
  201. raise NotADirectoryError(f"根目录不是有效文件夹:{root_dir_abs}")
  202. print(f"\n开始扫描根目录:{root_dir_abs}")
  203. print(f"排除文件:{exclude_files},排除子文件夹:{exclude_dirs}")
  204. # 3. 遍历根目录下的所有子文件夹(每个子文件夹=一个站源)
  205. for dir_name in os.listdir(root_dir_abs):
  206. # 3.1 过滤非文件夹、排除文件夹
  207. site_dir_path = os.path.join(root_dir_abs, dir_name)
  208. if not os.path.isdir(site_dir_path) or dir_name in exclude_dirs:
  209. continue
  210. site_source = dir_name # 核心:子文件夹名称 = 站源名称
  211. print(f"\n=== 开始处理站源:{site_source}(文件夹:{dir_name})===")
  212. # 3.2 遍历当前站源文件夹下的所有.py文件
  213. for filename in os.listdir(site_dir_path):
  214. # 过滤非.py文件、排除文件
  215. if not filename.endswith(".py") or filename in exclude_files:
  216. continue
  217. file_path = os.path.join(site_dir_path, filename)
  218. module_name = f"{site_source}_{os.path.splitext(filename)[0]}" # 模块名:站源_文件名(避免重复)
  219. # 3.3 动态导入模块
  220. try:
  221. spec = importlib.util.spec_from_file_location(module_name, file_path)
  222. if spec is None or spec.loader is None:
  223. print(f"警告:站源{site_source},无法创建模块规格,跳过文件:{filename}")
  224. continue
  225. module = importlib.util.module_from_spec(spec)
  226. spec.loader.exec_module(module)
  227. print(f"站源{site_source},成功加载模块:{module_name}(文件:{filename})")
  228. except Exception as e:
  229. print(f"警告:站源{site_source},加载模块{module_name}失败,跳过文件:{filename}")
  230. print(f"错误详情:{str(e)}")
  231. continue
  232. # 3.4 扫描模块中的模板并注册(绑定站源)
  233. self._scan_and_register_from_module(module, filename, site_source)
  234. # 4. 输出注册汇总
  235. print(f"\n=== 文件夹扫描完成 ===")
  236. print(f"已识别站源数量:{len(self.template_site_map)}")
  237. for site, tpl_ids in self.template_site_map.items():
  238. print(f" 站源{site}:{len(tpl_ids)}个模板(IDs:{tpl_ids})")
  239. print(f"总计注册模板数量:{len(self.registered_templates)}")
  240. def _scan_and_register_from_module(self, module: Any, filename: str, site_source: str) -> None:
  241. """
  242. 辅助方法:从模块中识别模板并注册(新增站源参数)
  243. """
  244. for member_name, member_obj in inspect.getmembers(module):
  245. # 筛选条件:AbstractTemplate子类 + 非基类本身
  246. if (inspect.isclass(member_obj)
  247. and issubclass(member_obj, AbstractTemplate)
  248. and member_obj != AbstractTemplate):
  249. # 实例化并注册(传入站源参数)
  250. try:
  251. template_instance = member_obj()
  252. self.register_template(template_instance, site_source) # 调用修改后的注册方法
  253. except Exception as e:
  254. print(f"警告:站源{site_source},实例化模板类{member_name}失败(文件:{filename})")
  255. print(f"错误详情:{str(e)}")
  256. continue
  257. # ------------------------------ 使用示例(按站源调用模板) ------------------------------
  258. if __name__ == "__main__":
  259. # 1. 初始化模板管理器
  260. manager = TemplateManager()
  261. # 2. 【关键】自动扫描根目录,按“子文件夹名称”绑定站源
  262. # 目录结构示例:
  263. # ./templates_root/
  264. # ├─ nanan/ (站源1:nanan)
  265. # │ ├─ bidding_place_tpl.py (模板:tpl_bidding_place_001)
  266. # │ └─ bidding_time_tpl.py (模板:tpl_bidding_time_001)
  267. # └─ xxs/ (站源2:xxs)
  268. # └─ bidding_place_tpl.py (模板:tpl_bidding_place_002)
  269. ROOT_TEMPLATE_DIR = "./templates" # 根目录(包含多个站源子文件夹)
  270. try:
  271. manager.auto_register_templates_from_dir(
  272. root_target_dir=ROOT_TEMPLATE_DIR,
  273. exclude_files=["test_tpl.py"], # 排除测试文件
  274. exclude_dirs=["temp_dir"] # 排除临时文件夹
  275. )
  276. except Exception as e:
  277. traceback.print_exc()
  278. print(f"自动注册失败:{str(e)}")
  279. exit(1)
  280. with open('d:/html/2.html', encoding='utf-8') as f:
  281. html = f.read()
  282. tree = etree.HTML(html)
  283. # 3. 准备测试数据(2条数据,分别对应不同站源)
  284. test_pre_data = [
  285. # 数据1:对应站源nanan
  286. {
  287. "text": """招标单位:南安市美林街道溪一村民委员会
  288. 代理单位:福建省协祥工程项目管理咨询有限公司
  289. 地区:南安市(福建省)
  290. 开标地点:其他场地
  291. 行业:市政
  292. 预算:323014.0元
  293. 项目状态:已公示""",
  294. "ai_extract_result": {"招标单位": "南安市美林街道溪一村民委员会"}
  295. },
  296. # 数据2:对应站源xxs
  297. {
  298. "text": """招标单位:XX市住建局
  299. 代理单位:XX工程咨询公司
  300. 地区:XX市(浙江省)
  301. 开标地点:市公共资源交易中心3楼
  302. 行业:房建
  303. 预算:1500000.0元""",
  304. "ai_extract_result": {"地区": "XX市(浙江省)"}
  305. },
  306. {
  307. "tree": tree,
  308. "html": html
  309. },
  310. ]
  311. test_general_result = [
  312. {"行业": "未知", "项目状态": "未知"},
  313. {"代理单位": "未知", "预算": "未知"},
  314. {}
  315. ]
  316. site_source = [
  317. '',
  318. '15151',
  319. 'uni_jilin',
  320. ''
  321. ]
  322. # # 4. 批量调用模板并输出结果
  323. # print("\n=== 所有注册模板信息 ===")
  324. # for idx, tpl_info in enumerate(manager.get_template_info_list(), 1):
  325. # print(f"{idx}. {tpl_info}")
  326. #
  327. # print("\n=== 批量提取结果 ===")
  328. # batch_result = manager.batch_call(test_pre_data, test_general_result, 'DX000002')
  329. # for idx, result in enumerate(batch_result, 1):
  330. # print(f"数据{idx}:{result}")
  331. import pandas as pd
  332. from collections import Counter
  333. import json
  334. from 分析表格内容 import 获取表格行列信息,获取行内文本,获取大纲文本信息,extract_table_headers
  335. # df = pd.read_csv('E:\实体识别数据/2023-08-24所有公告.csv')
  336. # df2 = pd.read_csv('E:\实体识别数据/2023-08-24所有公告html.csv')
  337. df = pd.read_csv('E:/模版提取/医院大学站源公告.csv')
  338. df2 = pd.read_csv('E:\模版提取/医院大学站源公告_html.csv')
  339. print('公告:', len(df), len(df2), len(set(df['docid'])&set(df2['docid'])))
  340. print(Counter(df['web_source_name']).most_common(100))
  341. print(Counter(df['web_source_no']).most_common(100))
  342. # df = df[df['web_source_name']=='哈尔滨工业大学']
  343. # df = df[df['web_source_no'].str.contains('XX0182')]
  344. df = df.merge(df2, on='docid', how='inner')
  345. print('公告在:', set(df['docid'])&set(df2['docid']))
  346. print('公告数量:', len(df), df['docid'].tolist()[:20],df.columns)
  347. # print(Counter(df['docchannel']).most_common())
  348. # print(df['web_source_no'].tolist()[:20])
  349. df.fillna('', inplace=True)
  350. # extractor = OutlineExtractor()
  351. #
  352. # df = pd.read_csv('E:/模版提取/中国南方电网-站源公告.csv')
  353. # df2 = pd.read_csv('E:/模版提取/中国南方电网-站源公告_html.csv')
  354. # df = df.merge(df2, on='docid', how='left')
  355. # # df = df[df['docid'].isin([673112494])]
  356. # print(Counter(df['docchannel']).most_common())
  357. import time
  358. t1 = time.time()
  359. rs_list = []
  360. n1 = n2 = 0
  361. not_extract_heads = []
  362. head_docid = {}
  363. for docid, html2,web_no in zip(df['docid'], df['dochtmlcon'], df['web_source_no']):
  364. tree = etree.HTML(html2)
  365. # outline = extractor.extract(html2)
  366. tables = get_html_table(html2)
  367. pre_data = {
  368. # "tree": tree,
  369. # "html": html2,
  370. # "outline": outline,
  371. "表格列表": tables
  372. }
  373. # with open('preprocessed_data.json', 'w', encoding='utf-8') as f:
  374. # text = json.dumps(pre_data, ensure_ascii=False, indent=2)
  375. # f.write(text)
  376. # print(text)
  377. if '-' in web_no:
  378. web_no = web_no.split('-')[0]
  379. # if web_no in ['Y00236', 'XX0182']:
  380. # continue
  381. result = manager.call_templates(pre_data, {}, web_no)
  382. # result = {}
  383. # for table in tables:
  384. # pre_data = {
  385. # "表格列表": [table]
  386. # }
  387. # if '-' in web_no:
  388. # web_no = web_no.split('-')[0]
  389. # result_tmp = manager.call_templates(pre_data, {}, web_no)
  390. # if result_tmp:
  391. # for k, v in result_tmp.items():
  392. # if k not in result:
  393. # result[k] = v
  394. # else:
  395. # if isinstance(v, list):
  396. # result[k].extend(v)
  397. # else:
  398. # result[k] = v
  399. # print('公告数:%d, 耗时:%.4f'%(len(df), time.time()-t1)) # 公告数:714, 耗时:1.6382 1.5857 1.6712 # 公告数:714, 耗时:1.7684 1.7891 1.8029
  400. # 耗时:1.9181 耗时:2.2976
  401. if len(result) == 0:
  402. print('docid: ', docid)
  403. # 获取表格行列信息(html2)
  404. # 获取行内文本(html2)
  405. # 获取大纲文本信息(html2)
  406. headers = extract_table_headers(html2)
  407. if len(headers) > 2:
  408. print('未提取表头类型:', '_'.join(headers))
  409. not_extract_heads.append('_'.join(headers))
  410. head_docid['_'.join(headers)] = docid
  411. n1 += 1
  412. else:
  413. print('提取结果:', docid, result)
  414. if len(result.get('招标信息', []))>1 and '标的' not in result.get('招标信息', [])[0] and '标包' not in result.get('招标信息', [])[0]:
  415. print('信息错误,缺少标包:', docid)
  416. if len(result.get('中标信息', []))>1 and '标的' not in result.get('中标信息', [])[0] and '标包' not in result.get('中标信息', [])[0]:
  417. print('信息错误,缺少标包:', docid)
  418. rs_list.append((docid, json.dumps(result, ensure_ascii=False, indent=2, sort_keys=True)))
  419. # print('docid: ', docid)
  420. n2 += 1
  421. print()
  422. print('比例:', n1, n2, n2/(n1+n2))
  423. # #
  424. # # df = pd.DataFrame(rs_list, columns=['docid', '结果'])
  425. # # df.to_excel('E:\待检查数据/中国南方电网-站源公告_模板提取结果.xlsx')
  426. # for it in Counter(not_extract_heads).most_common(50):
  427. # print(it, len(it[0].split('_')), head_docid[it[0]])
  428. # print(Counter([it for l in not_extract_heads for it in l.split('_') ]).most_common(100))