1#  Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
2#  See https://llvm.org/LICENSE.txt for license information.
3#  SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
4
5# Provide a convenient name for sub-packages to resolve the main C-extension
6# with a relative import.
7from .._mlir_libs import _mlir as _cext
8from typing import Sequence as _Sequence, Union as _Union
9
10__all__ = [
11    "equally_sized_accessor",
12    "extend_opview_class",
13    "get_default_loc_context",
14    "get_op_result_or_value",
15    "get_op_results_or_values",
16    "segmented_accessor",
17]
18
19
20def extend_opview_class(ext_module):
21  """Decorator to extend an OpView class from an extension module.
22
23  Extension modules can expose various entry-points:
24    Stand-alone class with the same name as a parent OpView class (i.e.
25    "ReturnOp"). A name-based match is attempted first before falling back
26    to a below mechanism.
27
28    def select_opview_mixin(parent_opview_cls):
29      If defined, allows an appropriate mixin class to be selected dynamically
30      based on the parent OpView class. Should return NotImplemented if a
31      decision is not made.
32
33  Args:
34    ext_module: A module from which to locate extensions. Can be None if not
35      available.
36
37  Returns:
38    A decorator that takes an OpView subclass and further extends it as
39    needed.
40  """
41
42  def class_decorator(parent_opview_cls: type):
43    if ext_module is None:
44      return parent_opview_cls
45    mixin_cls = NotImplemented
46    # First try to resolve by name.
47    try:
48      mixin_cls = getattr(ext_module, parent_opview_cls.__name__)
49    except AttributeError:
50      # Fall back to a select_opview_mixin hook.
51      try:
52        select_mixin = getattr(ext_module, "select_opview_mixin")
53      except AttributeError:
54        pass
55      else:
56        mixin_cls = select_mixin(parent_opview_cls)
57
58    if mixin_cls is NotImplemented or mixin_cls is None:
59      return parent_opview_cls
60
61    # Have a mixin_cls. Create an appropriate subclass.
62    try:
63
64      class LocalOpView(mixin_cls, parent_opview_cls):
65        pass
66    except TypeError as e:
67      raise TypeError(
68          f"Could not mixin {mixin_cls} into {parent_opview_cls}") from e
69    LocalOpView.__name__ = parent_opview_cls.__name__
70    LocalOpView.__qualname__ = parent_opview_cls.__qualname__
71    return LocalOpView
72
73  return class_decorator
74
75
76def segmented_accessor(elements, raw_segments, idx):
77  """
78  Returns a slice of elements corresponding to the idx-th segment.
79
80    elements: a sliceable container (operands or results).
81    raw_segments: an mlir.ir.Attribute, of DenseIntElements subclass containing
82        sizes of the segments.
83    idx: index of the segment.
84  """
85  segments = _cext.ir.DenseIntElementsAttr(raw_segments)
86  start = sum(segments[i] for i in range(idx))
87  end = start + segments[idx]
88  return elements[start:end]
89
90
91def equally_sized_accessor(elements, n_variadic, n_preceding_simple,
92                           n_preceding_variadic):
93  """
94  Returns a starting position and a number of elements per variadic group
95  assuming equally-sized groups and the given numbers of preceding groups.
96
97    elements: a sequential container.
98    n_variadic: the number of variadic groups in the container.
99    n_preceding_simple: the number of non-variadic groups preceding the current
100        group.
101    n_preceding_variadic: the number of variadic groups preceding the current
102        group.
103  """
104
105  total_variadic_length = len(elements) - n_variadic + 1
106  # This should be enforced by the C++-side trait verifier.
107  assert total_variadic_length % n_variadic == 0
108
109  elements_per_group = total_variadic_length // n_variadic
110  start = n_preceding_simple + n_preceding_variadic * elements_per_group
111  return start, elements_per_group
112
113
114def get_default_loc_context(location=None):
115  """
116  Returns a context in which the defaulted location is created. If the location
117  is None, takes the current location from the stack, raises ValueError if there
118  is no location on the stack.
119  """
120  if location is None:
121    # Location.current raises ValueError if there is no current location.
122    return _cext.ir.Location.current.context
123  return location.context
124
125
126def get_op_result_or_value(
127    arg: _Union[_cext.ir.OpView, _cext.ir.Operation, _cext.ir.Value, _cext.ir.OpResultList]
128) -> _cext.ir.Value:
129  """Returns the given value or the single result of the given op.
130
131  This is useful to implement op constructors so that they can take other ops as
132  arguments instead of requiring the caller to extract results for every op.
133  Raises ValueError if provided with an op that doesn't have a single result.
134  """
135  if isinstance(arg, _cext.ir.OpView):
136    return arg.operation.result
137  elif isinstance(arg, _cext.ir.Operation):
138    return arg.result
139  elif isinstance(arg, _cext.ir.OpResultList):
140    return arg[0]
141  else:
142    assert isinstance(arg, _cext.ir.Value)
143    return arg
144
145
146def get_op_results_or_values(
147    arg: _Union[_cext.ir.OpView, _cext.ir.Operation,
148                _Sequence[_Union[_cext.ir.OpView, _cext.ir.Operation, _cext.ir.Value]]]
149) -> _Union[_Sequence[_cext.ir.Value], _cext.ir.OpResultList]:
150  """Returns the given sequence of values or the results of the given op.
151
152  This is useful to implement op constructors so that they can take other ops as
153  lists of arguments instead of requiring the caller to extract results for
154  every op.
155  """
156  if isinstance(arg, _cext.ir.OpView):
157    return arg.operation.results
158  elif isinstance(arg, _cext.ir.Operation):
159    return arg.results
160  else:
161    return [get_op_result_or_value(element) for element in arg]
162