MCPcopy Create free account
hub / github.com/hankcs/HanLP / predict

Method predict

hanlp/components/mtl/multi_task_learning.py:459–568  ·  view source on GitHub ↗

Predict on data. Args: data: A sentence or a list of sentences. tasks: The tasks to predict. skip_tasks: The tasks to skip. resolved_tasks: The resolved tasks to override ``tasks`` and ``skip_tasks``. **kwargs: Not used. R

(self,
                data: Union[str, List[str]],
                tasks: Optional[Union[str, List[str]]] = None,
                skip_tasks: Optional[Union[str, List[str]]] = None,
                resolved_tasks=None,
                **kwargs)

Source from the content-addressed store, hash-verified

457 return MultiTaskModel(transformer_module, scalar_mixes, decoders, use_raw_hidden_states)
458
459 def predict(self,
460 data: Union[str, List[str]],
461 tasks: Optional[Union[str, List[str]]] = None,
462 skip_tasks: Optional[Union[str, List[str]]] = None,
463 resolved_tasks=None,
464 **kwargs) -> Document:
465 """Predict on data.
466
467 Args:
468 data: A sentence or a list of sentences.
469 tasks: The tasks to predict.
470 skip_tasks: The tasks to skip.
471 resolved_tasks: The resolved tasks to override ``tasks`` and ``skip_tasks``.
472 **kwargs: Not used.
473
474 Returns:
475 A :class:`~hanlp_common.document.Document`.
476 """
477 doc = Document()
478 target_tasks = resolved_tasks or self.resolve_tasks(tasks, skip_tasks)
479 if data == []:
480 for group in target_tasks:
481 for task_name in group:
482 doc[task_name] = []
483 return doc
484 flatten_target_tasks = [self.tasks[t] for group in target_tasks for t in group]
485 cls_is_bos = any([x.cls_is_bos for x in flatten_target_tasks])
486 sep_is_eos = any([x.sep_is_eos for x in flatten_target_tasks])
487 # Now build the dataloaders and execute tasks
488 first_task_name: str = list(target_tasks[0])[0]
489 first_task: Task = self.tasks[first_task_name]
490 encoder_transform, transform = self.build_transform(first_task)
491 # Override the tokenizer config of the 1st task
492 encoder_transform.sep_is_eos = sep_is_eos
493 encoder_transform.cls_is_bos = cls_is_bos
494 average_subwords = self.model.encoder.average_subwords
495 flat = first_task.input_is_flat(data)
496 if flat:
497 data = [data]
498 device = self.device
499 samples = first_task.build_samples(data, cls_is_bos=cls_is_bos, sep_is_eos=sep_is_eos)
500 dataloader = first_task.build_dataloader(samples, transform=transform, device=device)
501 results = defaultdict(list)
502 order = []
503 for batch in dataloader:
504 order.extend(batch[IDX])
505 # Run the first task, let it make the initial batch for the successors
506 output_dict = self.predict_task(first_task, first_task_name, batch, results, run_transform=True,
507 cls_is_bos=cls_is_bos, sep_is_eos=sep_is_eos)
508 # Run each task group in order
509 for group_id, group in enumerate(target_tasks):
510 # We could parallelize this in the future
511 for task_name in group:
512 if task_name == first_task_name:
513 continue
514 output_dict = self.predict_task(self.tasks[task_name], task_name, batch, results, output_dict,
515 run_transform=True, cls_is_bos=cls_is_bos, sep_is_eos=sep_is_eos)
516 if group_id == 0 and len(target_tasks) > 1:

Callers

nothing calls this directly

Calls 15

resolve_tasksMethod · 0.95
build_transformMethod · 0.95
predict_taskMethod · 0.95
DocumentClass · 0.90
reorderFunction · 0.90
extendMethod · 0.80
getMethod · 0.80
pad_dataMethod · 0.80
input_is_flatMethod · 0.45
build_samplesMethod · 0.45
build_dataloaderMethod · 0.45

Tested by

no test coverage detected