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)
| 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: |
nothing calls this directly
no test coverage detected