Test model with decoder
(model, data_loader, device, keywords_token, keywords_idxset,
args)
| 160 | |
| 161 | |
| 162 | def executor_test(model, data_loader, device, keywords_token, keywords_idxset, |
| 163 | args): |
| 164 | ''' Test model with decoder |
| 165 | ''' |
| 166 | assert args.get('test_dir', None) is not None, \ |
| 167 | 'Please config param: test_dir, to store score file' |
| 168 | score_abs_path = os.path.join(args['test_dir'], 'score.txt') |
| 169 | log_interval = args.get('log_interval', 10) |
| 170 | |
| 171 | model.eval() |
| 172 | infer_seconds = 0.0 |
| 173 | decode_seconds = 0.0 |
| 174 | with torch.no_grad(), open(score_abs_path, 'w', encoding='utf8') as fout: |
| 175 | for batch_idx, batch in enumerate(data_loader): |
| 176 | batch_start_time = datetime.datetime.now() |
| 177 | |
| 178 | keys, feats, target, feats_lengths, target_lengths = batch |
| 179 | feats = feats.to(device) |
| 180 | feats_lengths = feats_lengths.to(device) |
| 181 | if target_lengths is not None: |
| 182 | target_lengths = target_lengths.to(device) |
| 183 | num_utts = feats_lengths.size(0) |
| 184 | if num_utts == 0: |
| 185 | continue |
| 186 | |
| 187 | logits, _ = model(feats) |
| 188 | logits = logits.softmax(2) # (1, maxlen, vocab_size) |
| 189 | logits = logits.cpu() |
| 190 | |
| 191 | infer_end_time = datetime.datetime.now() |
| 192 | for i in range(len(keys)): |
| 193 | key = keys[i] |
| 194 | score = logits[i][:feats_lengths[i]] |
| 195 | hyps = ctc_prefix_beam_search(score, feats_lengths[i], |
| 196 | keywords_idxset) |
| 197 | hit_keyword = None |
| 198 | hit_score = 1.0 |
| 199 | # start = 0; end = 0 |
| 200 | for one_hyp in hyps: |
| 201 | prefix_ids = one_hyp[0] |
| 202 | # path_score = one_hyp[1] |
| 203 | prefix_nodes = one_hyp[2] |
| 204 | assert len(prefix_ids) == len(prefix_nodes) |
| 205 | for word in keywords_token.keys(): |
| 206 | lab = keywords_token[word]['token_id'] |
| 207 | offset = is_sublist(prefix_ids, lab) |
| 208 | if offset != -1: |
| 209 | hit_keyword = word |
| 210 | # start = prefix_nodes[offset]['frame'] |
| 211 | # end = prefix_nodes[offset+len(lab)-1]['frame'] |
| 212 | for idx in range(offset, offset + len(lab)): |
| 213 | hit_score *= prefix_nodes[idx]['prob'] |
| 214 | break |
| 215 | if hit_keyword is not None: |
| 216 | hit_score = math.sqrt(hit_score) |
| 217 | break |
| 218 | |
| 219 | if hit_keyword is not None: |
no test coverage detected
searching dependent graphs…