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