| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494 |
- # -*- coding: utf-8 -*-
- import re
- import os
- import importlib
- import inspect
- import traceback
- from lxml import etree
- from typing import Dict, List, Optional, Any
- from BiddingKG.dl.template_extract.abstract_template import AbstractTemplate
- from BiddingKG.dl.template_extract.out_line_extractor import OutlineExtractor
- from BiddingKG.dl.template_extract.table_extractor import get_html_table
- from BiddingKG.dl.common.Utils import log
- class TemplateManager:
- """
- 模板管理与调用类(新增站源属性绑定+按站源筛选调用)
- 核心特性:
- 1. 自动注册时从模板文件的父文件夹名称读取“站源”属性
- 2. 调用模板时支持传入站源参数,仅执行对应站源的模板
- 3. 保留原有优先级规则、结果融合逻辑
- """
- def __init__(self):
- # 存储已注册的模板:key=模板ID,value=AbstractTemplate子类实例
- self.registered_templates: Dict[str, AbstractTemplate] = {}
- # 缓存模板优先级映射(避免重复读取,提升筛选效率)
- self.template_priority_cache: Dict[str, int] = {}
- # 新增:缓存模板-站源映射(快速筛选站源对应的模板)
- self.template_site_map: Dict[str, List[str]] = {} # key=站源名,value=模板ID列表
- def register_template(self, template: AbstractTemplate, site_source: str) -> None:
- """
- 注册模板(新增站源参数,绑定模板与站源)
- :param template: AbstractTemplate子类实例
- :param site_source: 站源名称(从模板文件父文件夹名称获取)
- """
- if not isinstance(template, AbstractTemplate):
- raise TypeError("仅支持注册AbstractTemplate子类的模板实例")
- # 1. 处理模板覆盖逻辑
- template_id = template.template_id
- if template_id in self.registered_templates:
- print(f"模板ID={template_id}已存在,将覆盖旧版本(原站源:{self._get_template_site(template_id)} → 新站源:{site_source})")
- # 移除旧模板在站源映射中的关联
- old_site = self._get_template_site(template_id)
- if old_site and old_site in self.template_site_map:
- self.template_site_map[old_site].remove(template_id)
- if not self.template_site_map[old_site]: # 若站源下无模板,删除空列表
- del self.template_site_map[old_site]
- # 2. 注册模板核心信息
- self.registered_templates[template_id] = template
- self.template_priority_cache[template_id] = template.priority
- # 3. 绑定模板与站源(更新站源-模板ID映射)
- if site_source not in self.template_site_map:
- self.template_site_map[site_source] = []
- if template_id not in self.template_site_map[site_source]:
- self.template_site_map[site_source].append(template_id)
- # 4. 扩展模板实例属性:动态添加站源属性(便于后续查看)
- setattr(template, "site_source", site_source)
- print(f"模板注册成功:ID={template_id},站源={site_source}")
- def unregister_template(self, template_id: str) -> bool:
- """
- 注销模板(同步删除站源-模板映射)
- """
- if template_id not in self.registered_templates:
- print(f"模板ID={template_id}不存在,注销失败")
- return False
- # 1. 移除站源-模板映射
- site_source = self._get_template_site(template_id)
- if site_source and site_source in self.template_site_map:
- self.template_site_map[site_source].remove(template_id)
- if not self.template_site_map[site_source]:
- del self.template_site_map[site_source]
- # 2. 注销模板核心信息
- del self.registered_templates[template_id]
- del self.template_priority_cache[template_id]
- print(f"模板注销成功:ID={template_id},站源={site_source}")
- return True
- def _get_template_site(self, template_id: str) -> Optional[str]:
- """
- 辅助方法:通过模板ID获取对应的站源(从实例属性或站源映射反查)
- """
- if template_id not in self.registered_templates:
- return None
- # 从模板实例的动态属性获取站源(优先)
- template = self.registered_templates[template_id]
- if hasattr(template, "site_source"):
- return template.site_source
- # 从站源映射反查(兜底)
- for site, tpl_ids in self.template_site_map.items():
- if template_id in tpl_ids:
- return site
- return None
- def get_matched_templates(self, preprocessed_data: Dict[str, Any], site_source: Optional[str] = None) -> List[AbstractTemplate]:
- """
- 筛选匹配模板(新增站源筛选逻辑)
- :param site_source: 目标站源(若传入,仅筛选该站源的模板;若为None,筛选所有站源模板)
- :return: 按“站源→优先级”筛选后的可用模板列表
- """
- # 1. 第一步:按站源筛选模板(核心新增逻辑)
- if site_source and site_source in self.template_site_map:
- # 仅保留目标站源的模板ID
- site_tpl_ids = self.template_site_map[site_source]
- candidate_templates = [
- self.registered_templates[tpl_id]
- for tpl_id in site_tpl_ids
- ]
- else:
- # 无站源参数或站源不存在:候选模板为空
- # candidate_templates = list(self.registered_templates.values())
- candidate_templates = []
- if site_source and site_source not in self.template_site_map:
- print(f"警告:站源{site_source}无对应模板")
- # 2. 第二步:按“调用时机”筛选(原有逻辑)
- matched_templates = [
- tpl for tpl in candidate_templates
- if tpl.check_call_timing(preprocessed_data)
- ]
- # 3. 第三步:按优先级降序排序(原有逻辑)
- sorted_templates = sorted(
- matched_templates,
- key=lambda tpl: tpl.priority,
- reverse=True
- )
- return sorted_templates
- 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]:
- """
- 调用模板并融合结果(新增站源参数,仅调用对应站源模板)
- :param site_source: 目标站源(仅调用该站源的模板)
- """
- # 1. 初始化结果:先纳入通用提取结果(基础值)
- final_result = general_result.copy()
- # 2. 获取“指定站源+符合调用时机”的模板,按优先级执行(核心修改)
- sorted_templates = self.get_matched_templates(preprocessed_data, site_source)
- if not sorted_templates:
- sorted_templates = self.get_matched_templates(preprocessed_data, site_source='shared_template')
- if not sorted_templates:
- print(f"站源{site_source}(或所有站源)无可用模板,直接使用通用结果+AI结果")
- else:
- for tpl in sorted_templates:
- tpl_result = tpl.extract(preprocessed_data)
- # 用模板结果补充/覆盖(非空值才更新)
- for key, value in tpl_result.items():
- if value is not None and str(value).strip():
- final_result[key] = value
- # print(f"调用模板:ID={tpl.template_id},站源={tpl.site_source},提取结果:{tpl_result}")
- log("调用模板:%s,站源:%s,docid:%s, 提取结果:%s"%(tpl.template_id, site_source, docid, str(tpl_result)))
- # 3. 融合AI提取结果(最高优先级,原有逻辑)
- ai_result = preprocessed_data.get("ai_extract_result", {})
- for key, value in ai_result.items():
- if value is not None and str(value).strip():
- final_result[key] = value
- return final_result
- 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]]:
- """
- 批量调用模板(新增站源参数,批量筛选站源模板)
- :param site_source: 目标站源(所有批量数据共用一个站源)
- """
- if len(preprocessed_data_list) != len(general_result_list):
- raise ValueError("预处理数据列表与通用结果列表长度不匹配")
- # 批量调用时传入统一的站源参数
- batch_result = [
- self.call_templates(pre_data, general_result, site_source)
- for pre_data, general_result in zip(preprocessed_data_list, general_result_list)
- ]
- return batch_result
- def get_template_info_list(self, site_source: Optional[str] = None) -> List[Dict[str, Any]]:
- """
- 获取模板信息(支持按站源筛选)
- :return: 包含站源属性的模板信息列表
- """
- # 1. 按站源筛选模板
- if site_source and site_source in self.template_site_map:
- target_templates = [
- self.registered_templates[tpl_id]
- for tpl_id in self.template_site_map[site_source]
- ]
- else:
- target_templates = list(self.registered_templates.values())
- # 2. 构造信息字典(新增站源字段)
- template_info_list = []
- for tpl in target_templates:
- info = tpl.get_template_info()
- info["site_source"] = getattr(tpl, "site_source", "未知") # 新增站源信息
- template_info_list.append(info)
- return template_info_list
- # ------------------------------ 核心修改:自动注册时绑定站源(父文件夹名称) ------------------------------
- def auto_register_templates_from_dir(self, root_target_dir: str, exclude_files: Optional[List[str]] = None, exclude_dirs: Optional[List[str]] = None) -> None:
- """
- 扫描文件夹并自动注册模板(核心修改:从父文件夹名称读取站源)
- 目录结构要求:root_target_dir/[站源文件夹]/[模板文件.py]
- 例:templates/zhihu/extract_tpl.py → 站源=zhihu;templates/toutiao/tpl_bidding.py → 站源=toutiao
- :param root_target_dir: 根目录(包含多个站源子文件夹)
- :param exclude_files: 需排除的文件名列表(默认排除__init__.py)
- :param exclude_dirs: 需排除的子文件夹列表(如["test_dir", "temp"])
- """
- # 1. 初始化排除列表
- exclude_files = exclude_files or []
- if "__init__.py" not in exclude_files:
- exclude_files.append("__init__.py")
- exclude_dirs = exclude_dirs or [] # 新增:排除不需要的子文件夹
- # 2. 验证根目录有效性
- root_dir_abs = os.path.abspath(root_target_dir)
- if not os.path.isdir(root_dir_abs):
- raise NotADirectoryError(f"根目录不是有效文件夹:{root_dir_abs}")
- print(f"\n开始扫描根目录:{root_dir_abs}")
- print(f"排除文件:{exclude_files},排除子文件夹:{exclude_dirs}")
- # 3. 遍历根目录下的所有子文件夹(每个子文件夹=一个站源)
- for dir_name in os.listdir(root_dir_abs):
- # 3.1 过滤非文件夹、排除文件夹
- site_dir_path = os.path.join(root_dir_abs, dir_name)
- if not os.path.isdir(site_dir_path) or dir_name in exclude_dirs:
- continue
- site_source = dir_name # 核心:子文件夹名称 = 站源名称
- print(f"\n=== 开始处理站源:{site_source}(文件夹:{dir_name})===")
- # 3.2 遍历当前站源文件夹下的所有.py文件
- for filename in os.listdir(site_dir_path):
- # 过滤非.py文件、排除文件
- if not filename.endswith(".py") or filename in exclude_files:
- continue
- file_path = os.path.join(site_dir_path, filename)
- module_name = f"{site_source}_{os.path.splitext(filename)[0]}" # 模块名:站源_文件名(避免重复)
- # 3.3 动态导入模块
- try:
- spec = importlib.util.spec_from_file_location(module_name, file_path)
- if spec is None or spec.loader is None:
- print(f"警告:站源{site_source},无法创建模块规格,跳过文件:{filename}")
- continue
- module = importlib.util.module_from_spec(spec)
- spec.loader.exec_module(module)
- print(f"站源{site_source},成功加载模块:{module_name}(文件:{filename})")
- except Exception as e:
- print(f"警告:站源{site_source},加载模块{module_name}失败,跳过文件:{filename}")
- print(f"错误详情:{str(e)}")
- continue
- # 3.4 扫描模块中的模板并注册(绑定站源)
- self._scan_and_register_from_module(module, filename, site_source)
- # 4. 输出注册汇总
- print(f"\n=== 文件夹扫描完成 ===")
- print(f"已识别站源数量:{len(self.template_site_map)}")
- for site, tpl_ids in self.template_site_map.items():
- print(f" 站源{site}:{len(tpl_ids)}个模板(IDs:{tpl_ids})")
- print(f"总计注册模板数量:{len(self.registered_templates)}")
- def _scan_and_register_from_module(self, module: Any, filename: str, site_source: str) -> None:
- """
- 辅助方法:从模块中识别模板并注册(新增站源参数)
- """
- for member_name, member_obj in inspect.getmembers(module):
- # 筛选条件:AbstractTemplate子类 + 非基类本身
- if (inspect.isclass(member_obj)
- and issubclass(member_obj, AbstractTemplate)
- and member_obj != AbstractTemplate):
- # 实例化并注册(传入站源参数)
- try:
- template_instance = member_obj()
- self.register_template(template_instance, site_source) # 调用修改后的注册方法
- except Exception as e:
- print(f"警告:站源{site_source},实例化模板类{member_name}失败(文件:{filename})")
- print(f"错误详情:{str(e)}")
- continue
- # ------------------------------ 使用示例(按站源调用模板) ------------------------------
- if __name__ == "__main__":
- # 1. 初始化模板管理器
- manager = TemplateManager()
- # 2. 【关键】自动扫描根目录,按“子文件夹名称”绑定站源
- # 目录结构示例:
- # ./templates_root/
- # ├─ nanan/ (站源1:nanan)
- # │ ├─ bidding_place_tpl.py (模板:tpl_bidding_place_001)
- # │ └─ bidding_time_tpl.py (模板:tpl_bidding_time_001)
- # └─ xxs/ (站源2:xxs)
- # └─ bidding_place_tpl.py (模板:tpl_bidding_place_002)
- ROOT_TEMPLATE_DIR = "./templates" # 根目录(包含多个站源子文件夹)
- try:
- manager.auto_register_templates_from_dir(
- root_target_dir=ROOT_TEMPLATE_DIR,
- exclude_files=["test_tpl.py"], # 排除测试文件
- exclude_dirs=["temp_dir"] # 排除临时文件夹
- )
- except Exception as e:
- traceback.print_exc()
- print(f"自动注册失败:{str(e)}")
- exit(1)
- with open('d:/html/2.html', encoding='utf-8') as f:
- html = f.read()
- tree = etree.HTML(html)
- # 3. 准备测试数据(2条数据,分别对应不同站源)
- test_pre_data = [
- # 数据1:对应站源nanan
- {
- "text": """招标单位:南安市美林街道溪一村民委员会
- 代理单位:福建省协祥工程项目管理咨询有限公司
- 地区:南安市(福建省)
- 开标地点:其他场地
- 行业:市政
- 预算:323014.0元
- 项目状态:已公示""",
- "ai_extract_result": {"招标单位": "南安市美林街道溪一村民委员会"}
- },
- # 数据2:对应站源xxs
- {
- "text": """招标单位:XX市住建局
- 代理单位:XX工程咨询公司
- 地区:XX市(浙江省)
- 开标地点:市公共资源交易中心3楼
- 行业:房建
- 预算:1500000.0元""",
- "ai_extract_result": {"地区": "XX市(浙江省)"}
- },
- {
- "tree": tree,
- "html": html
- },
- ]
- test_general_result = [
- {"行业": "未知", "项目状态": "未知"},
- {"代理单位": "未知", "预算": "未知"},
- {}
- ]
- site_source = [
- '',
- '15151',
- 'uni_jilin',
- ''
- ]
- # # 4. 批量调用模板并输出结果
- # print("\n=== 所有注册模板信息 ===")
- # for idx, tpl_info in enumerate(manager.get_template_info_list(), 1):
- # print(f"{idx}. {tpl_info}")
- #
- # print("\n=== 批量提取结果 ===")
- # batch_result = manager.batch_call(test_pre_data, test_general_result, 'DX000002')
- # for idx, result in enumerate(batch_result, 1):
- # print(f"数据{idx}:{result}")
- import pandas as pd
- from collections import Counter
- import json
- from 分析表格内容 import 获取表格行列信息,获取行内文本,获取大纲文本信息,extract_table_headers
- # df = pd.read_csv('E:\实体识别数据/2023-08-24所有公告.csv')
- # df2 = pd.read_csv('E:\实体识别数据/2023-08-24所有公告html.csv')
- df = pd.read_csv('E:/模版提取/医院大学站源公告.csv')
- df2 = pd.read_csv('E:\模版提取/医院大学站源公告_html.csv')
- print('公告:', len(df), len(df2), len(set(df['docid'])&set(df2['docid'])))
- print(Counter(df['web_source_name']).most_common(100))
- print(Counter(df['web_source_no']).most_common(100))
- # df = df[df['web_source_name']=='哈尔滨工业大学']
- # df = df[df['web_source_no'].str.contains('XX0182')]
- df = df.merge(df2, on='docid', how='inner')
- print('公告在:', set(df['docid'])&set(df2['docid']))
- print('公告数量:', len(df), df['docid'].tolist()[:20],df.columns)
- # print(Counter(df['docchannel']).most_common())
- # print(df['web_source_no'].tolist()[:20])
- df.fillna('', inplace=True)
- # extractor = OutlineExtractor()
- #
- # df = pd.read_csv('E:/模版提取/中国南方电网-站源公告.csv')
- # df2 = pd.read_csv('E:/模版提取/中国南方电网-站源公告_html.csv')
- # df = df.merge(df2, on='docid', how='left')
- # # df = df[df['docid'].isin([673112494])]
- # print(Counter(df['docchannel']).most_common())
- import time
- t1 = time.time()
- rs_list = []
- n1 = n2 = 0
- not_extract_heads = []
- head_docid = {}
- for docid, html2,web_no in zip(df['docid'], df['dochtmlcon'], df['web_source_no']):
- tree = etree.HTML(html2)
- # outline = extractor.extract(html2)
- tables = get_html_table(html2)
- pre_data = {
- # "tree": tree,
- # "html": html2,
- # "outline": outline,
- "表格列表": tables
- }
- # with open('preprocessed_data.json', 'w', encoding='utf-8') as f:
- # text = json.dumps(pre_data, ensure_ascii=False, indent=2)
- # f.write(text)
- # print(text)
- if '-' in web_no:
- web_no = web_no.split('-')[0]
- # if web_no in ['Y00236', 'XX0182']:
- # continue
- result = manager.call_templates(pre_data, {}, web_no)
- # result = {}
- # for table in tables:
- # pre_data = {
- # "表格列表": [table]
- # }
- # if '-' in web_no:
- # web_no = web_no.split('-')[0]
- # result_tmp = manager.call_templates(pre_data, {}, web_no)
- # if result_tmp:
- # for k, v in result_tmp.items():
- # if k not in result:
- # result[k] = v
- # else:
- # if isinstance(v, list):
- # result[k].extend(v)
- # else:
- # result[k] = v
- # print('公告数:%d, 耗时:%.4f'%(len(df), time.time()-t1)) # 公告数:714, 耗时:1.6382 1.5857 1.6712 # 公告数:714, 耗时:1.7684 1.7891 1.8029
- # 耗时:1.9181 耗时:2.2976
- if len(result) == 0:
- print('docid: ', docid)
- # 获取表格行列信息(html2)
- # 获取行内文本(html2)
- # 获取大纲文本信息(html2)
- headers = extract_table_headers(html2)
- if len(headers) > 2:
- print('未提取表头类型:', '_'.join(headers))
- not_extract_heads.append('_'.join(headers))
- head_docid['_'.join(headers)] = docid
- n1 += 1
- else:
- print('提取结果:', docid, result)
- if len(result.get('招标信息', []))>1 and '标的' not in result.get('招标信息', [])[0] and '标包' not in result.get('招标信息', [])[0]:
- print('信息错误,缺少标包:', docid)
- if len(result.get('中标信息', []))>1 and '标的' not in result.get('中标信息', [])[0] and '标包' not in result.get('中标信息', [])[0]:
- print('信息错误,缺少标包:', docid)
- rs_list.append((docid, json.dumps(result, ensure_ascii=False, indent=2, sort_keys=True)))
- # print('docid: ', docid)
- n2 += 1
- print()
- print('比例:', n1, n2, n2/(n1+n2))
- # #
- # # df = pd.DataFrame(rs_list, columns=['docid', '结果'])
- # # df.to_excel('E:\待检查数据/中国南方电网-站源公告_模板提取结果.xlsx')
- # for it in Counter(not_extract_heads).most_common(50):
- # print(it, len(it[0].split('_')), head_docid[it[0]])
- # print(Counter([it for l in not_extract_heads for it in l.split('_') ]).most_common(100))
|