Skip to content

Commit 77d7e79

Browse files
authored
batch-level decoding in DuIE for memory issues (PaddlePaddle#253)
* fix duie out of memory when prediction * batch-level decoding
1 parent 76065a9 commit 77d7e79

4 files changed

Lines changed: 34 additions & 38 deletions

File tree

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
{"predicate": ["empty", "empty", "注册资本", "作者", "所属专辑", "歌手", "邮政编码", "主演", "上映时间", "上映时间", "饰演", "饰演", "国籍", "成立日期", "毕业院校", "作曲", "作词", "编剧", "导演", "面积", "占地面积", "总部地点", "制片人", "嘉宾", "简称", "主持人", "获奖", "获奖", "获奖", "获奖", "海拔", "出品公司", "配音", "配音", "所在城市", "号", "主角", "创始人", "父亲", "祖籍", "母亲", "朝代", "董事长", "人口数量", "妻子", "丈夫", "票房", "票房", "专业代码", "气候", "修业年限", "改编自", "官方语言", "首都", "主题曲", "校长", "代言人"], "subject_type": ["empty", "empty", "企业", "图书作品", "歌曲", "歌曲", "行政区", "影视作品", "影视作品", "影视作品", "娱乐人物", "娱乐人物", "人物", "机构", "人物", "歌曲", "歌曲", "影视作品", "影视作品", "行政区", "机构", "企业", "影视作品", "电视综艺", "机构", "电视综艺", "娱乐人物", "娱乐人物", "娱乐人物", "娱乐人物", "地点", "影视作品", "娱乐人物", "娱乐人物", "景点", "历史人物", "文学作品", "企业", "人物", "人物", "人物", "历史人物", "企业", "行政区", "人物", "人物", "影视作品", "影视作品", "学科专业", "行政区", "学科专业", "影视作品", "国家", "国家", "影视作品", "学校", "企业/品牌"], "object_type": ["empty", "empty", "Number", "人物", "音乐专辑", "人物", "Text", "人物", "Date_@value", "地点_inArea", "人物_@value", "影视作品_inWork", "国家", "Date", "学校", "人物", "人物", "人物", "人物", "Number", "Number", "地点", "人物", "人物", "Text", "人物", "奖项_@value", "作品_inWork", "Date_onDate", "Number_period", "Number", "企业", "人物_@value", "影视作品_inWork", "城市", "Text", "人物", "人物", "人物", "地点", "人物", "Text", "人物", "Number", "人物", "人物", "Number_@value", "地点_inArea", "Text", "气候", "Number", "作品", "语言", "城市", "歌曲", "人物", "人物"]}
Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1 @@
1+
{"O": 0, "I": 1, "注册资本": 2, "作者": 3, "所属专辑": 4, "歌手": 5, "邮政编码": 6, "主演": 7, "上映时间_@value": 8, "上映时间_inArea": 9, "饰演_@value": 10, "饰演_inWork": 11, "国籍": 12, "成立日期": 13, "毕业院校": 14, "作曲": 15, "作词": 16, "编剧": 17, "导演": 18, "面积": 19, "占地面积": 20, "总部地点": 21, "制片人": 22, "嘉宾": 23, "简称": 24, "主持人": 25, "获奖_@value": 26, "获奖_inWork": 27, "获奖_onDate": 28, "获奖_period": 29, "海拔": 30, "出品公司": 31, "配音_@value": 32, "配音_inWork": 33, "所在城市": 34, "号": 35, "主角": 36, "创始人": 37, "父亲": 38, "祖籍": 39, "母亲": 40, "朝代": 41, "董事长": 42, "人口数量": 43, "妻子": 44, "丈夫": 45, "票房_@value": 46, "票房_inArea": 47, "专业代码": 48, "气候": 49, "修业年限": 50, "改编自": 51, "官方语言": 52, "首都": 53, "主题曲": 54, "校长": 55, "代言人": 56}

‎examples/information_extraction/DuIE/run_duie.py‎

Lines changed: 25 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -88,47 +88,41 @@ def evaluate(model, criterion, data_loader, file_path, mode):
8888
predict_test.json and predict_test.json.zip \
8989
under args.data_path dir for later submission or evaluation.
9090
"""
91+
example_all = []
92+
with open(file_path, "r", encoding="utf-8") as fp:
93+
for line in fp:
94+
example_all.append(json.loads(line))
95+
id2spo_path = os.path.join(os.path.dirname(file_path), "id2spo.json")
96+
with open(id2spo_path, 'r', encoding='utf8') as fp:
97+
id2spo = json.load(fp)
98+
9199
model.eval()
92-
probs_all = None
93-
seq_len_all = None
94-
tok_to_orig_start_index_all = None
95-
tok_to_orig_end_index_all = None
96100
loss_all = 0
97101
eval_steps = 0
102+
formatted_outputs = []
103+
current_idx = 0
98104
for batch in tqdm(data_loader, total=len(data_loader)):
99105
eval_steps += 1
100106
input_ids, seq_len, tok_to_orig_start_index, tok_to_orig_end_index, labels = batch
101107
logits = model(input_ids=input_ids)
102-
mask = (input_ids != 0).logical_and((input_ids != 1)).logical_and(
103-
(input_ids != 2))
108+
mask = (input_ids != 0).logical_and((input_ids != 1)).logical_and((input_ids != 2))
104109
loss = criterion(logits, labels, mask)
105110
loss_all += loss.numpy().item()
106111
probs = F.sigmoid(logits)
107-
if probs_all is None:
108-
probs_all = probs.numpy()
109-
seq_len_all = seq_len.numpy()
110-
tok_to_orig_start_index_all = tok_to_orig_start_index.numpy()
111-
tok_to_orig_end_index_all = tok_to_orig_end_index.numpy()
112-
else:
113-
probs_all = np.append(probs_all, probs.numpy(), axis=0)
114-
seq_len_all = np.append(seq_len_all, seq_len.numpy(), axis=0)
115-
tok_to_orig_start_index_all = np.append(
116-
tok_to_orig_start_index_all,
117-
tok_to_orig_start_index.numpy(),
118-
axis=0)
119-
tok_to_orig_end_index_all = np.append(
120-
tok_to_orig_end_index_all,
121-
tok_to_orig_end_index.numpy(),
122-
axis=0)
112+
logits_batch = probs.numpy()
113+
seq_len_batch = seq_len.numpy()
114+
tok_to_orig_start_index_batch = tok_to_orig_start_index.numpy()
115+
tok_to_orig_end_index_batch = tok_to_orig_end_index.numpy()
116+
formatted_outputs.extend(decoding(example_all[current_idx: current_idx+len(logits)],
117+
id2spo,
118+
logits_batch,
119+
seq_len_batch,
120+
tok_to_orig_start_index_batch,
121+
tok_to_orig_end_index_batch))
122+
current_idx = current_idx+len(logits)
123123
loss_avg = loss_all / eval_steps
124124
print("eval loss: %f" % (loss_avg))
125125

126-
id2spo_path = os.path.join(os.path.dirname(file_path), "id2spo.json")
127-
with open(id2spo_path, 'r', encoding='utf8') as fp:
128-
id2spo = json.load(fp)
129-
formatted_outputs = decoding(file_path, id2spo, probs_all, seq_len_all,
130-
tok_to_orig_start_index_all,
131-
tok_to_orig_end_index_all)
132126
if mode == "predict":
133127
predict_file_path = os.path.join(args.data_path, 'predictions.json')
134128
else:
@@ -228,7 +222,7 @@ def do_train():
228222
optimizer.clear_grad()
229223
loss_item = loss.numpy().item()
230224
global_step += 1
231-
225+
232226
if global_step % logging_steps == 0 and rank == 0:
233227
print(
234228
"epoch: %d / %d, steps: %d / %d, loss: %f, speed: %.2f step/s"
@@ -270,7 +264,7 @@ def do_train():
270264

271265
def do_predict():
272266
paddle.set_device(args.device)
273-
267+
274268
# Reads label_map.
275269
label_map_path = os.path.join(args.data_path, "predicate2id.json")
276270
if not (os.path.exists(label_map_path) and os.path.isfile(label_map_path)):
@@ -313,6 +307,7 @@ def do_predict():
313307

314308

315309
if __name__ == "__main__":
310+
316311
if args.do_train:
317312
do_train()
318313
elif args.do_predict:

‎examples/information_extraction/DuIE/utils.py‎

Lines changed: 7 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -42,19 +42,18 @@ def find_entity(text_raw, id_, predictions, tok_to_orig_start_index,
4242
return list(set(entity_list))
4343

4444

45-
def decoding(file_path, id2spo, logits_all, seq_len_all,
46-
tok_to_orig_start_index_all, tok_to_orig_end_index_all):
45+
def decoding(example_batch,
46+
id2spo,
47+
logits_batch,
48+
seq_len_batch,
49+
tok_to_orig_start_index_batch,
50+
tok_to_orig_end_index_batch):
4751
"""
4852
model output logits -> formatted spo (as in data set file)
4953
"""
50-
example_all = []
51-
with open(file_path, "r", encoding="utf-8") as fp:
52-
for line in fp:
53-
example_all.append(json.loads(line))
54-
5554
formatted_outputs = []
5655
for (i, (example, logits, seq_len, tok_to_orig_start_index, tok_to_orig_end_index)) in \
57-
enumerate(zip(example_all, logits_all, seq_len_all, tok_to_orig_start_index_all, tok_to_orig_end_index_all)):
56+
enumerate(zip(example_batch, logits_batch, seq_len_batch, tok_to_orig_start_index_batch, tok_to_orig_end_index_batch)):
5857

5958
logits = logits[1:seq_len +
6059
1] # slice between [CLS] and [SEP] to get valid logits

0 commit comments

Comments
 (0)