本文共 1425 字,大约阅读时间需要 4 分钟。
import kashgarifrom data_tools import clean_dataimport reimport tqdmdef cut_text(text, length): textArr = re.findall(f'.{length}', text) remainder = text[len(textArr) * length:] if remainder: textArr.append(remainder) return textArrdef clean_data(source_file, target_file, ner_model): data = DataReader().read_conll_format_file(source_file) for i, (text, labels) in enumerate(data): if len(text) <= 100: prediction = ner_model.predict([text]) ner = prediction[0] else: chunks = cut_text(''.join(text), 100) predictions = [] for chunk in chunks: prediction = ner_model.predict([chunk]) predictions.extend(prediction[0]) ner = predictions for j, (token, label) in enumerate(zip(text, labels)): if ner[j].startswith('B') or ner[j].startswith('I'): if labels[j] == 'O': labels[j] = ner[j] with open(target_file, 'a', encoding='utf-8') as f: f.write('\n'.join( [f"{token} {label}" for token, label in zip(text, labels)] )) data, _ = DataReader().read_conll_format_file(source_file) print(f"数据清洗完成:{data == original_data}, 数据标签长度一致") import kashgarifrom data_tools import clean_datatime_ner = kashgari.utils.load_model('time_ner.h5')clean_data('./data/example.dev', 'example.dev', time_ner) 转载地址:http://kjofk.baihongyu.com/