-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathE_conll.py
More file actions
146 lines (131 loc) · 5.32 KB
/
Copy pathE_conll.py
File metadata and controls
146 lines (131 loc) · 5.32 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
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
# encoding:utf-8
'''
@Author: catnlp
@Email: wk_nlp@163.com
@Time: 2018/5/2 14:14
'''
from E_util.E_config import Config
from E_util.E_helpers import *
import os
import sys
import argparse
import random
import torch
import numpy as np
os.environ["CUDA_VISIBLE_DEVICES"] = '2'
seed_num = 100
random.seed(seed_num)
torch.manual_seed(seed_num)
np.random.seed(seed_num)
if __name__ == '__main__':
parser = argparse.ArgumentParser(description='Tuning with NER')
parser.add_argument('--wordemb', help='Embedding for words', default='glove')
parser.add_argument('--charemb', help='Embedding for chars', default='None')
parser.add_argument('--status', choices=['train', 'test', 'decode'], help='update algorithm', default='train')
parser.add_argument('--savemodel', default='E_model/E_group/conll/') # catnlp
parser.add_argument('--savedset', help='Dir of saved data setting', default='E_model/E_group/conll/train_bmes1.dset') # catnlp
parser.add_argument('--dataset', help='Dir of datset', default='data/E_group/conll/data_conll.txt')
parser.add_argument('--train', default='train.tsv')
parser.add_argument('--devel', default='devel.tsv')
parser.add_argument('--test', default='test.tsv')
parser.add_argument('--gpu', default='True')
parser.add_argument('--seg', default='True')
parser.add_argument('--extendalphabet', default='True')
# parser.add_argument('--raw', default='data/E_group/conll2003/test.bmes')
# parser.add_argument('--loadmodel', default='E_model/conll2003/train_bmes/conll2003_train_bmes')
# parser.add_argument('--output', default='data/E_group/conll2003/decode_train_bmes.txt')
args = parser.parse_args()
# train_file = args.train
# dev_file = args.dev
# test_file = args.test
# raw_file = args.raw
# model_dir = args.loadmodel
dset_dir = args.savedset
# output_file = args.output
if args.seg.lower() == 'true':
seg = True
else:
seg = False
status = args.status.lower()
save_model_dir = args.savemodel
if args.gpu.lower() == 'false':
gpu = False
else:
gpu = torch.cuda.is_available()
print('Seed num: ', seed_num)
print('GPU available: ', gpu)
print('Status: ', status)
print('Seg: ', seg)
# print('Raw file: ', raw_file)
if status == 'train':
print('Model saved to: ', save_model_dir)
sys.stdout.flush()
if status == 'train':
emb = args.wordemb.lower()
print('Word Embedding: ', emb)
if emb == 'glove':
emb_file = 'data/embedding/glove.6B.100d.txt'
else:
emb_file = None
char_emb_file = args.charemb.lower()
print('Char Embedding: ', char_emb_file)
name = 'BaseLSTM' # catnlp
config = Config()
config.optim = 'SGD'
config.lr = 0.015
config.iteration = 200
config.hidden_dim = 200
# config.clip = True
config.number_normalized = True
config.gpu = gpu
config.word_features = name
print('Word features: ', config.word_features)
count = 0
with open(args.dataset, 'r') as f:
for line in f:
print(line)
count += 1
line = line.strip()
if not line:
break
train_file = os.path.join(line, args.train)
devel_file = os.path.join(line, args.devel)
test_file = os.path.join(line, args.test)
print('Train file: ', train_file)
print('Dev file: ', devel_file)
print('Test file: ', test_file)
data_initialization(config, train_file, devel_file, test_file)
# config.fix_alphabet()
config.num_corpus = count
with open(args.dataset, 'r') as f:
for line in f:
line = line.strip()
train_file = os.path.join(line, args.train)
devel_file = os.path.join(line, args.devel)
test_file = os.path.join(line, args.test)
config.generate_instance(train_file, 'train')
config.generate_instance(devel_file, 'dev')
config.generate_instance(test_file, 'test')
if emb_file:
print('load word emb file...norm: ', config.norm_word_emb)
config.build_word_pretain_emb(emb_file)
if char_emb_file != 'none':
print('load char emb file...norm: ', config.norm_char_emb)
config.build_char_pretrain_emb(char_emb_file)
for label in config.label_alphabet.instances:
print(label)
name = 'E_conll-kaggle-2003_bio'
train_model(config, name, dset_dir, save_model_dir, seg)
# elif status == 'test':
# data = load_data_setting(dset_dir)
# data.generate_instance(dev_file, 'dev')
# load_model_decode(model_dir, data, 'dev', gpu, seg)
# data.generate_instance(test_file, 'test')
# load_model_decode(model_dir, data, 'test', gpu, seg)
# elif status == 'decode':
# data = load_data_setting(dset_dir)
# data.generate_instance(raw_file, 'raw')
# decode_results = load_model_decode(model_dir, data, 'raw', gpu, seg)
# data.write_decoded_results(output_file, decode_results, 'raw')
else:
print('Invalid argument! Please use valid arguments! (train/test/decode)')