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
5from typing import Callable, Dict, List, Sequence, Tuple, Union
6
7from .....ir import *
8
9from .... import func
10from .... import linalg
11from .... import math
12from .... import arith
13from .... import complex
14from ...._ods_common import get_op_result_or_value as _get_op_result_or_value, get_op_results_or_values as _get_op_results_or_values
15
16from .scalar_expr import *
17from .config import *
18from .comprehension import *
19import numpy as np
20
21__all__ = [
22    "emit_generic_structured_op",
23    "emit_named_structured_op",
24    "ValueList",
25]
26
27# Type aliases.
28ValueList = Union[Sequence[Value], OpResultList]
29
30
31def isa(cls: Type, ty: Type):
32  try:
33    cls(ty)
34    return True
35  except ValueError:
36    return False
37
38
39def prepare_common_structured_op(op_config: LinalgStructuredOpConfig,
40                                 *ins: Value, outs: ValueList,
41                                 **attrs: Union[Sequence[int], TypeFnType]):
42  all_arg_defs = op_config.ordered_operands
43  in_arg_defs = [
44      d for d in all_arg_defs
45      if d.kind in [OperandKind.SCALAR, OperandKind.INPUT_TENSOR]
46  ]
47  out_arg_defs = [
48      d for d in all_arg_defs if d.kind == OperandKind.OUTPUT_TENSOR
49  ]
50  index_attr_arg_defs = [
51      d for d in all_arg_defs if d.kind == OperandKind.INDEX_ATTR
52  ]
53  fn_attr_arg_defs = [
54      d for d in all_arg_defs if d.kind in [
55          OperandKind.UNARY_FN_ATTR, OperandKind.BINARY_FN_ATTR,
56          OperandKind.TYPE_FN_ATTR
57      ]
58  ]
59
60  # Verify outs is a sequence or a list of results.
61  if not isinstance(outs, (Sequence, OpResultList)):
62    raise ValueError(f"Expected named argument outs to have type Sequence or "
63                     f"OpResultLis but got {type(outs)}")
64
65  # Arity validation.
66  if len(ins) != len(in_arg_defs):
67    raise ValueError(f"Expected {len(in_arg_defs)} inputs but got "
68                     f"{len(ins)} for {op_config}")
69  if outs and len(outs) != len(out_arg_defs):
70    raise ValueError(f"Expected {len(out_arg_defs)} outputs but got "
71                     f"{len(outs)} for {op_config}")
72
73  # Compute a replacement list for all index attribute symbols.
74  expressions = []  # type: Sequence[AffineExpr]
75  replacements = []  # type: Sequence[AffineExpr]
76  for index_attr in index_attr_arg_defs:
77    index_attr_vals = index_attr.operand_def.default_indices
78    if index_attr.name in attrs:
79      index_attr_vals = attrs.get(index_attr.name)
80    assert index_attr_vals, "Index attribute has no value"
81    if not all(isinstance(value, int) for value in index_attr_vals):
82      raise ValueError(f"Attribute {index_attr.name} needs to be of type "
83                       f"Sequence[int] but got {type(index_attr_vals)}")
84    results = index_attr.index_attr_map.results  # type: AffineExprList
85    if len(index_attr_vals) != len(results):
86      raise ValueError(f"Attribute {index_attr.name} has length {len(results)} "
87                       f"but got {len(index_attr_vals)} values")
88    for expr, value in zip(results, index_attr_vals):
89      expressions.append(expr)
90      replacements.append(AffineConstantExpr.get(value))
91
92  # Replace all index attribute symbols by their value.
93  # TODO: Add support for shape symbols.
94  indexing_maps = []  # type: Sequence[AffineMap]
95  for curr in op_config.indexing_maps:
96    for expression, replacement in zip(expressions, replacements):
97      curr = curr.replace(expression, replacement, curr.n_dims, curr.n_symbols)
98    indexing_maps.append(curr)
99
100  # TODO: Linalg verification does not currently allow symbols.
101  # Compress them for now and verify none are left.
102  indexing_maps = AffineMap.compress_unused_symbols(indexing_maps,
103                                                    Context.current)
104  if any(indexing_map.n_symbols != 0 for indexing_map in indexing_maps):
105    raise ValueError(f"Expected indexing_maps to use no symbols after "
106                     f"replacement and compression but got {indexing_maps}")
107
108  outs, out_types = _infer_structured_outs(op_config, in_arg_defs, ins,
109                                           out_arg_defs, outs)
110
111  result_types = [t for t in out_types if isa(RankedTensorType, t)]
112
113  # Initialize the type dictionary with the predefined types.
114  type_mapping = dict()  # type: Dict[str, Type]
115  type_mapping["F32"] = F32Type.get()
116  type_mapping["F64"] = F64Type.get()
117  type_mapping["I32"] = IntegerType.get_signless(32)
118  type_mapping["I64"] = IntegerType.get_signless(64)
119
120  # Extract type vars for input/output based types.
121  block_arg_types = list()  # type: List[Type]
122  for arg_def, arg_element_type in zip(in_arg_defs + out_arg_defs,
123                                       _get_types_from_values(*ins, *outs)):
124    _add_type_mapping(arg_def, arg_element_type, type_mapping, block_arg_types)
125
126  # Emit the generic op.
127  # TODO: Support emission of pure memref form.
128  indexing_maps_attr = ArrayAttr.get(
129      [AffineMapAttr.get(am) for am in indexing_maps])
130  iterator_types_attr = ArrayAttr.get(
131      [StringAttr.get(s) for s in op_config.iterator_types])
132
133  # Compute the index attributes used when emitting a named structured op.
134  index_attrs = {}  # type: Dict[str, DenseElementAttr]
135  for index_attr in index_attr_arg_defs:
136    index_attr_vals = attrs.get(index_attr.name)
137    # Only forward attributes set to a non-default value.
138    if index_attr_vals:
139      array = np.array(index_attr_vals, dtype=np.int64)
140      index_attrs[index_attr.name] = DenseElementsAttr.get(array)
141
142  # Compute the function attribute mapping.
143  fn_attr_mapping = {}
144  for fn_attr in fn_attr_arg_defs:
145    attr_val = fn_attr.operand_def.default_fn
146    attr_kind = fn_attr.kind
147    if fn_attr.name in attrs:
148      fn = attrs.get(fn_attr.name)
149      if attr_kind == OperandKind.UNARY_FN_ATTR:
150        if not isinstance(fn, UnaryFnType):
151          raise ValueError(f"Attribute {fn_attr.name} needs to be of type "
152                           f"UnaryFnType but got {type(attr_val)}")
153      elif attr_kind == OperandKind.BINARY_FN_ATTR:
154        if not isinstance(fn, BinaryFnType):
155          raise ValueError(f"Attribute {fn_attr.name} needs to be of type "
156                           f"BinaryFnType but got {type(attr_val)}")
157      else:
158        if not isinstance(fn, TypeFnType):
159          raise ValueError(f"Attribute {fn_attr.name} needs to be of type "
160                           f"TypeFnType but got {type(attr_val)}")
161      attr_val = fn.fn_name
162    assert attr_val, "Function attribute has no value"
163    fn_attr_mapping[fn_attr.name] = (attr_val, attr_kind)
164
165  return (all_arg_defs, in_arg_defs, out_arg_defs, outs, result_types,
166          type_mapping, indexing_maps_attr, iterator_types_attr, index_attrs,
167          fn_attr_mapping, block_arg_types)
168
169
170def emit_generic_structured_op(op_config: LinalgStructuredOpConfig, *ins: Value,
171                               outs: ValueList, **attrs: Sequence[int]):
172  all_arg_defs, in_arg_defs, out_arg_defs, outs, result_types, type_mapping, \
173  indexing_maps_attr, iterator_types_attr, index_attrs, fn_attr_mapping, \
174  block_arg_types = \
175     prepare_common_structured_op(op_config, *ins, outs = outs, **attrs)
176
177  # An operation that accesses only scalars and scalar/rank zero tensors is
178  # rank polymorhpic. We implement rank polymorphism by generating different
179  # indexing maps and iterators that match the rank of the first output tensor.
180  # An operation is rank polymorphic if the iteration domain has rank zero.
181  if not iterator_types_attr:
182    rank = ShapedType(outs[0].type).rank
183    iterator_types_attr = ArrayAttr.get([StringAttr.get("parallel")] * rank)
184    scalar_map = AffineMap.get(rank, 0, [])
185    tensor_map = AffineMap.get_identity(rank)
186    indexing_maps = []
187    for arg_def in all_arg_defs:
188      if arg_def.operand_def.kind == OperandKind.SCALAR:
189        indexing_maps.append(scalar_map)
190      if arg_def.operand_def.is_tensor():
191        idx = arg_def.operand_def.registered_index
192        if idx < len(ins) and ShapedType(ins[idx].type).rank == 0:
193          indexing_maps.append(scalar_map)
194        else:
195          indexing_maps.append(tensor_map)
196    indexing_maps_attr = ArrayAttr.get(
197        [AffineMapAttr.get(am) for am in indexing_maps])
198
199  generic_op = linalg.GenericOp(
200      result_tensors=result_types,
201      inputs=ins,
202      outputs=outs,
203      indexing_maps=indexing_maps_attr,
204      iterator_types=iterator_types_attr,
205      doc=None,  # TODO: Make optional.
206      library_call=None)  # TODO: Make optional.
207
208  # Construct the body.
209  block_arg_names = _get_operand_def_names(*in_arg_defs, *out_arg_defs)
210  block = generic_op.regions[0].blocks.append(*block_arg_types)
211  block_arg_mapping = dict(zip(block_arg_names, block.arguments))
212  with InsertionPoint(block):
213    body_builder = _BodyBuilder(type_mapping, block_arg_mapping,
214                                fn_attr_mapping)
215    for assignment in op_config.assignments:
216      body_builder.assign(assignment)
217    body_builder.yield_outputs(*_get_operand_def_names(*out_arg_defs))
218
219  if len(result_types) == 1:
220    return generic_op.result
221  else:
222    return generic_op.results
223
224
225def emit_named_structured_op(op_config: LinalgStructuredOpConfig, op_name: str,
226                             op_class_name: str, *ins: Value, outs: ValueList,
227                             **attrs: Sequence[int]):
228  all_arg_defs, in_arg_defs, out_arg_defs, outs, result_types, type_mapping, \
229  indexing_maps_attr, iterator_types_attr, index_attrs, fn_attr_mapping, \
230  block_arg_types = \
231     prepare_common_structured_op(op_config, *ins, outs = outs, **attrs)
232
233  # If we get here, there must exist a builtin class `op_class_name`.
234  ctx = Context.current
235  fully_qualified_name = "linalg." + op_name
236  if (not ctx.is_registered_operation(fully_qualified_name) or
237      not op_class_name in linalg.__dict__.keys()):
238    raise NotImplementedError(
239        f"Unknown named op_name / op_class_name: {op_name} / {op_class_name}")
240
241  # Set the index attributes used to compute the indexing maps.
242  named_op = getattr(linalg, op_class_name)(ins, outs, result_types)
243  for name, value in index_attrs.items():
244    named_op.operation.attributes[name] = value
245
246  # Compute the function attributes by combining operand kind and function name.
247  for name, (fn_name, kind) in fn_attr_mapping.items():
248    assert kind.name.lower().endswith("_attr")
249    enum_name = kind.name.lower()[:-5]
250    named_op.operation.attributes[name] = Attribute.parse(
251        f"#linalg.{enum_name}<{fn_name}>")
252
253  linalg.fill_builtin_region(named_op.operation)
254
255  if len(result_types) == 1:
256    return named_op.result
257  else:
258    return named_op.results
259
260
261class _BodyBuilder:
262  """Constructs a structured op body by evaluating assignments."""
263
264  def __init__(self, type_mapping: Dict[str, Type],
265               block_arg_mapping: Dict[str, Value], fn_attr_mapping: Dict[str,
266                                                                          str]):
267    self.type_mapping = type_mapping
268    self.block_arg_mapping = block_arg_mapping
269    self.fn_attr_mapping = fn_attr_mapping
270    self.yield_mapping = dict()  # type: Dict[str, Value]
271
272  def assign(self, assignment: ScalarAssign):
273    if assignment.arg in self.yield_mapping:
274      raise ValueError(
275          f"Multiple assignments to the same argument are forbidden: "
276          f"{assignment}")
277    self.yield_mapping[assignment.arg] = self.expression(assignment.value)
278
279  def expression(self, expr: ScalarExpression) -> Value:
280    if expr.scalar_arg:
281      try:
282        return self.block_arg_mapping[expr.scalar_arg.arg]
283      except KeyError:
284        raise ValueError(f"Argument {expr.scalar_arg.arg} is not bound for "
285                         f"this structured op.")
286    elif expr.scalar_const:
287      value_attr = Attribute.parse(expr.scalar_const.value)
288      return arith.ConstantOp(value_attr.type, value_attr).result
289    elif expr.scalar_index:
290      dim_attr = IntegerAttr.get(
291          IntegerType.get_signless(64), expr.scalar_index.dim)
292      return linalg.IndexOp(dim_attr).result
293    elif expr.scalar_fn:
294      kind = expr.scalar_fn.kind.name.lower()
295      fn_name = expr.scalar_fn.fn_name
296      if expr.scalar_fn.attr_name:
297        fn_name, _ = self.fn_attr_mapping[expr.scalar_fn.attr_name]
298      fn = self._get_function(f"_{kind}_{fn_name}")
299      operand_values = [
300          self.expression(operand) for operand in expr.scalar_fn.operands
301      ]
302      if expr.scalar_fn.kind == FunctionKind.TYPE:
303        operand_values = [expr.scalar_fn.type_var.name] + operand_values
304      return fn(*operand_values)
305    raise NotImplementedError(f"Unimplemented scalar body expression: {expr}")
306
307  def yield_outputs(self, *output_names: str):
308    output_values = []
309    for n in output_names:
310      try:
311        output_values.append(self.yield_mapping[n])
312      except KeyError:
313        raise ValueError(f"Body assignments do not assign all outputs: "
314                         f"missing '{n}'")
315    linalg.YieldOp(output_values)
316
317  def _get_function(self, fn_name: str) -> Callable:
318    try:
319      fn = getattr(self, f"{fn_name}")
320    except AttributeError:
321      raise ValueError(f"Function '{fn_name}' is not a known function")
322    return fn
323
324  def _cast(self,
325            type_var_name: str,
326            operand: Value,
327            is_unsigned_cast: bool = False) -> Value:
328    try:
329      to_type = self.type_mapping[type_var_name]
330    except KeyError:
331      raise ValueError(f"Unbound type variable '{type_var_name}' ("
332                       f"expected one of {self.type_mapping.keys()}")
333    if operand.type == to_type:
334      return operand
335    if _is_integer_type(to_type):
336      return self._cast_to_integer(to_type, operand, is_unsigned_cast)
337    elif _is_floating_point_type(to_type):
338      return self._cast_to_floating_point(to_type, operand, is_unsigned_cast)
339
340  def _cast_to_integer(self, to_type: Type, operand: Value,
341                       is_unsigned_cast: bool) -> Value:
342    to_width = IntegerType(to_type).width
343    operand_type = operand.type
344    if _is_floating_point_type(operand_type):
345      if is_unsigned_cast:
346        return arith.FPToUIOp(to_type, operand).result
347      return arith.FPToSIOp(to_type, operand).result
348    if _is_index_type(operand_type):
349      return arith.IndexCastOp(to_type, operand).result
350    # Assume integer.
351    from_width = IntegerType(operand_type).width
352    if to_width > from_width:
353      if is_unsigned_cast:
354        return arith.ExtUIOp(to_type, operand).result
355      return arith.ExtSIOp(to_type, operand).result
356    elif to_width < from_width:
357      return arith.TruncIOp(to_type, operand).result
358    raise ValueError(f"Unable to cast body expression from {operand_type} to "
359                     f"{to_type}")
360
361  def _cast_to_floating_point(self, to_type: Type, operand: Value,
362                              is_unsigned_cast: bool) -> Value:
363    operand_type = operand.type
364    if _is_integer_type(operand_type):
365      if is_unsigned_cast:
366        return arith.UIToFPOp(to_type, operand).result
367      return arith.SIToFPOp(to_type, operand).result
368    # Assume FloatType.
369    to_width = _get_floating_point_width(to_type)
370    from_width = _get_floating_point_width(operand_type)
371    if to_width > from_width:
372      return arith.ExtFOp(to_type, operand).result
373    elif to_width < from_width:
374      return arith.TruncFOp(to_type, operand).result
375    raise ValueError(f"Unable to cast body expression from {operand_type} to "
376                     f"{to_type}")
377
378  def _type_cast_signed(self, type_var_name: str, operand: Value) -> Value:
379    return self._cast(type_var_name, operand, False)
380
381  def _type_cast_unsigned(self, type_var_name: str, operand: Value) -> Value:
382    return self._cast(type_var_name, operand, True)
383
384  def _unary_exp(self, x: Value) -> Value:
385    if _is_floating_point_type(x.type):
386      return math.ExpOp(x).result
387    raise NotImplementedError("Unsupported 'exp' operand: {x}")
388
389  def _unary_log(self, x: Value) -> Value:
390    if _is_floating_point_type(x.type):
391      return math.LogOp(x).result
392    raise NotImplementedError("Unsupported 'log' operand: {x}")
393
394  def _unary_abs(self, x: Value) -> Value:
395    if _is_floating_point_type(x.type):
396      return math.AbsOp(x).result
397    raise NotImplementedError("Unsupported 'abs' operand: {x}")
398
399  def _unary_ceil(self, x: Value) -> Value:
400    if _is_floating_point_type(x.type):
401      return math.CeilOp(x).result
402    raise NotImplementedError("Unsupported 'ceil' operand: {x}")
403
404  def _unary_floor(self, x: Value) -> Value:
405    if _is_floating_point_type(x.type):
406      return math.FloorOp(x).result
407    raise NotImplementedError("Unsupported 'floor' operand: {x}")
408
409  def _unary_negf(self, x: Value) -> Value:
410    if _is_floating_point_type(x.type):
411      return arith.NegFOp(x).result
412    if _is_complex_type(x.type):
413      return complex.NegOp(x).result
414    raise NotImplementedError("Unsupported 'negf' operand: {x}")
415
416  def _binary_add(self, lhs: Value, rhs: Value) -> Value:
417    if _is_floating_point_type(lhs.type):
418      return arith.AddFOp(lhs, rhs).result
419    if _is_integer_type(lhs.type) or _is_index_type(lhs.type):
420      return arith.AddIOp(lhs, rhs).result
421    if _is_complex_type(lhs.type):
422      return complex.AddOp(lhs, rhs).result
423    raise NotImplementedError("Unsupported 'add' operands: {lhs}, {rhs}")
424
425  def _binary_sub(self, lhs: Value, rhs: Value) -> Value:
426    if _is_floating_point_type(lhs.type):
427      return arith.SubFOp(lhs, rhs).result
428    if _is_integer_type(lhs.type) or _is_index_type(lhs.type):
429      return arith.SubIOp(lhs, rhs).result
430    if _is_complex_type(lhs.type):
431      return complex.SubOp(lhs, rhs).result
432    raise NotImplementedError("Unsupported 'sub' operands: {lhs}, {rhs}")
433
434  def _binary_mul(self, lhs: Value, rhs: Value) -> Value:
435    if _is_floating_point_type(lhs.type):
436      return arith.MulFOp(lhs, rhs).result
437    if _is_integer_type(lhs.type) or _is_index_type(lhs.type):
438      return arith.MulIOp(lhs, rhs).result
439    if _is_complex_type(lhs.type):
440      return complex.MulOp(lhs, rhs).result
441    raise NotImplementedError("Unsupported 'mul' operands: {lhs}, {rhs}")
442
443  def _binary_max_signed(self, lhs: Value, rhs: Value) -> Value:
444    if _is_floating_point_type(lhs.type):
445      return arith.MaxFOp(lhs, rhs).result
446    if _is_integer_type(lhs.type) or _is_index_type(lhs.type):
447      return arith.MaxSIOp(lhs, rhs).result
448    raise NotImplementedError("Unsupported 'max' operands: {lhs}, {rhs}")
449
450  def _binary_max_unsigned(self, lhs: Value, rhs: Value) -> Value:
451    if _is_floating_point_type(lhs.type):
452      return arith.MaxFOp(lhs, rhs).result
453    if _is_integer_type(lhs.type) or _is_index_type(lhs.type):
454      return arith.MaxUIOp(lhs, rhs).result
455    raise NotImplementedError(
456        "Unsupported 'max_unsigned' operands: {lhs}, {rhs}")
457
458  def _binary_min_signed(self, lhs: Value, rhs: Value) -> Value:
459    if _is_floating_point_type(lhs.type):
460      return arith.MinFOp(lhs, rhs).result
461    if _is_integer_type(lhs.type) or _is_index_type(lhs.type):
462      return arith.MinSIOp(lhs, rhs).result
463    raise NotImplementedError("Unsupported 'min' operands: {lhs}, {rhs}")
464
465  def _binary_min_unsigned(self, lhs: Value, rhs: Value) -> Value:
466    if _is_floating_point_type(lhs.type):
467      return arith.MinFOp(lhs, rhs).result
468    if _is_integer_type(lhs.type) or _is_index_type(lhs.type):
469      return arith.MinUIOp(lhs, rhs).result
470    raise NotImplementedError(
471        "Unsupported 'min_unsigned' operands: {lhs}, {rhs}")
472
473
474def _infer_structured_outs(
475    op_config: LinalgStructuredOpConfig,
476    in_arg_defs: Sequence[OperandDefConfig], ins: Sequence[Value],
477    out_arg_defs: Sequence[OperandDefConfig],
478    outs: Union[Sequence[Value], OpResultList]) -> Tuple[ValueList, List[Type]]:
479  """Infers implicit outs and output types.
480
481  Respects existing contents of outs if not empty.
482
483  Returns:
484    normalized outs, output types
485  """
486  # If outs were explicitly provided, we accept them verbatim.
487  if outs:
488    return outs, [out.type for out in outs]
489
490  raise NotImplementedError(f"Output tensor inference not yet supported for "
491                            "structured ops")
492
493
494def _get_types_from_values(*values: Value) -> Sequence[Type]:
495  types = []
496  for v in values:
497    types.append(v.type)
498  return types
499
500
501def _get_operand_def_names(*operand_configs: OperandDefConfig) -> Sequence[str]:
502  return [odc.operand_def.name for odc in operand_configs]
503
504
505def _add_type_mapping(operand_config: OperandDefConfig, operand_type: Type,
506                      type_mapping: Dict[str, Type],
507                      block_arg_types: Sequence[Type]):
508  element_or_self_type = operand_type
509  # Get the element type for tensor operands and the type itself for scalars.
510  if operand_config.shape_map:
511    try:
512      element_or_self_type = ShapedType(operand_type).element_type
513    except Exception as e:
514      raise ValueError(f"Expected ShapedType but got {operand_type}") from e
515  name = operand_config.type_var.name
516  if name in type_mapping:
517    if type_mapping[name] != element_or_self_type:
518      raise ValueError(f"Cannot overwrite type mapping {name} = "
519                       f"{type_mapping[name]} by type {element_or_self_type}")
520  type_mapping[name] = element_or_self_type
521  block_arg_types.append(element_or_self_type)
522
523
524def _is_complex_type(t: Type) -> bool:
525  return ComplexType.isinstance(t)
526
527
528def _is_floating_point_type(t: Type) -> bool:
529  # TODO: Create a FloatType in the Python API and implement the switch
530  # there.
531  return (F64Type.isinstance(t) or F32Type.isinstance(t) or
532          F16Type.isinstance(t) or BF16Type.isinstance(t))
533
534
535def _is_integer_type(t: Type) -> bool:
536  return IntegerType.isinstance(t)
537
538
539def _is_index_type(t: Type) -> bool:
540  return IndexType.isinstance(t)
541
542
543def _get_floating_point_width(t: Type) -> int:
544  # TODO: Create a FloatType in the Python API and implement the switch
545  # there.
546  if F64Type.isinstance(t):
547    return 64
548  if F32Type.isinstance(t):
549    return 32
550  if F16Type.isinstance(t):
551    return 16
552  if BF16Type.isinstance(t):
553    return 16
554  raise NotImplementedError(f"Unhandled floating point type switch {t}")
555