channel_bert.py 35 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591592593594595596597598599600601602603604605606607608609610611612613614615616617618619620621622623624625626627628629630631632633634635636637638639640641642643644645646647648649650651652653654655656657658659660661662663664665666667668669670671672673674675676677678679680681682683684685686687688689690691692693694695696697698699700701702703704705706707708709710711712713714715716717718719
  1. # coding: UTF-8
  2. import torch
  3. import torch.nn as nn
  4. import torch.nn.functional as F
  5. import time
  6. import re
  7. import os
  8. import transformers
  9. from transformers import ElectraTokenizer
  10. import numpy as np
  11. import json
  12. # device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
  13. device = torch.device("cpu") # 线上用CPU
  14. class PositionalEncoding(nn.Module):
  15. def __init__(self,dim_hid):
  16. super(PositionalEncoding,self).__init__()
  17. base_array = np.array([np.power(10000,2*(hid_j//2)/dim_hid) for hid_j in range(dim_hid)])
  18. self.base_tensor = torch.from_numpy(base_array).to(torch.float32).to(device) #[1,D]
  19. def forward(self,x):
  20. # x(B,N,d)
  21. B,N,d = x.shape
  22. pos = torch.arange(N).unsqueeze(-1).to(torch.float32).to(device) #[N,1]
  23. pos = pos/self.base_tensor
  24. pos = pos.unsqueeze(0)
  25. pos[:,:,0::2] = torch.sin(pos[:,:,0::2])
  26. pos[:,:,1::2] = torch.cos(pos[:,:,1::2])
  27. return x+pos
  28. class ScaledDotProductAttention(nn.Module):
  29. ''' Scaled Dot-Product Attention '''
  30. def __init__(self, temperature, attn_dropout=0.1):
  31. super().__init__()
  32. self.temperature = temperature
  33. self.dropout = nn.Dropout(attn_dropout)
  34. def forward(self, q, k, v, mask=None):
  35. # print(q.shape,k.shape)
  36. attn = torch.matmul(q / self.temperature, k.transpose(2, 3))
  37. if mask is not None:
  38. attn = attn.masked_fill(mask == 0, -1e9)
  39. # t1 = time.time()
  40. attn = self.dropout(torch.softmax(attn, dim=-1))
  41. # print('cost',time.time()-t1) # 主要时间花费
  42. output = torch.matmul(attn, v)
  43. return output, attn
  44. class MultiHeadAttention(nn.Module):
  45. ''' Multi-Head Attention module '''
  46. def __init__(self, n_head, d_model, d_k, d_v, dropout=0.1):
  47. super().__init__()
  48. self.n_head = n_head
  49. self.d_k = d_k
  50. self.d_v = d_v
  51. self.w_qs = nn.Linear(d_model, n_head * d_k, bias=False)
  52. self.w_ks = nn.Linear(d_model, n_head * d_k, bias=False)
  53. self.w_vs = nn.Linear(d_model, n_head * d_v, bias=False)
  54. self.fc = nn.Linear(n_head * d_v, d_model, bias=False)
  55. self.attention = ScaledDotProductAttention(temperature=d_k ** 0.5)
  56. self.dropout = nn.Dropout(dropout)
  57. self.layer_norm = nn.LayerNorm(d_model, eps=1e-6)
  58. self.rotaryEmbedding = RotaryEmbedding(d_k)
  59. def forward(self, q, k, v, mask=None):
  60. d_k, d_v, n_head = self.d_k, self.d_v, self.n_head
  61. sz_b, len_q, len_k, len_v = q.size(0), q.size(1), k.size(1), v.size(1)
  62. residual = q
  63. # Pass through the pre-attention projection: b x lq x (n*dv)
  64. # Separate different heads: b x lq x n x dv
  65. q = self.w_qs(q).view(sz_b, len_q, n_head, d_k)
  66. k = self.w_ks(k).view(sz_b, len_k, n_head, d_k)
  67. v = self.w_vs(v).view(sz_b, len_v, n_head, d_v)
  68. # Transpose for attention dot product: b x n x lq x dv
  69. q, k, v = q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)
  70. # RoPE embed
  71. target_tensor = torch.zeros((q.size(0), q.size(2)))
  72. position_ids = torch.arange(q.size(2), dtype=torch.long).unsqueeze(0).expand_as(target_tensor)
  73. _cos, _sin = self.rotaryEmbedding(q, position_ids)
  74. q, k = apply_rotary_pos_emb(q, k, _cos, _sin)
  75. if mask is not None:
  76. # mask = mask.unsqueeze(1) # For head axis broadcasting.
  77. mask = mask.unsqueeze(1).unsqueeze(2) # For head axis broadcasting.
  78. q, attn = self.attention(q, k, v, mask=mask)
  79. #q (sz_b,n_head,N=len_q,d_k)
  80. #k (sz_b,n_head,N=len_k,d_k)
  81. #v (sz_b,n_head,N=len_v,d_v)
  82. # Transpose to move the head dimension back: b x lq x n x dv
  83. # Combine the last two dimensions to concatenate all the heads together: b x lq x (n*dv)
  84. q = q.transpose(1, 2).contiguous().view(sz_b, len_q, -1)
  85. #q (sz_b,len_q,n_head,N * d_k)
  86. q = self.dropout(self.fc(q))
  87. q += residual
  88. q = self.layer_norm(q)
  89. return q, attn
  90. class PositionwiseFeedForward(nn.Module):
  91. ''' A two-feed-forward-layer module '''
  92. def __init__(self, d_in, d_hid, dropout=0.1):
  93. super().__init__()
  94. self.w_1 = nn.Linear(d_in, d_hid) # position-wise
  95. self.w_2 = nn.Linear(d_hid, d_in) # position-wise
  96. self.layer_norm = nn.LayerNorm(d_in, eps=1e-6)
  97. self.dropout = nn.Dropout(dropout)
  98. def forward(self, x):
  99. residual = x
  100. x = self.w_2(torch.relu(self.w_1(x)))
  101. x = self.dropout(x)
  102. x += residual
  103. x = self.layer_norm(x)
  104. return x
  105. class RotaryEmbedding(nn.Module):
  106. def __init__(self, dim, max_position_embeddings=2048, base=10000, device=None):
  107. super().__init__()
  108. self.dim = dim # it is set to the head_dim
  109. self.max_position_embeddings = max_position_embeddings
  110. self.base = base
  111. # Calculate the theta according to the formula theta_i = base^(2i/dim) where i = 0, 1, 2, ..., dim // 2
  112. inv_freq = 1.0 / (self.base ** (torch.arange(0, self.dim, 2, dtype=torch.int64).float() / self.dim))
  113. self.register_buffer("inv_freq", tensor=inv_freq, persistent=False)
  114. @torch.no_grad()
  115. def forward(self, x, position_ids, seq_len=None):
  116. # x: [bs, num_attention_heads, seq_len, head_size]
  117. self.inv_freq = self.inv_freq.to(device)
  118. position_ids = position_ids.to(device)
  119. # Copy the inv_freq tensor for batch in the sequence
  120. # inv_freq_expanded: [Batch_Size, Head_Dim // 2, 1]
  121. inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1)
  122. # position_ids_expanded: [Batch_Size, 1, Seq_Len]
  123. position_ids_expanded = position_ids[:, None, :].float()
  124. # Multiply each theta by the position (which is the argument of the sin and cos functions)
  125. # freqs: [Batch_Size, Head_Dim // 2, 1] @ [Batch_Size, 1, Seq_Len] --> [Batch_Size, Seq_Len, Head_Dim // 2]
  126. freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2)
  127. # emb: [Batch_Size, Seq_Len, Head_Dim]
  128. emb = torch.cat((freqs, freqs), dim=-1)
  129. # cos, sin: [Batch_Size, Seq_Len, Head_Dim]
  130. cos = emb.cos()
  131. sin = emb.sin()
  132. return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
  133. def rotate_half(x):
  134. # Build the [-x2, x1, -x4, x3, ...] tensor for the sin part of the positional encoding.
  135. x1 = x[..., : x.shape[-1] // 2] # Takes the first half of the last dimension
  136. x2 = x[..., x.shape[-1] // 2 :] # Takes the second half of the last dimension
  137. return torch.cat((-x2, x1), dim=-1)
  138. def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1):
  139. cos = cos.unsqueeze(unsqueeze_dim) # Add the head dimension
  140. sin = sin.unsqueeze(unsqueeze_dim) # Add the head dimension
  141. # Apply the formula (34) of the Rotary Positional Encoding paper.
  142. q_embed = (q * cos) + (rotate_half(q) * sin)
  143. k_embed = (k * cos) + (rotate_half(k) * sin)
  144. return q_embed, k_embed
  145. class EncoderLayer(nn.Module):
  146. ''' Compose with two layers '''
  147. def __init__(self, d_model, d_inner, n_head, d_k, d_v, dropout=0.1):
  148. super(EncoderLayer, self).__init__()
  149. self.slf_attn = MultiHeadAttention(n_head, d_model, d_k, d_v, dropout=dropout)
  150. self.pos_ffn = PositionwiseFeedForward(d_model, d_inner, dropout=dropout)
  151. def forward(self, enc_input, slf_attn_mask=None):
  152. enc_output, enc_slf_attn = self.slf_attn(
  153. enc_input, enc_input, enc_input, mask=slf_attn_mask)
  154. enc_output = self.pos_ffn(enc_output)
  155. return enc_output, enc_slf_attn
  156. class Encoder(nn.Module):
  157. ''' A encoder model with self attention mechanism. '''
  158. def __init__(
  159. self, n_src_vocab, d_word_vec, n_layers, n_head, d_k, d_v,
  160. d_model, d_inner, pad_idx, dropout=0.1, n_position=200, scale_emb=False,embedding=None):
  161. super().__init__()
  162. if embedding is None:
  163. self.src_word_emb = nn.Embedding(n_src_vocab, d_word_vec, padding_idx=pad_idx)
  164. else:
  165. self.src_word_emb = nn.Embedding(n_src_vocab, d_word_vec,padding_idx=pad_idx,
  166. # _weight=torch.from_numpy(embedding ,))
  167. # _weight=torch.tensor(embedding ,dtype=torch.float64).to(device))
  168. _weight=torch.tensor(embedding))
  169. self.position_enc = PositionalEncoding(d_word_vec)
  170. self.dropout = nn.Dropout(p=dropout)
  171. self.layer_stack = nn.ModuleList([
  172. EncoderLayer(d_model, d_inner, n_head, d_k, d_v, dropout=dropout)
  173. for _ in range(n_layers)])
  174. self.layer_norm = nn.LayerNorm(d_model, eps=1e-6)
  175. self.scale_emb = scale_emb
  176. self.d_model = d_model
  177. def forward(self, src_seq, src_mask, return_attns=False):
  178. enc_slf_attn_list = []
  179. # -- Forward
  180. enc_output = self.src_word_emb(src_seq)
  181. if self.scale_emb:
  182. enc_output *= self.d_model ** 0.5
  183. # enc_output = self.dropout(self.position_enc(enc_output))
  184. enc_output = self.dropout(enc_output)
  185. enc_output = self.layer_norm(enc_output)
  186. for enc_layer in self.layer_stack:
  187. enc_output, enc_slf_attn = enc_layer(enc_output, slf_attn_mask=src_mask)
  188. enc_slf_attn_list += [enc_slf_attn] if return_attns else []
  189. if return_attns:
  190. return enc_output, enc_slf_attn_list
  191. return enc_output
  192. class bidiBert(nn.Module):
  193. def __init__(self, n_src_vocab, d_word_vec, n_layers, n_head, d_k, d_v,
  194. d_model, d_inner, pad_idx,n_class,embedding = None):
  195. super(bidiBert, self).__init__()
  196. self.encoder = Encoder(n_src_vocab, d_word_vec, n_layers, n_head, d_k, d_v,
  197. d_model, d_inner, pad_idx,embedding = embedding)
  198. # self.src_word_emb = nn.Embedding(n_src_vocab, d_word_vec, padding_idx=pad_idx,
  199. # # _weight=torch.from_numpy(embedding ,))
  200. # # _weight=torch.tensor(embedding ,dtype=torch.float64).to(device))
  201. # _weight=torch.tensor(embedding))
  202. # self.encoder = nn.LSTM(128, 256, 2,
  203. # bidirectional=True, batch_first=True, dropout=0.3)
  204. self.dropout = nn.Dropout(p=0.1)
  205. self.pooler = nn.Linear(d_inner,d_inner)
  206. self.liner = nn.Linear(d_inner,n_class)
  207. # self.liner = nn.Linear(256*4,n_class)
  208. self.avg_pool1 = nn.AdaptiveAvgPool1d(1) # Max Pooling: nn.AdaptiveMaxPool1d(1)
  209. self.avg_pool2 = nn.AdaptiveAvgPool1d(1)
  210. def forward(self, input_title,input_doctext):
  211. # out = self.encoder(inputs, attention_mask)
  212. # input_title = self.src_word_emb(input_title)
  213. # input_title,_ = self.encoder(input_title)
  214. # input_title = self.encoder(src_seq=input_title[0], src_mask=input_title[1])
  215. # input_title = input_title[:, 1, :]
  216. # input_title = torch.mean(input_title,dim=-2)
  217. # input_title = self.avg_pool1(input_title.transpose(1, 2)).squeeze(-1)
  218. # input_title = torch.tanh(self.pooler(input_title[:,0]))
  219. # input_doctext = self.src_word_emb(input_doctext)
  220. # input_doctext,_ = self.encoder(input_doctext)
  221. input_doctext = self.encoder(src_seq=input_doctext[0], src_mask=input_doctext[1])
  222. # input_doctext = input_doctext[:, 1, :]
  223. # input_doctext = torch.mean(input_doctext,dim=-2)
  224. # input_doctext = self.avg_pool2(input_doctext.transpose(1, 2)).squeeze(-1)
  225. input_doctext = self.pooler(input_doctext[:,0])
  226. input_doctext = self.dropout(input_doctext)
  227. input_doctext = torch.tanh(input_doctext)
  228. # print('size:',input_title.size(),input_doctext.size())
  229. # out = torch.cat((input_title, input_doctext), dim=-1)
  230. out = input_doctext
  231. # bs, n, m = out.size()
  232. # out = out.view(bs,n*m)
  233. out = self.liner(out)
  234. out = F.softmax(out, dim=-1)
  235. return out
  236. phone = re.compile('1[3-9][0-9][-—-―]?\d{4}[-—-―]?\d{4}|'
  237. '\+86.?1[3-9]\d{9}|'
  238. # '0[^0]\d{1,2}[-—-―][1-9]\d{6,7}/[1-9]\d{6,10}|'
  239. '0[1-9]\d{1,2}[-—-―][2-9]\d{6}\d?[-—-―]\d{1,4}|'
  240. '0[1-9]\d{1,2}[-—-―]{0,2}[2-9]\d{6}\d?(?=1[3-9]\d{9})|'
  241. '0[1-9]\d{1,2}[-—-―]{0,2}[2-9]\d{6}\d?(?=0[1-9]\d{1,2}[-—-―]?[2-9]\d{6}\d?)|'
  242. '0[1-9]\d{1,2}[-—-―]{0,2}[2-9]\d{6}\d?(?=[2-9]\d{6,7})|'
  243. '0[1-9]\d{1,2}[-—-―]{0,2}[2-9]\d{6}\d?|'
  244. '[\(|\(]0[1-9]\d{1,2}[\)|\)]-?[2-9]\d{6}\d?-?\d{,4}|'
  245. '400\d{7}转\d{1,4}|'
  246. '[2-9]\d{6,7}')
  247. def text_process(text):
  248. text = text.strip()
  249. text = re.sub(r'[\000-\010]|[\013-\014]|[\016-\037]',"",text) # 非法字符
  250. text = re.sub("extractJson:|fullTextSeg:","",text)
  251. # text = re.sub("[??]{1,}", "", text)
  252. text = re.sub("[??]{2,}", "", text)
  253. text = re.sub(r'(http[s]?://|www\.)(?:[a-zA-Z]|[0-9]|[$-_@.&+]|[!*\\(\\),]|(?:%[0-9a-fA-F][0-9a-fA-F]))+', "", text) # 网站
  254. text = re.sub(r'[a-zA-Z0-9_.+-]+@[a-zA-Z0-9-]+\.[a-zA-Z0-9-.]+', "", text)# 邮箱
  255. text = re.sub('[0-9a-zA-Z@#*%=?&~()_|<>/.(){}【】{}\[\]\-]{6,}', "", text) # 编号
  256. text = re.sub(phone,"",text) # 号码
  257. text = re.sub(r'\b\d[-.\s]?\d{3}[-.\s]?\d{4}\b', "", text) # 座机
  258. text = re.sub('0[1-9]\d{1,2}[-—-―]{0,2}[2-9]\d{6}\d?', "", text) # 座机
  259. text = re.sub('1(3[0-9]|4[01456879]|5[0-35-9]|6[2567]|7[0-8]|8[0-9]|9[0-35-9])\d{8}', "", text)# 手机号
  260. text = re.sub('&?nbsp;?|&?ensp;?|&?emsp;?', "", text)
  261. text = re.sub('\\\\n|\\\\r|\\\\t', "", text)
  262. # text = re.sub("\s+", "", text)
  263. text = re.sub("\s+", " ", text)
  264. # 优化部分未识别表达
  265. text = re.sub("中止", "终止", text)
  266. text = re.sub("遴选", "招标", text)
  267. text = re.sub("合同段", "", text)
  268. return text
  269. label2class_dict = {
  270. 0: 51, 1:52 , 2:101,
  271. 3:102, 4:103, 5:105,
  272. 6:114, 7:118, 8:119,
  273. 9:120, 10:121, 11:122
  274. }
  275. def channel_predict(title,text,min_text_len=None):
  276. if globals().get("channel_pytorch_model") is None or globals().get("channel_tokenizer") is None:
  277. # config
  278. config = {
  279. # 'n_src_vocab': len(vocab),
  280. 'd_word_vec': 128,
  281. 'n_layers': 3,
  282. 'n_head': 3,
  283. 'd_k': 128,
  284. 'd_v': 128,
  285. 'd_model': 128,
  286. 'd_inner': 128,
  287. 'pad_idx': 0
  288. }
  289. # n_src_vocab = config['n_src_vocab']
  290. d_word_vec = config['d_word_vec']
  291. n_layers = config['n_layers']
  292. n_head = config['n_head']
  293. d_k = config['d_k']
  294. d_v = config['d_v']
  295. d_model = config['d_model']
  296. d_inner = config['d_inner']
  297. pad_idx = config['pad_idx']
  298. n_class = 12
  299. # tokenizer
  300. base_model_name = os.path.abspath(os.path.dirname(__file__)) + "/pytorch_model/tokenizer"
  301. tokenizer = ElectraTokenizer.from_pretrained(base_model_name)
  302. n_src_vocab = len(tokenizer.get_vocab())
  303. # 实例化模型
  304. model_path = os.path.abspath(os.path.dirname(__file__)) + '/pytorch_model/channel.pth'
  305. model = bidiBert(n_src_vocab, d_word_vec, n_layers, n_head, d_k, d_v,
  306. d_model, d_inner, pad_idx, n_class, embedding=None)
  307. model.to(device)
  308. model_state = torch.load(model_path, map_location=device)
  309. model_state_dict = model.state_dict()
  310. pretrained_state_dict = model_state
  311. # missing_keys = set(model_state_dict.keys()) - set(pretrained_state_dict.keys())
  312. unexpected_keys = set(pretrained_state_dict.keys()) - set(model_state_dict.keys())
  313. add_kv = []
  314. for k, v in model_state.items():
  315. if k in unexpected_keys:
  316. # model_state[k.replace("module.","")] = v
  317. add_kv.append([k.replace("module.", ""), v])
  318. for i in add_kv:
  319. model_state[i[0]] = i[1]
  320. for k in list(unexpected_keys):
  321. del model_state[k]
  322. model.load_state_dict(model_state)
  323. # 将模型设置为评估模式
  324. model.eval()
  325. globals()["channel_pytorch_model"] = model
  326. globals()["channel_tokenizer"] = tokenizer
  327. else:
  328. model = globals().get("channel_pytorch_model")
  329. tokenizer = globals().get("channel_tokenizer")
  330. # process text
  331. if title in text:
  332. text = text.replace(title, '', 1)
  333. text = text.lstrip(",")
  334. text = text.lstrip("。")
  335. if "##attachment##" in text:
  336. main_text,attachment_text = text.split("##attachment##",maxsplit=1)
  337. # print('main_text',main_text)
  338. if len(main_text)>=500: # 正文有足够的内容时不需要使用附件预测
  339. text = main_text
  340. text = re.sub("##attachment##。?","",text)
  341. text = text_process(text)
  342. if min_text_len is None:
  343. min_text_len = 200
  344. if len(text)<=min_text_len:
  345. # 正文内容过短时,不预测
  346. return
  347. # elif len(text)<=150:
  348. # # 正文内容过短时,重复正文
  349. # text = text * 2
  350. text = text[:2000]
  351. title = text_process(title)
  352. title = title[:100]
  353. text = "公告标题:" + title + "。" + "公告内容:" + text
  354. text = text[:2000]
  355. # print('predict text:',text)
  356. # to torch data
  357. text = [text]
  358. text_max_len = 2000
  359. # text = [tokenizer.encode_plus(
  360. # _t,
  361. # add_special_tokens=True, # 添加特殊标记,如[CLS]和[SEP]
  362. # max_length=text_max_len, # 设置最大长度
  363. # padding='max_length', # 填充到最大长度
  364. # truncation=True, # 截断超过最大长度的文本
  365. # return_attention_mask=True, # 返回attention_mask
  366. # return_tensors='pt' # 返回PyTorch张量
  367. # ) for _t in text]
  368. # text = [torch.LongTensor(np.array([_t['input_ids'].numpy()[0] for _t in text])).to(device),
  369. # torch.LongTensor(np.array([_t['attention_mask'].numpy()[0] for _t in text])).to(device)]
  370. text = [tokenizer.encode_plus(
  371. _t,
  372. add_special_tokens=True, # 添加特殊标记,如[CLS]和[SEP]
  373. max_length=text_max_len, # 设置最大长度
  374. padding='max_length', # 填充到最大长度
  375. truncation=True, # 截断超过最大长度的文本
  376. return_attention_mask=True, # 返回attention_mask
  377. # return_tensors='pt' # 返回PyTorch张量
  378. return_tensors=None #不返回PyTorch张量
  379. ) for _t in text]
  380. text = [torch.LongTensor(np.array([_t['input_ids'] for _t in text])).to(device),
  381. torch.LongTensor(np.array([_t['attention_mask'] for _t in text])).to(device)]
  382. # predict
  383. with torch.no_grad():
  384. outputs = model(None, text)
  385. predic = torch.max(outputs.data, 1)[1].cpu().numpy()
  386. pred_prob = torch.max(outputs.data, 1)[0].cpu().numpy()
  387. # print('pred_prob',pred_prob)
  388. if pred_prob>0.5:
  389. pred_label = predic[0]
  390. pred_class = label2class_dict[pred_label]
  391. else:
  392. return
  393. # print('check rule before',pred_class)
  394. # check rule
  395. if pred_class==101 and re.search("((资格|资质)(审查|预审|后审|审核)|资审)结果(公告|公示)?|(资质|资格)(预审|后审)公示|资审及业绩公示",title): # 纠正部分‘资审结果’模型错误识别为中标
  396. pred_class = 105
  397. elif pred_class==122 and re.search("验收服务",title):
  398. pred_class = None
  399. # elif pred_class==118 and re.search("重新招标",title): #重新招标类公告,因之前公告的废标原因而错识别为废标公告
  400. # pred_class = 52
  401. return pred_class
  402. class_dict = {51: '公告变更',
  403. 52: '招标公告',
  404. 101: '中标信息',
  405. 102: '招标预告',
  406. 103: '招标答疑',
  407. 104: '招标文件',
  408. 105: '资审结果',
  409. 106: '法律法规',
  410. 107: '新闻资讯',
  411. 108: '拟建项目',
  412. 109: '展会推广',
  413. 110: '企业名录',
  414. 111: '企业资质',
  415. 112: '全国工程',
  416. 113: '业主采购',
  417. 114: '采购意向',
  418. 115: '拍卖出让',
  419. 116: '土地矿产',
  420. 117: '产权交易',
  421. 118: '废标公告',
  422. 119: '候选人公示',
  423. 120: '合同公告',
  424. 121: '开标记录',
  425. 122: '验收合同'
  426. }
  427. tenderee_type = ['公告变更','招标公告','招标预告','招标答疑','资审结果','采购意向']
  428. win_type = ['中标信息','废标公告','候选人公示','合同公告','开标记录','验收合同']
  429. def is_contain_winner(extract_json):
  430. if re.search('win_tenderer', extract_json):
  431. return True
  432. else:
  433. return False
  434. def merge_channel(list_articles,channel_dic,original_docchannel,page_time,prem={},web_source_no="",is_finnal=False):
  435. def merge_rule(title,text,docchannel,pred_channel,channel_dic,original_docchannel):
  436. front_text_len = len(text)//3 if len(text)>300 else 100
  437. front_text = text[:front_text_len]
  438. pred_channel = class_dict[pred_channel]
  439. # print('pred_channel',pred_channel,'docchannel',docchannel,'original_docchannel',original_docchannel)
  440. if pred_channel == docchannel:
  441. channel_dic['docchannel']['use_original_docchannel'] = 0
  442. else:
  443. if pred_channel in ['采购意向','招标预告'] and docchannel in ['采购意向','招标预告']:
  444. merge_res = '采购意向' if re.search("意向|意愿",title) or re.search("意向|意愿",front_text) else "招标预告"
  445. channel_dic['docchannel']['docchannel'] = merge_res
  446. channel_dic['docchannel']['use_original_docchannel'] = 0
  447. elif pred_channel in ['公告变更','招标答疑'] and docchannel in ['公告变更','招标答疑']:
  448. channel_dic['docchannel']['docchannel'] = docchannel
  449. channel_dic['docchannel']['use_original_docchannel'] = 0
  450. elif pred_channel=='公告变更' and docchannel in ['中标信息','废标公告','候选人公示','合同公告']: #中标类的变更还是中标类公告
  451. channel_dic['docchannel']['docchannel'] = docchannel
  452. channel_dic['docchannel']['use_original_docchannel'] = 0
  453. elif docchannel=='公告变更' and pred_channel in ['中标信息','废标公告','候选人公示','合同公告']:
  454. channel_dic['docchannel']['docchannel'] = pred_channel
  455. channel_dic['docchannel']['use_original_docchannel'] = 0
  456. elif docchannel in ['中标信息','候选人公示'] and pred_channel in ['中标信息','候选人公示']:
  457. if re.search('候选人(变更)?公[告示]|评标(结果)?(公[告示]|报告)|评审结果', title):
  458. channel_dic['docchannel']['docchannel'] = '候选人公示'
  459. channel_dic['docchannel']['use_original_docchannel'] = 0
  460. else:
  461. if original_docchannel in [101,119]:
  462. channel_dic['docchannel']['docchannel'] = class_dict.get(original_docchannel, '原始类别')
  463. channel_dic['docchannel']['use_original_docchannel'] = 1
  464. else:
  465. channel_dic['docchannel']['docchannel'] = pred_channel
  466. channel_dic['docchannel']['use_original_docchannel'] = 0
  467. elif docchannel in ['中标信息','候选人公示'] and pred_channel=='开标记录':
  468. re_text_len = max(500,len(text)//3)
  469. re_text = text[:re_text_len]
  470. if re.search('开标记录|截标信息|开标安排|开标数据表|开标信息|开标情况|开标一览表|开标结果',re_text):
  471. channel_dic['docchannel']['docchannel'] = '开标记录'
  472. channel_dic['docchannel']['use_original_docchannel'] = 0
  473. elif pred_channel=='招标答疑' and docchannel in ['招标预告','招标公告']: #招标答疑公告修正
  474. if re.search("招标答疑|澄清.?修改|答疑.?澄清|澄清.?答疑|(澄清|答疑)(公告|公示)",title+text[:100]):
  475. channel_dic['docchannel']['docchannel'] = '招标答疑'
  476. channel_dic['docchannel']['use_original_docchannel'] = 0
  477. elif docchannel=='招标答疑' and pred_channel in ['招标预告','招标公告']:
  478. if re.search("招标答疑|澄清.?修改|答疑.?澄清|澄清.?答疑|(澄清|答疑)(公告|公示)",title+text[:100]):
  479. channel_dic['docchannel']['docchannel'] = '招标答疑'
  480. channel_dic['docchannel']['use_original_docchannel'] = 0
  481. # 以上规则都没用上,使用和original_docchannel相同的预测结果
  482. if 'use_original_docchannel' not in channel_dic['docchannel']:
  483. if docchannel==class_dict.get(original_docchannel, '原始类别') or pred_channel==class_dict.get(original_docchannel, '原始类别'):
  484. # 其中一个结果和original_docchannel相同
  485. channel_dic['docchannel']['docchannel'] = class_dict.get(original_docchannel, '原始类别')
  486. channel_dic['docchannel']['use_original_docchannel'] = 0
  487. if 'use_original_docchannel' not in channel_dic['docchannel']:
  488. original_type = class_dict.get(original_docchannel, '原始类别')
  489. if pred_channel in tenderee_type and docchannel in tenderee_type and original_type not in tenderee_type:
  490. # pred_channel和docchannel都是同一(招标/中标)类型时,original_docchannel不一致时不使用原网类型
  491. channel_dic['docchannel']['use_original_docchannel'] = 0
  492. elif pred_channel in win_type and docchannel in win_type and original_type not in win_type:
  493. # pred_channel和docchannel都是同一(招标/中标)类型时,original_docchannel不一致时不使用原网类型
  494. channel_dic['docchannel']['use_original_docchannel'] = 0
  495. else:
  496. channel_dic = {'docchannel': {'doctype': '采招数据',
  497. 'docchannel': original_type,
  498. 'life_docchannel': original_type}}
  499. channel_dic['docchannel']['use_original_docchannel'] = 1
  500. return channel_dic
  501. article = list_articles[0]
  502. title = article.title
  503. text = article.content
  504. doctype = channel_dic['docchannel']['doctype']
  505. docchannel = channel_dic['docchannel']['docchannel']
  506. # print('doctype',doctype,'docchannel',docchannel,'original_docchannel',original_docchannel)
  507. compare_type = ['公告变更','招标公告','中标信息','招标预告','招标答疑','资审结果','采购意向','废标公告','候选人公示',
  508. '合同公告','开标记录','验收合同']
  509. prem_json = json.dumps(prem, ensure_ascii=False)
  510. contain_winner = is_contain_winner(prem_json)
  511. # 仅比较部分数据
  512. special_pattern = re.compile("单一来源|直接采购|单源直采|直采")
  513. if doctype=='采招数据' and docchannel in compare_type:
  514. if not re.search(special_pattern,title) and not re.search(special_pattern,text[:100]):
  515. pred = channel_predict(title, text)
  516. # print(text[:2000], '\n pred_res1', pred)
  517. if pred is not None and original_docchannel: # 无original_docchannel时不进行对比校正
  518. channel_dic = merge_rule(title,text,docchannel,pred,channel_dic,original_docchannel)
  519. elif doctype=='采招数据' and docchannel=="":
  520. if is_finnal:
  521. pred = channel_predict(title, text,min_text_len=100)
  522. else:
  523. pred = channel_predict(title, text)
  524. # print(text[:2000], '\n pred_res2', pred)
  525. if pred is not None:
  526. pred = class_dict[pred]
  527. channel_dic['docchannel']['docchannel'] = pred
  528. channel_dic['docchannel']['life_docchannel'] = pred
  529. channel_dic['docchannel']['use_original_docchannel'] = 0
  530. # print('channel_dic1',channel_dic)
  531. if is_finnal:
  532. # 公告分类后验规则
  533. if doctype=="采招数据":
  534. if contain_winner and original_docchannel in [101,119,120] and \
  535. channel_dic['docchannel']['docchannel'] in ['招标公告','招标预告','采购意向']:
  536. if class_dict.get(original_docchannel):
  537. channel_dic['docchannel']['docchannel'] = class_dict.get(original_docchannel)
  538. channel_dic['docchannel']['use_original_docchannel'] = 1
  539. elif contain_winner and docchannel in ['中标信息','候选人公示','合同公告'] and \
  540. channel_dic['docchannel']['docchannel'] in ['招标公告','招标预告','采购意向']:
  541. channel_dic['docchannel']['docchannel'] = docchannel
  542. channel_dic['docchannel']['use_original_docchannel'] = 0
  543. elif not contain_winner and ((channel_dic['docchannel'].get('docchannel')=='' and original_docchannel in [101,119]) or
  544. channel_dic['docchannel'].get('docchannel') in ['中标信息','候选人公示']) and re.search("开标.?记录|开标.?安排|开标.?时间|开标.?(地址|地点)|开标一览表|开标日程|评审专家公示|评委名单公示|开标数据表",title+text[:150]) and len(text)<=300:
  545. # 识别公告内容较短的开标安排(时间、地点),例:798134361,790800148
  546. channel_dic['docchannel']['docchannel'] = "开标记录"
  547. channel_dic['docchannel']['use_original_docchannel'] = 0
  548. elif channel_dic['docchannel']['docchannel'] in ['招标公告','招标预告'] and re.search("第[1-9一二三]次(变更|更正|更改|修改)|((变更|更正|更改|修改)(事项)?|延期|暂停)(招标|采购)?的?(公告|公示|通知)|变更$|更正$",title):
  549. # 对标题关键词明显的招标公告进行修正
  550. channel_dic['docchannel']['docchannel'] = "公告变更"
  551. channel_dic['docchannel']['use_original_docchannel'] = 0
  552. # elif not contain_winner and original_docchannel in [52] and channel_dic['docchannel'].get('docchannel') in ['招标公告'] and \
  553. # re.search("开标.?记录|开标.?安排|开标.?时间|开标.?(地址|地点)|开标一览表|开标日程|评审专家公示|评委名单公示|开标数据表",title+text[:150]) and len(text)<=300:
  554. # # 识别公告内容较短的开标安排(时间、地点),例:799590061
  555. # channel_dic['docchannel']['docchannel'] = "开标记录"
  556. # channel_dic['docchannel']['use_original_docchannel'] = 0
  557. # elif not contain_winner and original_docchannel in [52,102,114] and \
  558. # channel_dic['docchannel']['docchannel'] in ['中标信息','合同公告']:
  559. # channel_dic['docchannel']['docchannel'] = class_dict.get(original_docchannel)
  560. # channel_dic['docchannel']['use_original_docchannel'] = 0
  561. # elif not contain_winner and docchannel in ['招标公告','招标预告','采购意向'] and \
  562. # channel_dic['docchannel']['docchannel'] in ['中标信息','合同公告']:
  563. # channel_dic['docchannel']['docchannel'] = docchannel
  564. # '招标预告'类 规则纠正,规则排除部分站源
  565. if channel_dic['docchannel']['doctype']=='采招数据' and channel_dic['docchannel']['docchannel'] in ["招标公告","公告变更"] and web_source_no not in ['DX000027-1']:
  566. if "##attachment##" in text:
  567. main_text, attachment_text = text.split("##attachment##", maxsplit=1)
  568. else:
  569. main_text = text
  570. main_text = text_process(main_text)
  571. # if re.search("采购实施月份|采购月份|预计(招标|采购|发标|发包)(时间|月份)|招标公告预计发布时间",main_text[:max(500,len(main_text)//2)]):
  572. if re.search("采购实施月份|采购月份|(计划|预计|预期)(招标|采购|发标|发包)(时间|月份)|(招标公告|资格预审公告|招标公告[(\(]资格预审公告[)\)])预计发布时间|预计(招标公告|资格预审公告|招标公告[(\(]资格预审公告[)\)])发布时间",main_text):
  573. front_text_len = len(main_text) // 3 if len(main_text) > 300 else 100
  574. front_text = main_text[:front_text_len]
  575. if re.search("意向|意愿",title) or re.search("意向|意愿",front_text):
  576. channel_dic['docchannel']['docchannel'] = "采购意向"
  577. else:
  578. channel_dic['docchannel']['docchannel'] = "招标预告"
  579. channel_dic['docchannel']['use_original_docchannel'] = 0
  580. # '招标预告'类规则纠正,有开标/截标时间的改为'招标公告'
  581. if channel_dic['docchannel']['doctype']=='采招数据' and channel_dic['docchannel']['docchannel']=="招标预告":
  582. time_bidopen = prem.get("time_bidopen","")
  583. time_bidclose = prem.get("time_bidclose","")
  584. time_getFileEnd = prem.get("time_getFileEnd","")
  585. time_registrationEnd = prem.get("time_registrationEnd","")
  586. for _time in [time_bidopen,time_bidclose,time_getFileEnd,time_registrationEnd]:
  587. if _time and page_time and _time[:10]>=page_time:
  588. # print("change docchannel by time",_time,page_time)
  589. channel_dic['docchannel']['docchannel'] = "招标公告"
  590. channel_dic['docchannel']['use_original_docchannel'] = 0
  591. break
  592. # print("original channel_dic",channel_dic)
  593. # docchannel预测为空时,补充原网类别
  594. if channel_dic['docchannel'].get('docchannel') == '' and is_finnal is True:
  595. if channel_dic['docchannel'].get('doctype') in ("采招数据","拍卖出让","产权交易","土地矿产"):
  596. if channel_dic['docchannel'].get('life_docchannel') and channel_dic['docchannel']['life_docchannel'] not in ("拍卖出让","产权交易","土地矿产","原始类别"):
  597. channel_dic['docchannel']['docchannel'] = channel_dic['docchannel']['life_docchannel']
  598. channel_dic['docchannel']['use_original_docchannel'] = 1
  599. # print('channel_dic2',channel_dic)
  600. return channel_dic
  601. if __name__ == '__main__':
  602. title = '关于【2024年四好农村路大中村药红路、空坦路延伸段设计服务】无效项目的公示'
  603. text = '''关于【2024年四好农村路大中村药红路、空坦路延伸段设计服务】无效项目的公示 点击查看招标公告 关于【2024年四好农村路大中村药红路、空坦路延伸段设计服务】无效项目的公示 项目名称 2024年四好农村路大中村药红路、空坦路延伸段设计服务, 采购人 重庆市巴南区人民政府莲花街道办事处, 选取方式 直接选取, 是否重新发布招标公告 是 ,无效类型 项目取消, 无效原因 资质设置错误,附件已盖章上传 ,无效时间 2024-10-21 ,公示附件 大中村设计变更.jpg'''
  604. pred_class = channel_predict(title,text)
  605. print(pred_class)
  606. # pred_class2 = channel_predict(title,text)
  607. # print(pred_class2)
  608. pass