66import numpy as np
77from tqdm ._tqdm_notebook import tqdm_notebook
88from tqdm import tqdm
9+ import json
910
1011
1112class 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
248288def 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):
260301class 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 )
0 commit comments