data_util.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323
  1. #!/usr/bin/python3
  2. # -*- coding: utf-8 -*-
  3. # @Author : bidikeji
  4. # @Time : 2021/1/13 0013 14:19
  5. # DEPRECATED(Phase1): 本文件中的 PostgreSQL 硬编码连接(host=192.168.*,
  6. # password=postgres 等)将在后续 training/ Phase 迁移到 BiddingKG.dl.infra.db。
  7. # 迁移完成前可临时使用:from BiddingKG.dl.infra.db import get_connection;
  8. # conn = get_connection("<dbname>")
  9. # 详见 ARCHITECTURE.md 第 11 章 Phase 1 与 REFACTOR_LOG.md。
  10. import re
  11. import os
  12. import math
  13. import json
  14. import random
  15. import numpy as np
  16. import pandas as pd
  17. from BiddingKG.dl.common.Utils import getVocabAndMatrix,getModel_word,viterbi_decode, load
  18. tag2id = {'S':0,'B-pro':1, 'I-pro':2, 'E-pro':3, 'B-rea':4, 'I-rea':5, 'E-rea':6}
  19. id_to_tag = {v:k for k,v in tag2id.items()}
  20. path1 = os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))+"/interface/codename_vocab.pk"
  21. path2 = os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))+"/interface/codename_w2v_matrix.pk"
  22. vocab = load(path1)
  23. matrix = load(path2)
  24. max_id = len(vocab)
  25. word2id = {k: v for v, k in enumerate(vocab)}
  26. def df2data(df):
  27. import pandas as pd
  28. import json
  29. datas = []
  30. for idx in df.index:
  31. docid = df.loc[idx, 'docid']
  32. text = df.loc[idx, 'text']
  33. # string = list(text)
  34. tags = [0]*len(text)
  35. labels = json.loads(df.loc[idx, 'label'])
  36. for label in labels:
  37. _, _, begin, end, _ = re.split('\s',label)
  38. begin = int(begin)
  39. end = int(end)
  40. if end-begin>=2:
  41. tags[begin]=1
  42. tags[end-1]=3
  43. for i in range(begin+1,end-1):
  44. tags[i]=2
  45. # datas.append([string, tags])
  46. text_sentence = []
  47. ids_sentence = []
  48. tag_sentence = []
  49. for i in range(len(text)):
  50. text_sentence.append(text[i])
  51. # ids_sentence.append(word2id.get(text[i], max_id))
  52. ids_sentence.append(word2id.get(text[i], word2id.get('<unk>')))
  53. tag_sentence.append(tags[i])
  54. if text[i] in ['。','!']:
  55. if text_sentence:
  56. # if len(text_sentence) > 100:
  57. if len(text_sentence)>5 and len(text_sentence)<1000:
  58. datas.append([text_sentence, ids_sentence,tag_sentence])
  59. else:
  60. print('单句小于5或大于1000,句子长度为:%d,文章ID:%s'%(len(text_sentence), docid))
  61. text_sentence = []
  62. ids_sentence = []
  63. tag_sentence = []
  64. if text_sentence:
  65. # if len(text_sentence) > 5:
  66. if len(text_sentence) > 5 and len(text_sentence) < 1000:
  67. datas.append([text_sentence, ids_sentence, tag_sentence])
  68. else:
  69. print('单句小于5或大于1000,句子长度为:%d,文章ID:%s' % (len(text_sentence), docid))
  70. return datas
  71. def find_kw_from_text(kw, s):
  72. '''
  73. 输入关键词及句子信息,返回句子中关键词的所有出现位置
  74. :param kw: 关键词
  75. :param s: 文本
  76. :return:
  77. '''
  78. begin = s.find(kw, 0)
  79. kws = []
  80. while begin!=-1:
  81. end = begin + len(kw)
  82. # print(s[begin:end])
  83. kws.append((begin, end))
  84. begin = s.find(kw, end)
  85. return kws
  86. def get_feature(text, lbs):
  87. '''
  88. 输入文章预处理后文本内容及产品名称列表,返回句子列表,数字化句子列表,数字化标签列表
  89. :param text: 文本内容
  90. :param lbs: 产品名称列表
  91. :return:
  92. '''
  93. lbs = sorted(set(lbs), key=lambda x: len(x), reverse=True)
  94. sentences = []
  95. ids_list = []
  96. tags_list = []
  97. for sentence in text.split('。'):
  98. if len(sentence) < 5:
  99. continue
  100. if len(sentence) > 1000:
  101. sentence = sentence[:1000]
  102. tags = [0] * len(sentence)
  103. # ids = [word2id.get(word, max_id) for word in sentence]
  104. ids = [word2id.get(word, word2id.get('<unk>')) for word in sentence]
  105. for lb in lbs:
  106. kw_indexs = find_kw_from_text(lb, sentence)
  107. for indexs in kw_indexs:
  108. b, e = indexs
  109. if tags[b] == 0 and tags[e - 1] == 0:
  110. tags[b] = 1
  111. tags[e - 1] = 3
  112. for i in range(b+1, e - 1):
  113. tags[i] = 2
  114. sentences.append(list(sentence))
  115. ids_list.append(ids)
  116. tags_list.append(tags)
  117. return sentences, ids_list, tags_list
  118. def dfsearchlb(df):
  119. datas = []
  120. for i in df.index:
  121. text = df.loc[i, 'text']
  122. lbs = json.loads(df.loc[i, 'lbset'])
  123. sentences, ids_list, tags_list = get_feature(text, lbs)
  124. for sen, ids, tags in zip(sentences, ids_list, tags_list):
  125. datas.append([sen, ids, tags])
  126. return datas
  127. def get_label_data():
  128. import psycopg2
  129. conn = psycopg2.connect(dbname='iepy_product', user='postgres', password='postgres', host='192.168.2.101')
  130. cursor = conn.cursor()
  131. sql = "select human_identifier, text from corpus_iedocument where edittime NOTNULL AND jump_signal=0 \
  132. and creation_date > to_timestamp('2021-01-14 00:00:00','yyyy-MM-dd HH24:mi:ss');"
  133. cursor.execute(sql)
  134. writer = open('label_data.txt', 'w', encoding='utf-8')
  135. datas = []
  136. for row in cursor.fetchall():
  137. docid = row[0]
  138. text = row[1]
  139. # string = list(text)
  140. tags = [0]*len(text)
  141. sql_lb = "select b.value from brat_bratannotation as b where document_id = '{}' and b.value like 'T%product%';".format(docid)
  142. cursor.execute(sql_lb)
  143. for row_lb in cursor.fetchall():
  144. label = row_lb[0]
  145. _, _, begin, end, _ = re.split('\s',label)
  146. begin = int(begin)
  147. end = int(end)
  148. if end-begin>=2:
  149. tags[begin]=1
  150. tags[end-1]=3
  151. for i in range(begin+1,end-1):
  152. tags[i]=2
  153. # datas.append([string, tags])
  154. text_sentence = []
  155. ids_sentence = []
  156. tag_sentence = []
  157. for i in range(len(text)):
  158. text_sentence.append(text[i])
  159. # ids_sentence.append(word2id.get(text[i], max_id))
  160. ids_sentence.append(word2id.get(text[i], word2id.get('<unk>')))
  161. tag_sentence.append(tags[i])
  162. writer.write("%s\t%s\n"%(text[i],tags[i]))
  163. if text[i] in ['。','?','!',';']:
  164. writer.write('\n')
  165. if text_sentence:
  166. if len(text_sentence) > 100:
  167. # if len(text_sentence)>5 and len(text_sentence)<1000:
  168. datas.append([text_sentence, ids_sentence,tag_sentence])
  169. elif len(text_sentence) > 5:
  170. continue
  171. else:
  172. print('单句小于5或大于100,句子长度为:%d,文章ID:%s'%(len(text_sentence), docid))
  173. text_sentence = []
  174. ids_sentence = []
  175. tag_sentence = []
  176. if text_sentence:
  177. if len(text_sentence) > 5:
  178. # if len(text_sentence) > 5 and len(text_sentence) < 1000:
  179. datas.append([text_sentence, ids_sentence, tag_sentence])
  180. else:
  181. print('单句小于5或大于100,句子长度为:%d,文章ID:%s' % (len(text_sentence), docid))
  182. writer.close()
  183. return datas
  184. def input_from_line(line):
  185. string = list(line)
  186. # ids = [word2id.get(k, max_id) for k in string]
  187. ids = [word2id.get(k, word2id.get('<unk>')) for k in string]
  188. tags = []
  189. return [[string], [ids], [tags]]
  190. def process_data(sentences):
  191. '''
  192. 字符串数字化并统一长度
  193. :param sentences: 文章分句字符串列表['招标公告','招标代理']
  194. :return: 数字化后的统一长度
  195. '''
  196. maxLen = max([len(sentence) for sentence in sentences])
  197. # tags = [[word2id.get(k, max_id) for k in sentence] for sentence in sentences]
  198. tags = [[word2id.get(k, word2id.get('<unk>')) for k in sentence] for sentence in sentences]
  199. pad_tags = [tag[:maxLen]+[0]*(maxLen-len(tag)) for tag in tags]
  200. return pad_tags
  201. def get_ner(BIE_tag):
  202. ner = set()
  203. for it in re.finditer('BI*E',BIE_tag):
  204. ner.add((it.start(),it.end()))
  205. return ner
  206. def decode(logits, lengths, matrix):
  207. paths = []
  208. small = -1000.0
  209. # start = np.asarray([[small]*4+[0]]) # 只有产品
  210. start = np.asarray([[small]*7+[0]]) # 产品及失败原因
  211. for score, length in zip(logits, lengths):
  212. score = score[:length]
  213. pad = small * np.ones([length, 1])
  214. logits = np.concatenate([score, pad], axis=1)
  215. logits = np.concatenate([start, logits], axis=0)
  216. path, _ = viterbi_decode(logits, matrix)
  217. paths.append(path[1:])
  218. return paths
  219. def result_to_json(line, tags):
  220. result = []
  221. ner = []
  222. tags = ''.join([str(it) for it in tags])
  223. for it in re.finditer("12*3", tags):
  224. start = it.start()
  225. end = it.end()
  226. ner.append([line[start:end], (start, end)])
  227. # for it in re.finditer("45*6", tags):
  228. # start = it.start()
  229. # end = it.end()
  230. # ner.append([line[start:end], (start, end)])
  231. result.append([line, ner])
  232. # print(tags)
  233. return result
  234. class BatchManager(object):
  235. def __init__(self, data, batch_size):
  236. self.batch_data = self.sort_and_pad(data, batch_size)
  237. self.len_data = len(self.batch_data)
  238. def sort_and_pad(self, data, batch_size):
  239. num_batch = int(math.ceil(len(data)/batch_size))
  240. sorted_data = sorted(data, key=lambda x:len(x[0]))
  241. print('最小句子长度:%d;最大句子长度:%d' % (len(sorted_data[0][0]), len(sorted_data[-1][0]))) # 临时增加打印句子长度
  242. batch_data = list()
  243. for i in range(num_batch):
  244. batch_data.append(self.pad_data(sorted_data[i*int(batch_size):(i+1)*int(batch_size)]))
  245. return batch_data
  246. @staticmethod
  247. def pad_data(data):
  248. strings = []
  249. chars = []
  250. targets = []
  251. max_length = max([len(sentence[0]) for sentence in data])
  252. for line in data:
  253. string, char, target = line
  254. padding = [0]*(max_length-len(string))
  255. strings.append(string + padding)
  256. chars.append(char + padding)
  257. targets.append(target + padding)
  258. return [strings, chars, targets]
  259. def iter_batch(self, shuffle=False):
  260. if shuffle:
  261. random.shuffle(self.batch_data)
  262. for idx in range(self.len_data):
  263. yield self.batch_data[idx]
  264. def 获取原始标注数据():
  265. import psycopg2
  266. import json
  267. conn = psycopg2.connect(dbname='iepy_product', user='postgres', password='postgres', host='192.168.2.103')
  268. cursor = conn.cursor()
  269. sql = "select human_identifier, text from corpus_iedocument where edittime NOTNULL AND jump_signal=0 ;"
  270. cursor.execute(sql)
  271. writer = open('label_data.txt', 'w', encoding='utf-8')
  272. datas = []
  273. for row in cursor.fetchall():
  274. docid = row[0]
  275. text = row[1]
  276. sql_lb = "select b.value from brat_bratannotation as b where document_id = '{}' and b.value like 'T%product%';".format(docid)
  277. cursor.execute(sql_lb)
  278. rows = cursor.fetchall()
  279. print('len(rows)', len(rows))
  280. datas.append((docid, text, json.dumps(rows, ensure_ascii=False), len(rows)))
  281. df = pd.DataFrame(datas, columns=['docid', 'text', 'rows', 'product_num'])
  282. df.to_excel('data/产品数据自己人标注的原始数据.xlsx')
  283. if __name__=="__main__":
  284. # import os
  285. import pickle
  286. # with open('data/dev_data2.pkl', 'rb') as f:
  287. # dev_data = pickle.load(f)
  288. # print(len(dev_data))
  289. # print(os.path.exists('data/testdata.xlsx'))
  290. # df = pd.read_excel('data/testdata.xlsx')
  291. # print(len(df))
  292. # data_test = df2data(df)
  293. # print(len(data_test), len(data_test[0][0]))
  294. # 获取原始标注数据()
  295. df = pd.read_excel('data/产品数据自己人标注的原始数据.xlsx')
  296. with open('data/dev_data2.pkl', 'rb') as f:
  297. dev_data = pickle.load(f)
  298. print(len(set(df['docid'])))
  299. print('')