MCPcopy Create free account
hub / github.com/modelscope/modelscope / executor_test

Function executor_test

modelscope/trainers/audio/kws_utils/batch_utils.py:162–242  ·  view source on GitHub ↗

Test model with decoder

(model, data_loader, device, keywords_token, keywords_idxset,
                  args)

Source from the content-addressed store, hash-verified

160
161
162def 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:

Callers 1

evaluateMethod · 0.85

Calls 10

ctc_prefix_beam_searchFunction · 0.85
is_sublistFunction · 0.85
infoMethod · 0.80
getMethod · 0.45
evalMethod · 0.45
toMethod · 0.45
sizeMethod · 0.45
softmaxMethod · 0.45
keysMethod · 0.45
writeMethod · 0.45

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…