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