A `Dataset` that maps a function over elements in its input in parallel.
| 3550 | |
| 3551 | |
| 3552 | class ParallelMapDataset(UnaryDataset): |
| 3553 | """A `Dataset` that maps a function over elements in its input in parallel.""" |
| 3554 | |
| 3555 | def __init__(self, |
| 3556 | input_dataset, |
| 3557 | map_func, |
| 3558 | num_parallel_calls, |
| 3559 | use_inter_op_parallelism=True, |
| 3560 | preserve_cardinality=False, |
| 3561 | use_legacy_function=False): |
| 3562 | """See `Dataset.map()` for details.""" |
| 3563 | self._input_dataset = input_dataset |
| 3564 | self._use_inter_op_parallelism = use_inter_op_parallelism |
| 3565 | self._map_func = StructuredFunctionWrapper( |
| 3566 | map_func, |
| 3567 | self._transformation_name(), |
| 3568 | dataset=input_dataset, |
| 3569 | use_legacy_function=use_legacy_function) |
| 3570 | self._num_parallel_calls = ops.convert_to_tensor( |
| 3571 | num_parallel_calls, dtype=dtypes.int32, name="num_parallel_calls") |
| 3572 | self._preserve_cardinality = preserve_cardinality |
| 3573 | variant_tensor = gen_dataset_ops.parallel_map_dataset( |
| 3574 | input_dataset._variant_tensor, # pylint: disable=protected-access |
| 3575 | self._map_func.function.captured_inputs, |
| 3576 | f=self._map_func.function, |
| 3577 | num_parallel_calls=self._num_parallel_calls, |
| 3578 | use_inter_op_parallelism=self._use_inter_op_parallelism, |
| 3579 | preserve_cardinality=self._preserve_cardinality, |
| 3580 | **self._flat_structure) |
| 3581 | super(ParallelMapDataset, self).__init__(input_dataset, variant_tensor) |
| 3582 | |
| 3583 | def _functions(self): |
| 3584 | return [self._map_func] |
| 3585 | |
| 3586 | @property |
| 3587 | def element_spec(self): |
| 3588 | return self._map_func.output_structure |
| 3589 | |
| 3590 | def _transformation_name(self): |
| 3591 | return "Dataset.map()" |
| 3592 | |
| 3593 | |
| 3594 | class FlatMapDataset(UnaryDataset): |
no outgoing calls
no test coverage detected