Skip to content

Commit a4f76bf

Browse files
committed
add pos encoding
1 parent 95dbb5c commit a4f76bf

4 files changed

Lines changed: 198 additions & 52 deletions

File tree

modules/data/bert_data.py

Lines changed: 74 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
import numpy as np
77
from tqdm._tqdm_notebook import tqdm_notebook
88
from tqdm import tqdm
9+
import json
910

1011

1112
class InputFeatures(object):
@@ -16,11 +17,13 @@ def __init__(
1617
# Bert data
1718
bert_tokens, input_ids, input_mask, input_type_ids,
1819
# Origin data
19-
tokens, labels, labels_ids, labels_mask, tok_map, cls=None, cls_idx=None):
20+
tokens, labels, labels_ids, labels_mask, tok_map, cls=None, cls_idx=None, meta=None):
2021
"""
2122
Data has the following structure.
2223
data[0]: list, tokens ids
2324
data[1]: list, tokens mask
25+
data[2]: list, tokens type ids (for bert)
26+
data[3]: list, tokens meta info (if meta is not None)
2427
...
2528
data[-2]: list, labels mask
2629
data[-1]: list, labels ids
@@ -34,6 +37,10 @@ def __init__(
3437
self.data.append(input_mask)
3538
self.input_type_ids = input_type_ids
3639
self.data.append(input_type_ids)
40+
# Meta data
41+
self.meta = meta
42+
if meta is not None:
43+
self.data.append(meta)
3744
# Origin data
3845
self.tokens = tokens
3946
self.labels = labels
@@ -63,21 +70,25 @@ def __init__(self, data_set, shuffle, cuda, **kwargs):
6370

6471
def collate_fn(self, data):
6572
res = []
66-
token_ml = max(map(lambda x: sum(x.data[1]), data))
67-
label_ml = max(map(lambda x: sum(x.data[-2]), data))
68-
sorted_idx = np.argsort(list(map(lambda x: sum(x.data[1]), data)))[::-1]
73+
token_ml = max(map(lambda x_: sum(x_.data[1]), data))
74+
label_ml = max(map(lambda x_: sum(x_.data[-2]), data))
75+
sorted_idx = np.argsort(list(map(lambda x_: sum(x_.data[1]), data)))[::-1]
6976
for idx in sorted_idx:
7077
f = data[idx]
7178
example = []
72-
for x in f.data[:-2]:
79+
for idx_, x in enumerate(f.data[:-2]):
7380
if isinstance(x, list):
7481
x = x[:token_ml]
7582
example.append(x)
7683
example.append(f.data[-2][:label_ml])
7784
example.append(f.data[-1][:label_ml])
7885
res.append(example)
79-
res = list(zip(*res))
80-
res = [torch.LongTensor(x) for x in res]
86+
res = []
87+
for idx, x in enumerate(zip(*res)):
88+
if data[0].meta is not None and idx == 3:
89+
res.append(torch.FloatTensor(x))
90+
else:
91+
res.append(torch.LongTensor(x))
8192
if self.cuda:
8293
res = [t.cuda() for t in res]
8394
return res
@@ -95,9 +106,9 @@ def __init__(self, data_set, cuda, **kwargs):
95106

96107
def collate_fn(self, data):
97108
res = []
98-
token_ml = max(map(lambda x: sum(x.data[1]), data))
99-
label_ml = max(map(lambda x: sum(x.data[-2]), data))
100-
sorted_idx = np.argsort(list(map(lambda x: sum(x.data[1]), data)))[::-1]
109+
token_ml = max(map(lambda x_: sum(x_.data[1]), data))
110+
label_ml = max(map(lambda x_: sum(x_.data[-2]), data))
111+
sorted_idx = np.argsort(list(map(lambda x_: sum(x_.data[1]), data)))[::-1]
101112
for idx in sorted_idx:
102113
f = data[idx]
103114
example = []
@@ -108,35 +119,55 @@ def collate_fn(self, data):
108119
example.append(f.data[-2][:label_ml])
109120
example.append(f.data[-1][:label_ml])
110121
res.append(example)
111-
res = list(zip(*res))
112-
res = [torch.LongTensor(x) for x in res]
122+
res = []
123+
for idx, x in enumerate(zip(*res)):
124+
if data[0].meta is not None and idx == 3:
125+
res.append(torch.FloatTensor(x))
126+
else:
127+
res.append(torch.LongTensor(x))
113128
sorted_idx = torch.LongTensor(list(sorted_idx))
114129
if self.cuda:
115130
res = [t.cuda() for t in res]
116131
sorted_idx = sorted_idx.cuda()
117132
return res, sorted_idx
118133

119134

120-
def get_data(df, tokenizer, label2idx=None, max_seq_len=424, pad="<pad>", cls2idx=None, is_cls=False):
135+
def get_data(
136+
df, tokenizer, label2idx=None, max_seq_len=424, pad="<pad>", cls2idx=None,
137+
is_cls=False, is_meta=False):
121138
if label2idx is None:
122139
label2idx = {pad: 0, '[CLS]': 1, '[SEP]': 2}
123140
features = []
141+
all_args = []
124142
if is_cls:
125143
# Use joint model
126144
if cls2idx is None:
127145
cls2idx = dict()
128-
zip_args = zip(df["1"].tolist(), df["0"].tolist(), df["2"].tolist())
146+
all_args.extend([df["1"].tolist(), df["0"].tolist(), df["2"].tolist()])
129147
else:
130-
zip_args = zip(df["1"].tolist(), df["0"].tolist())
148+
all_args.extend([df["1"].tolist(), df["0"].tolist()])
149+
if is_meta:
150+
all_args.append(df["3"].tolist())
131151
total = len(df["0"].tolist())
132152
cls = None
133-
134-
for args in tqdm_notebook(enumerate(zip_args), total=total, leave=False):
153+
meta = None
154+
for args in tqdm_notebook(enumerate(zip(*all_args)), total=total, leave=False):
135155
if is_cls:
136-
idx, (text, labels, cls) = args
156+
if is_meta:
157+
idx, (text, labels, cls, meta) = args
158+
else:
159+
idx, (text, labels, cls) = args
137160
else:
138-
idx, (text, labels) = args
161+
if is_meta:
162+
idx, (text, labels, meta) = args
163+
else:
164+
idx, (text, labels) = args
165+
139166
tok_map = []
167+
meta_tokens = []
168+
if is_meta:
169+
meta = json.loads(meta)
170+
meta_tokens.append([0] * len(meta[0]))
140171
bert_tokens = []
141172
bert_labels = []
142173
bert_tokens.append("[CLS]")
@@ -147,7 +178,7 @@ def get_data(df, tokenizer, label2idx=None, max_seq_len=424, pad="<pad>", cls2id
147178
pad_idx = label2idx[pad]
148179
assert len(orig_tokens) == len(labels)
149180
prev_label = ""
150-
for orig_token, label in zip(orig_tokens, labels):
181+
for idx_, (orig_token, label) in enumerate(zip(orig_tokens, labels)):
151182
prefix = "B_"
152183
if label != "O":
153184
label = label.split("_")[1]
@@ -161,12 +192,15 @@ def get_data(df, tokenizer, label2idx=None, max_seq_len=424, pad="<pad>", cls2id
161192
if max_seq_len - 1 < len(bert_tokens) + len(cur_tokens):
162193
break
163194

195+
if is_meta:
196+
meta_tokens.extend(meta[idx_] * len(cur_tokens))
164197
bert_tokens.extend(cur_tokens)
165198
bert_label = [prefix + label] + ["I_" + label] * (len(cur_tokens) - 1)
166199
bert_labels.extend(bert_label)
167200
bert_tokens.append("[SEP]")
168201
bert_labels.append("[SEP]")
169-
202+
if is_meta:
203+
meta_tokens.append([0] * len(meta[0]))
170204
orig_tokens = ["[CLS]"] + orig_tokens + ["[SEP]"]
171205

172206
input_ids = tokenizer.convert_tokens_to_ids(bert_tokens)
@@ -187,6 +221,8 @@ def get_data(df, tokenizer, label2idx=None, max_seq_len=424, pad="<pad>", cls2id
187221
labels_ids.append(pad_idx)
188222
labels_mask.append(0)
189223
tok_map.append(-1)
224+
if is_meta:
225+
meta_tokens.append([0] * len(meta[0]))
190226
# assert len(input_ids) == len(bert_labels_ids)
191227
input_type_ids = [0] * len(input_ids)
192228
# For joint model
@@ -195,7 +231,8 @@ def get_data(df, tokenizer, label2idx=None, max_seq_len=424, pad="<pad>", cls2id
195231
if cls not in cls2idx:
196232
cls2idx[cls] = len(cls2idx)
197233
cls_idx = cls2idx[cls]
198-
234+
if is_meta:
235+
meta = meta_tokens
199236
features.append(InputFeatures(
200237
# Bert data
201238
bert_tokens=bert_tokens,
@@ -210,7 +247,9 @@ def get_data(df, tokenizer, label2idx=None, max_seq_len=424, pad="<pad>", cls2id
210247
tok_map=tok_map,
211248
# Joint data
212249
cls=cls,
213-
cls_idx=cls_idx
250+
cls_idx=cls_idx,
251+
# Meta data
252+
meta=meta
214253
))
215254
assert len(input_ids) == len(input_mask)
216255
assert len(input_ids) == len(input_type_ids)
@@ -222,20 +261,21 @@ def get_data(df, tokenizer, label2idx=None, max_seq_len=424, pad="<pad>", cls2id
222261
return features, label2idx
223262

224263

225-
def get_bert_data_loaders(train, valid, vocab_file, batch_size=16, cuda=True, is_cls=False, do_lower_case=False, max_seq_len=424):
264+
def get_bert_data_loaders(train, valid, vocab_file, batch_size=16, cuda=True, is_cls=False,
265+
do_lower_case=False, max_seq_len=424, is_meta=False):
226266
train = pd.read_csv(train)
227267
valid = pd.read_csv(valid)
228268

229269
cls2idx = None
230270

231271
tokenizer = tokenization.FullTokenizer(vocab_file=vocab_file, do_lower_case=do_lower_case)
232-
train_f, label2idx = get_data(train, tokenizer, is_cls=is_cls, max_seq_len=max_seq_len)
272+
train_f, label2idx = get_data(train, tokenizer, is_cls=is_cls, max_seq_len=max_seq_len, is_meta=is_meta)
233273
if is_cls:
234274
label2idx, cls2idx = label2idx
235275
train_dl = DataLoaderForTrain(
236276
train_f, batch_size=batch_size, shuffle=True, cuda=cuda)
237277
valid_f, label2idx = get_data(
238-
valid, tokenizer, label2idx, cls2idx=cls2idx, is_cls=is_cls, max_seq_len=max_seq_len)
278+
valid, tokenizer, label2idx, cls2idx=cls2idx, is_cls=is_cls, max_seq_len=max_seq_len, is_meta=is_meta)
239279
if is_cls:
240280
label2idx, cls2idx = label2idx
241281
valid_dl = DataLoaderForTrain(
@@ -248,8 +288,9 @@ def get_bert_data_loaders(train, valid, vocab_file, batch_size=16, cuda=True, is
248288
def get_bert_data_loader_for_predict(path, learner):
249289
df = pd.read_csv(path)
250290
f, _ = get_data(df, tokenizer=learner.data.tokenizer,
251-
label2idx=learner.data.label2idx, cls2idx=learner.data.cls2idx, is_cls=learner.data.is_cls,
252-
max_seq_len=learner.data.max_seq_len)
291+
label2idx=learner.data.label2idx, cls2idx=learner.data.cls2idx,
292+
is_cls=learner.data.is_cls,
293+
max_seq_len=learner.data.max_seq_len, is_meta=learner.data.is_meta)
253294
dl = DataLoaderForPredict(
254295
f, batch_size=learner.data.batch_size, shuffle=False,
255296
cuda=True)
@@ -260,13 +301,14 @@ def get_bert_data_loader_for_predict(path, learner):
260301
class BertNerData(object):
261302

262303
def __init__(self, train_dl, valid_dl, tokenizer, label2idx, max_seq_len=424,
263-
cls2idx=None, batch_size=16, cuda=True):
304+
cls2idx=None, batch_size=16, cuda=True, is_meta=False):
264305
self.train_dl = train_dl
265306
self.valid_dl = valid_dl
266307
self.tokenizer = tokenizer
267308
self.label2idx = label2idx
268309
self.cls2idx = cls2idx
269310
self.batch_size = batch_size
311+
self.is_meta = is_meta
270312
self.cuda = cuda
271313
self.id2label = sorted(label2idx.keys(), key=lambda x: label2idx[x])
272314
self.is_cls = False
@@ -277,7 +319,8 @@ def __init__(self, train_dl, valid_dl, tokenizer, label2idx, max_seq_len=424,
277319

278320
@classmethod
279321
def create(cls,
280-
train_path, valid_path, vocab_file, batch_size=16, cuda=True, is_cls=False, data_type="bert_cased", max_seq_len=424):
322+
train_path, valid_path, vocab_file, batch_size=16, cuda=True, is_cls=False,
323+
data_type="bert_cased", max_seq_len=424, is_meta=False):
281324
if ipython_info():
282325
global tqdm_notebook
283326
tqdm_notebook = tqdm
@@ -290,5 +333,5 @@ def create(cls,
290333
else:
291334
raise NotImplementedError("No requested mode :(.")
292335
return cls(*fn(
293-
train_path, valid_path, vocab_file, batch_size, cuda, is_cls, do_lower_case, max_seq_len),
336+
train_path, valid_path, vocab_file, batch_size, cuda, is_cls, do_lower_case, max_seq_len, is_meta),
294337
batch_size=batch_size, cuda=cuda)

modules/layers/encoders.py

Lines changed: 54 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
1-
import torch
21
from torch import nn
2+
import torch
33

44

55
class BertBiLSTMEncoder(nn.Module):
@@ -102,3 +102,56 @@ def create(cls, embeddings, hidden_dim=128, rnn_layers=1, use_cuda=True):
102102
model = cls(
103103
embeddings=embeddings, hidden_dim=hidden_dim, rnn_layers=rnn_layers, use_cuda=use_cuda)
104104
return model
105+
106+
107+
class BertMetaBiLSTMEncoder(nn.Module):
108+
109+
def __init__(self, embeddings, meta_dim,
110+
hidden_dim=128, rnn_layers=1, use_cuda=True):
111+
super(BertMetaBiLSTMEncoder, self).__init__()
112+
self.embeddings = embeddings
113+
self.hidden_dim = hidden_dim
114+
self.rnn_layers = rnn_layers
115+
self.use_cuda = use_cuda
116+
self.lstm = nn.LSTM(
117+
self.embeddings.embedding_dim, hidden_dim // 2,
118+
rnn_layers, batch_first=True, bidirectional=True)
119+
self.hidden = None
120+
if use_cuda:
121+
self.cuda()
122+
self.init_weights()
123+
self.meta_dim = meta_dim
124+
self.output_dim = hidden_dim + meta_dim
125+
126+
def init_weights(self):
127+
# for p in self.lstm.parameters():
128+
# nn.init.xavier_normal(p)
129+
pass
130+
131+
def forward(self, batch):
132+
input, input_mask = batch[0], batch[1]
133+
output = torch.cat((self.embeddings(*batch), batch[3]), axis=-1)
134+
# output = self.dropout(output)
135+
lens = input_mask.sum(-1)
136+
output = nn.utils.rnn.pack_padded_sequence(
137+
output, lens.tolist(), batch_first=True)
138+
output, self.hidden = self.lstm(output)
139+
output, _ = nn.utils.rnn.pad_packed_sequence(output, batch_first=True)
140+
return output, self.hidden
141+
142+
def get_n_trainable_params(self):
143+
pp = 0
144+
for p in list(self.parameters()):
145+
if p.requires_grad:
146+
num = 1
147+
for s in list(p.size()):
148+
num = num * s
149+
pp += num
150+
return pp
151+
152+
@classmethod
153+
def create(cls, embeddings, meta_dim, hidden_dim=128, rnn_layers=1, use_cuda=True):
154+
model = cls(
155+
embeddings=embeddings, meta_dim=meta_dim,
156+
idden_dim=hidden_dim, rnn_layers=rnn_layers, use_cuda=use_cuda)
157+
return model

0 commit comments

Comments
 (0)