(
x, axes, starts, ends, steps, decrease_axes, none_axes, use_strided_slice
)
| 715 | |
| 716 | |
| 717 | def get_tensor_with_basic_indexing( |
| 718 | x, axes, starts, ends, steps, decrease_axes, none_axes, use_strided_slice |
| 719 | ): |
| 720 | from .dygraph.base import in_to_static_mode |
| 721 | |
| 722 | out_is_view = False |
| 723 | if in_to_static_mode() and hasattr(x, "is_view_var"): |
| 724 | x.is_view_var = True |
| 725 | |
| 726 | if len(axes) == 0: |
| 727 | out = x |
| 728 | else: |
| 729 | out_is_view = True |
| 730 | op_type = "strided_slice" if use_strided_slice else "slice" |
| 731 | inputs = {'Input': [x]} |
| 732 | attrs = { |
| 733 | 'axes': axes, |
| 734 | 'starts': [], |
| 735 | 'ends': [], |
| 736 | 'decrease_axis': decrease_axes, |
| 737 | } |
| 738 | if use_strided_slice: |
| 739 | attrs['strides'] = [] |
| 740 | infer_flags = [1] * len(axes) |
| 741 | deal_attrs( |
| 742 | attrs, starts, "starts", "StartsTensorList", inputs, infer_flags |
| 743 | ) |
| 744 | deal_attrs(attrs, ends, "ends", "EndsTensorList", inputs, infer_flags) |
| 745 | deal_attrs( |
| 746 | attrs, steps, "strides", "StridesTensorList", inputs, infer_flags |
| 747 | ) |
| 748 | attrs['infer_flags'] = infer_flags |
| 749 | |
| 750 | from . import in_dynamic_or_pir_mode, in_pir_mode |
| 751 | |
| 752 | if in_dynamic_or_pir_mode(): |
| 753 | if "StartsTensorList" in inputs.keys(): |
| 754 | st = inputs['StartsTensorList'] |
| 755 | else: |
| 756 | st = attrs['starts'] |
| 757 | if "EndsTensorList" in inputs.keys(): |
| 758 | end = inputs['EndsTensorList'] |
| 759 | else: |
| 760 | end = attrs['ends'] |
| 761 | if "StridesTensorList" in inputs.keys(): |
| 762 | stride = inputs['StridesTensorList'] |
| 763 | else: |
| 764 | stride = attrs['strides'] |
| 765 | if use_strided_slice: |
| 766 | # TODO(zoooo0820): support strided_slice_array until PIR API is ready |
| 767 | if in_pir_mode(): |
| 768 | if isinstance(st, (list, tuple)): |
| 769 | if paddle.utils._contain_var(st): |
| 770 | st = paddle.utils.get_int_tensor_list(st) |
| 771 | if isinstance(end, (list, tuple)): |
| 772 | if paddle.utils._contain_var(end): |
| 773 | end = paddle.utils.get_int_tensor_list(end) |
| 774 | if isinstance(stride, (list, tuple)): |
no test coverage detected