Predict one query
(env)
| 222 | |
| 223 | |
| 224 | def 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 | |
| 266 | class DDParser(object): |
no test coverage detected