See `Dataset.concatenate()` for details.
(self, input_dataset, dataset_to_concatenate)
| 2941 | """A `Dataset` that concatenates its input with given dataset.""" |
| 2942 | |
| 2943 | def __init__(self, input_dataset, dataset_to_concatenate): |
| 2944 | """See `Dataset.concatenate()` for details.""" |
| 2945 | self._input_dataset = input_dataset |
| 2946 | self._dataset_to_concatenate = dataset_to_concatenate |
| 2947 | |
| 2948 | output_types = get_legacy_output_types(input_dataset) |
| 2949 | if output_types != get_legacy_output_types(dataset_to_concatenate): |
| 2950 | raise TypeError( |
| 2951 | "Two datasets to concatenate have different types %s and %s" % |
| 2952 | (output_types, get_legacy_output_types(dataset_to_concatenate))) |
| 2953 | |
| 2954 | output_classes = get_legacy_output_classes(input_dataset) |
| 2955 | if output_classes != get_legacy_output_classes(dataset_to_concatenate): |
| 2956 | raise TypeError( |
| 2957 | "Two datasets to concatenate have different classes %s and %s" % |
| 2958 | (output_classes, get_legacy_output_classes(dataset_to_concatenate))) |
| 2959 | |
| 2960 | input_shapes = get_legacy_output_shapes(self._input_dataset) |
| 2961 | output_shapes = nest.pack_sequence_as(input_shapes, [ |
| 2962 | ts1.most_specific_compatible_shape(ts2) |
| 2963 | for (ts1, ts2) in zip( |
| 2964 | nest.flatten(input_shapes), |
| 2965 | nest.flatten(get_legacy_output_shapes( |
| 2966 | self._dataset_to_concatenate))) |
| 2967 | ]) |
| 2968 | |
| 2969 | self._structure = structure.convert_legacy_structure( |
| 2970 | output_types, output_shapes, output_classes) |
| 2971 | |
| 2972 | self._input_datasets = [input_dataset, dataset_to_concatenate] |
| 2973 | # pylint: disable=protected-access |
| 2974 | variant_tensor = gen_dataset_ops.concatenate_dataset( |
| 2975 | input_dataset._variant_tensor, dataset_to_concatenate._variant_tensor, |
| 2976 | **self._flat_structure) |
| 2977 | # pylint: enable=protected-access |
| 2978 | super(ConcatenateDataset, self).__init__(variant_tensor) |
| 2979 | |
| 2980 | def _inputs(self): |
| 2981 | return self._input_datasets |
nothing calls this directly
no test coverage detected