1from abc
import ABC, abstractmethod
5from operator
import itemgetter
47 return self.
_mlir.pop()
76 use_accelerator =
False
77 use_memory_manager =
False
82 target_variable_name =
None
83 mlir_symbol_table = {}
84 mlir_memref_types = {}
88 type_map.update({
'void':
'void'})
89 type_map.update({
'int':
'i64'})
90 type_map.update({
'double':
'f64'})
91 type_map.update({
'double*':
'memref<?xf64>'})
92 type_map.update({
'tarch::la::Vector<2, double>':
'memref<?xi64>'})
93 type_map.update({
'tarch::la::Vector<3, double>':
'memref<?xi64>'})
94 type_map.update({
'CellData<double, double>&':
'!llvm.ptr'})
95 type_map.update({
'const FluxFunctor&':
'(!llvm.ptr, !llvm.ptr, !llvm.ptr, f64, f64, i32, !llvm.ptr) -> ()'})
96 type_map.update({
'const SourceFunctor&':
'(!llvm.ptr, !llvm.ptr, !llvm.ptr, f64, f64, !llvm.ptr) -> ()'})
97 type_map.update({
'const NonconservativeProductFunctor&':
'(!llvm.ptr, !llvm.ptr, !llvm.ptr, !llvm.ptr, f64, f64, i32, !llvm.ptr) -> ()'})
98 type_map.update({
'const MaxEigenvalueFunctorInPlace&':
'(!llvm.ptr, !llvm.ptr, !llvm.ptr, f64, f64, i32, !llvm.ptr) -> ()'})
99 type_map.update({
'const MaxEigenvalueFunctor&':
'(!llvm.ptr, !llvm.ptr, !llvm.ptr, f64, f64, i32, !llvm.ptr) -> ()'})
100 type_map.update({
'peano4::utils::LoopPlacement':
'i32'})
104 type_map.update({
'flux':
'(memref<?xf64>, !llvm.ptr, !llvm.ptr, f64, f64, i32, memref<?xf64>) -> ()'})
105 type_map.update({
'nonconservativeProduct':
'(memref<?xf64>, memref<?xf64>, !llvm.ptr, !llvm.ptr, f64, f64, i32, memref<?xf64>) -> ()'})
106 type_map.update({
'sourceTerm':
'(memref<?xf64>, !llvm.ptr, !llvm.ptr, f64, f64, memref<?xf64>) -> ()'})
107 type_map.update({
'maxEigenvalue':
'(memref<?xf64>, !llvm.ptr, !llvm.ptr, f64, f64, i32, memref<?xf64>) -> ()'})
110 type_map.update({
'source':
'memref<?xmemref<?xf64>>'})
111 type_map.update({
'maxEigenvaluesX':
'memref<?xmemref<?xf64>>'})
112 type_map.update({
'maxEigenvaluesY':
'memref<?xmemref<?xf64>>'})
113 type_map.update({
'maxEigenvaluesZ':
'memref<?xmemref<?xf64>>'})
114 type_map.update({
'maxEigenvaluesX2':
'memref<?xmemref<?xf64>>'})
115 type_map.update({
'maxEigenvaluesY2':
'memref<?xmemref<?xf64>>'})
116 type_map.update({
'maxEigenvaluesZ2':
'memref<?xmemref<?xf64>>'})
117 type_map.update({
'maxEigenvaluesXY':
'memref<?xmemref<?xf64>>'})
118 type_map.update({
'maxEigenvaluesPerVolume':
'memref<?xmemref<?xf64>>'})
122 type_map.update({
'lambdaLeft':
'memref<?xmemref<?xf64>>'})
123 type_map.update({
'lambdaRight':
'memref<?xmemref<?xf64>>'})
124 type_map.update({
'lambdaBottom':
'memref<?xmemref<?xf64>>'})
125 type_map.update({
'lambdaTop':
'memref<?xmemref<?xf64>>'})
126 type_map.update({
'lambdaNear':
'memref<?xmemref<?xf64>>'})
127 type_map.update({
'lambdaFar':
'memref<?xmemref<?xf64>>'})
128 type_map.update({
'fluxLeft':
'memref<?xmemref<?xf64>>'})
129 type_map.update({
'fluxRight':
'memref<?xmemref<?xf64>>'})
130 type_map.update({
'fluxBottom':
'memref<?xmemref<?xf64>>'})
131 type_map.update({
'fluxTop':
'memref<?xmemref<?xf64>>'})
132 type_map.update({
'fluxNear':
'memref<?xmemref<?xf64>>'})
133 type_map.update({
'fluxFar':
'memref<?xmemref<?xf64>>'})
138 type_map.update({
'flux_x':
'memref<?xmemref<?xf64>>'})
139 type_map.update({
'flux_y':
'memref<?xmemref<?xf64>>'})
140 type_map.update({
'flux_z':
'memref<?xmemref<?xf64>>'})
141 type_map.update({
'average_x':
'memref<?xmemref<?xf64>>'})
142 type_map.update({
'average_y':
'memref<?xmemref<?xf64>>'})
143 type_map.update({
'average_z':
'memref<?xmemref<?xf64>>'})
144 type_map.update({
'delta_x':
'memref<?xmemref<?xf64>>'})
145 type_map.update({
'delta_y':
'memref<?xmemref<?xf64>>'})
146 type_map.update({
'delta_z':
'memref<?xmemref<?xf64>>'})
147 type_map.update({
'nonconservative_x':
'memref<?xmemref<?xf64>>'})
148 type_map.update({
'nonconservative_y':
'memref<?xmemref<?xf64>>'})
149 type_map.update({
'nonconservative_z':
'memref<?xmemref<?xf64>>'})
152 type_map.update({
'memref<?xf64>':
'void*'})
153 type_map.update({
'!llvm.ptr':
'void*'})
154 type_map.update({
'f64':
'double'})
155 type_map.update({
'i32':
'int'})
156 type_map.update({
'i64':
'long'})
157 type_map.update({
'i1':
'bool'})
160 descriptor2DType =
"!llvm.struct<(ptr, ptr, i64, array<2 x i64>, array<2 x i64>)>"
161 descriptor1DType =
"!llvm.struct<(ptr, ptr, i64, array<1 x i64>, array<1 x i64>)>"
164 patchTypes = [descriptor2DType, descriptor2DType, descriptor1DType, descriptor1DType,
165 descriptor1DType, descriptor1DType, descriptor1DType,
'i64',
'i64',
'i64',
'ptr',
'ptr']
167 struct_map.update({
'patchData': [
168 (
'QIn' , f
"{descriptor2DType}|||!llvm.struct<({', '.join(patchTypes)})>"),
169 (
'QOut' , f
"{descriptor2DType}|||!llvm.struct<({', '.join(patchTypes)})>"),
170 (
'cellCentre' , f
"{descriptor1DType}|||!llvm.struct<({', '.join(patchTypes)})>"),
171 (
'cellSize' , f
"{descriptor1DType}|||!llvm.struct<({', '.join(patchTypes)})>"),
172 (
't' , f
"{descriptor1DType}|||!llvm.struct<({', '.join(patchTypes)})>"),
173 (
'dt' , f
"{descriptor1DType}|||!llvm.struct<({', '.join(patchTypes)})>"),
174 (
'id' , f
"{descriptor1DType}|||!llvm.struct<({', '.join(patchTypes)})>"),
175 (
'numberOfCells' , f
"i64|||!llvm.struct<({', '.join(patchTypes)})>"),
176 (
'memoryLocation' , f
"i64|||!llvm.struct<({', '.join(patchTypes)})>"),
177 (
'targetDevice' , f
"i64|||!llvm.struct<({', '.join(patchTypes)})>"),
178 (
'QOut_legacy' , f
"!llvm.ptr|||!llvm.struct<({', '.join(patchTypes)})>"),
179 (
'maxEigenvalue' , f
"memref<?xf64>|||!llvm.struct<({', '.join(patchTypes)})>")]})
183 for struct_name, fields
in struct_map.items():
184 for field_name, field_info
in fields:
185 target_type = field_info.split(
'|||')[0]
187 if target_type == descriptor2DType:
188 type_map.update({field_name:
'memref<?x?xf64>'})
189 elif target_type == descriptor1DType:
193 if field_name
in [
'cellCentre',
'cellSize',
'h',
'x']:
194 type_map.update({field_name:
'memref<?xi64>'})
196 type_map.update({field_name:
'memref<?xf64>'})
197 elif target_type ==
'i64':
198 type_map.update({field_name:
'i64'})
199 elif target_type ==
'!llvm.ptr':
201 type_map.update({field_name:
'!llvm.ptr'})
204 type_map.update({field_name: target_type})
213 if parent
is not None:
240 if isinstance(type_string, str):
241 if type_string ==
"int":
243 elif type_string ==
"double":
245 elif type_string ==
"bool":
251 indent =
' ' * indent_level * Node.spaces_per_tab
256 inner_memref_id = f
"%{Node.get_mlir_id()}"
257 mlir_code += indent + f
"{inner_memref_id} = memref.load {memref_mlir}[{row_index}] {{name = \"{memref_mlir.strip('%')}\"}} : memref<?xmemref<?xf64>>\n"
260 ptr_idx_id = f
"%{Node.get_mlir_id()}"
261 mlir_code += indent + f
"{ptr_idx_id} = memref.extract_aligned_pointer_as_index {inner_memref_id} : memref<?xf64> -> index {{name = \"{memref_mlir.strip('%')}\"}}\n"
264 element_size_bytes = f
"%{Node.get_mlir_id()}"
265 mlir_code += indent + f
"{element_size_bytes} = arith.constant 8 : index\n"
267 offset_bytes_id = f
"%{Node.get_mlir_id()}"
268 mlir_code += indent + f
"{offset_bytes_id} = arith.muli {col_index}, {element_size_bytes} : index\n"
271 ptr_with_offset = f
"%{Node.get_mlir_id()}"
272 mlir_code += indent + f
"{ptr_with_offset} = arith.addi {ptr_idx_id}, {offset_bytes_id} : index\n"
275 ptr_i64 = f
"%{Node.get_mlir_id()}"
276 mlir_code += indent + f
"{ptr_i64} = arith.index_cast {ptr_with_offset} : index to i64\n"
278 result_ptr = f
"%{Node.get_mlir_id()}"
279 mlir_code += indent + f
"{result_ptr} = llvm.inttoptr {ptr_i64} : i64 to !llvm.ptr\n"
282 col_width_id = f
"%{Node.get_mlir_id()}"
283 one_const = f
"%{Node.get_mlir_id()}"
284 mlir_code += indent + f
"{one_const} = arith.constant 1 : index\n"
285 mlir_code += indent + f
"{col_width_id} = memref.dim {memref_mlir}, {one_const} : memref<?x?xf64>\n"
288 row_offset = f
"%{Node.get_mlir_id()}"
289 mlir_code += indent + f
"{row_offset} = arith.muli {row_index}, {col_width_id} : index\n"
292 elem_offset = f
"%{Node.get_mlir_id()}"
293 mlir_code += indent + f
"{elem_offset} = arith.addi {row_offset}, {col_index} : index\n"
296 element_size_bytes = f
"%{Node.get_mlir_id()}"
297 mlir_code += indent + f
"{element_size_bytes} = arith.constant 8 : index\n"
299 offset_bytes_id = f
"%{Node.get_mlir_id()}"
300 mlir_code += indent + f
"{offset_bytes_id} = arith.muli {elem_offset}, {element_size_bytes} : index\n"
303 base_ptr_idx = f
"%{Node.get_mlir_id()}"
304 mlir_code += indent + f
"{base_ptr_idx} = memref.extract_aligned_pointer_as_index {memref_mlir} : memref<?x?xf64> -> index {{name = \"{memref_mlir.strip('%')}\"}}\n"
307 ptr_with_offset = f
"%{Node.get_mlir_id()}"
308 mlir_code += indent + f
"{ptr_with_offset} = arith.addi {base_ptr_idx}, {offset_bytes_id} : index\n"
311 ptr_i64 = f
"%{Node.get_mlir_id()}"
312 mlir_code += indent + f
"{ptr_i64} = arith.index_cast {ptr_with_offset} : index to i64\n"
314 result_ptr = f
"%{Node.get_mlir_id()}"
315 mlir_code += indent + f
"{result_ptr} = llvm.inttoptr {ptr_i64} {{name = \"{memref_mlir.strip('%')}\"}} : i64 to !llvm.ptr\n"
317 return mlir_code, result_ptr
324 desc_id = f
"%{cls.get_mlir_id()}"
325 desc_0_id = f
"%{cls.get_mlir_id()}"
326 desc_1_id = f
"%{cls.get_mlir_id()}"
327 desc_2_id = f
"%{cls.get_mlir_id()}"
328 desc_3_id = f
"%{cls.get_mlir_id()}"
329 desc_4_id = f
"%{cls.get_mlir_id()}"
333 memref_id = target_var_name
335 memref_id = f
"%{cls.get_mlir_id()}"
338 c0_i64_id = f
"%{cls.get_mlir_id()}"
339 c1_i64_id = f
"%{cls.get_mlir_id()}"
343 if target_type.startswith(
'memref<?x?'):
345 struct_type =
'!llvm.struct<(ptr, ptr, i64, array<2 x i64>, array<2 x i64>)>'
349 struct_type =
'!llvm.struct<(ptr, ptr, i64, array<1 x i64>, array<1 x i64>)>'
353 if target_type.startswith(
'memref<?xmemref')
or size_hint == 1:
355 elif target_type.startswith(
'memref<?xi64>')
and size_hint == 2:
357 elif target_type.startswith(
'memref<?xi64>'):
365 mlir_str += indent + f
"{c0_i64_id} = arith.constant 0 : i64\n"
366 mlir_str += indent + f
"{c1_i64_id} = arith.constant 1 : i64\n"
370 size_id = f
"%{cls.get_mlir_id()}"
371 size_i64_id = f
"%{cls.get_mlir_id()}"
372 mlir_str += indent + f
"{size_id} = arith.constant {size_val} : index\n"
373 mlir_str += indent + f
"{size_i64_id} = arith.index_cast {size_id} : index to i64\n"
374 size_ref = size_i64_id
379 mlir_str += indent + f
"{desc_id} = llvm.mlir.undef : {struct_type}\n"
380 mlir_str += indent + f
"{desc_0_id} = llvm.insertvalue {ptr_id}, {desc_id}[0] : {struct_type}\n"
381 mlir_str += indent + f
"{desc_1_id} = llvm.insertvalue {ptr_id}, {desc_0_id}[1] : {struct_type}\n"
382 mlir_str += indent + f
"{desc_2_id} = llvm.insertvalue {c0_i64_id}, {desc_1_id}[2] : {struct_type}\n"
383 mlir_str += indent + f
"{desc_3_id} = llvm.insertvalue {size_ref}, {desc_2_id}[3, 0] : {struct_type}\n"
384 if num_dimensions == 2:
386 desc_3_1_id = f
"%{cls.get_mlir_id()}"
387 mlir_str += indent + f
"{desc_3_1_id} = llvm.insertvalue {size_ref}, {desc_3_id}[3, 1] : {struct_type}\n"
388 mlir_str += indent + f
"{desc_4_id} = llvm.insertvalue {c1_i64_id}, {desc_3_1_id}[4, 0] : {struct_type}\n"
389 desc_final_id = f
"%{cls.get_mlir_id()}"
390 mlir_str += indent + f
"{desc_final_id} = llvm.insertvalue {c1_i64_id}, {desc_4_id}[4, 1] : {struct_type}\n"
393 mlir_str += indent + f
"{desc_4_id} = llvm.insertvalue {c1_i64_id}, {desc_3_id}[4, 0] : {struct_type}\n"
394 desc_final_id = desc_4_id
395 mlir_str += indent + f
"{memref_id} = builtin.unrealized_conversion_cast {desc_final_id} : {struct_type} to {target_type} {{name = \"{memref_id.strip('%')}\"}}\n"
397 return mlir_str, memref_id
400 def create_memref_from_extracted_descriptor(cls, alloc_ptr_id, aligned_ptr_id, offset_id, sizes_0_id, sizes_1_id, strides_0_id, strides_1_id, target_type, struct_type, indent_level=0, target_var_name=None, is_2d=False):
405 memref_id = target_var_name
407 memref_id = f
"%{cls.get_mlir_id()}"
410 desc_id = f
"%{cls.get_mlir_id()}"
411 desc_0_id = f
"%{cls.get_mlir_id()}"
412 desc_1_id = f
"%{cls.get_mlir_id()}"
413 desc_2_id = f
"%{cls.get_mlir_id()}"
414 desc_3_id = f
"%{cls.get_mlir_id()}"
415 desc_4_id = f
"%{cls.get_mlir_id()}"
417 c1_i64_id = f
"%{cls.get_mlir_id()}"
422 mlir_str += indent + f
"{c1_i64_id} = arith.constant 1 : i64\n"
426 mlir_str += indent + f
"{desc_id} = llvm.mlir.undef : {struct_type}\n"
427 mlir_str += indent + f
"{desc_0_id} = llvm.insertvalue {alloc_ptr_id}, {desc_id}[0] : {struct_type}\n"
428 mlir_str += indent + f
"{desc_1_id} = llvm.insertvalue {aligned_ptr_id}, {desc_0_id}[1] : {struct_type}\n"
429 mlir_str += indent + f
"{desc_2_id} = llvm.insertvalue {offset_id}, {desc_1_id}[2] : {struct_type}\n"
430 mlir_str += indent + f
"{desc_3_id} = llvm.insertvalue {sizes_0_id}, {desc_2_id}[3, 0] : {struct_type}\n"
435 desc_3_1_id = f
"%{cls.get_mlir_id()}"
436 mlir_str += indent + f
"{desc_3_1_id} = llvm.insertvalue {sizes_1_id}, {desc_3_id}[3, 1] : {struct_type}\n"
437 mlir_str += indent + f
"{desc_4_id} = llvm.insertvalue {strides_0_id}, {desc_3_1_id}[4, 0] : {struct_type}\n"
438 desc_final_id = f
"%{cls.get_mlir_id()}"
439 mlir_str += indent + f
"{desc_final_id} = llvm.insertvalue {strides_1_id}, {desc_4_id}[4, 1] : {struct_type}\n"
442 mlir_str += indent + f
"{desc_4_id} = llvm.insertvalue {strides_0_id}, {desc_3_id}[4, 0] : {struct_type}\n"
443 desc_final_id = desc_4_id
445 mlir_str += indent + f
"{memref_id} = builtin.unrealized_conversion_cast {desc_final_id} : {struct_type} to {target_type} {{name = \"{memref_id.strip('%')}\"}}\n"
447 return mlir_str, memref_id
469class Statement(Node):
503 return ' ' * indent_level * Node.spaces_per_tab +
'void'
506 return ' ' * indent_level * Node.spaces_per_tab +
'void'
509 return ' ' * indent_level * Node.spaces_per_tab +
'void'
512 return ' ' * indent_level * Node.spaces_per_tab +
'void'
520 def __init__(self, id, argument_type: Type, namespace =
None, is_function =
False):
527 if Node.use_accelerator ==
True:
528 if type(self.
_type)
is TCustom
and self.
_type._type.find(
"CellData") != -1:
529 index = self.
_type._type.find(
"CellData")
530 self.
_type._type = self.
_type._type[0:index] +
"CopyCellDataGPU" + self.
_type._type[index + 8:]
543 return f
"%{self.id} : {self._type.print_mlir(0)}"
546 return ' ' * indent_level * Node.spaces_per_tab + f
"Argument: {self.id}"
553 def __init__(self, id, return_type: Type, template =
None, namespaces = [], stateless =
False, top_level =
False):
562 FunctionDefinition._stateless = stateless
571 if statement
is not None:
572 self.
_body.append(statement)
579 Node.type_map.update({argument.id: argument._namespace})
580 if Node.type_map.get(argument._type.print_cpp())
is None:
581 Node.type_map.update({argument._type.print_cpp(): argument._namespace})
584 namespace_header = f
"""namespace {"::".join(self._namespaces)} {{"""
585 namespace_footer =
"}"
588 template_strings = [argument.print_cpp()
for argument
in self.
_template]
589 template_string =
"template <" +
",".join(template_strings) +
">\n"
591 current_indent = indent_level + 1
592 argument_prints = [argument.print_cpp()
for argument
in self.
_arguments]
594 return f
"""{namespace_header}
595{template_string}{' ' * current_indent * Node.spaces_per_tab}{self._return_type.print_cpp()} {self.id}({', '.join(argument_prints)});
600 namespace_header = f
"""namespace {"::".join(self._namespaces)} {{"""
601 namespace_footer =
"}"
604 template_strings = [argument.print_cpp()
for argument
in self.
_template]
605 template_string =
"template <" +
",".join(template_strings) +
">\n"
607 current_indent = indent_level + 1
608 statement_prints = [statement.print_cpp(current_indent + 1)
for statement
in self.
_body]
609 argument_prints = [argument.print_cpp()
for argument
in self.
_arguments]
611 return f
"""{namespace_header}
612{template_string}{' ' * current_indent * Node.spaces_per_tab}{self._return_type.print_cpp()} {self.id}({', '.join(argument_prints)}) {{
613{os.linesep.join(statement_prints)}
614{' ' * current_indent * Node.spaces_per_tab}}}
619 namespace_header = f
"""namespace {"::".join(self._namespaces)} {{"""
620 namespace_footer =
"}"
623 template_strings = [argument.print_omp()
for argument
in self.
_template]
624 template_string =
"template <" +
",".join(template_strings) +
">\n"
626 current_indent = indent_level + 1
627 statement_prints = [statement.print_omp(current_indent + 1)
for statement
in self.
_body]
628 argument_prints = [argument.print_omp()
for argument
in self.
_arguments]
630 return f
"""{namespace_header}
631{template_string}{' ' * current_indent * Node.spaces_per_tab}{self._return_type.print_omp()} {self.id}({', '.join(argument_prints)}) {{
632{os.linesep.join(statement_prints)}
633{' ' * current_indent * Node.spaces_per_tab}}}
638 namespace_header = f
"""namespace {"::".join(self._namespaces)} {{"""
639 namespace_footer =
"}"
642 template_strings = [argument.print_sycl()
for argument
in self.
_template]
643 template_string =
"template <" +
",".join(template_strings) +
">\n"
645 current_indent = indent_level + 1
646 statement_prints = [statement.print_sycl(current_indent + 1)
for statement
in self.
_body]
647 argument_prints = [argument.print_sycl()
for argument
in self.
_arguments]
648 statement_prints.insert(0,
"::sycl::queue& queue = tarch::accelerator::getSYCLQueue(targetDevice);")
649 statement_prints.insert(1,
"size_t range0, range1, range2;")
651 data_copy_prints = []
652 for dataBlockCreation
in FunctionDefinition._syclDataToCopy:
653 if len(dataBlockCreation._dataBlock._iteration_range) > 1:
655 for d
in dataBlockCreation._dataBlock._iteration_range[0:-1]:
656 step_size = step_size * (d[1] - d[0])
657 size = step_size * (dataBlockCreation._dataBlock._iteration_range[-1][1] - dataBlockCreation._dataBlock._iteration_range[-1][0])
659 return f
"""{namespace_header}
660{template_string}{' ' * current_indent * Node.spaces_per_tab}{self._return_type.print_sycl()} {self.id}({', '.join(argument_prints)}) {{
661{os.linesep.join(statement_prints)}
662{os.linesep.join(data_copy_prints)}
663{' ' * current_indent * Node.spaces_per_tab}}}
671 Node.reset_mlir_ids()
680 mlir_signature = argument._namespace
683 if argument._is_function:
686 func_decl = f
"func.func private @{argument.id}_bridge{mlir_signature}"
687 globals_mlir.append(
' ' * (indent_level) * Node.spaces_per_tab + func_decl)
693 "MLIR code generation requires template parameters to determine bridge function signatures. "
694 "Ensure the function has template parameters that specify the required functors."
698 function_global_loads = []
702 new_id = Node.get_mlir_id()
703 globals_mlir.append(
' ' * (indent_level) * Node.spaces_per_tab +f
"llvm.mlir.global external @{argument.id}() : {argument._type.print_mlir(0)}")
704 function_global_loads.append(
' ' * (indent_level + 1) * Node.spaces_per_tab + f
"%{new_id} = llvm.mlir.addressof @{argument.id} : !llvm.ptr")
707 arg_mlir_type = argument._type.print_mlir(0)
710 if arg_mlir_type
and arg_mlir_type.strip() !=
'' and len(arg_mlir_type.strip()) > 0:
711 function_global_loads.append(
' ' * (indent_level + 1) * Node.spaces_per_tab + f
"%{argument.id} = llvm.load %{new_id} {{name = \"{argument.id.strip('%')}\"}} : !llvm.ptr -> {arg_mlir_type}")
714 function_global_loads.append(
' ' * (indent_level + 1) * Node.spaces_per_tab + f
"// Skipping load for {argument.id} (empty type: '{arg_mlir_type}')")
718 Node.context.push_block()
719 argument_prints = [argument.print_mlir(indent_level).strip(
'')
for argument
in self.
_arguments]
720 Node.context.push_block()
722 statement_prints = [statement.print_mlir(indent_level + 1)
for statement
in self.
_body]
723 statements = Node.context.pop_block().get_mlir()
724 statements += statement_prints
727 return_type_str = f
" -> ({return_type})" if return_type !=
'void' else ''
731{os.linesep.join(globals_mlir)}
732 func.func @{self.id + "_omp" if Node.use_mlir_omp else self.id}({', '.join(argument_prints)}){return_type_str} {{
733{os.linesep.join(function_global_loads)}
734{os.linesep.join(statements)}
740 Node.context.push_block()
741 argument_prints = [argument.print_mlir(indent_level)
for argument
in self.
_arguments]
742 Node.context.push_block()
743 statement_prints = [statement.print_mlir(indent_level + 1)
for statement
in self.
_body]
744 statements = Node.context.pop_block().get_mlir()
745 statements += statement_prints
747 return_type_str = f
" -> ({return_type})" if return_type !=
'void' else ''
749 func.func @{self.id}({', '.join(argument_prints)}){return_type_str} {{
750{os.linesep.join(statements)}
756 statement_prints = [statement.print_tree(indent_level + 1)
for statement
in self.
_body]
757 argument_prints = [argument.print_tree(indent_level + 1)
for argument
in self.
_arguments]
759FunctionDefinition: {self.id}:
760{os.linesep.join(argument_prints)}
761{os.linesep.join(statement_prints)}"""
764 namespace_header = f
"""namespace {"::".join(self._namespaces)} {{"""
765 namespace_footer =
"}"
767 function_call_string = self.
id
769 template_strings = [argument.print_cpp()
for argument
in self.
_template]
770 template_string =
"template <" +
",".join(template_strings) +
">\n"
771 function_call_string +=
"<" +
",".join([template.split(
' ')[1]
for template
in template_strings]) +
">"
774 arguments.append(
Argument(
"measurement",
TCustom(
"tarch::timing::Measurement&")))
776 current_indent = indent_level + 1
777 argument_prints = [argument.print_cpp()
for argument
in arguments]
779 return f
"""{namespace_header}
780{template_string}{' ' * current_indent * Node.spaces_per_tab}{self._return_type.print_cpp()} {self.id}({', '.join(argument_prints)}) {{
781tarch::timing::Watch watch("{"::".join(self._namespaces)}", "{self.id}", false, true);
782{function_call_string}({",".join([argument.print_cpp().split(' ')[-1] for argument in self._arguments])});
784measurement.setValue(watch.getCalendarTime());
785{' ' * current_indent * Node.spaces_per_tab}}}
790 namespace_header = f
"""namespace {"::".join(self._namespaces)} {{"""
791 namespace_footer =
"}"
794 template_strings = [argument.print_cpp()
for argument
in self.
_template]
795 template_string =
"template <" +
",".join(template_strings) +
">\n"
798 arguments.append(
Argument(
"measurement",
TCustom(
"tarch::timing::Measurement&")))
800 current_indent = indent_level + 1
801 argument_prints = [argument.print_cpp()
for argument
in arguments]
803 return f
"""{namespace_header}
804{template_string}{' ' * current_indent * Node.spaces_per_tab}{self._return_type.print_cpp()} {self.id}({', '.join(argument_prints)});
814 return f
"""{self._output.print_cpp()} << {self._statement.print_cpp()};"""
817 return f
"""{self._output.print_omp()} << {self._statement.print_omp()};"""
820 return f
"""{self._output.print_sycl()} << {self._statement.print_sycl()};"""
837 if self.
type is None:
839 self.
type._parent = self
848 return Name.variables[self.
id].
index(i)
851 return Name.variables[self.
id].
index(i, j)
869 if type(rhs)
is Name:
870 return (Name.variables[self.
id] > Name.variables[rhs.id])
871 return (Name.variables[self.
id] > rhs)
874 if type(rhs)
is Name:
875 if rhs.id
in Name.variables
and self.
id in Name.variables:
876 return (self.
_value == Name.variables[rhs.id])
880 if self.
id in Name.variables:
881 return (Name.variables[self.
id] == rhs)
889 return ' ' * indent_level * Node.spaces_per_tab + self.
id
892 return ' ' * indent_level * Node.spaces_per_tab + self.
id
895 return ' ' * indent_level * Node.spaces_per_tab + self.
id
902 return ' ' * indent_level * Node.spaces_per_tab +
"Name:" + self.
id
907class Integer(Expression):
916 if type(rhs)
is Integer:
922 if type(rhs)
is Integer:
924 if type(rhs)
is UnaryOperation
and rhs._operation ==
"-":
927 if type(rhs)
is Integer:
944 if type(rhs)
is Name:
945 if rhs.id
in Name.variables:
946 return (self.
_value < Name.variables[rhs.id])
949 return (self.
_value < rhs)
952 if type(rhs)
is Name:
953 if rhs.id
in Name.variables:
954 return (self.
_value > Name.variables[rhs.id])
957 return (self.
_value > rhs)
960 if type(rhs)
is Integer:
961 return (self.
_value == rhs._value)
969 struct, field = self.
_string.split(
'.')
970 if struct
in Node.struct_map:
971 struct_list = Node.struct_map[struct]
973 for item
in struct_list:
977 target_type = mlir_type.split(
'|||')[0]
if '|||' in mlir_type
else mlir_type
979 if target_type ==
'i64':
982 elif target_type ==
'i32':
984 elif target_type ==
'f64':
991 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_value)
993 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_string)
997 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_value)
999 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_string)
1003 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_value)
1005 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_string)
1008 indent =
' ' * indent_level * Node.spaces_per_tab
1013 self.
mlir_id(Node.get_mlir_id())
1014 mlir_str =
"".join(Node.context.get_block().pop_mlir())
1015 mlir_str += indent + f
"%{self.mlir_id()} = arith.constant {self._value} : {self.get_type().print_mlir(0)}\n"
1016 Node.context.push_block()
1017 Node.context.get_block().mlir_append(mlir_str)
1018 id = f
"%{self.mlir_id()}"
1028 struct, field = x.split(
'.')
1029 if struct
in Node.struct_map:
1030 struct_list = Node.struct_map[struct]
1031 idx = [item[0]
for item
in struct_list].index(field)
1032 mlir_type = struct_list[idx][1]
1034 mlir_type_struct = mlir_type.split(
'|||')[1]
if '|||' in mlir_type
else mlir_type
1036 mlir_str = f
"llvm.getelementptr %{struct}[0, {idx}] : (!llvm.ptr) -> !llvm.ptr, {mlir_type_struct}"
1042 return ' ' * indent_level * Node.spaces_per_tab +
"Integer: " + str(self.
_value)
1044 return ' ' * indent_level * Node.spaces_per_tab +
"Integer: " + str(self.
_string)
1059 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_value)
1062 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_value)
1065 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_value)
1072 return ' ' * indent_level * Node.spaces_per_tab +
"Boolean: " + str(self.
_value)
1086 return ' ' * indent_level * Node.spaces_per_tab + self.
_value
1089 return ' ' * indent_level * Node.spaces_per_tab + self.
_value
1092 return ' ' * indent_level * Node.spaces_per_tab + self.
_value
1095 indent =
' ' * indent_level * Node.spaces_per_tab
1099 struct, field = x.split(
'.')
1100 if struct
in Node.struct_map:
1101 struct_list = Node.struct_map[struct]
1102 idx = [item[0]
for item
in struct_list].index(field)
1103 mlir_type_target, mlir_type_struct = struct_list[idx][1].split(
'|||')
1106 struct_val_id = f
"%{Node.get_mlir_id()}"
1109 if field ==
"maxEigenvalue":
1110 field_id = f
"%{field}s"
1112 field_id = f
"%{Node.get_mlir_id()}"
1115 mlir_str = indent + f
"{struct_val_id} = llvm.load %{struct} {{name = \"{struct.strip('%')}\"}}: !llvm.ptr -> {mlir_type_struct}\n"
1119 if not mlir_type_target.startswith(
"!llvm.struct<"):
1121 ptr_field_id = f
"%{field}{'_ptr' if field != 'maxEigenvalue' else '_ptr'}"
1122 mlir_str += indent + f
"{ptr_field_id} = llvm.extractvalue {struct_val_id}[{idx}] : {mlir_type_struct}\n"
1123 field_id = ptr_field_id
1126 mlir_str += indent + f
"{field_id} = llvm.extractvalue {struct_val_id}[{idx}] : {mlir_type_struct}\n"
1131 if mlir_type_target.startswith(
"!llvm.struct<"):
1133 ptr_getelementptr_id = f
"%{Node.get_mlir_id()}"
1134 ptr_load_id = f
"%{Node.get_mlir_id()}"
1135 mlir_str += indent + f
"{ptr_getelementptr_id} = llvm.getelementptr %{struct}[0, {idx}] : (!llvm.ptr) -> !llvm.ptr, {mlir_type_struct}\n"
1136 mlir_str += indent + f
"{ptr_load_id} = llvm.load {ptr_getelementptr_id} {{name = \"{struct.strip('%')}\"}}: !llvm.ptr -> {mlir_type_target}\n"
1139 alloc_ptr_id = f
"%{Node.get_mlir_id()}"
1140 aligned_ptr_id = f
"%{Node.get_mlir_id()}"
1141 offset_id = f
"%{Node.get_mlir_id()}"
1142 sizes_0_id = f
"%{Node.get_mlir_id()}"
1143 strides_0_id = f
"%{Node.get_mlir_id()}"
1145 mlir_str += indent + f
"{alloc_ptr_id} = llvm.extractvalue {ptr_load_id}[0] : {mlir_type_target}\n"
1146 mlir_str += indent + f
"{aligned_ptr_id} = llvm.extractvalue {ptr_load_id}[1] : {mlir_type_target}\n"
1147 mlir_str += indent + f
"{offset_id} = llvm.extractvalue {ptr_load_id}[2] : {mlir_type_target}\n"
1148 mlir_str += indent + f
"{sizes_0_id} = llvm.extractvalue {ptr_load_id}[3, 0] : {mlir_type_target}\n"
1149 mlir_str += indent + f
"{strides_0_id} = llvm.extractvalue {ptr_load_id}[4, 0] : {mlir_type_target}\n"
1152 is_2d =
"array<2 x i64>" in mlir_type_target
1159 sizes_1_id = f
"%{Node.get_mlir_id()}"
1160 strides_1_id = f
"%{Node.get_mlir_id()}"
1161 mlir_str += indent + f
"{sizes_1_id} = llvm.extractvalue {ptr_load_id}[3, 1] : {mlir_type_target}\n"
1162 mlir_str += indent + f
"{strides_1_id} = llvm.extractvalue {ptr_load_id}[4, 1] : {mlir_type_target}\n"
1166 if field
in Node.type_map:
1167 type_entry = Node.type_map[field]
1169 if isinstance(type_entry, str):
1170 target_memref_type = type_entry
1173 target_memref_type = type_entry.print_mlir(0).strip()
1178 target_memref_type =
"memref<?x?xf64>"
1181 target_memref_type =
"memref<?xf64>"
1184 conversion_mlir, memref_id = Node.create_memref_from_extracted_descriptor(
1195 target_var_name=Node.target_variable_name,
1197 mlir_str += conversion_mlir
1198 final_id = memref_id
1199 elif 'memref' in mlir_type_target
or mlir_type_target ==
'!llvm.ptr':
1202 clean_target_type = mlir_type_target.strip(
'()')
1204 conversion_mlir, memref_id = Node.create_memref_from_ptr(
1208 target_var_name=Node.target_variable_name)
1209 mlir_str += conversion_mlir
1210 final_id = memref_id
1213 Node.context.get_block().mlir_append(mlir_str)
1218 return ' ' * indent_level * Node.spaces_per_tab +
"String: " + self.
_value
1222 def __init__(self, value, string = None, reference = False):
1237 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_value)
1239 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_string)
1243 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_value)
1245 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_string)
1249 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_value)
1251 return ' ' * indent_level * Node.spaces_per_tab + str(self.
_string)
1254 indent =
' ' * indent_level * Node.spaces_per_tab
1259 self.
mlir_id(Node.get_mlir_id())
1260 mlir_str =
"".join(Node.context.get_block().pop_mlir())
1261 mlir_str += indent + f
"%{self.mlir_id()} = arith.constant {self._value} : {self.get_type().print_mlir(0)}\n"
1262 Node.context.push_block()
1263 Node.context.get_block().mlir_append(mlir_str)
1264 id = f
"%{self.mlir_id()}"
1271 return ' ' * indent_level * Node.spaces_per_tab +
"Float: " + str(self.
_value)
1275 def __init__(self, iteration_range, internal, requires_memory_allocation, id = None, underlying_type = String(
"double")):
1292 for i
in range(len(offset)):
1293 if type(offset[i])
is list:
1294 if offset[i][0]
is None:
1296 if offset[i][1]
is None:
1300 output._iteration_range[i][1] = offset[i][1]
1301 output._iteration_range[i][0] = offset[i][0]
1303 raise Exception(
"Invalid index")
1305 for i
in range(len(output._offset), len(output._iteration_range)):
1306 output._offset.append(
Integer(0))
1312 for i
in range(0, len(dimensions) - 1):
1314 if Node.use_accelerator ==
False and type(self)
is not FaceDataBlock:
1315 index_list = indices[0:-1]
1317 index_list = indices
1319 if type(dimensions[-len(index_list)][1] - dimensions[-len(index_list)][0])
is Integer
and (dimensions[-len(index_list)][1] - dimensions[-len(index_list)][0])._value == 1:
1320 index = self.
_offset[offset_start_index]
1322 if type(self.
_offset[offset_start_index + 0].
get_type())
is TDataBlock:
1325 index = index_list[0] + self.
_offset[offset_start_index + 0]
1327 index = factors[-len(index_list)] * index
1328 for i
in range(1, len(index_list)):
1329 if type(self.
_offset[offset_start_index + i].
get_type())
is TDataBlock:
1332 temp = index_list[i] + self.
_offset[offset_start_index + i]
1333 index = index + factors[-len(index_list) + i] * temp
1340 if Node.use_accelerator ==
False:
1353 return ' ' * indent_level * Node.spaces_per_tab + self.
id
1359 return ' ' * indent_level * Node.spaces_per_tab + self.
id
1365 return ' ' * indent_level * Node.spaces_per_tab + self.
id
1372 return "%" + self.
id
1376 return ' ' * indent_level * Node.spaces_per_tab + f
"""DataBlock:
1377{self._internal.print_tree(indent_level + 1)}"""
1379 return ' ' * indent_level * Node.spaces_per_tab +
"DataBlock: " + self.
id
1382FaceDataBlock distinguishes itself from normal DataBlocks by the fact that its internal array is always 1d.
1385 def __init__(self, iteration_range, internal, requires_memory_allocation, id = None, underlying_type = String(
"double")):
1386 super().
__init__(iteration_range, internal, requires_memory_allocation, id, underlying_type)
1392 for i
in range(len(offset)):
1393 if type(offset[i])
is list:
1394 if offset[i][0]
is None:
1396 if offset[i][1]
is None:
1400 output._iteration_range[i][1] = offset[i][1]
1401 output._iteration_range[i][0] = offset[i][0]
1403 raise Exception(
"Invalid index")
1405 for i
in range(len(output._offset), len(output._iteration_range)):
1406 output._offset.append(
Integer(0))
1424 for i
in range(len(self.
_dataBlock._memory_range) - 2, -1, -1):
1427 for i
in range(len(self.
_dataBlock._memory_range) - 1, -1, -1):
1428 index.append(loops[i].get_iteration_variable())
1432 for i
in range(0, len(loops) - 1):
1433 loops[i].add_statement(loops[i + 1])
1436 for i
in range(0, len(self.
_dataBlock._memory_range) - 2):
1439 for loop
in loops[::-1]:
1444 return f
"""log.open("{self._filename.print_cpp()}");
1445{self._loop.print_cpp()}
1462 def __init__(self, value: Expression, index: Expression):
1471 return isinstance(self.
_value, Subscript)
1477 return ' ' * indent_level * Node.spaces_per_tab + f
"{self._value.print_cpp()}[{self._index.print_cpp()}]"
1480 return ' ' * indent_level * Node.spaces_per_tab + f
"{self._value.print_omp()}[{self._index.print_omp()}]"
1483 return ' ' * indent_level * Node.spaces_per_tab + f
"{self._value.print_sycl()}[{self._index.print_sycl()}]"
1486 indent =
' ' * indent_level * Node.spaces_per_tab
1489 while isinstance(tmp, Subscript):
1495 type_str = value_type.print_cpp().strip()
1497 mlir_type = Node.type_map.get(type_str)
1498 if mlir_type
is None:
1500 mlir_type =
'' + value_type.print_mlir(0)
1502 mlir_str =
"".join(Node.context.get_block().pop_mlir())
1506 mlir_str +=
"".join(Node.context.get_block().pop_mlir())
1508 if not getattr(self.
_index,
'_index_type',
False):
1509 new_index = f
"%{Node.get_mlir_id()}"
1512 mlir_type_cast = index_expr_type.print_mlir(0).strip()
1513 mlir_str += indent + new_index + f
" = arith.index_cast {index.strip(' ')}: {mlir_type_cast} to index\n"
1519 if isinstance(self.
_value, String):
1525 if isinstance(base_type, TDataBlock):
1526 memref_type = base_type.print_mlir(0).strip()
1529 if "x?x" in memref_type:
1532 result_id = f
"%{Node.get_mlir_id()}"
1533 mlir_str += indent + f
"{result_id} = memref.load {base_mlir}[%0, {self._new_index}] {{name = \"{base_mlir.strip('%')}\"}} : {memref_type}\n"
1534 Node.context.push_block()
1535 Node.context.get_block().mlir_append(mlir_str)
1539 result_id = f
"%{Node.get_mlir_id()}"
1540 mlir_str += indent + f
"{result_id} = memref.load {base_mlir}[{self._new_index}] {{name = \"{base_mlir.strip('%')}\"}} : {memref_type}\n"
1541 Node.context.push_block()
1542 Node.context.get_block().mlir_append(mlir_str)
1548 Node.context.push_block()
1549 Node.context.get_block().mlir_append(mlir_str)
1556 result_id = f
"%{Node.get_mlir_id()}"
1559 outer_base = base_expr.print_mlir(indent_level)
1562 if outer_base
in Node.mlir_memref_types:
1563 actual_memref_type = Node.mlir_memref_types[outer_base]
1566 base_expr_type = base_expr.get_type()
1567 if isinstance(base_expr_type, TDataBlock):
1568 actual_memref_type = base_expr_type.print_mlir(0).strip()
1570 actual_memref_type =
"memref<?x?xf64>"
1574 if isinstance(base_expr, Name):
1575 base_name = base_expr.id
1576 elif hasattr(base_expr,
'id'):
1578 base_name = base_expr.id
1580 if base_name
and base_name
in Node.type_map
and _is_nested_memref(Node.type_map[base_name]):
1581 actual_memref_type = Node.type_map[base_name]
1586 while isinstance(current, Subscript):
1587 indices.insert(0, current._new_index)
1588 current = current._value
1591 indices_str =
", ".join(indices)
1592 mlir_str =
"".join(Node.context.get_block().pop_mlir())
1596 if isinstance(actual_memref_type, str):
1597 is_nested = actual_memref_type.count(
"memref<") > 1
1598 actual_memref_type_str = actual_memref_type
1601 is_nested = actual_memref_type.is_nested_memref()
1602 actual_memref_type_str = actual_memref_type.print_mlir(0).strip()
1604 if is_nested
and len(indices) >= 2:
1607 inner_memref_id = f
"%{Node.get_mlir_id()}"
1609 if actual_memref_type_str.startswith(
"memref<?x")
and actual_memref_type_str.endswith(
">>"):
1610 inner_type = actual_memref_type_str[9:-1]
1612 inner_type =
"memref<?xf64>"
1614 mlir_str += indent + f
"{inner_memref_id} = memref.load {outer_base}[{indices[0]}] {{name = \"{outer_base.strip('%')}\"}} : {actual_memref_type_str}\n"
1616 remaining_indices =
", ".join(indices[1:])
1617 mlir_str += indent + f
"{result_id} = memref.load {inner_memref_id}[{remaining_indices}] {{name = \"{outer_base.strip('%')}\"}} : {inner_type}\n"
1620 mlir_str += indent + f
"{result_id} = memref.load {outer_base}[{indices_str}] {{name = \"{outer_base.strip('%')}\"}} : {actual_memref_type_str}\n"
1622 Node.context.push_block()
1623 Node.context.get_block().mlir_append(mlir_str)
1631 return f
"""Subscript:
1632{self._value.print_tree(indent_level + 1)}
1633{self._index.print_tree(indent_level + 1)}"""
1654 indent =
' ' * indent_level * Node.spaces_per_tab
1655 mlir_str =
"".join(Node.context.get_block().pop_mlir())
1664 base_expr.hoist(
True)
1665 base_mlir = base_expr.print_mlir(indent_level)
1667 mlir_str +=
"".join(Node.context.get_block().pop_mlir())
1671 index_expr.hoist(
True)
1672 index_mlir = index_expr.print_mlir(indent_level)
1673 mlir_str +=
"".join(Node.context.get_block().pop_mlir())
1676 new_index = f
"%{Node.get_mlir_id()}"
1677 mlir_type = (
'index' if hasattr(index_expr,
'_index_type')
and index_expr._index_type
else 'i32')
1678 mlir_str += indent + new_index + f
" = arith.index_cast {index_mlir.strip(' ')}: {mlir_type} to index\n"
1684 base_type = base_expr.get_type()
1685 if isinstance(base_type, TDataBlock):
1686 result_id = f
"%{Node.get_mlir_id()}"
1689 if isinstance(base_expr, Subscript):
1692 if isinstance(base_expr._value, Subscript):
1693 result_id = base_mlir
1696 base_mlir_name = base_mlir.lstrip(
'%')
1697 is_nested_memref =
False
1698 if base_mlir_name
in Node.type_map
and _is_nested_memref(Node.type_map[base_mlir_name]):
1699 is_nested_memref =
True
1701 outer_index = base_expr._new_index
1704 ptr_code, result_id = Node.generate_pointer_to_2d_element(
1705 base_mlir, outer_index, new_index, is_nested_memref, indent_level
1708 mlir_str += ptr_code
1710 element_type =
"f64"
1711 mlir_str += indent + f
"{result_id} = memref.subview {base_mlir}[{new_index}][1][1] : memref<?x{element_type}> to memref<?x{element_type}>\n"
1714 Node.context.push_block()
1715 Node.context.get_block().mlir_append(mlir_str)
1722 return ' ' * indent_level * Node.spaces_per_tab + f
"""Reference:
1723{self._expression.print_tree(indent_level + 1)}"""
1727 def __init__(self, elements, element_type, string = None):
1734 element_prints = [element.print_cpp()
for element
in self.
_elements]
1735 return ' ' * indent_level * Node.spaces_per_tab + f
"""tarch::la::Vector<{len(self._elements)}, {self._type.print_cpp()}>{{{", ".join(element_prints)}}}"""
1739 element_prints = [element.print_cpp()
for element
in self.
_elements]
1740 return ' ' * indent_level * Node.spaces_per_tab + f
"""tarch::la::Vector<{len(self._elements)}, {self._type.print_cpp()}>{{{", ".join(element_prints)}}}"""
1743 element_prints = [element.print_cpp()
for element
in self.
_elements]
1744 return ' ' * indent_level * Node.spaces_per_tab + f
"""tarch::la::Vector<{len(self._elements)}, {self._type.print_cpp()}>{{{", ".join(element_prints)}}}"""
1755 return TCustom(f
"tarch::la::Vector<{len(self._elements)}, {self._type.print_cpp()}>")
1766 diagonal_index = index[1]
1768 for i
in range(2, len(index) - 1):
1769 factor = factor * (dimensions[i - 1][1] - dimensions[i - 1][0])
1770 diagonal_index = diagonal_index + factor * index[i]
1774 if self.
_id is not None:
1780 if self.
_id is not None:
1792 return ' ' * indent_level * Node.spaces_per_tab +
"DiagonalMatrix"
1799 def __init__(self, dimensions, internal, index_dimensions):
1805 def index(self, row_index, column_index):
1809 if self.
_id is not None:
1815 if self.
_id is not None:
1827 return ' ' * indent_level * Node.spaces_per_tab +
"Matrix"
1836 if type(matrix)
is Name:
1843 lhs_matrix = self.
_matrix._internal._arguments[0]
1844 if type(lhs_matrix)
is Name:
1845 lhs_matrix = Name.variables[lhs_matrix.id]
1847 rhs_matrix = self.
_matrix._internal._arguments[1]
1848 if type(rhs_matrix)
is Name:
1849 rhs_matrix = Name.variables[rhs_matrix.id]
1852 internal_loop =
For([
Integer(0), rhs_matrix._width])
1853 index = self.
_assignment.get_iteration_variable() * rhs_matrix._width + internal_loop.get_iteration_variable()
1855 if type(lhs_matrix._internal)
is Float
or type(lhs_matrix._internal)
is Integer:
1856 lhs_index = lhs_matrix._internal
1857 rhs_index =
Subscript(rhs_matrix, internal_loop.get_iteration_variable())
1858 if type(rhs_matrix._internal)
is Float
or type(rhs_matrix._internal)
is Integer:
1859 rhs_index = rhs_matrix._internal
1862 internal_loop.close_scope()
1866 return f
"""{self._memory_allocation.print_cpp(indent_level)}
1867{self._assignment.print_cpp(indent_level)}"""
1870 return f
"""{self._memory_allocation.print_omp(indent_level)}
1871{self._assignment.print_omp(indent_level)}"""
1880 return ' ' * indent_level * Node.spaces_per_tab +
"DiagonalKroneckerProduct"
1889 if type(matrix)
is Name:
1901 return f
"""{self._memory_allocation.print_cpp(indent_level)}
1902{self._assignment.print_cpp(indent_level)}"""
1905 return f
"""{self._memory_allocation.print_omp(indent_level)}
1906{self._assignment.print_omp(indent_level)}"""
1915 return ' ' * indent_level * Node.spaces_per_tab +
"DiagonalKroneckerProduct"
1924 if type(matrix)
is Name:
1935 lhs_matrix = self.
_matrix._internal._arguments[0]
1936 if type(lhs_matrix)
is Name:
1937 lhs_matrix = Name.variables[lhs_matrix.id]
1939 rhs_matrix = self.
_matrix._internal._arguments[1]
1940 if type(rhs_matrix)
is Name:
1941 rhs_matrix = Name.variables[rhs_matrix.id]
1943 if type(lhs_matrix.get_type())
is TMatrix:
1944 lhs_dimensions = lhs_matrix._dimensions
1946 lhs_dimensions = [lhs_matrix._width, lhs_matrix._width]
1948 if type(rhs_matrix.get_type())
is TMatrix:
1949 rhs_dimensions = rhs_matrix._dimensions
1951 rhs_dimensions = [rhs_matrix._width, rhs_matrix._width]
1954 if type(lhs_matrix.get_type())
is TMatrix:
1955 loops.append(
For([
Integer(0), lhs_matrix._dimensions[0]]))
1956 loops.append(
For([
Integer(0), lhs_matrix._dimensions[1]]))
1957 lhs_index = loops[0].get_iteration_variable() * lhs_dimensions[1] + loops[1].get_iteration_variable()
1958 startIndex = loops[0].get_iteration_variable() * lhs_dimensions[1] * rhs_dimensions[1] * rhs_dimensions[0] + loops[1].get_iteration_variable() * rhs_dimensions[1]
1960 loops.append(
For([
Integer(0), lhs_matrix._width]))
1961 lhs_index = loops[0].get_iteration_variable()
1962 startIndex = loops[0].get_iteration_variable() * lhs_dimensions[1] * rhs_dimensions[1] * rhs_dimensions[0] + loops[0].get_iteration_variable() * rhs_dimensions[1]
1964 if type(rhs_matrix.get_type())
is TMatrix:
1965 loops.append(
For([
Integer(0), rhs_matrix._dimensions[0]]))
1966 loops.append(
For([
Integer(0), rhs_matrix._dimensions[1]]))
1967 rhs_index = loops[-2].get_iteration_variable() * rhs_dimensions[1] + loops[-1].get_iteration_variable()
1968 index = startIndex + loops[-1].get_iteration_variable() + loops[-2].get_iteration_variable() * lhs_dimensions[1] * rhs_dimensions[1]
1970 loops.append(
For([
Integer(0), rhs_matrix._width]))
1971 rhs_index = loops[-1].get_iteration_variable()
1972 index = startIndex + loops[-1].get_iteration_variable() + loops[-1].get_iteration_variable() * lhs_dimensions[1] * rhs_dimensions[1]
1974 lhs_subscript =
Subscript(lhs_matrix, lhs_index)
1975 if type(lhs_matrix._internal)
is Float
or type(lhs_matrix._internal)
is Integer:
1976 lhs_subscript = lhs_matrix._internal
1977 rhs_subscript =
Subscript(rhs_matrix, rhs_index)
1978 if type(rhs_matrix._internal)
is Float
or type(rhs_matrix._internal)
is Integer:
1979 rhs_subscript = rhs_matrix._internal
1983 for i
in range(len(loops) - 1):
1984 loops[i].add_statement(loops[i + 1])
1985 loops[i].close_scope()
1986 loops[-1].close_scope()
1990 return f
"""{self._memory_allocation.print_cpp(indent_level)}
1991{self._initialisation.print_cpp(indent_level)}
1992{self._assignment.print_cpp(indent_level)}"""
1995 return f
"""{self._memory_allocation.print_omp(indent_level)}
1996{self._initialisation.print_omp(indent_level)}
1997{self._assignment.print_omp(indent_level)}"""
2006 return ' ' * indent_level * Node.spaces_per_tab +
"KroneckerProduct"
2013 def __init__(self, value: Expression, index: Expression):
2022 return ' ' * indent_level * Node.spaces_per_tab + f
"{self._value.print_cpp()}({self._index.print_cpp()})"
2025 return ' ' * indent_level * Node.spaces_per_tab + f
"{self._value.print_omp()}({self._index.print_omp()})"
2028 return ' ' * indent_level * Node.spaces_per_tab + f
"{self._value.print_sycl()}({self._index.print_sycl()})"
2031 indent =
' ' * indent_level * Node.spaces_per_tab
2034 if isinstance(self.
_value, Subscript):
2035 array_base = self.
_value._value
2036 array_index = self.
_value._index
2037 array_base_mlir = array_base.print_mlir(indent_level)
2038 array_index_mlir = array_index.print_mlir(indent_level)
2040 if not getattr(array_index,
'_index_type',
False):
2041 array_index_cast = f
"%{Node.get_mlir_id()}"
2042 array_index_type = array_index.get_type()
2043 array_index_mlir_type = array_index_type.print_mlir(0).strip()
2044 mlir_str += indent + f
"{array_index_cast} = arith.index_cast {array_index_mlir} : {array_index_mlir_type} to index\n"
2046 array_index_cast = array_index_mlir
2048 raw_ptr = f
"%{Node.get_mlir_id()}"
2049 mlir_str += indent + f
"{raw_ptr} = memref.load {array_base_mlir}[{array_index_cast}] {{name = \"{array_base_mlir.strip('%')}\"}} : memref<?xi64>\n"
2051 vec_ptr = f
"%{Node.get_mlir_id()}"
2052 mlir_str += indent + f
"{vec_ptr} = llvm.inttoptr {raw_ptr} {{name = \"{array_base_mlir.strip('%')}\"}} : i64 to !llvm.ptr\n"
2056 if isinstance(self.
_index, Integer):
2057 func_param_i32 = f
"%{Node.get_mlir_id()}"
2058 mlir_str += indent + f
"{func_param_i32} = arith.constant {int(self._index._value)} : i32\n"
2059 elif not getattr(self.
_index,
'_index_type',
False):
2060 func_param_i32 = f
"%{Node.get_mlir_id()}"
2061 mlir_str += indent + f
"{func_param_i32} = arith.index_cast {func_param_mlir} : index to i32\n"
2063 func_param_i32 = func_param_mlir
2065 result = f
"%{Node.get_mlir_id()}"
2067 func_param_i64 = f
"%{Node.get_mlir_id()}"
2068 mlir_str += indent + f
"{func_param_i64} = arith.extsi {func_param_i32} : i32 to i64\n"
2070 velem_ptr = f
"%{Node.get_mlir_id()}"
2071 mlir_str += indent + f
"{velem_ptr} = llvm.getelementptr {vec_ptr}[{func_param_i64}] {{name = \"{array_base_mlir.strip('%')}\"}} : (!llvm.ptr, i64) -> !llvm.ptr, f64\n"
2073 mlir_str += indent + f
"{result} = llvm.load {velem_ptr} {{name = \"{array_base_mlir.strip('%')}\"}} : !llvm.ptr -> f64\n"
2074 Node.context.get_block().mlir_append(mlir_str)
2082{self._value.print_tree(indent_level + 1)}
2083{self._index.print_tree(indent_level + 1)}"""
2096 return ' ' * indent_level * Node.spaces_per_tab + f
"({self._lhs.print_cpp()} {self._operation} {self._rhs.print_cpp()})"
2099 return ' ' * indent_level * Node.spaces_per_tab + f
"({self._lhs.print_omp()} {self._operation} {self._rhs.print_omp()})"
2102 return ' ' * indent_level * Node.spaces_per_tab + f
"({self._lhs.print_sycl()} {self._operation} {self._rhs.print_sycl()})"
2106 if op_type_class
is TFloat:
2108 "<":
"arith.cmpf olt",
2109 ">":
"arith.cmpf ogt",
2110 "<=":
"arith.cmpf ole",
2111 ">=":
"arith.cmpf oge",
2112 "==":
"arith.cmpf oeq",
2113 "!=":
"arith.cmpf one"
2117 "<":
"arith.cmpi slt",
2118 ">":
"arith.cmpi sgt",
2119 "<=":
"arith.cmpi sle",
2120 ">=":
"arith.cmpi sge",
2121 "==":
"arith.cmpi eq",
2122 "!=":
"arith.cmpi ne"
2124 return op_map.get(operation,
None)
2127 op_type_class = TDataBlock.get_mlir_type_for_operation(self.
_lhs, self.
_rhs)
2129 mlir_operation = Comparison.get_mlir_comparison_op(self.
_operation, op_type_class)
2132 if mlir_operation
is None:
2133 raise ValueError(f
"Unsupported comparison operation: {self._operation}")
2135 const_zero_id =
None
2138 if TDataBlock.needs_memref_load(self.
_lhs):
2139 const_zero_id = TDataBlock.ensure_const_zero(const_zero_id, indent_level)
2140 lhs_mlir = TDataBlock.generate_memref_load(self.
_lhs, lhs_mlir, const_zero_id, indent_level)
2143 if TDataBlock.needs_memref_load(self.
_rhs):
2144 const_zero_id = TDataBlock.ensure_const_zero(const_zero_id, indent_level)
2145 rhs_mlir = TDataBlock.generate_memref_load(self.
_rhs, rhs_mlir, const_zero_id, indent_level)
2147 return ' ' * indent_level * Node.spaces_per_tab + f
"""{mlir_operation} {lhs_mlir}, {rhs_mlir} : {result_type}"""
2151Comparison {self._operation}:
2152{self._lhs.print_tree(indent_level + 1)}
2153{self._rhs.print_tree(indent_level + 1)}"""
2162 if type(self.
_lhs)
is BinaryOperation:
2163 if self.
_lhs._operation ==
"+":
2168 elif self.
_lhs._operation ==
"-":
2171 elif self.
_lhs._operation ==
"*":
2177 if type(self.
_rhs)
is BinaryOperation:
2178 if self.
_rhs._operation ==
"+":
2183 elif self.
_rhs._operation ==
"-":
2186 elif self.
_rhs._operation ==
"*":
2192 if self.
_operation ==
"-" and type(self.
_rhs)
is UnaryOperation
and self.
_rhs._operation ==
"-":
2195 if self.
_operation ==
"+" and type(self.
_rhs)
is UnaryOperation
and self.
_rhs._operation ==
"-":
2220 op_type_class = TDataBlock.get_mlir_type_for_operation(self.
_lhs, self.
_rhs)
2221 return op_type_class()
2224 return ' ' * indent_level * Node.spaces_per_tab + f
"({self._lhs.print_cpp()} {self._operation} {self._rhs.print_cpp()})"
2227 return ' ' * indent_level * Node.spaces_per_tab + f
"({self._lhs.print_omp()} {self._operation} {self._rhs.print_omp()})"
2230 return ' ' * indent_level * Node.spaces_per_tab + f
"({self._lhs.print_sycl()} {self._operation} {self._rhs.print_sycl()})"
2234 if op_type_class
is TFloat:
2248 return op_map.get(operation,
None)
2252 indent =
' ' * indent_level * Node.spaces_per_tab
2253 operand_type = operand.get_type()
2255 if type(operand_type)
is TDataBlock:
2256 operand_type = operand_type.get_element_type()
2258 if isinstance(operand_type, TInteger)
and result_type == TFloat:
2259 new_op = f
"%{Node.get_mlir_id()}"
2260 Node.context.get_block().mlir_append(indent + f
"{new_op} = arith.sitofp {operand_mlir.strip(' ')} : i32 to f64\n")
2262 elif isinstance(operand_type, TLong)
and result_type == TFloat:
2263 new_op = f
"%{Node.get_mlir_id()}"
2264 Node.context.get_block().mlir_append(indent + f
"{new_op} = arith.sitofp {operand_mlir.strip(' ')} : i64 to f64\n")
2266 elif isinstance(operand_type, TFloat)
and result_type == TInteger:
2267 new_op = f
"%{Node.get_mlir_id()}"
2268 Node.context.get_block().mlir_append(indent + f
"{new_op} = arith.fptosi {operand_mlir.strip(' ')} : f64 to i32\n")
2270 elif isinstance(operand_type, TFloat)
and result_type == TLong:
2271 new_op = f
"%{Node.get_mlir_id()}"
2272 Node.context.get_block().mlir_append(indent + f
"{new_op} = arith.fptosi {operand_mlir.strip(' ')} : f64 to i64\n")
2274 elif isinstance(operand_type, TInteger)
and result_type == TLong:
2275 new_op = f
"%{Node.get_mlir_id()}"
2276 Node.context.get_block().mlir_append(indent + f
"{new_op} = arith.extsi {operand_mlir.strip(' ')} : i32 to i64\n")
2278 elif isinstance(operand_type, TLong)
and result_type == TInteger:
2279 new_op = f
"%{Node.get_mlir_id()}"
2280 Node.context.get_block().mlir_append(indent + f
"{new_op} = arith.trunci {operand_mlir.strip(' ')} : i64 to i32\n")
2282 elif operand._index_type
and result_type == TFloat:
2283 new_op = f
"%{Node.get_mlir_id()}"
2284 Node.context.get_block().mlir_append(indent + f
"{new_op} = arith.index_cast {operand_mlir.strip(' ')} : index to f64\n")
2286 elif operand._index_type
and result_type == TInteger:
2287 new_op = f
"%{Node.get_mlir_id()}"
2288 Node.context.get_block().mlir_append(indent + f
"{new_op} = arith.index_cast {operand_mlir.strip(' ')} : index to i32\n")
2290 elif operand._index_type
and result_type == TLong:
2291 new_op = f
"%{Node.get_mlir_id()}"
2292 Node.context.get_block().mlir_append(indent + f
"{new_op} = arith.index_cast {operand_mlir.strip(' ')} : index to i64\n")
2298 op_type_class = TDataBlock.get_mlir_type_for_operation(self.
_lhs, self.
_rhs)
2299 return BinaryOperation.get_mlir_arithmetic_op(self.
_operation, op_type_class)
2302 indent =
' ' * indent_level * Node.spaces_per_tab
2306 op_type_class = TDataBlock.get_mlir_type_for_operation(self.
_lhs, self.
_rhs)
2307 result_type = op_type_class.print_mlir(0)
2309 op_type_class = TDataBlock.get_mlir_type_for_operation(self.
_lhs, self.
_rhs)
2312 const_zero_id =
None
2317 if TDataBlock.needs_memref_load(self.
_lhs):
2318 const_zero_id = TDataBlock.ensure_const_zero(const_zero_id, indent_level)
2319 lhs = TDataBlock.generate_memref_load(self.
_lhs, lhs, const_zero_id, indent_level)
2321 lhs = BinaryOperation.cast_operand(lhs, self.
_lhs, op_type_class, indent_level)
2326 if TDataBlock.needs_memref_load(self.
_rhs):
2327 const_zero_id = TDataBlock.ensure_const_zero(const_zero_id, indent_level)
2328 rhs = TDataBlock.generate_memref_load(self.
_rhs, rhs, const_zero_id, indent_level)
2330 rhs = BinaryOperation.cast_operand(rhs, self.
_rhs, op_type_class, indent_level)
2333 for mlir
in Node.context.pop_block().get_mlir():
2334 mlir_str += indent + mlir.strip(
' ')
2336 self.
mlir_id(f
"%{Node.get_mlir_id()}")
2338 mlir_str += indent + f
"""{self.mlir_id()} = {self.get_mlir_op()} {lhs}, {rhs} : {result_type}\n"""
2339 Node.context.push_block()
2340 Node.context.get_block().mlir_append(mlir_str)
2346BinaryOperation {self._operation}:
2347{self._lhs.print_tree(indent_level + 1)}
2348{self._rhs.print_tree(indent_level + 1)}"""
2376 return f
"({self._operation} {self._operand.print_cpp()})"
2379 return f
"({self._operation} {self._operand.print_omp()})"
2382 return f
"({self._operation} {self._operand.print_sycl()})"
2385 indent =
' ' * indent_level * Node.spaces_per_tab
2389 if isinstance(operand_type, TDataBlock):
2390 element_type_class = operand_type.get_element_type()
2391 operand_type = element_type_class()
if element_type_class
else operand_type
2393 if isinstance(operand_type, TFloat):
2399 mlir_str =
"".join(Node.context.get_block().pop_mlir())
2400 neg_one_id = f
"%{Node.get_mlir_id()}"
2401 mlir_str += indent + f
"{neg_one_id} = arith.constant -1.0 : f64\n"
2403 self.
mlir_id(f
"%{Node.get_mlir_id()}")
2404 mlir_str += indent + f
"{self.mlir_id()} = arith.mulf {op_mlir}, {neg_one_id} : f64\n"
2405 Node.context.push_block()
2406 Node.context.get_block().mlir_append(mlir_str)
2414 mlir_str =
"".join(Node.context.get_block().pop_mlir())
2415 neg_one_id = f
"%{Node.get_mlir_id()}"
2416 mlir_str += indent + f
"{neg_one_id} = arith.constant -1 : i32\n"
2418 self.
mlir_id(f
"%{Node.get_mlir_id()}")
2419 mlir_str += indent + f
"{self.mlir_id()} = arith.muli {op_mlir}, {neg_one_id} : i32\n"
2420 Node.context.push_block()
2421 Node.context.get_block().mlir_append(mlir_str)
2426UnaryOperation {self._operation}:
2427{self._operand.print_tree(indent_level + 1)}"""
2432 if type(dataBlock)
is Name:
2466class DataBlockComparison(Expression):
2468 if type(lhs)
is Name
and type(lhs.get_type())
is TDataBlock:
2473 if type(rhs)
is Name
and type(rhs.get_type())
is TDataBlock:
2489 if type(self.
_rhs._memory_range[0])
is not BinaryOperation
and type(self.
_memory_range[0])
is not BinaryOperation
and self.
_rhs._memory_range[0][1] > self.
_memory_range[0][1]:
2526 def __init__(self, operation, lhs, rhs, useFunctionSyntax=False):
2528 if type(lhs)
is Name
and type(lhs.get_type())
is TDataBlock:
2533 if type(rhs)
is Name
and type(rhs.get_type())
is TDataBlock:
2547 elif type(self.
_rhs._iteration_range[0][1])
is not BinaryOperation
and type(self.
_iteration_range[0][1])
is not BinaryOperation
and self.
_rhs._iteration_range[0][1] > self.
_iteration_range[0][1]:
2583 lhs_is_double =
False
2584 rhs_is_double =
False
2586 lhs_is_double = (str(self.
_lhs.
get_type().get_single()) ==
"double")
2588 lhs_is_double = (str(self.
_lhs.
get_type()) ==
"double")
2591 rhs_is_double = (str(self.
_rhs.
get_type().get_single()) ==
"double")
2593 rhs_is_double = (str(self.
_rhs.
get_type()) ==
"double")
2595 if (lhs_is_double
or rhs_is_double):
2625 if type(dataBlock)
is Name
and type(dataBlock.get_type())
is TDataBlock:
2664 if type(dataBlock)
is Name:
2674 for i
in range(len(self.
_dataBlock._iteration_range) - 2, -1, -1):
2678 for loop
in self.
_loops[::-1]:
2679 index.append(loop.get_iteration_variable())
2691 if Node.use_accelerator:
2692 return f
"""{' ' * indent_level * Node.spaces_per_tab}omp_target_memset(maxEigenvalues, 0, N * sizeof(double), targetDevice);
2693{' ' * indent_level * Node.spaces_per_tab}for (int {self._loop._iteration_variable.print_cpp()} = {self._loop._iteration_range[0].print_cpp()}; {self._loop._iteration_variable.print_cpp()} < {self._loop._iteration_range[1].print_cpp()}; {self._loop._iteration_variable.print_cpp()}++) {{
2694{' ' * (indent_level + 1) * Node.spaces_per_tab}#pragma omp target teams loop collapse(Dimensions) reduction(max:{self._outputVariable.multidimensional_index([self._loop.get_iteration_variable()]).print_omp()})
2695{self._loop._statements[1].print_omp(indent_level + 1)}
2696{' ' * indent_level * Node.spaces_per_tab}}}
2702 ranges = [f
"range{i} = {loop.get_interval_size().print_sycl()};" for i, loop
in enumerate(self.
_loops[0:3])]
2703 indexing = [f
"int {loop.get_iteration_variable().print_sycl()} = index[{i}];" for i, loop
in enumerate(self.
_loops[0:3])]
2704 statement_prints = [statement.print_sycl(indent_level + 1)
for statement
in self.
_loops[2]._statements]
2705 return os.linesep.join(ranges) +
"\n" +
' ' * indent_level * Node.spaces_per_tab + f
"""queue.submit([&](::sycl::handler& handler) {{handler.memset({self._outputVariable.print_sycl()}, 0, sizeof(double) * {self._loop.get_interval_size().print_sycl()}); }}).wait();
2706{' ' * indent_level * Node.spaces_per_tab}queue.submit([&](::sycl::handler& handler) {{
2707{' ' * indent_level * Node.spaces_per_tab}handler.parallel_for(::sycl::range<3>{{range0, range1, range2}}, [=](::sycl::item<3> index) {{
2708{os.linesep.join(indexing)}
2709{self._loops[3].print_sycl(indent_level + 1)}
2710{' ' * (indent_level + 1) * Node.spaces_per_tab}}});
2711{' ' * indent_level * Node.spaces_per_tab}}}).wait();"""
2717 return f
"""DataBlockUnaryMax:
2718{self._dataBlock.print_tree(indent_level + 1)}"""
2752 return ' ' * indent_level * Node.spaces_per_tab + self.
_type
2755 return ' ' * indent_level * Node.spaces_per_tab + self.
_type
2758 return ' ' * indent_level * Node.spaces_per_tab + self.
_type
2761 mlir_type = Node.get_mlir_type(self.
_type)
2762 if mlir_type
is not None:
2763 return ' ' * indent_level * Node.spaces_per_tab + mlir_type
2766 type_entry = Node.type_map[f
"{self._type}"]
2767 if isinstance(type_entry, str):
2768 return ' ' * indent_level * Node.spaces_per_tab + type_entry
2771 return ' ' * indent_level * Node.spaces_per_tab + type_entry.print_mlir(0)
2774 return ' ' * indent_level * Node.spaces_per_tab +
"TCustom: " + self.
_type
2778 return isinstance(self.
_type, str)
and (
'xmemref<' in self.
_type or self.
_type.count(
'memref<') > 1)
2785 return ' ' * indent_level * Node.spaces_per_tab +
"int"
2788 return ' ' * indent_level * Node.spaces_per_tab +
"int"
2791 return ' ' * indent_level * Node.spaces_per_tab +
"int"
2794 return ' ' * indent_level * Node.spaces_per_tab +
"i32"
2797 return ' ' * indent_level * Node.spaces_per_tab +
"TInteger"
2801 """Type for 64-bit integers (long) used in MLIRCellData struct fields"""
2803 return ' ' * indent_level * Node.spaces_per_tab +
"long"
2806 return ' ' * indent_level * Node.spaces_per_tab +
"long"
2809 return ' ' * indent_level * Node.spaces_per_tab +
"long"
2812 return ' ' * indent_level * Node.spaces_per_tab +
"i64"
2815 return ' ' * indent_level * Node.spaces_per_tab +
"TLong"
2820 return ' ' * indent_level * Node.spaces_per_tab +
"bool"
2823 return ' ' * indent_level * Node.spaces_per_tab +
"bool"
2826 return ' ' * indent_level * Node.spaces_per_tab +
"bool"
2829 return ' ' * indent_level * Node.spaces_per_tab +
"i1"
2832 return ' ' * indent_level * Node.spaces_per_tab +
"TBoolean"
2839 return ' ' * indent_level * Node.spaces_per_tab +
"double" + (
"&" if self.
_reference else "")
2842 return ' ' * indent_level * Node.spaces_per_tab +
"double" + (
"&" if self.
_reference else "")
2845 return ' ' * indent_level * Node.spaces_per_tab +
"double" + (
"&" if self.
_reference else "")
2848 return ' ' * indent_level * Node.spaces_per_tab +
"f64"
2851 return ' ' * indent_level * Node.spaces_per_tab +
"TFloat"
2856 return ' ' * indent_level * Node.spaces_per_tab +
"const char*"
2859 return ' ' * indent_level * Node.spaces_per_tab +
"const char*"
2862 return ' ' * indent_level * Node.spaces_per_tab +
"const char*"
2868 return ' ' * indent_level * Node.spaces_per_tab +
"TString"
2877 return ' ' * indent_level * Node.spaces_per_tab + f
'double**'
2880 return ' ' * indent_level * Node.spaces_per_tab + f
'memref<?xmemref<?x{self._element_type.print_mlir(0).strip()}>>'
2883 return ' ' * indent_level * Node.spaces_per_tab + f
'double**'
2886 return ' ' * indent_level * Node.spaces_per_tab + f
'double**'
2902 return ' ' * indent_level * Node.spaces_per_tab +
'void*'
2905 return ' ' * indent_level * Node.spaces_per_tab + self.
_signature
2908 return ' ' * indent_level * Node.spaces_per_tab +
'void*'
2911 return ' ' * indent_level * Node.spaces_per_tab +
'void*'
2915 def __init__(self, element_type=TFloat(), dimensions=1):
2921 return ' ' * indent_level * Node.spaces_per_tab +
'void*'
2926 return ' ' * indent_level * Node.spaces_per_tab + f
'memref<{dims}x{self._element_type.print_mlir(0).strip()}>'
2929 return ' ' * indent_level * Node.spaces_per_tab +
'void*'
2932 return ' ' * indent_level * Node.spaces_per_tab +
'void*'
2945 if isinstance(operand, VectorIndex):
2948 operand_type = operand.get_type()
2949 if isinstance(operand_type, TDataBlock):
2950 if (isinstance(operand_type._dimensions, int)
and operand_type._dimensions == 1)
or \
2951 (isinstance(operand_type._dimensions, list)
and len(operand_type._dimensions) == 1):
2954 if isinstance(operand, Subscript)
and operand.is_nested_subscript():
2961 indent =
' ' * indent_level * Node.spaces_per_tab
2962 load_id = f
"%{Node.get_mlir_id()}"
2963 operand_type = operand.get_type()
2965 if isinstance(operand_type, TDataBlock):
2967 if isinstance(operand, Subscript)
and operand.is_nested_subscript():
2970 element_type_class = operand_type.get_element_type()
2971 if element_type_class
is not None:
2972 element_type = element_type_class().
print_mlir(0)
2974 element_type =
"f64"
2976 load_str = indent + f
"{load_id} = memref.load {operand_mlir}[{const_zero_id}] {{name = \"{operand_mlir.strip('%')}\"}} : memref<?x{element_type}>\n"
2977 Node.context.get_block().mlir_append(load_str)
2984 if const_zero_id
is None:
2985 const_zero_id = f
"%{Node.get_mlir_id()}"
2986 indent =
' ' * indent_level * Node.spaces_per_tab
2987 Node.context.get_block().mlir_append(indent + f
"{const_zero_id} = arith.constant 0 : index\n")
2988 return const_zero_id
3006 lhs_operand_type = lhs.get_type()
3007 if type(lhs_operand_type)
is TDataBlock:
3008 lhs_type = lhs_operand_type.get_element_type()
3010 lhs_type = type(lhs_operand_type)
3012 rhs_operand_type = rhs.get_type()
3013 if type(rhs_operand_type)
is TDataBlock:
3014 rhs_type = rhs_operand_type.get_element_type()
3016 rhs_type = type(rhs_operand_type)
3019 if lhs_type
is TFloat
or rhs_type
is TFloat:
3021 elif lhs_type
is TLong
or rhs_type
is TLong:
3023 elif lhs_type
is TInteger
and rhs_type
is TInteger:
3025 elif lhs_type
is TBoolean
and rhs_type
is TBoolean:
3043 mlir_element_type =
"i64"
3045 mlir_element_type = element_type_class().
print_mlir(0)
3047 return element_type_class
or TFloat, is_vector, mlir_element_type
3068 dims.append(str(dim))
3070 dims.append(dim.print_mlir(indent_level))
3072 Node.context.push_block()
3074 op.print_mlir(indent_level)
3075 dims.append(
''.join(Node.context.pop_block().get_mlir()))
3084 if type_name
is None:
3086 type_name = Node.type_map[f
"{self._underlying_type}"]
3090 if isinstance(type_name, str)
and type_name.startswith(
"memref<"):
3091 return ' ' * indent_level * Node.spaces_per_tab + type_name
3114 dimension_list = [
"?"] * actual_dims
3115 memref_dims =
"x".join(dimension_list)
3117 return ' ' * indent_level * Node.spaces_per_tab + f
"""memref<{memref_dims}x{type_name}>"""
3120 return ' ' * indent_level * Node.spaces_per_tab +
"TDataBlock"
3124 if isinstance(type_map_entry, str):
3126 return 'xmemref<' in type_map_entry
3127 elif hasattr(type_map_entry,
'is_nested_memref'):
3129 return type_map_entry.is_nested_memref()
3148 return ' ' * indent_level * Node.spaces_per_tab +
"TMatrix"
3165 return ' ' * indent_level * Node.spaces_per_tab +
"TDiagonalMatrix"
3182 return ' ' * indent_level * Node.spaces_per_tab +
"TVector"
3193 if statement
is not None:
3197 statement_prints = [statement.print_cpp(indent_level + 1)
for statement
in self.
_statements]
3198 return ' ' * indent_level * Node.spaces_per_tab + f
"""if constexpr ({self._boolean.print_cpp()}) {{
3199{os.linesep.join(statement_prints)}
3200""" +
' ' * indent_level * Node.spaces_per_tab +
"}"
3203 statement_prints = [statement.print_omp(indent_level + 1)
for statement
in self.
_statements]
3204 return ' ' * indent_level * Node.spaces_per_tab + f
"""if constexpr ({self._boolean.print_omp()}) {{
3205{os.linesep.join(statement_prints)}
3206""" +
' ' * indent_level * Node.spaces_per_tab +
"}"
3209 statement_prints = [statement.print_sycl(indent_level + 1)
for statement
in self.
_statements]
3210 return ' ' * indent_level * Node.spaces_per_tab + f
"""if constexpr ({self._boolean.print_sycl()}) {{
3211{os.linesep.join(statement_prints)}
3212""" +
' ' * indent_level * Node.spaces_per_tab +
"}"
3215 indent =
' ' * indent_level * Node.spaces_per_tab
3216 Node.context.push_block()
3217 statement_prints = [statement.print_mlir(indent_level + 1)
for statement
in self.
_statements]
3218 body_statements =
''.join(Node.context.pop_block().get_mlir())
3219 body_statements +=
''.join(statement_prints)
3220 return indent + f
"""scf.if {self._boolean.print_mlir(indent_level)} {{
3221{body_statements}""" + indent +
"}"
3225 statement_prints = [statement.print_tree(indent_level + 1)
for statement
in self.
_statements]
3226 return ' ' * indent_level * Node.spaces_per_tab + f
"""If:
3227{os.linesep.join(statement_prints)}"""
3230 _inuse_iteration_variables = set()
3231 _iteration_variable_names = [
'i',
'j',
'k',
'l',
'n',
'm',
'a',
'b',
'c',
'd']
3233 def __init__(self, iteration_range, iteration_variable_name = None, use_scheduler = False):
3239 if iteration_variable_name ==
None:
3240 for variable_name
in For._iteration_variable_names:
3241 if variable_name
not in For._inuse_iteration_variables:
3244 For._inuse_iteration_variables.add(variable_name)
3248 For._inuse_iteration_variables.update({iteration_variable_name : self.
_iteration_variable})
3263 statement_prints = [statement.print_cpp(indent_level + 1)
for statement
in self.
_statements]
3265 return ' ' * indent_level * Node.spaces_per_tab + f
"""parallelForWithSchedulerInstructions({self._iteration_variable.print_cpp()}, {self._iteration_range[1].print_cpp()}, loopParallelism) {{
3266{os.linesep.join(statement_prints)}
3267{' ' * indent_level * Node.spaces_per_tab}}}
3270 return ' ' * indent_level * Node.spaces_per_tab + f"""for (int {self._iteration_variable.print_cpp()} = {self._iteration_range[0].print_cpp()}; {self._iteration_variable.print_cpp()} < {self._iteration_range[1].print_cpp()}; {self._iteration_variable.print_cpp()}++) {{
3271{os.linesep.join(statement_prints)}
3272""" + ' ' * indent_level * Node.spaces_per_tab + "}"
3275 def print_omp(self, indent_level = 0):
3276 statement_prints = [statement.print_omp(indent_level + 1) for statement in self._statements]
3277 return ' ' * indent_level * Node.spaces_per_tab + f"""for (int {self._iteration_variable.print_omp()} = {self._iteration_range[0].print_omp()}; {self._iteration_variable.print_omp()} < {self._iteration_range[1].print_omp()}; {self._iteration_variable.print_omp()}++) {{
3278{os.linesep.join(statement_prints)}
3279""" + ' ' * indent_level * Node.spaces_per_tab + "}"
3281 def print_sycl(self, indent_level = 0):
3282 statement_prints = [statement.print_sycl(indent_level + 1) for statement in self._statements]
3283 return ' ' * indent_level * Node.spaces_per_tab + f"""for (int {self._iteration_variable.print_sycl()} = {self._iteration_range[0].print_sycl()}; {self._iteration_variable.print_sycl()} < {self._iteration_range[1].print_sycl()}; {self._iteration_variable.print_sycl()}++) {{
3284{os.linesep.join(statement_prints)}
3285""" + ' ' * indent_level * Node.spaces_per_tab + "}"
3287 def print_mlir(self, indent_level=0):
3288 indent = ' ' * indent_level * Node.spaces_per_tab
3291 if isinstance(self._iteration_range[0], list) or isinstance(self._iteration_range[0], tuple):
3292 dims = len(self._iteration_range)
3297 for d in range(dims):
3298 var = For._iteration_variable_names[d]
3299 iter_vars.append(f"%{var}")
3300 self._iteration_range[d][0].hoist(True)
3301 self._iteration_range[d][1].hoist(True)
3302 if isinstance(self._iteration_range[d][0], Integer):
3303 mlir_id = f"%{Node.get_mlir_id()}"
3304 bounds_casting += indent + f"{mlir_id} = arith.constant {self._iteration_range[d][0].print_mlir(indent_level)} : index\n"
3307 lwb = self._iteration_range[d][0].print_mlir(indent_level).strip()
3308 if isinstance(self._iteration_range[d][1], Integer):
3309 mlir_id = f"%{Node.get_mlir_id()}"
3310 bounds_casting += indent + f"{mlir_id} = arith.constant {self._iteration_range[d][1].print_mlir(indent_level)} : index\n"
3313 upb = self._iteration_range[d][1].print_mlir(indent_level).strip()
3316 mlir_id = f"%{Node.get_mlir_id()}"
3317 bounds_casting += indent + f"{mlir_id} = arith.constant 1 : index\n"
3318 steps.append(mlir_id)
3320 self._iteration_range[0].hoist(True)
3321 self._iteration_range[1].hoist(True)
3322 if not self._iteration_variable._index_type:
3323 self._iteration_variable.hoist(True)
3324 mlir_str = self._iteration_variable.print_mlir(indent_level).strip()
3325 mlir_id = f"%{Node.get_mlir_id()}"
3326 bounds_casting += ''.join(Node.context.pop_block().get_mlir())
3327 Node.context.push_block()
3328 bounds_casting += indent + f"{mlir_id} = arith.index_cast {mlir_str} : {self._iteration_variable.get_type().print_mlir(0)} to index\n"
3329 iter_vars = [mlir_id]
3331 iter_vars = [self._iteration_variable.print_mlir(indent_level).strip()]
3333 if not self._iteration_range[0]._index_type:
3334 self._iteration_range[0].hoist(True)
3335 mlir_str = self._iteration_range[0].print_mlir(indent_level).strip()
3336 mlir_id = f"%{Node.get_mlir_id()}"
3337 bounds_casting += ''.join(Node.context.pop_block().get_mlir())
3338 Node.context.push_block()
3339 bounds_casting += indent + f"{mlir_id} = arith.index_cast {mlir_str} : {self._iteration_range[0].get_type().print_mlir(indent_level).strip()} to index\n"
3342 lwbs += [self._iteration_range[0].print_mlir(indent_level).strip()]
3344 if not self._iteration_range[1]._index_type:
3345 self._iteration_range[1].hoist(True)
3346 mlir_str = self._iteration_range[1].print_mlir(indent_level).strip()
3347 mlir_id = f"%{Node.get_mlir_id()}"
3348 bounds_casting += ''.join(Node.context.pop_block().get_mlir())
3349 Node.context.push_block()
3350 bounds_casting += indent + f"{mlir_id} = arith.index_cast {mlir_str} : {self._iteration_range[1].get_type().print_mlir(indent_level).strip()} to index\n"
3353 upbs += [self._iteration_range[1].print_mlir(indent_level).strip()]
3355 bounds_casting += ''.join(Node.context.pop_block().get_mlir())
3356 mlir_id = f"%{Node.get_mlir_id()}"
3357 bounds_casting += indent + f"{mlir_id} = arith.constant 1 : index\n"
3361 Node.context.push_block()
3362 statement_prints = [statement.print_mlir(indent_level + 1) for statement in self._statements]
3363 statements = ''.join(Node.context.pop_block().get_mlir())
3364 statements += ''.join(statement_prints)
3366 # Check if loop contains memory allocation (nested memref allocation must be sequential)
3367 # Use scf.for for allocation loops, scf.parallel for computation loops
3368 has_allocation = any(isinstance(stmt, MemoryAllocation) for stmt in self._statements)
3371 # Use sequential scf.for for memory allocation (cannot parallelize)
3373 f"{indent}scf.for {', '.join(iter_vars)} = "
3374 f"{', '.join(lwbs)} to {', '.join(upbs)} step {', '.join(steps)} {{\n"
3377 # Use scf.parallel for computation loops (safe to parallelize)
3379 f"{indent}scf.parallel ({', '.join(iter_vars)}) = "
3380 f"({', '.join(lwbs)}) to ({', '.join(upbs)}) step ({', '.join(steps)}) {{\n"
3382 loop_footer = "\n" + indent + "}\n"
3384 return bounds_casting + loop_header + statements + loop_footer
3386 def print_tree(self, indent_level = 0):
3387 statement_prints = [statement.print_tree(indent_level + 1) for statement in self._statements]
3388 return ' ' * indent_level * Node.spaces_per_tab + f"""For:
3389{os.linesep.join(statement_prints)}"""
3392class Comment(Statement):
3393 def __init__(self, comment):
3395 self._comment = comment
3400 def print_cpp(self, indent_level = 0):
3401 return ' ' * indent_level * Node.spaces_per_tab + "//" + self._comment[1:]
3403 def print_omp(self, indent_level = 0):
3404 return ' ' * indent_level * Node.spaces_per_tab + "//" + self._comment[1:]
3406 def print_sycl(self, indent_level = 0):
3407 return ' ' * indent_level * Node.spaces_per_tab + "//" + self._comment[1:]
3409 def print_mlir(self, indent_level = 0):
3410 return ' ' * indent_level * Node.spaces_per_tab + "//" + self._comment[1:]
3412 def print_tree(self, indent_level=0):
3413 return ' ' * indent_level * Node.spaces_per_tab + "Comment"
3416class FunctionCall(Statement):
3417 def __init__(self, id, arguments, is_offloadable = False):
3419 self._arguments = arguments
3420 self._is_offloadable = is_offloadable
3422 def add_argument(self, argument: Expression):
3423 self._arguments.append(argument)
3425 def print_cpp(self, indent_level=0):
3426 argument_prints = [argument.print_cpp() for argument in self._arguments]
3427 return ' ' * indent_level * Node.spaces_per_tab + f"""{self.id}({", ".join(argument_prints)});"""
3429 def print_omp(self, indent_level=0):
3430 argument_prints = [argument.print_omp() for argument in self._arguments]
3431 return ' ' * indent_level * Node.spaces_per_tab + f"""{self.id}({", ".join(argument_prints)});"""
3433 def print_sycl(self, indent_level=0):
3434 argument_prints = [argument.print_sycl() for argument in self._arguments]
3435 return ' ' * indent_level * Node.spaces_per_tab + f"""{self.id}({", ".join(argument_prints)});"""
3437 def print_mlir(self, indent_level=0):
3438 indent = ' ' * indent_level * Node.spaces_per_tab
3439 Node.context.push_block()
3441 # Detect if this is a bridge function call
3442 is_bridge_call = self.id in ['sourceTerm', 'flux', 'nonconservativeProduct']
3444 # Process arguments and handle array subscripts
3445 argument_prints = []
3446 for i, argument in enumerate(self._arguments):
3447 argument.hoist(True)
3449 # Check if this is a Reference to a 2D subscript for a bridge call
3450 is_reference_to_2d_subscript = (
3452 isinstance(argument, Reference) and
3453 isinstance(argument._expression, Subscript) and
3454 isinstance(argument._expression._value, Subscript)
3457 if is_reference_to_2d_subscript:
3458 # This is Reference(Subscript[Subscript[2D_memref]])
3459 # Extract the nested structure
3460 outer_subscript = argument._expression # Subscript[Subscript[...], col_index]
3461 inner_subscript = outer_subscript._value # Subscript[2D_memref, row_index]
3462 memref_2d = inner_subscript._value
3464 outer_subscript.hoist(True)
3465 outer_mlir = outer_subscript.print_mlir(indent_level)
3466 mlir_str_from_outer = "".join(Node.context.get_block().pop_mlir())
3467 Node.context.get_block().mlir_append(mlir_str_from_outer)
3469 # Get MLIR representations
3470 memref_2d.hoist(True)
3471 memref_2d_mlir = memref_2d.print_mlir(indent_level).strip(' ')
3473 # Get the base memref variable name (strip % prefix)
3474 base_mlir_name = memref_2d_mlir.lstrip('%')
3475 is_nested_base = False
3476 if base_mlir_name in Node.type_map and _is_nested_memref(Node.type_map[base_mlir_name]):
3477 is_nested_base = True
3479 # Use unified method from Node to generate pointer arithmetic
3480 ptr_code, final_ptr = Node.generate_pointer_to_2d_element(
3481 memref_2d_mlir, inner_subscript._new_index, outer_subscript._new_index, is_nested_base, indent_level
3483 Node.context.get_block().mlir_append(ptr_code)
3484 arg_mlir = final_ptr
3486 # Not a 2D subscript for bridge call - use standard approach
3487 arg_mlir = argument.print_mlir(indent_level).strip(' ')
3489 # Regular subscript (array indexing) that needs loading
3490 if not is_reference_to_2d_subscript and isinstance(argument, Subscript):
3491 # Get the base array and index
3492 base_value = arg_mlir # This is the array reference
3493 index = argument._new_index # This is the converted index
3494 base_type = argument._value.get_type()
3496 if isinstance(base_type, TDataBlock):
3497 # Check if this is a multi-dimensional (2D+) DataBlock
3498 if isinstance(base_type._dimensions, int) and base_type._dimensions >= 2:
3499 # This is a flat multi-dimensional memref (like memref<?x?xf64>)
3500 # When passed to a function expecting memref<?xf64>, just pass it as-is
3501 arg_mlir = f"{base_value}[{index}]"
3503 # Single-dimensional DataBlock
3504 element_type_class, is_vector, mlir_element_type = base_type.get_array_element_info()
3507 # This is an array of vectors - generate a load to get the vector pointer as i64, then cast to !llvm.ptr
3508 result_id = f"%{Node.get_mlir_id()}"
3509 load_op = indent + f"{result_id} = memref.load {base_value}[{index}] {{name = \"{base_value.strip('%')}\"}} : memref<?x{mlir_element_type}>\n"
3510 Node.context.get_block().mlir_append(load_op)
3511 # Now cast the i64 to !llvm.ptr
3512 ptr_id = f"%{Node.get_mlir_id()}"
3513 cast_op = indent + f"{ptr_id} = llvm.inttoptr {result_id} : {mlir_element_type} to !llvm.ptr\n"
3514 Node.context.get_block().mlir_append(cast_op)
3517 # Regular array of scalars - generate a load to get the scalar value
3518 result_id = f"%{Node.get_mlir_id()}"
3519 load_op = indent + f"{result_id} = memref.load {base_value}[{index}] {{name = \"{base_value.strip('%')}\"}} : memref<?x{mlir_element_type}>\n"
3520 Node.context.get_block().mlir_append(load_op)
3521 arg_mlir = result_id
3523 # Handle non-TDataBlock types (like nested memrefs)
3524 if str(base_type).startswith("memref<?xmemref<?xf64>>"):
3525 # This is an array of vectors (stored as memref of memrefs)
3526 result_id = f"%{Node.get_mlir_id()}"
3527 load_op = indent + f"{result_id} = memref.load {base_value}[{index}] {{name = \"{base_value.strip('%')}\"}} : memref<?xmemref<?xf64>>\n"
3528 Node.context.get_block().mlir_append(load_op)
3529 arg_mlir = result_id
3531 # For other types, just pass the reference
3532 arg_mlir = f"{base_value}[{index}]"
3534 argument_prints.append(arg_mlir)
3536 argument_block = Node.context.pop_block().get_mlir()
3537 for op in argument_block:
3538 Node.context.get_block().mlir_append(' ' * indent_level * Node.spaces_per_tab + op.strip(' '))
3540 # NOTE: we lookup the function pointer name in type_map to get the type
3541 id = self.id[0].upper() + self.id[1:]
3542 type_entry = Node.type_map[self.id]
3543 # Handle both Type objects and strings
3544 if isinstance(type_entry, str):
3545 fn_type = type_entry
3547 fn_type = type_entry.print_mlir(0).strip()
3549 bridge_function_name = f"{self.id}_bridge"
3550 call_statement = f"""func.call @{bridge_function_name}({", ".join(argument_prints)}) : {fn_type.strip(' ')}"""
3552 Node.context.get_block().mlir_append(indent + call_statement)
3555 def print_tree(self, indent_level=0):
3556 argument_prints = [argument.print_tree() for argument in self._arguments]
3557 return ' ' * indent_level * Node.spaces_per_tab + f"""FunctionCall:
3558{' ' * (indent_level + 1) * Node.spaces_per_tab + ",".join(argument_prints)}"""
3561class Max(Statement):
3562 def __init__(self, lhs, rhs):
3567 def print_cpp(self, indent_level=0):
3568 return ' ' * indent_level * Node.spaces_per_tab + f"""std::max({self._lhs.print_cpp()}, {self._rhs.print_cpp()});"""
3570 def print_omp(self, indent_level=0):
3571 return ' ' * indent_level * Node.spaces_per_tab + f"""std::max({self._lhs.print_omp()}, {self._rhs.print_omp()});"""
3573 def print_sycl(self, indent_level=0):
3574 return ' ' * indent_level * Node.spaces_per_tab + f"""::sycl::max({self._lhs.print_sycl()}, {self._rhs.print_sycl()});"""
3576 def print_mlir(self, indent_level=0):
3577 indent = ' ' * indent_level * Node.spaces_per_tab
3579 def unpack_nested_memref_access(expr, indent_level):
3580 indent = ' ' * indent_level * Node.spaces_per_tab
3582 if not (isinstance(expr, Subscript) and isinstance(expr._value, Subscript)):
3583 return expr.print_mlir(indent_level).strip(' ')
3585 outer_subscript = expr._value
3587 if not hasattr(outer_subscript, '_new_index'):
3588 outer_subscript.print_mlir(indent_level)
3589 if not hasattr(expr, '_new_index'):
3590 expr.print_mlir(indent_level)
3592 outer_index = getattr(outer_subscript, '_new_index', outer_subscript._index.print_mlir(indent_level))
3593 inner_index = getattr(expr, '_new_index', expr._index.print_mlir(indent_level))
3595 inner_memref_id = f"%{Node.get_mlir_id()}"
3597 # Get the outer memref's MLIR variable name and check type_map
3598 outer_memref_var = outer_subscript._value.print_mlir(indent_level).strip(' ')
3600 # Try to determine the outer memref type
3601 # First, check if this variable is registered in type_map
3603 # Get the cpp name to check type_map
3604 outer_var_cpp = outer_subscript._value.print_cpp().strip()
3605 if outer_var_cpp in Node.type_map:
3606 type_entry = Node.type_map[outer_var_cpp]
3607 # Handle both Type objects and strings
3608 if isinstance(type_entry, str):
3609 outer_memref_type = type_entry
3611 # Type object - call print_mlir()
3612 outer_memref_type = type_entry.print_mlir(0).strip()
3614 # Fallback to default nested memref type
3615 outer_memref_type = "memref<?xmemref<?xf64>>"
3617 outer_memref_type = "memref<?xmemref<?xf64>>"
3619 mlir_str = indent + f"{inner_memref_id} = memref.load {outer_memref_var}[{outer_index}] {{name = \"{outer_memref_var.strip('%')}\"}} : {outer_memref_type}\n"
3620 value_id = f"%{Node.get_mlir_id()}"
3621 mlir_str += indent + f"{value_id} = memref.load {inner_memref_id}[{inner_index}] {{name = \"{outer_memref_var.strip('%')}\"}} : memref<?xf64>\n"
3622 Node.context.get_block().mlir_append(mlir_str)
3626 Node.context.push_block()
3628 if type(self._lhs.get_type()) is TDataBlock:
3629 self._lhs.hoist(True)
3630 if type(self._rhs.get_type()) is TDataBlock:
3631 self._rhs.hoist(True)
3633 const_zero_id = None
3634 if isinstance(self._lhs, Subscript) and isinstance(self._lhs._value, Subscript):
3635 lhs = unpack_nested_memref_access(self._lhs, indent_level)
3637 lhs = self._lhs.print_mlir(indent_level).strip(' ')
3638 if TDataBlock.needs_memref_load(self._lhs):
3639 const_zero_id = TDataBlock.ensure_const_zero(const_zero_id, indent_level)
3640 lhs = TDataBlock.generate_memref_load(self._lhs, lhs, const_zero_id, indent_level)
3642 if isinstance(self._rhs, Subscript) and isinstance(self._rhs._value, Subscript):
3643 rhs = unpack_nested_memref_access(self._rhs, indent_level)
3645 rhs = self._rhs.print_mlir(indent_level).strip(' ')
3646 if TDataBlock.needs_memref_load(self._rhs):
3647 const_zero_id = TDataBlock.ensure_const_zero(const_zero_id, indent_level)
3648 rhs = TDataBlock.generate_memref_load(self._rhs, rhs, const_zero_id, indent_level)
3650 argument_block = Node.context.pop_block().get_mlir()
3651 for op in argument_block:
3652 Node.context.get_block().mlir_append(op)
3654 # Use the unified type determination and operation mapping from TDataBlock
3655 op_type_class = TDataBlock.get_mlir_type_for_operation(self._lhs, self._rhs)
3656 result_type = op_type_class.print_mlir(0)
3658 if op_type_class is TInteger:
3659 mlir_op = "arith.maxsi"
3661 mlir_op = "arith.maximumf" # Default to float
3663 new_id = f"%{Node.get_mlir_id()}"
3664 mlir_str = indent + f"{new_id} = {mlir_op} {lhs}, {rhs} : {result_type}\n"
3665 Node.context.get_block().mlir_append(mlir_str)
3669 def print_tree(self, indent_level=0):
3670 return ' ' * indent_level * Node.spaces_per_tab + f"""Max:
3671{self._lhs.print_tree(indent_level + 1)}
3672{self._rhs.print_tree(indent_level + 1)}"""
3675class MemoryAllocation(Statement):
3676 memoryAllocated = []
3677 allocated_types = {} # Map from MLIR variable name to allocated memref type string
3679 def __init__(self, name: Name, object_type: Type, dimensions, specify_type = True, add_to_stack=True):
3681 self._type = object_type
3682 if type(dimensions[0]) is list:
3683 self._dimensions = [dimension[1] for dimension in dimensions]
3685 self._dimensions = dimensions
3687 self._specify_type = specify_type
3689 if add_to_stack == True:
3690 MemoryAllocation.memoryAllocated[-1].append(self)
3692 if len(self._dimensions) > 1:
3693 self._size = self._dimensions[-2]
3694 for i in range(len(self._dimensions) - 3, -1, -1):
3695 self._size = self._size * self._dimensions[i]
3697 self._loop = For([Integer(0), self._dimensions[-1]])
3698 self._loop.add_statement(MemoryAllocation(Subscript(self._name, self._loop.get_iteration_variable()), self._type, [self._size], False, False))
3699 self._loop.close_scope()
3701 def print_cpp(self, indent_level=0):
3702 if len(self._dimensions) == 1:
3703 return ' ' * indent_level * Node.spaces_per_tab + f"""{self._type.print_cpp() + "* " if self._specify_type else ""}{self._name.print_cpp()} = new {self._type.print_cpp()}[{self._dimensions[0].print_cpp()}];"""
3705 if Node.use_accelerator == False and type(Name.variables[self._name.id]) is not FaceDataBlock:
3706 return ' ' * indent_level * Node.spaces_per_tab + f"""{self._type.print_cpp()}** {self._name.print_cpp()} = new {self._type.print_cpp()}*[{self._dimensions[-1].print_cpp()}];
3707{self._loop.print_cpp(indent_level)}"""
3709 return ' ' * indent_level * Node.spaces_per_tab + f"""{self._type.print_cpp()}* {self._name.print_cpp()} = new {self._type.print_cpp()}[{(self._size * self._dimensions[-1]).print_cpp()}];"""
3712 def print_omp(self, indent_level=0):
3713 if Node.use_accelerator == False:
3714 if len(self._dimensions) == 1:
3715 return ' ' * indent_level * Node.spaces_per_tab + f"""{self._type.print_omp() + "* " if self._specify_type else ""}{self._name.print_omp()} = new {self._type.print_omp()}[{self._dimensions[0].print_omp()}];"""
3717 return ' ' * indent_level * Node.spaces_per_tab + f"""{self._type.print_omp()}** {self._name.print_omp()} = new {self._type.print_omp()}*[{self._dimensions[-1].print_omp()}];
3718{self._loop.print_omp(indent_level)}"""
3720 if len(self._dimensions) == 1:
3721 if Node.use_memory_manager == True:
3722 return ' ' * indent_level * Node.spaces_per_tab + f"""{self._type.print_omp() + "* " if self._specify_type else ""}{self._name.print_omp()} = tarch::accelerator::GPUMemoryManager::getInstance().allocate<{self._type.print_omp()}>({self._dimensions[0].print_omp()}, targetDevice);"""
3724 return ' ' * indent_level * Node.spaces_per_tab + f"""{self._type.print_omp() + "* " if self._specify_type else ""}{self._name.print_omp()} = ({self._type.print_omp()}*)omp_target_alloc(sizeof({self._type.print_omp()}) * {self._dimensions[0].print_omp()}, targetDevice);"""
3726 if Node.use_memory_manager == True:
3727 return ' ' * indent_level * Node.spaces_per_tab + f"""{self._type.print_omp()}* {self._name.print_omp()} = tarch::accelerator::GPUMemoryManager::getInstance().allocate<{self._type.print_omp()}>({(self._size * self._dimensions[-1]).print_omp()}, targetDevice);"""
3729 return ' ' * indent_level * Node.spaces_per_tab + f"""{self._type.print_omp()}* {self._name.print_omp()} = ({self._type.print_omp()}*)omp_target_alloc(sizeof({self._type.print_omp()}) * {(self._size * self._dimensions[-1]).print_omp()}, targetDevice);"""
3731 def print_sycl(self, indent_level=0):
3732 if len(self._dimensions) == 1:
3733 return ' ' * indent_level * Node.spaces_per_tab + f"""{self._type.print_sycl() + "* " if self._specify_type else ""}{self._name.print_sycl()} = new {self._type.print_sycl()}[{self._dimensions[0].print_sycl()}];"""
3735 return ' ' * indent_level * Node.spaces_per_tab + f"""{self._type.print_sycl()}** {self._name.print_sycl()} = new {self._type.print_sycl()}*[{self._dimensions[-1].print_sycl()}];
3736{self._loop.print_sycl(indent_level)}"""
3738 def print_mlir(self, indent_level=0):
3739 indent = ' ' * indent_level * Node.spaces_per_tab
3741 # Comprehensive tracing for MemoryAllocation
3742 name_mlir = self._name.print_mlir(0)
3743 name_type = self._name.get_type()
3744 is_datablock = isinstance(name_type, TDataBlock)
3745 var_name_for_map = self._name.print_cpp().strip()
3746 in_type_map = var_name_for_map in Node.type_map
3749 if len(self._dimensions) == 1:
3750 # Process the dimension expression first and collect any intermediate operations
3751 Node.context.push_block()
3752 index_val = self._dimensions[0].print_mlir(indent_level)
3753 intermediate_ops = Node.context.pop_block().get_mlir()
3755 # Now build the allocation operations
3756 index_id = f"%{Node.get_mlir_id()}"
3757 decl = ''.join(intermediate_ops)
3759 # Add the index cast and allocation with correct source type
3760 index_type = self._dimensions[0].get_type()
3761 index_mlir_type = index_type.print_mlir(0).strip()
3762 decl += indent + f"{index_id} = arith.index_cast {index_val} : {index_mlir_type} to index\n"
3763 new_id = f"%{Node.get_mlir_id()}"
3765 # For 1D allocation, we should always generate 1D memref type
3766 # Get the element type correctly
3767 name_type = self._name.get_type()
3768 if isinstance(name_type, TDataBlock):
3769 # Get the element type of the DataBlock
3770 element_type_class = name_type.get_element_type()
3771 element_type = element_type_class().print_mlir(0)
3772 mlir_type = f"memref<?x{element_type}>"
3774 mlir_type = name_type.print_mlir(0)
3776 decl += indent + f"{new_id} = memref.alloc({index_id}) {{name = \"{name_mlir.strip('%')}\"}} : {mlir_type}\n"
3777 # Track the allocated type for deallocation
3778 MemoryAllocation.allocated_types[new_id] = mlir_type
3780 # Check if this is an inner allocation that needs to be stored in an outer array
3781 # This happens when self._name is a Subscript (indicating arr[i] = allocated_inner_array)
3782 if isinstance(self._name, Subscript) and not self._specify_type:
3783 # This is an inner allocation for nested memref, add the store operation
3784 # Get the outer array name and index
3785 outer_array = self._name._value # This should be the outer array name
3786 outer_index = self._name._index
3788 # Process the outer array and index expressions to collect any additional operations
3789 Node.context.push_block()
3790 outer_array_mlir = outer_array.print_mlir(indent_level)
3791 outer_index_mlir = outer_index.print_mlir(indent_level)
3792 additional_ops = Node.context.pop_block().get_mlir()
3794 # Add any additional operations
3795 decl += ''.join(additional_ops)
3797 # When storing a memref into a nested memref array, we need to use memref<?xmemref<?xf64>>
3798 type_mlir = Node.get_mlir_type(self._type)
3799 if type_mlir is None:
3800 type_entry = Node.type_map[self._type.print_cpp()]
3801 # Handle both Type objects and strings
3802 if isinstance(type_entry, str):
3803 type_mlir = type_entry.strip(' ')
3805 # Type object - call print_mlir()
3806 type_mlir = type_entry.print_mlir(0).strip()
3807 outer_type = f"memref<?xmemref<?x{type_mlir}>>"
3809 # Add the store operation
3810 decl += indent + f"memref.store {new_id}, {outer_array_mlir}[{outer_index_mlir}] {{name = \"{outer_array_mlir.strip('%')}\"}} : {outer_type}"
3814 # For multi-dimensional arrays, check type_map first (ordering is important)
3815 # This allows overriding default DataBlock behavior for nested memref types
3816 var_name_for_type_check = self._name.print_cpp().strip()
3817 in_type_map = var_name_for_type_check in Node.type_map
3819 # Use nested memref if in type_map, otherwise check if DataBlock flat
3821 # Variable has explicit type specification - use nested memref path
3822 is_datablock = False # Force nested path
3823 actual_dims_calc = -1
3825 # For multi-dimensional arrays, check if this is a DataBlock
3826 # If so, use flat memref with the correct number of dimensions instead of nested memref
3827 is_datablock = isinstance(self._name.get_type(), TDataBlock)
3829 # Calculate actual_dims based on DataBlock iteration_range
3831 datablock_obj = self._name.get_type()
3833 if hasattr(datablock_obj, '_iteration_range'):
3834 actual_dims_calc = len(datablock_obj._iteration_range)
3835 elif isinstance(datablock_obj._dimensions, int):
3836 actual_dims_calc = 1 if datablock_obj._dimensions == 1 else 2
3837 elif isinstance(datablock_obj._dimensions, list):
3838 actual_dims_calc = len(datablock_obj._dimensions) if len(datablock_obj._dimensions) > 0 else 1
3840 actual_dims_calc = 1
3842 actual_dims_calc = -1 # Not a DataBlock
3844 if is_datablock and actual_dims_calc >= 2:
3845 # Multi-dimensional DataBlock case: allocate as flat N-D memref
3846 element_type_class = self._name.get_type().get_element_type()
3847 element_type = element_type_class().print_mlir(0)
3849 # We need actual_dims_calc dimensions from self._dimensions
3850 # For a 2D array: _dimensions might have 2, 3, or 4 elements (depends on iteration_range structure)
3851 # Take the first actual_dims_calc elements
3852 dim_exprs = self._dimensions[:actual_dims_calc]
3853 if len(dim_exprs) < actual_dims_calc:
3854 # Fallback: if we don't have enough dimensions, use what we have
3855 dim_exprs = self._dimensions
3857 intermediate_decl = ""
3860 for i, dim_expr in enumerate(dim_exprs):
3861 # Evaluate the dimension expression - this appends its operations to context
3862 dim_val = dim_expr.print_mlir(indent_level)
3863 dim_id = f"%{Node.get_mlir_id()}"
3864 dim_ids.append(dim_id)
3865 dim_type = dim_expr.get_type()
3866 dim_mlir_type = dim_type.print_mlir(0).strip()
3868 # Generate the index cast (will be appended after dimension expressions)
3869 intermediate_decl += indent + f"{dim_id} = arith.index_cast {dim_val} : {dim_mlir_type} to index\n"
3871 # These must come before the index_cast operations
3872 dim_expr_ops = Node.context.pop_block().get_mlir()
3873 intermediate_decl = ''.join(dim_expr_ops) + intermediate_decl
3875 # Re-push a new block for remaining operations
3876 Node.context.push_block()
3878 # Process name and allocate
3879 name_mlir = self._name.print_mlir(indent_level)
3881 # Build memref type: memref<?x?x...xf64> with actual_dims_calc '?'s
3882 dimension_list = ["?"] * actual_dims_calc
3883 memref_dims = "x".join(dimension_list)
3884 alloc_args = ", ".join(dim_ids)
3886 # Allocate as flat N-D memref
3887 alloc_line = indent + f"{name_mlir} = memref.alloc({alloc_args}) {{name = \"{name_mlir.strip('%')}\"}} : memref<{memref_dims}x{element_type}>\n"
3889 # Track the allocated memref type so it can be used in load/store operations
3890 memref_type_full = f"memref<{memref_dims}x{element_type}>"
3891 Node.mlir_memref_types[name_mlir] = memref_type_full
3892 MemoryAllocation.allocated_types[name_mlir] = memref_type_full
3894 return intermediate_decl + alloc_line
3896 # For non-DataBlock multi-dimensional arrays, use nested memref with loop
3897 # Use the existing loop structure created in the constructor
3898 size = self._dimensions[-2]
3899 for i in range(len(self._dimensions) - 3, -1, -1):
3900 size = size * self._dimensions[i]
3902 # Process the outer dimension expression first and collect any intermediate operations
3903 Node.context.push_block()
3904 outer_size_val = self._dimensions[-1].print_mlir(indent_level)
3905 intermediate_ops = Node.context.pop_block().get_mlir()
3907 # Build the allocation operations
3908 decl = ''.join(intermediate_ops)
3910 # Add the index cast and allocation with correct source type
3911 outer_size_id = f"%{Node.get_mlir_id()}"
3912 outer_size_type = self._dimensions[-1].get_type()
3913 outer_size_mlir_type = outer_size_type.print_mlir(0).strip()
3914 decl += indent + f"{outer_size_id} = arith.index_cast {outer_size_val} : {outer_size_mlir_type} to index\n"
3916 # Determine the correct type for the outer array
3917 # Check if this variable name is in type_map (e.g., maxEigenvaluesX, maxEigenvaluesY)
3918 var_name = self._name.print_cpp().strip()
3920 # NOTE: We need to check if the variable name is in type_map first
3921 if var_name in Node.type_map:
3922 # Use the explicitly registered type from type_map
3923 type_entry = Node.type_map[var_name]
3925 # Handle both Type objects and strings
3926 if isinstance(type_entry, str):
3927 outer_type = type_entry
3929 # Type object (e.g., TNestedMemref)
3930 outer_type = type_entry.print_mlir(0).strip()
3931 elif isinstance(self._name.get_type(), TDataBlock):
3932 # Get the element type correctly for the inner memref
3933 element_type_class = self._name.get_type().get_element_type()
3934 element_type = element_type_class().print_mlir(0)
3935 outer_type = f"memref<?xmemref<?x{element_type}>>"
3937 # Fallback for non-DataBlock types
3938 name_type_mlir = self._name.get_type().print_mlir(0)
3940 if name_type_mlir is None:
3941 type_entry = Node.type_map[self._name.get_type().print_cpp()]
3942 # Handle both Type objects and strings
3943 if isinstance(type_entry, str):
3944 name_type_mlir = type_entry.strip(' ')
3947 name_type_mlir = type_entry.print_mlir(0).strip()
3948 outer_type = f"memref<?x{name_type_mlir}>"
3950 # Process the name expression to collect any additional operations
3951 Node.context.push_block()
3952 name_mlir = self._name.print_mlir(indent_level)
3953 name_ops = Node.context.pop_block().get_mlir()
3955 # Add any name-related operations
3956 decl += ''.join(name_ops)
3957 decl += indent + f"{name_mlir} = memref.alloc({outer_size_id}) {{name = \"{name_mlir.strip('%')}\"}} : {outer_type}\n"
3958 # Track the allocated type for deallocation
3959 MemoryAllocation.allocated_types[name_mlir] = outer_type
3961 # Use the existing loop structure for inner allocations
3962 loop_mlir = self._loop.print_mlir(indent_level)
3964 return decl + loop_mlir
3966 def print_tree(self, indent_level=0):
3967 return ' ' * indent_level * Node.spaces_per_tab + f"""MemoryAllocation:
3968{self._name.print_tree(indent_level + 1)}"""
3971class MemoryDeallocation(Statement):
3972 def __init__(self, allocation: MemoryAllocation):
3974 self._dimensions = allocation._dimensions
3975 self._name = allocation._name
3976 self._allocated_type = None # Will store the MLIR type string when known
3978 if len(self._dimensions) > 1:
3979 size = self._dimensions[-2]
3980 for i in range(len(self._dimensions) - 3, -1, -1):
3981 size = size * self._dimensions[i]
3983 self._loop = For([Integer(0), self._dimensions[-1]])
3984 self._loop.add_statement(MemoryDeallocation(MemoryAllocation(Subscript(self._name, self._loop.get_iteration_variable()), None, [size], False, False)))
3985 self._loop.close_scope()
3987 def print_cpp(self, indent_level=0):
3988 if len(self._dimensions) == 1:
3989 return ' ' * indent_level * Node.spaces_per_tab + f"""delete[] {self._name.print_cpp()};"""
3991 if Node.use_accelerator == False and type(Name.variables[self._name.id]) is not FaceDataBlock:
3992 return f"""{self._loop.print_cpp(indent_level)}
3993{' ' * indent_level * Node.spaces_per_tab}delete[] {self._name.print_cpp()};"""
3995 return f"""{' ' * indent_level * Node.spaces_per_tab}delete[] {self._name.print_cpp()};"""
3997 def print_omp(self, indent_level=0):
3998 if Node.use_accelerator == False:
3999 if len(self._dimensions) == 1:
4000 return ' ' * indent_level * Node.spaces_per_tab + f"""delete[] {self._name.print_omp()};"""
4002 return f"""{self._loop.print_omp(indent_level)}
4003{' ' * indent_level * Node.spaces_per_tab}delete[] {self._name.print_omp()};"""
4005 if Node.use_memory_manager:
4006 return ' ' * indent_level * Node.spaces_per_tab + f"""tarch::accelerator::GPUMemoryManager::getInstance().free({self._name.print_omp()}, targetDevice);"""
4008 return ' ' * indent_level * Node.spaces_per_tab + f"""omp_target_free({self._name.print_omp()}, targetDevice);"""
4010 def print_sycl(self, indent_level=0):
4011 return ' ' * indent_level * Node.spaces_per_tab + f"""::sycl::free({self._name.print_sycl()}, queue);"""
4013 # return f"""{' ' * indent_level * Node.spaces_per_tab}::sycl::free({self._name.print_sycl()}, queue);"""
4014# return f"""{self._loop.print_sycl(indent_level)}
4015#{' ' * indent_level * Node.spaces_per_tab}::sycl::free({self._name.print_sycl()}, queue);"""
4017 def print_mlir(self, indent_level=0):
4018 indent = ' ' * indent_level * Node.spaces_per_tab
4020 # Get the MLIR variable name
4021 name_mlir = self._name.print_mlir(indent_level)
4023 # Use the stored allocated type if available
4024 if name_mlir in MemoryAllocation.allocated_types:
4025 name_type_mlir = MemoryAllocation.allocated_types[name_mlir]
4026 elif self._allocated_type is not None:
4027 # Use the allocated type if it was set explicitly
4028 name_type_mlir = self._allocated_type
4030 # Fallback: try to reconstruct the type from the allocation dimensions
4031 # For nested memrefs (multi-dimensional arrays), we need memref<?xmemref<?xf64>>
4032 # For 1D arrays, we need memref<?xf64>
4033 name_element_type = self._name.get_type()
4034 element_type = name_element_type.print_mlir(0).strip()
4036 if len(self._dimensions) == 1:
4037 name_type_mlir = f"memref<?x{element_type}>"
4039 # For multidimensional, assume nested memref
4040 name_type_mlir = f"memref<?xmemref<?x{element_type}>>"
4042 return "\n" + indent + f"""memref.dealloc {name_mlir} : {name_type_mlir}\n"""
4044 def print_tree(self, indent_level=0):
4045 return ' ' * indent_level * Node.spaces_per_tab + f"""MemoryDeallocation:
4046{self._name.print_tree(indent_level + 1)}"""
4049class Construction(Statement):
4050 def __init__(self, name: Name, expression: Expression):
4053 self._name.on_left(True)
4054 self._expression = expression
4055 Name.variables.update({self._name.id: expression})
4057 def print_cpp(self, indent_level = 0):
4058 return ' ' * indent_level * Node.spaces_per_tab + f"{self._name.get_type().print_cpp()} {self._name.print_cpp()} = {self._expression.print_cpp()};"
4060 def print_omp(self, indent_level = 0):
4061 return ' ' * indent_level * Node.spaces_per_tab + f"{self._name.get_type().print_omp()} {self._name.print_omp()} = {self._expression.print_omp()};"
4063 def print_sycl(self, indent_level = 0):
4064 return ' ' * indent_level * Node.spaces_per_tab + f"{self._name.get_type().print_sycl()} {self._name.print_sycl()} = {self._expression.print_sycl()};"
4066 def print_mlir(self, indent_level = 0):
4067 indent = ' ' * indent_level * Node.spaces_per_tab
4069 if self._expression._string is None:
4070 expr_mlir = self._expression.print_mlir(indent_level)
4072 # Get any operations that were added to the context
4073 mlir_str = ''.join(Node.context.get_block().pop_mlir())
4074 mlir_str += indent + f"{self._name.print_mlir(indent_level)} = {expr_mlir} : {self._name.get_type().print_mlir(indent_level).strip(' ')}\n"
4076 # Get another variable id and load the pointer into that, then load the value into the actual variable (self)
4077 new_id = Node.get_mlir_id()
4078 mlir_str = indent + f"%{new_id} = {self._expression.print_mlir(indent_level)}\n"
4080 # Use the expression's type, not the name's type
4081 expr_type = self._expression.get_type().print_mlir(indent_level).strip(' ')
4082 mlir_str += indent + f"{self._name.print_mlir(indent_level)} = llvm.load %{new_id} {{name = \"{self._name.print_mlir(indent_level).strip('%')}\"}} : !llvm.ptr -> {expr_type}\n"
4084 Node.context.get_block().mlir_append(mlir_str)
4087 def print_tree(self, indent_level=0):
4090{self._name.print_tree(indent_level + 1)}
4091{self._expression.print_tree(indent_level + 1)}"""
4094class DataBlockConstructionFromExisting(Statement):
4095 def __init__(self, name: Name, dataBlock: DataBlock):
4098 self._dataBlock = dataBlock
4099 self._dataBlock.id = name.id
4100 Name.variables.update({self._name.id: self._dataBlock})
4101 FunctionDefinition._syclDataToCopy.append(self)
4103 def print_cpp(self, indent_level = 0):
4104 type_print = self._name.get_type().print_cpp()
4105 if len(self._dataBlock._iteration_range) > 1 and Node.use_accelerator == False and type(self._dataBlock) is not FaceDataBlock:
4107 return ' ' * indent_level * Node.spaces_per_tab + f"{type_print} {self._name.print_cpp()} = {self._dataBlock._internal.print_cpp()};"
4109 def print_omp(self, indent_level = 0):
4110 type_print = self._name.get_type().print_omp()
4111 if len(self._dataBlock._iteration_range) > 1 and Node.use_accelerator == False and type(self._dataBlock) is not FaceDataBlock:
4113 return ' ' * indent_level * Node.spaces_per_tab + f"{type_print} {self._name.print_omp()} = {self._dataBlock._internal.print_omp()};"
4115 def print_sycl(self, indent_level = 0):
4116 type_print = self._name.get_type().print_sycl()
4117 if len(self._dataBlock._iteration_range) > 1 and Node.use_accelerator == False and type(self._dataBlock) is not FaceDataBlock:
4119 return ' ' * indent_level * Node.spaces_per_tab + f"{type_print} {self._name.print_sycl()} = {self._dataBlock._internal.print_sycl()};"
4120# step_size = Integer(1)
4121# for d in self._dataBlock._iteration_range[0:-1]:
4122# step_size = step_size * (d[1] - d[0])
4123# size = step_size * (self._dataBlock._iteration_range[-1][1] - self._dataBlock._iteration_range[-1][0])
4124# return f"""{' ' * indent_level * Node.spaces_per_tab}auto {self._name.print_sycl()} = ::sycl::malloc_shared<double>({size.print_sycl()}, queue);
4125#{' ' * indent_level * Node.spaces_per_tab}for (int z = 0; z < {(self._dataBlock._iteration_range[-1][1]-self._dataBlock._iteration_range[-1][0]).print_sycl()}; z++){{
4126#""" + (f"""{' ' * (indent_level + 1) * Node.spaces_per_tab}queue.memcpy(&{self._name.print_sycl()}[{step_size.print_sycl()} * z], {self._dataBlock._internal.print_sycl()}[z], {step_size.print_sycl()} * sizeof({self._dataBlock._underlying_type.print_sycl()})).wait();""" if self._dataBlock._internal.print_sycl()[-4:] != "QOut" else """""") + f"""
4127#{' ' * indent_level * Node.spaces_per_tab}}}
4129# for d in self._dataBlock._iteration_range[0:-1]:
4130# size = size * (d[1] - d[0])
4131# return f"""{' ' * indent_level * Node.spaces_per_tab}auto {self._name.print_sycl()} = ::sycl::malloc_shared<double*>({(self._dataBlock._iteration_range[-1][1]-self._dataBlock._iteration_range[-1][0]).print_sycl()}, queue);
4132#{' ' * indent_level * Node.spaces_per_tab}for (int z = 0; z < {(self._dataBlock._iteration_range[-1][1]-self._dataBlock._iteration_range[-1][0]).print_sycl()}; z++){{
4133#{' ' * (indent_level + 1) * Node.spaces_per_tab}{self._name.print_sycl()}[z] = ::sycl::malloc_shared<double>({size.print_sycl()}, queue);
4134#""" + (f"""{' ' * (indent_level + 1) * Node.spaces_per_tab}queue.memcpy({self._name.print_sycl()}[z], {self._dataBlock._internal.print_sycl()}[z], {size.print_sycl()} * sizeof({self._dataBlock._underlying_type.print_sycl()})).wait();""" if self._dataBlock._internal.print_sycl()[-4:] != "QOut" else """""") + f"""
4135#{' ' * indent_level * Node.spaces_per_tab}}}
4137 return ' ' * indent_level * Node.spaces_per_tab + f"{type_print} {self._name.print_sycl()} = {self._dataBlock._internal.print_sycl()};"
4139 def print_mlir(self, indent_level = 0):
4140 indent = ' ' * indent_level * Node.spaces_per_tab
4141 type_print = self._name.get_type().print_mlir(indent_level)
4144 for i in range(len(self._dataBlock._iteration_range) - 3, -1, -1):
4145 dims.append(self._dataBlock._iteration_range[i])
4147 if len(self._dataBlock._iteration_range) > 1:
4148 type_print = "!llvm.ptr " #+ f"{self._dataBlock._underlying_type} dims: {dims})"
4150 # NOTE: We encode structs as Strings, so check here if they are a struct and generate accordingly
4151 # Set the target variable name so String.print_mlir() can use it
4152 target_var_name = self._name.print_mlir(indent_level)
4153 Node.target_variable_name = target_var_name
4155 # Store the mapping from DSL name to MLIR variable name for later use in assignments
4156 # TODO: This is clumsy - it is assuming the exact flow here...
4157 dsl_name = self._dataBlock._internal._value
4158 Node.mlir_symbol_table[dsl_name] = target_var_name
4160 # Generate the struct field access operations with the target variable name
4161 struct_mlir = self._dataBlock._internal.print_mlir(indent_level).strip(' ')
4162 mlir_str = ''.join(Node.context.pop_block().get_mlir())
4164 # Clear the target variable name - we've finished generating the MLIR for this Node
4165 Node.target_variable_name = None
4170 def print_tree(self, indent_level=0):
4172DataBlockConstructionFromExisting:
4173{self._name.print_tree(indent_level + 1)}
4174{self._dataBlock._internal.print_tree(indent_level + 1)}"""
4177class DataBlockConstructionFromOperation:
4178 def __init__(self, name: Name, dataBlockOperation):
4181 self._dataBlock = DataBlock(dataBlockOperation._iteration_range, None, False, self._name.id, underlying_type=dataBlockOperation.get_type().get_single())
4182 Name.variables.update({self._name.id: self._dataBlock})
4183 if type(dataBlockOperation) is DataBlockBinaryOperation or type(dataBlockOperation) is DataBlockUnaryOperation or type(dataBlockOperation) is DataBlockMax or type(dataBlockOperation) is DataBlockComparison or type(dataBlockOperation._internal) is String:
4184 self._operation = dataBlockOperation
4186 self._operation = dataBlockOperation._internal
4188 self.memoryAllocation = MemoryAllocation(self._name, self._dataBlock.get_type().get_single(), self._dataBlock._memory_range)
4189 self.assignment = DataBlockAssignment(self._name, self._operation)
4191 def print_cpp(self, indent_level = 0):
4192 return f"""{self.memoryAllocation.print_cpp(indent_level)}
4193{self.assignment.print_cpp(indent_level)}"""
4195 def print_omp(self, indent_level = 0):
4196 return f"""{self.memoryAllocation.print_omp(indent_level)}
4197{self.assignment.print_omp(indent_level)}"""
4199 def print_sycl(self, indent_level = 0):
4200 return f"""{self.memoryAllocation.print_sycl(indent_level)}
4201{self.assignment.print_sycl(indent_level)}"""
4203 def print_mlir(self, indent_level = 0):
4204 return f"""{self.memoryAllocation.print_mlir(indent_level)}
4205{self.assignment.print_mlir(indent_level)}"""
4207 def print_tree(self, indent_level = 0):
4208 return ' ' * indent_level * Node.spaces_per_tab + f"""DataBlockConstructionFromOperation:
4209{self._name.print_tree(indent_level + 1)}
4210{self._dataBlock.print_tree(indent_level + 1)}"""
4213class DataBlockConstructionFromMatrixVectorProduct:
4214 def __init__(self, name: Name, matrix, vector):
4216 if type(vector) is Name:
4217 self._vector = Name.variables[vector.id]
4219 self._vector = vector
4221 if type(matrix) is Name:
4222 self._matrix = Name.variables[matrix.id]
4224 self._matrix = matrix
4226 self._dataBlock = DataBlock(self._vector._iteration_range, None, False, self._name.id, underlying_type=vector.get_type().get_single())
4227 Name.variables.update({self._name.id: self._dataBlock})
4229 self.memoryAllocation = MemoryAllocation(self._name, self._dataBlock.get_type().get_single(), self._dataBlock._memory_range)
4231 element_width = self._vector._iteration_range[0][1] - self._vector._iteration_range[0][0]
4234 loops.append(For([Integer(0), self._dataBlock._iteration_range[-1][1] - self._dataBlock._iteration_range[-1][0]]))
4235 loops.append(For([Integer(0), self._matrix._dimensions[0]]))
4236 loops.append(For([Integer(0), self._dataBlock._iteration_range[0][1] - self._dataBlock._iteration_range[0][0]]))
4238 if Node.use_accelerator:
4240 for [a, b] in self._dataBlock._memory_range[0:-1]:
4241 size = size * (b - a)
4242 index = size * loops[0].get_iteration_variable() + (loops[1].get_iteration_variable() * element_width + loops[2].get_iteration_variable())
4243 loops[2].add_statement(Assignment(Subscript(self._dataBlock, index), Float(0.0)))
4245 loops[2].add_statement(Assignment(Subscript(Subscript(self._dataBlock, loops[0].get_iteration_variable()), loops[1].get_iteration_variable() * element_width + loops[2].get_iteration_variable()), Float(0.0)))
4247 loops.append(For([Integer(0), self._matrix._dimensions[1]]))
4249 if Node.use_accelerator:
4251 for [a, b] in self._dataBlock._memory_range[0:-1]:
4252 size = size * (b - a)
4253 index = size * loops[0].get_iteration_variable() + (loops[1].get_iteration_variable() * element_width + loops[2].get_iteration_variable())
4254 loops[3].add_statement(Assignment(Subscript(self._dataBlock, index), BinaryOperation("+", Subscript(self._dataBlock, index), BinaryOperation("*", matrix.index(loops[1].get_iteration_variable(), loops[3].get_iteration_variable()), Subscript(self._vector, loops[0].get_iteration_variable() * size + element_width * loops[3].get_iteration_variable() + loops[2].get_iteration_variable())))))
4256 loops[3].add_statement(Assignment(Subscript(Subscript(self._dataBlock, loops[0].get_iteration_variable()), loops[1].get_iteration_variable() * element_width + loops[2].get_iteration_variable()), BinaryOperation("+", Subscript(Subscript(self._dataBlock, loops[0].get_iteration_variable()), loops[1].get_iteration_variable() * element_width + loops[2].get_iteration_variable()), BinaryOperation("*", matrix.index(loops[1].get_iteration_variable(), loops[3].get_iteration_variable()), Subscript(Subscript(self._vector, loops[0].get_iteration_variable()), element_width * loops[3].get_iteration_variable() + loops[2].get_iteration_variable())))))
4260 loops[2].add_statement(loops[3])
4261 loops[1].add_statement(loops[2])
4262 loops[0].add_statement(loops[1])
4263 loops[3].close_scope()
4264 loops[2].close_scope()
4265 loops[1].close_scope()
4266 loops[0].close_scope()
4267 self.assignment = loops[0]
4269 def print_cpp(self, indent_level = 0):
4270 return f"""{self.memoryAllocation.print_cpp(indent_level)}
4271{self.assignment.print_cpp(indent_level)}"""
4273 def print_omp(self, indent_level = 0):
4274 return f"""{self.memoryAllocation.print_omp(indent_level)}
4275{self.assignment.print_omp(indent_level)}"""
4277 def print_sycl(self, indent_level = 0):
4278 return f"""{self.memoryAllocation.print_sycl(indent_level)}
4279{self.assignment.print_sycl(indent_level)}"""
4281 def print_mlir(self, indent_level = 0):
4282 return f"""{self.memoryAllocation.print_mlir(indent_level)}
4283{self.assignment.print_mlir(indent_level)}"""
4285 def print_tree(self, indent_level = 0):
4286 return ' ' * indent_level * Node.spaces_per_tab + f"""DataBlockConstructionFromOperation:
4287{self._name.print_tree(indent_level + 1)}
4288{self._dataBlock.print_tree(indent_level + 1)}"""
4291class DataBlockConstructionFromDotProduct:
4292 def __init__(self, name: Name, matrix, vector):
4294 if type(vector) is Name:
4295 self._vector = Name.variables[vector.id]
4297 self._vector = vector
4299 if type(matrix) is Name:
4300 self._matrix = Name.variables[matrix.id]
4302 self._matrix = matrix
4304 iteration_range = self._vector._iteration_range
4305 iteration_range[-1][1] = Integer(1)
4306 self._dataBlock = DataBlock(self._vector._iteration_range, None, False, self._name.id, underlying_type=vector.get_type().get_single())
4307 Name.variables.update({self._name.id: self._dataBlock})
4309 self.memoryAllocation = MemoryAllocation(self._name, self._dataBlock.get_type().get_single(), self._dataBlock._memory_range)
4313 for [a, b] in self._dataBlock._iteration_range[-1::-1]:
4314 loops.append(For([Integer(0), b - a]))
4316 for loop in loops[::-1]:
4317 index.append(loop.get_iteration_variable())
4319 loops[-1].add_statement(Assignment(self._dataBlock.multidimensional_index(index), Float(0.0)))
4320 for i in range(len(loops) - 1):
4321 loops[i].add_statement(loops[i + 1])
4323 loops.append(For([Integer(0), self._matrix._width]))
4324 rhs_index = index.copy()
4325 rhs_index[-1] = loops[-1].get_iteration_variable()
4326 loops[-1].add_statement(Assignment(self._dataBlock.multidimensional_index(index), BinaryOperation("+", self._dataBlock.multidimensional_index(index), BinaryOperation("*", Subscript(self._matrix, rhs_index[-1]), self._vector.multidimensional_index(rhs_index)))))
4327 loops[-2].add_statement(loops[-1])
4331 self.assignment = loops[0]
4333 def print_cpp(self, indent_level = 0):
4334 return f"""{self.memoryAllocation.print_cpp(indent_level)}
4335{self.assignment.print_cpp(indent_level)}"""
4337 def print_omp(self, indent_level = 0):
4338 return f"""{self.memoryAllocation.print_omp(indent_level)}
4339{self.assignment.print_omp(indent_level)}"""
4341 def print_sycl(self, indent_level = 0):
4342 return f"""{self.memoryAllocation.print_sycl(indent_level)}
4343{self.assignment.print_sycl(indent_level)}"""
4345 def print_mlir(self, indent_level = 0):
4346 return f"""{self.memoryAllocation.print_mlir(indent_level)}
4347{self.assignment.print_mlir(indent_level)}"""
4349 def print_tree(self, indent_level = 0):
4350 return ' ' * indent_level * Node.spaces_per_tab + f"""DataBlockConstructionFromOperation:
4351{self._name.print_tree(indent_level + 1)}
4352{self._dataBlock.print_tree(indent_level + 1)}"""
4355class DataBlockConstructionFromFunction(Statement):
4356 def __init__(self, name: Name, dataBlock: DataBlock):
4359 self._dataBlock = dataBlock
4360 self._dataBlock.id = name.id
4361 Name.variables.update({self._name.id: self._dataBlock})
4363 self.memoryAllocation = MemoryAllocation(self._name, self._dataBlock.get_type().get_single(), self._dataBlock._memory_range)
4364 self.initialisation = For(self._dataBlock._memory_range[-1])
4367 for i in range(len(self._dataBlock._iteration_range) - 2, 0, -1):
4368 loop_bound = self._dataBlock._memory_range[i][1] - self._dataBlock._memory_range[i][0]
4369 inner_loop = For([Integer(0), loop_bound])
4370 inner_loops.append(inner_loop)
4371 if type(self._dataBlock._internal) is not FunctionCall:
4372 inner_loop = For(self._dataBlock._memory_range[0])
4373 inner_loops.append(inner_loop)
4376 for i in range(len(inner_loops)):
4377 index.append(inner_loops[-(i + 1)].get_iteration_variable())
4378 index.append(self.initialisation.get_iteration_variable())
4380 output_offset = [Integer(0) for i in self._dataBlock._iteration_range[1:]]
4381 if type(self._dataBlock._internal._arguments[0]) is Index:
4382 input_offset = [0 for i in range(1, len(self._dataBlock._iteration_range))]
4384 input_offset = [self._dataBlock._iteration_range[i][0] - Name.variables[self._dataBlock._internal._arguments[0].id]._iteration_range[i][0] for i in range(1, len(self._dataBlock._iteration_range))]
4386 if type(self._dataBlock._internal) is FunctionCall:
4387 functionCall = FunctionCall(self._dataBlock._internal.id, [])
4388 functionCall._is_offloadable = self._dataBlock._internal._is_offloadable
4390 if type(self._dataBlock._internal._arguments[0]) is Index:
4391 functionCall.add_argument(Vector(index[0:-1], TInteger()))
4393 functionCall.add_argument(Reference(self._dataBlock._internal._arguments[0].multidimensional_index(index, 1)))
4395 for argument in self._dataBlock._internal._arguments[1:]:
4396 if type(argument.get_type()) is TDataBlock:
4397 offset = [i[0] for i in Name.variables[argument.id]._iteration_range[1:]]
4398 if len(Name.variables[argument.id]._iteration_range) == 1 or Name.variables[argument.id]._iteration_range[0][1] == Integer(1):
4399 functionCall.add_argument(Name.variables[argument.id].multidimensional_index(index, 1))
4401 functionCall.add_argument(Reference(Name.variables[argument.id].multidimensional_index(index, 1)))
4403 functionCall.add_argument(argument)
4405 functionCall.add_argument(Reference(self._dataBlock.multidimensional_index(index, 1)))
4406 if functionCall._is_offloadable:
4407 functionCall.add_argument(String("Solver::Offloadable::Yes"))
4409 functionCall = Assignment(self._dataBlock.multidimensional_index(index), self._dataBlock._internal)
4411 self.initialisation.add_statement(inner_loops[0])
4413 for i in range (0, len(inner_loops) - 1):
4414 inner_loops[i].add_statement(inner_loops[i + 1])
4415 inner_loops[-1].add_statement(functionCall)
4417 for i in range(len(inner_loops) - 1, -1, -1):
4418 inner_loops[i].close_scope()
4419 self.initialisation.close_scope()
4420 self._loops = inner_loops
4421 self._loops.insert(0, self.initialisation)
4423 def print_cpp(self, indent_level = 0):
4424 return f"""{self.memoryAllocation.print_cpp(indent_level)}
4425{self.initialisation.print_cpp(indent_level)}"""
4427 def print_omp(self, indent_level = 0):
4428 if Node.use_accelerator == True:
4429 omp_pragma = "#pragma omp target teams loop collapse(Dimensions + 1) device(targetDevice)"
4431 omp_pragma = "#pragma omp parallel for simd collapse(Dimensions + 1)"
4432 return f"""{self.memoryAllocation.print_omp(indent_level)}
4433{' ' * indent_level * Node.spaces_per_tab + omp_pragma}
4434{self.initialisation.print_omp(indent_level)}"""
4436 def print_sycl(self, indent_level = 0):
4437 ranges = [f"range{i} = {loop.get_interval_size().print_sycl()};" for i, loop in enumerate(self._loops[0:3])]
4438 indexing = [f"int {loop.get_iteration_variable().print_sycl()} = index[{i}];" for i, loop in enumerate(self._loops[0:3])]
4439 statement_prints = [statement.print_sycl(indent_level + 1) for statement in self._loops[2]._statements]
4440 return self.memoryAllocation.print_sycl(indent_level) + os.linesep.join(ranges) + "\n" + ' ' * indent_level * Node.spaces_per_tab + f"""queue.submit([&](::sycl::handler& handler) {{
4441{' ' * indent_level * Node.spaces_per_tab}handler.parallel_for(::sycl::range<3>{{range0, range1, range2}}, [=](::sycl::item<3> index) {{
4442{os.linesep.join(indexing)}
4443{os.linesep.join(statement_prints)}
4444{' ' * (indent_level + 1) * Node.spaces_per_tab}}});
4445{' ' * indent_level * Node.spaces_per_tab}}}).wait();"""
4448 def print_mlir(self, indent_level = 0):
4449 Node.context.push_block()
4450 mem_alloc = self.memoryAllocation.print_mlir(indent_level)
4451 init = self.initialisation.print_mlir(indent_level)
4452 alloc_block = Node.context.pop_block().get_mlir()
4453 for op in alloc_block:
4454 Node.context.get_block().mlir_append(op)
4456 return f"""{mem_alloc}
4459 def print_tree(self, indent_level = 0):
4460 return ' ' * indent_level * Node.spaces_per_tab + f"""DataBlockConstructionFromFunction:
4461{self._name.print_tree(indent_level + 1)}
4462{self._dataBlock.print_tree(indent_level + 1)}"""
4464class Assignment(Statement):
4465 def __init__(self, lhs: Expression, rhs: Expression):
4468 self._lhs.on_left(True)
4470 self._rhs.on_left(False) # Just to be sure...
4472 def print_cpp(self, indent_level = 0):
4473 return ' ' * indent_level * Node.spaces_per_tab + f"{self._lhs.print_cpp()} = {self._rhs.print_cpp()};"
4475 def print_omp(self, indent_level = 0):
4476 return ' ' * indent_level * Node.spaces_per_tab + f"{self._lhs.print_omp()} = {self._rhs.print_omp()};"
4478 def print_sycl(self, indent_level = 0):
4479 return ' ' * indent_level * Node.spaces_per_tab + f"{self._lhs.print_sycl()} = {self._rhs.print_sycl()};"
4481 def print_mlir(self, indent_level = 0):
4482 indent = ' ' * indent_level * Node.spaces_per_tab
4484 Node.context.push_block()
4485 rhs = self._rhs.print_mlir(indent_level)
4489 rhs_casting_mlir = ''
4491 if isinstance(self._rhs, Subscript):
4492 mlir_list = Node.context.get_block().get_mlir()
4493 for i in range(self._rhs._dimensions):
4494 rhs_dims.append('?')
4495 if len(mlir_list) > 0:
4496 rhs_indices.insert(0, self._rhs._new_index)
4497 rhs_casting_mlir += mlir_list.pop()
4499 mlir_str += ''.join(Node.context.pop_block().get_mlir())
4500 mlir_str += rhs_casting_mlir
4502 # Assume scalar here
4503 tmp = ''.join(Node.context.pop_block().get_mlir())
4506 Node.context.push_block()
4507 lhs = self._lhs.print_mlir(indent_level)
4511 lhs_casting_mlir = ''
4513 if isinstance(self._lhs, Subscript):
4514 mlir_list = Node.context.get_block().get_mlir()
4515 for i in range(self._lhs._dimensions):
4516 lhs_dims.append('?')
4517 if len(mlir_list) > 0:
4518 popped = mlir_list.pop()
4519 lhs_indices.insert(0, self._lhs._new_index)
4520 lhs_casting_mlir += popped
4522 mlir_str += ''.join(Node.context.pop_block().get_mlir())
4523 mlir_str += lhs_casting_mlir
4525 tmp = ''.join(Node.context.pop_block().get_mlir())
4528 # Check if LHS is a flat 2D memref (Subscript[Subscript] on String-based field)
4529 lhs_is_flat_2d = False
4530 lhs_memref_name = None
4531 all_lhs_indices = []
4533 if isinstance(self._lhs, Subscript) and isinstance(self._lhs._value, Subscript):
4534 # Check if the base is a String-based field
4535 base_base = self._lhs._value._value
4536 if isinstance(base_base, DataBlock) and isinstance(base_base._internal, String):
4537 # This is a String-based field which is now a flat 2D memref
4538 lhs_is_flat_2d = True
4539 dsl_name = base_base._internal._value
4540 lhs_memref_name = Node.mlir_symbol_table[dsl_name]
4542 outer_sub = self._lhs
4543 inner_sub = self._lhs._value
4545 if hasattr(inner_sub, '_new_index'):
4546 all_lhs_indices.append(inner_sub._new_index)
4547 if hasattr(outer_sub, '_new_index'):
4548 all_lhs_indices.append(outer_sub._new_index)
4553 lhs_index_str = '[' + ','.join(all_lhs_indices) + ']'
4554 lhs_type = f"memref<?x?xf64>"
4555 mlir_str += indent + f"memref.store {rhs_value}, {lhs_memref_name}{lhs_index_str} {{name = \"{lhs_memref_name.strip('%')}\"}} : {lhs_type}\n"
4556 elif isinstance(self._lhs, Subscript) and isinstance(self._lhs._value, Subscript):
4557 base_base = self._lhs._value._value
4559 # If base is not a String field, it's a true nested memref (e.g. lambdaLeft[i][j])
4560 if not (isinstance(base_base, DataBlock) and isinstance(base_base._internal, String)):
4561 # This is a nested memref store (e.g., lambdaLeft[i][j] = value)
4563 if isinstance(self._rhs, Float):
4564 const_id = f"%{Node.get_mlir_id()}"
4565 mlir_str += indent + f"{const_id} = arith.constant {self._rhs._value} : f64\n"
4566 rhs_value = const_id
4567 elif isinstance(self._rhs, Integer):
4568 const_id = f"%{Node.get_mlir_id()}"
4569 mlir_str += indent + f"{const_id} = arith.constant {self._rhs._value} : i32\n"
4570 rhs_value = const_id
4571 elif isinstance(self._rhs, Subscript):
4572 rhs_id = f"%{Node.get_mlir_id()}"
4573 rhs_index_str = '[' + ','.join(rhs_indices) + ']'
4574 rhs_type = "memref<?xf64>" if not rhs_dims else f"memref<{ 'x'.join(rhs_dims) }xf64>"
4575 mlir_str += indent + f"{rhs_id} = memref.load {rhs}{rhs_index_str} {{name = \"{rhs.strip('%')}\"}} : {rhs_type.strip(' ')}\n"
4578 # Extract indices from nested subscript chain
4579 # lhs._value is the inner subscript, lhs is the outer subscript
4580 outer_index = self._lhs._new_index
4581 inner_memref_var = self._lhs._value._value.print_mlir(indent_level) # Base memref name
4583 # Load the inner memref from the outer nested structure
4584 inner_memref_id = f"%{Node.get_mlir_id()}"
4585 # Get the actual outer memref type from type_map
4586 base_var = self._lhs._value._value
4587 type_entry = Node.type_map[base_var.id]
4589 if isinstance(type_entry, str):
4590 outer_memref_type = type_entry
4592 outer_memref_type = type_entry.print_mlir(0).strip()
4594 # Load inner memref with first index
4595 inner_sub_index = self._lhs._value._new_index
4596 mlir_str += indent + f"{inner_memref_id} = memref.load {inner_memref_var}[{inner_sub_index}] {{name = \"{inner_memref_var.strip('%')}\"}} : {outer_memref_type}\n"
4598 # Store directly into the inner memref with the second index
4599 inner_memref_type = "memref<?xf64>" # Default inner arrays are 1D
4600 mlir_str += indent + f"memref.store {rhs_value}, {inner_memref_id}[{outer_index}] {{name = \"{inner_memref_var.strip('%')}\"}} : {inner_memref_type}\n"
4602 # Handle regular flat memref assignment (non-nested)
4603 rhs_id = f"%{Node.get_mlir_id()}"
4605 lhs_index_str = '[' + ','.join(lhs_indices) + ']'
4606 # Ensure at least one dimension for memref type
4608 lhs_type = f"memref<{ 'x'.join(lhs_dims) }xf64>"
4610 lhs_type = "memref<?xf64>"
4611 rhs_index_str = '[' + ','.join(rhs_indices) + ']'
4613 rhs_type = f"memref<{ 'x'.join(rhs_dims) }x f64>"
4615 rhs_type = "memref<?xf64>"
4617 if isinstance(self._rhs, Subscript):
4618 mlir_str += indent + f"{rhs_id} = memref.load {rhs}{rhs_index_str} {{name = \"{rhs.strip('%')}\"}} : {rhs_type.strip(' ')}\n"
4623 if isinstance(self._rhs, Float):
4624 const_id = f"%{Node.get_mlir_id()}"
4625 mlir_str += indent + f"{const_id} = arith.constant {self._rhs._value} : f64\n"
4626 rhs_value = const_id
4627 elif isinstance(self._rhs, Integer):
4628 const_id = f"%{Node.get_mlir_id()}"
4629 mlir_str += indent + f"{const_id} = arith.constant {self._rhs._value} : i32\n"
4630 rhs_value = const_id
4632 if isinstance(self._lhs, Subscript):
4633 mlir_str += indent + f"memref.store {rhs_value}, {lhs}{lhs_index_str} {{name = \"{lhs.strip('%')}\"}} : {lhs_type.strip(' ')}\n"
4635 lhs_type_obj = self._lhs.get_type()
4636 lhs_mlir_type = lhs_type_obj.print_mlir(indent_level).strip(' ')
4637 mlir_str += indent + f"memref.store {rhs_value}, {lhs} {{name = \"{lhs.strip('%')}\"}} : {lhs_mlir_type}\n"
4641 def print_tree(self, indent_level = 0):
4642 return ' ' * indent_level * Node.spaces_per_tab + f"""Assignment:
4643{self._lhs.print_tree(indent_level + 1)}
4644{self._rhs.print_tree(indent_level + 1)}"""
4647class DataBlockAssignment(Statement):
4648 def __init__(self, lhs: Expression, rhs: Expression):
4650 if type(lhs) is Name:
4651 self._lhs = Name.variables[lhs.id]
4652 self._lhs._id = lhs.id
4656 if type(rhs) is Name:
4657 self._rhs = Name.variables[rhs.id]
4661 self._lhs.on_left(True)
4662 self._rhs.on_left(False) # Just to be sure...
4667 self._loops.append(For([Integer(0), self._lhs._iteration_range[-1][1] - self._lhs._iteration_range[-1][0]]))
4668 for i in range(len(self._lhs._memory_range) - 2, -1, -1):
4669 self._loops.append(For([Integer(0), self._lhs._iteration_range[i][1] - self._lhs._iteration_range[i][0]]))
4671 for i in range(len(self._lhs._memory_range) - 1, -1, -1):
4672 index.append(self._loops[i].get_iteration_variable())
4674 if type(self._rhs.get_type()) is TDataBlock:
4676 for i in range(len(self._lhs._iteration_range)):
4677 lhs_start = self._lhs._iteration_range[i][0]
4678 rhs_start = self._rhs._iteration_range[i][0]
4679 if type(rhs_start.get_type()) is TDataBlock:
4680 rhs_start = Subscript(rhs_start, lhs_start)
4681 rhs_offset.append(lhs_start - rhs_start)
4683 rhs_offset = [Integer(0) for i in self._lhs._iteration_range]
4685 lhs_offset = [Integer(0) for i in self._lhs._iteration_range]
4687 self._loops[-1].add_statement(Assignment(self._lhs.multidimensional_index(index, dimensions=self._lhs._memory_range), self._rhs.multidimensional_index(index, dimensions=self._lhs._memory_range)))
4688 for i in range(0, len(self._loops) - 1):
4689 self._loops[i].add_statement(self._loops[i + 1])
4691 for loop in self._loops[::-1]:
4693 self._loop = self._loops[0]
4694 self._numDimensions = len(self._loops)
4696 def print_cpp(self, indent_level = 0):
4697 return self._loop.print_cpp(indent_level)
4699 def print_omp(self, indent_level = 0):
4700 if Node.use_accelerator == True:
4701 omp_pragma = "#pragma omp target teams loop collapse(Dimensions + 1) device(targetDevice)"
4703 omp_pragma = "#pragma omp parallel for simd collapse(Dimensions + 1)"
4704 return ' ' * indent_level * Node.spaces_per_tab + omp_pragma + f"""
4705{self._loop.print_omp(indent_level)}"""
4707 def print_sycl(self, indent_level = 0):
4708 #return self._loop.print_sycl(indent_level)
4709 ranges = [f"range{i} = {loop.get_interval_size().print_sycl()};" for i, loop in enumerate(self._loops[0:3])]
4710 indexing = [f"int {loop.get_iteration_variable().print_sycl()} = index[{i}];" for i, loop in enumerate(self._loops[0:3])]
4711 statement_prints = [statement.print_sycl(indent_level + 1) for statement in self._loops[2]._statements]
4712 return os.linesep.join(ranges) + "\n" + ' ' * indent_level * Node.spaces_per_tab + f"""queue.submit([&](::sycl::handler& handler) {{
4713{' ' * indent_level * Node.spaces_per_tab}handler.parallel_for(::sycl::range<3>{{range0, range1, range2}}, [=](::sycl::item<3> index) {{
4714{os.linesep.join(indexing)}
4715{os.linesep.join(statement_prints)}
4716{' ' * (indent_level + 1) * Node.spaces_per_tab}}});
4717{' ' * indent_level * Node.spaces_per_tab}}}).wait();"""
4720 def print_mlir(self, indent_level = 0):
4721 return self._loop.print_mlir(indent_level)
4723 def print_tree(self, indent_level = 0):
4724 return ' ' * indent_level * Node.spaces_per_tab + f"""DataBlockAssignment:
4725{self._lhs.print_tree(indent_level + 1)}
4726{self._rhs.print_tree(indent_level + 1)}"""
print_cpp(self, indent_level=0)
print_tree(self, indent_level=0)
print_omp(self, indent_level=0)
print_mlir(self, indent_level=0)
__init__(self, id, Type argument_type, namespace=None, is_function=False)
print_sycl(self, indent_level=0)
print_tree(self, indent_level=0)
print_sycl(self, indent_level=0)
print_cpp(self, indent_level=0)
__init__(self, operation, lhs, rhs)
cast_operand(operand_mlir, operand, result_type, indent_level)
print_mlir(self, indent_level=0)
multidimensional_index(self, index, start_index=0, dimensions=[])
get_mlir_arithmetic_op(operation, op_type_class)
print_omp(self, indent_level=0)
print_cpp(self, indent_level=0)
print_mlir(self, indent_level=0)
print_tree(self, indent_level=0)
print_sycl(self, indent_level=0)
print_omp(self, indent_level=0)
print_cpp(self, indent_level=0)
print_sycl(self, indent_level=0)
get_mlir_comparison_op(operation, op_type_class)
print_omp(self, indent_level=0)
print_tree(self, indent_level=0)
__init__(self, operation, lhs, rhs)
print_mlir(self, indent_level=0)
requires_memory_allocation
print_mlir(self, indent_level=0)
print_cpp(self, indent_level=0)
multidimensional_index(self, index, start_index=0, dimensions=[])
print_sycl(self, indent_level=0)
print_tree(self, indent_level=0)
print_omp(self, indent_level=0)
__init__(self, operation, lhs, rhs, useFunctionSyntax=False)
multidimensional_index(self, index, start_index=0, dimensions=[])
__init__(self, operation, lhs, rhs)
requires_memory_allocation
print_omp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_mlir(self, indent_level=0)
print_tree(self, indent_level=0)
print_cpp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_mlir(self, indent_level=0)
__init__(self, dataBlock, index)
print_omp(self, indent_level=0)
print_tree(self, indent_level=0)
print_cpp(self, indent_level=0)
multidimensional_index(self, index, start_index=0, dimensions=[])
print_omp(self, indent_level=0)
print_tree(self, indent_level=0)
print_sycl(self, indent_level=0)
print_cpp(self, indent_level=0)
print_mlir(self, indent_level=0)
multidimensional_index(self, index, start_index=0, dimensions=[])
print_cpp(self, indent_level=0)
print_tree(self, indent_level=0)
print_omp(self, indent_level=0)
print_sycl(self, indent_level=0)
__init__(self, dataBlock)
print_mlir(self, indent_level=0)
set_output_variable(self, outputVariable)
print_tree(self, indent_level=0)
print_cpp(self, indent_level=0)
multidimensional_index(self, index, start_index=0, dimensions=[])
print_mlir(self, indent_level=0)
requires_memory_allocation
print_omp(self, indent_level=0)
__init__(self, operation, dataBlock)
print_sycl(self, indent_level=0)
print_cpp(self, indent_level=0)
print_omp(self, indent_level=0)
__init__(self, iteration_range, internal, requires_memory_allocation, id=None, underlying_type=String("double"))
print_mlir(self, indent_level=0)
print_tree(self, indent_level=0)
linearise_index(self, indices, dimensions, offset_start_index=0)
requires_memory_allocation
print_sycl(self, indent_level=0)
multidimensional_index(self, index_list, start_index=0, dimensions=[])
print_sycl(self, indent_level=0)
__init__(self, id, matrix)
print_mlir(self, indent_level=0)
print_tree(self, indent_level=0)
print_cpp(self, indent_level=0)
print_omp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_cpp(self, indent_level=0)
print_tree(self, indent_level=0)
print_omp(self, indent_level=0)
__init__(self, id, matrix)
print_mlir(self, indent_level=0)
print_sycl(self, indent_level=0)
print_tree(self, indent_level=0)
print_omp(self, indent_level=0)
print_cpp(self, indent_level=0)
print_mlir(self, indent_level=0)
__init__(self, width, internal, index_dimensions)
multidimensional_index(self, index, start_index=0, dimensions=[])
multidimensional_index(self, index_list, start_index=0, dimensions=[])
__init__(self, iteration_range, internal, requires_memory_allocation, id=None, underlying_type=String("double"))
multidimensional_index(self, index_list, start_index=0, dimensions=[])
print_mlir(self, indent_level=0)
__init__(self, value, string=None, reference=False)
print_cpp(self, indent_level=0)
print_tree(self, indent_level=0)
print_sycl(self, indent_level=0)
print_omp(self, indent_level=0)
print_cpp(self, indent_level=0)
get_iteration_variable(self)
add_statement(self, statement)
__init__(self, iteration_range, iteration_variable_name=None, use_scheduler=False)
print_definition_with_timer(self, indent_level=0)
print_cpp(self, indent_level=0)
print_declaration(self, indent_level=0)
add_argument(self, Argument argument)
add_statement(self, Statement statement)
print_declaration_with_timer(self, indent_level=0)
top_level(self, top=None)
print_tree(self, indent_level=0)
print_sycl(self, indent_level=0)
__init__(self, id, Type return_type, template=None, namespaces=[], stateless=False, top_level=False)
print_mlir(self, indent_level=0)
print_omp(self, indent_level=0)
print_omp(self, indent_level=0)
print_cpp(self, indent_level=0)
print_tree(self, indent_level=0)
add_statement(self, statement)
print_mlir(self, indent_level=0)
print_sycl(self, indent_level=0)
print_cpp(self, indent_level=0)
print_omp(self, indent_level=0)
print_mlir(self, indent_level=0)
print_tree(self, indent_level=0)
__init__(self, value, string=None)
print_sycl(self, indent_level=0)
__init__(self, dataBlock, filename)
print_mlir(self, indent_level=0)
print_sycl(self, indent_level=0)
print_tree(self, indent_level=0)
print_cpp(self, indent_level=0)
print_omp(self, indent_level=0)
print_omp(self, indent_level=0)
print_mlir(self, indent_level=0)
__init__(self, statementToLog, outputStream)
print_tree(self, indent_level=0)
print_sycl(self, indent_level=0)
print_cpp(self, indent_level=0)
print_mlir(self, indent_level=0)
__init__(self, id, matrix)
print_sycl(self, indent_level=0)
print_tree(self, indent_level=0)
print_omp(self, indent_level=0)
print_cpp(self, indent_level=0)
index(self, row_index, column_index)
__init__(self, dimensions, internal, index_dimensions)
print_cpp(self, indent_level=0)
print_tree(self, indent_level=0)
print_mlir(self, indent_level=0)
print_sycl(self, indent_level=0)
print_omp(self, indent_level=0)
multidimensional_index(self, index_list, start_index=0, dimensions=[])
print_sycl(self, indent_level=0)
print_cpp(self, indent_level=0)
print_mlir(self, indent_level=0)
print_omp(self, indent_level=0)
print_tree(self, indent_level=0)
__init__(self, id, type=None)
create_memref_from_ptr(cls, ptr_id, target_type, indent_level=0, size_hint=None, target_var_name=None)
print_omp(self, indent_level=0)
print_cpp(self, indent_level=0)
print_sycl(self, indent_level=0)
parent(self, parent=None)
print_tree(self, indent_level=0)
generate_pointer_to_2d_element(memref_mlir, row_index, col_index, is_nested, indent_level=0)
print_mlir(self, indent_level=0)
create_memref_from_extracted_descriptor(cls, alloc_ptr_id, aligned_ptr_id, offset_id, sizes_0_id, sizes_1_id, strides_0_id, strides_1_id, target_type, struct_type, indent_level=0, target_var_name=None, is_2d=False)
get_mlir_type(type_string)
print_omp(self, indent_level=0)
print_cpp(self, indent_level=0)
print_mlir(self, indent_level=0)
__init__(self, Expression expression)
print_tree(self, indent_level=0)
print_sycl(self, indent_level=0)
print_mlir(self, indent_level=0)
print_omp(self, indent_level=0)
print_cpp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_tree(self, indent_level=0)
print_tree(self, indent_level=0)
print_mlir(self, indent_level=0)
is_nested_subscript(self)
print_cpp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_omp(self, indent_level=0)
__init__(self, Expression value, Expression index)
print_mlir(self, indent_level=0)
print_tree(self, indent_level=0)
print_omp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_cpp(self, indent_level=0)
print_tree(self, indent_level=0)
print_sycl(self, indent_level=0)
__init__(self, type_name)
print_mlir(self, indent_level=0)
print_omp(self, indent_level=0)
print_cpp(self, indent_level=0)
generate_memref_load(operand, operand_mlir, const_zero_id, indent_level)
print_cpp(self, indent_level=0)
get_mlir_type_for_operation(lhs, rhs)
print_omp(self, indent_level=0)
needs_memref_load(operand)
__init__(self, dimensions, underlying_type)
get_array_element_info(self)
print_tree(self, indent_level=0)
print_sycl(self, indent_level=0)
print_mlir(self, indent_level=0)
ensure_const_zero(const_zero_id, indent_level)
print_cpp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_omp(self, indent_level=0)
print_tree(self, indent_level=0)
print_mlir(self, indent_level=0)
print_omp(self, indent_level=0)
print_cpp(self, indent_level=0)
print_tree(self, indent_level=0)
print_sycl(self, indent_level=0)
print_mlir(self, indent_level=0)
__init__(self, reference=False)
print_cpp(self, indent_level=0)
print_omp(self, indent_level=0)
__init__(self, signature_string)
print_mlir(self, indent_level=0)
print_sycl(self, indent_level=0)
print_tree(self, indent_level=0)
print_omp(self, indent_level=0)
print_mlir(self, indent_level=0)
print_sycl(self, indent_level=0)
print_cpp(self, indent_level=0)
Type for 64-bit integers (long) used in MLIRCellData struct fields.
print_mlir(self, indent_level=0)
print_cpp(self, indent_level=0)
print_tree(self, indent_level=0)
print_omp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_cpp(self, indent_level=0)
print_omp(self, indent_level=0)
print_mlir(self, indent_level=0)
print_sycl(self, indent_level=0)
print_tree(self, indent_level=0)
print_cpp(self, indent_level=0)
print_omp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_mlir(self, indent_level=0)
__init__(self, element_type=TFloat())
print_omp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_mlir(self, indent_level=0)
print_cpp(self, indent_level=0)
__init__(self, element_type=TFloat(), dimensions=1)
print_tree(self, indent_level=0)
print_mlir(self, indent_level=0)
print_omp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_cpp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_mlir(self, indent_level=0)
print_cpp(self, indent_level=0)
print_tree(self, indent_level=0)
print_omp(self, indent_level=0)
print_omp(self, indent_level=0)
print_cpp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_mlir(self, indent_level=0)
print_tree(self, indent_level=0)
print_sycl(self, indent_level=0)
__init__(self, operation, operand)
print_tree(self, indent_level=0)
print_omp(self, indent_level=0)
print_mlir(self, indent_level=0)
print_cpp(self, indent_level=0)
print_cpp(self, indent_level=0)
__init__(self, Expression value, Expression index)
print_mlir(self, indent_level=0)
print_tree(self, indent_level=0)
print_omp(self, indent_level=0)
print_sycl(self, indent_level=0)
print_sycl(self, indent_level=0)
print_cpp(self, indent_level=0)
__init__(self, elements, element_type, string=None)
print_omp(self, indent_level=0)
print_mlir(self, indent_level=0)
print_tree(self, indent_level=0)
_is_nested_memref(type_map_entry)