Views
No views yet
1import torch
2import jieba
3import numpy as np
4from classifier import BertForMaskClassification
5from transformers import AutoTokenizer, AutoConfig, BertForTokenClassification
6
7label_list = ["O","COMMA","PERIOD","COLON"]
8
9label2punct = {
10 "COMMA": ",",
11 "PERIOD": "。",
12 "COLON":":",
13}
14
15model_name_or_path = "pmp-h256"
16
17tokenizer = AutoTokenizer.from_pretrained(model_name_or_path)
18model = BertForMaskClassification.from_pretrained(model_name_or_path)
19device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
20
21def punct(text):
22
23 tokenize_words = jieba.lcut(''.join(text))
24 mask_tokens = []
25 for word in tokenize_words:
26 mask_tokens.extend(word)
27 mask_tokens.append("[MASK]")
28 tokenized_inputs = tokenizer(mask_tokens,is_split_into_words=True, return_tensors="pt")
29 with torch.no_grad():
30 logits = model(**tokenized_inputs).logits
31 predictions = logits.argmax(-1).tolist()
32 predictions = predictions[0]
33 tokens = tokenizer.convert_ids_to_tokens(tokenized_inputs["input_ids"][0])
34
35 result =[]
36 print(tokens)
37 print(predictions)
38 for token, prediction in zip(tokens, predictions):
39 if token =="[CLS]" or token =="[SEP]":
40 continue
41 if token == "[MASK]":
42 label = label_list[prediction]
43 if label != "O":
44 punct = label2punct[label]
45 result.append(punct)
46 else:
47 result.append(token)
48
49 return "".join(result)
50
51text = '肝浊音界正常肝上界位于锁骨中线第五肋间移动浊音阴性肾区无叩痛'
52print(punct(text))
53
54# 肝浊音界正常,肝上界位于锁骨中线第五肋间,移动浊音阴性,肾区无叩痛。