Views
No views yet
1# pip install torch pytorch-lightning
2pip install zhpr1from zhpr.predict import DocumentDataset,merge_stride,decode_pred
2from transformers import AutoModelForTokenClassification,AutoTokenizer
3from torch.utils.data import DataLoader
4
5def predict_step(batch,model,tokenizer):
6 batch_out = []
7 batch_input_ids = batch
8
9 encodings = {'input_ids': batch_input_ids}
10 output = model(**encodings)
11
12 predicted_token_class_id_batch = output['logits'].argmax(-1)
13 for predicted_token_class_ids, input_ids in zip(predicted_token_class_id_batch, batch_input_ids):
14 out=[]
15 tokens = tokenizer.convert_ids_to_tokens(input_ids)
16
17 # compute the pad start in input_ids
18 # and also truncate the predict
19 # print(tokenizer.decode(batch_input_ids))
20 input_ids = input_ids.tolist()
21 try:
22 input_id_pad_start = input_ids.index(tokenizer.pad_token_id)
23 except:
24 input_id_pad_start = len(input_ids)
25 input_ids = input_ids[:input_id_pad_start]
26 tokens = tokens[:input_id_pad_start]
27
28 # predicted_token_class_ids
29 predicted_tokens_classes = [model.config.id2label[t.item()] for t in predicted_token_class_ids]
30 predicted_tokens_classes = predicted_tokens_classes[:input_id_pad_start]
31
32 for token,ner in zip(tokens,predicted_tokens_classes):
33 out.append((token,ner))
34 batch_out.append(out)
35 return batch_out
36
37if __name__ == "__main__":
38 window_size = 256
39 step = 200
40 text = "維基百科是維基媒體基金會運營的一個多語言的百科全書目前是全球網路上最大且最受大眾歡迎的參考工具書名列全球二十大最受歡迎的網站特點是自由內容自由編輯與自由著作權"
41 dataset = DocumentDataset(text,window_size=window_size,step=step)
42 dataloader = DataLoader(dataset=dataset,shuffle=False,batch_size=5)
43
44 model_name = 'p208p2002/zh-wiki-punctuation-restore'
45 model = AutoModelForTokenClassification.from_pretrained(model_name)
46 tokenizer = AutoTokenizer.from_pretrained(model_name)
47
48 model_pred_out = []
49 for batch in dataloader:
50 batch_out = predict_step(batch,model,tokenizer)
51 for out in batch_out:
52 model_pred_out.append(out)
53
54 merge_pred_result = merge_stride(model_pred_out,step)
55 merge_pred_result_deocde = decode_pred(merge_pred_result)
56 merge_pred_result_deocde = ''.join(merge_pred_result_deocde)
57 print(merge_pred_result_deocde)維基百科是維基媒體基金會運營的一個多語言的百科全書,目前是全球網路上最大且最受大眾歡迎的參考工具書,名列全球二十大最受歡迎的網站,特點是自由內容、自由編輯與自由著作權。