Restore SparseTensors after dequeue in batch, batch_join, etc.
(stored_list, sparse_info_list)
| 598 | |
| 599 | |
| 600 | def _restore_sparse_tensors(stored_list, sparse_info_list): |
| 601 | """Restore SparseTensors after dequeue in batch, batch_join, etc.""" |
| 602 | received_sequence = isinstance(stored_list, collections_abc.Sequence) |
| 603 | if not received_sequence: |
| 604 | stored_list = (stored_list,) |
| 605 | tensors = [ |
| 606 | _restore_sparse(sparse_map_op=info.map_op, |
| 607 | sparse_handles=array_ops.squeeze(s, [1]), |
| 608 | rank=tensor_shape.dimension_value(info.rank + 1)) |
| 609 | if info.sparse else s |
| 610 | for (s, info) in zip(stored_list, sparse_info_list)] |
| 611 | has_st = any(isinstance(x, sparse_tensor.SparseTensor) for x in tensors) |
| 612 | if has_st: |
| 613 | t_values = [ |
| 614 | x.values if isinstance(x, sparse_tensor.SparseTensor) |
| 615 | else x |
| 616 | for x in tensors] |
| 617 | with_deps = lambda x: control_flow_ops.with_dependencies(t_values, x) |
| 618 | ensure_restore_tensors = [ |
| 619 | sparse_tensor.SparseTensor(indices=with_deps(x.indices), |
| 620 | values=with_deps(x.values), |
| 621 | dense_shape=with_deps(x.dense_shape)) |
| 622 | if isinstance(x, sparse_tensor.SparseTensor) |
| 623 | else with_deps(x) |
| 624 | for x in tensors] |
| 625 | else: |
| 626 | ensure_restore_tensors = tensors |
| 627 | return ensure_restore_tensors if received_sequence else tensors[0] |
| 628 | |
| 629 | |
| 630 | def _validate(tensor_list): |
no test coverage detected