MCPcopy Create free account
hub / github.com/baidu/DDParser / predict_query

Function predict_query

ddparser/run.py:224–263  ·  view source on GitHub ↗

Predict one query

(env)

Source from the content-addressed store, hash-verified

222
223
224def predict_query(env):
225 """Predict one query"""
226 args = env.args
227 logging.info("Load the model")
228 model = load(args.model_path)
229 model.eval()
230 lac_mode = "seg" if args.feat != "pos" else "lac"
231 lac = LAC.LAC(mode=lac_mode)
232 if args.prob:
233 env.fields = env.fields._replace(PHEAD=Field("prob"))
234
235 while True:
236 query = input()
237 if isinstance(query, six.text_type):
238 pass
239 else:
240 query = query.decode("utf-8")
241 if not query:
242 logging.info("quit!")
243 return
244 if len(query) > 200:
245 logging.info("The length of the query should be less than 200!")
246 continue
247 start = datetime.datetime.now()
248 lac_results = lac.run([query])
249 predicts = Corpus.load_lac_results(lac_results, env.fields)
250 dataset = TextDataset(predicts, [env.WORD, env.FEAT])
251 # set the data loader
252 dataset.loader = batchify(dataset, args.batch_size, use_multiprocess=False, sequential_sampler=True)
253 pred_arcs, pred_rels, pred_probs = epoch_predict(env, args, model, dataset.loader)
254 predicts.head = pred_arcs
255 predicts.deprel = pred_rels
256 if args.prob:
257 predicts.prob = pred_probs
258 predicts._print()
259 total_time = datetime.datetime.now() - start
260 logging.info("{}s elapsed, {:.2f} Sents/s, {:.2f} ms/Sents".format(
261 total_time,
262 len(dataset) / total_time.total_seconds(),
263 total_time.total_seconds() / len(dataset) * 1000))
264
265
266class DDParser(object):

Callers 1

run.pyFile · 0.85

Calls 8

loadFunction · 0.90
FieldClass · 0.90
TextDatasetClass · 0.90
batchifyFunction · 0.90
epoch_predictFunction · 0.90
evalMethod · 0.80
load_lac_resultsMethod · 0.80
_printMethod · 0.80

Tested by

no test coverage detected