-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathselect_data.py
More file actions
64 lines (49 loc) · 2 KB
/
Copy pathselect_data.py
File metadata and controls
64 lines (49 loc) · 2 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
from tqdm.notebook import tqdm
from flair.datasets import ColumnCorpus
from flair.data import Sentence
from flair.embeddings import TransformerWordEmbeddings
from flair.models import SequenceTagger
from flair.trainers import ModelTrainer
from collections import defaultdict
import argparse
import pdb
parser = argparse.ArgumentParser(description="Main Program.")
parser.add_argument('--teacher_path', type=str, required=True, help="teacher model path")
parser.add_argument('--unlabeled_path', type=str, required=True, help="unlabeled data path")
parser.add_argument('--save_path', type=str, required=True, help="save path")
parser.add_argument('--number', type=int, default=50000, required=True, help="how many samples do you want?")
args = parser.parse_args()
tagger = SequenceTagger.load(f'{args.teacher_path}')
path = args.unlabeled_path
data = []
with open(path, 'r') as f:
for idx, line in tqdm(enumerate(f)):
if len(line.split(' ')) < 10: continue
sentence = Sentence(line)
tagger.predict(sentence)
if len(sentence.get_labels('ner')) < 2:
continue
data.append(sentence)
if data and len(data) % (args.number * 1.5 // 100) == 0:
partion = len(data) / (args.number * 1.5) * 100
print(f'{partion}% sentences have been tagged.')
if data and len(data) % (args.number * 1.5) == 0:
break
ief = defaultdict(int)
for sent in data:
for entity in sent.to_dict('ner')['entities']:
ief[entity['labels'][0].value] += 1
ief = {k: len(data)/v for k, v in ief.items()}
for i in range(len(data)):
sent = data[i]
ef = 0
for entity in sent.to_dict('ner')['entities']:
ef += ief[entity['labels'][0].value]
data[i] = (sent, ef / len(sent))
data.sort(key=lambda a: -a[1])
with open(args.save_path, 'w') as fw:
for sentence in data[:args.number]:
for i in sentence[0]:
text, tag = i.text, i.get_tag('ner').value
fw.write(f'{text} {tag}\n')
fw.write('\n')