Peano
Loading...
Searching...
No Matches
SyntaxTree.py
Go to the documentation of this file.
1from abc import ABC, abstractmethod
2import os
3import copy
4from enum import Enum
5from operator import itemgetter
6
7# We need to track the context of the code generation to support MLIR
8class Context:
9
10 def __init__(self):
11 # A stack of blocks e.g. for loops where we can store hoisted MLIR operations
12 self._blocks = [ Context.Block() ]
13
14 def push_block(self):
15 self._blocks.append(Context.Block())
16
17 def pop_block(self):
18 if self._blocks:
19 return self._blocks.pop()
20 else:
21 return Context.Block()
22
23 def get_block(self):
24 if not self._blocks:
25 self.push_block()
26 return self._blocks[-1]
27
28 def clear_block(self):
29 if self._blocks:
30 self._blocks[-1] = Context.Block()
31
32 # We need to store block-level MLIR that has been hoisted and ensure we don't
33 # repeat casting operators for indices within a block e.g. loops
34 class Block:
35
36 def __init__(self):
37 # We use this to hoist MLIR operations out of array indexing, function calls etc.
38 self._mlir = []
39 # We use a set to store the ids that we have already processed within the block
40 self._ids = {}
41
42 def push_mlir(self):
43 self._mlir.append([])
44
45 def pop_mlir(self):
46 if self._mlir:
47 return self._mlir.pop()
48 else:
49 return []
50
51 def get_mlir(self):
52 if not self._mlir:
53 self.push_mlir()
54 return self._mlir[-1]
55
56 def mlir_append(self, op):
57 self.get_mlir().append(op)
58
59 # def add_id(self, id):
60 # self._ids.update(id)
61
62 # def remove_id(self, id):
63 # self._ids.remove(id)
64
65 # def get_id(self, id):
66 # return self._ids.get(id)
67
68 # def find_id(self, id):
69 # return id in self._ids
70
71
72
73class Node(ABC):
74 output_format = None
75 spaces_per_tab = 4
76 use_accelerator = False
77 use_memory_manager = False
78 mlir_id_count = 0 # We need a variable to keep track of the MLIR ids %0, %1 etc.
79 type_map = {}
80 struct_map = {}
81 context = Context()
82 target_variable_name = None # For tracking when we're generating a specific variable name
83 mlir_symbol_table = {} # Maps DSL names (e.g., "patchData.QOut") to MLIR variable names (e.g., "%output")
84 mlir_memref_types = {} # Maps MLIR variable names (e.g., "%output") to their memref types (e.g., "memref<?x?xf64>")
85
86 # Initialise the type, struct and function maps
87 # NOTE: Populated lazily after Type classes are defined (see __init_type_map below)
88 type_map.update({'void': 'void'}) # Will be replaced with TVoid() in Phase 2a
89 type_map.update({'int': 'i64'}) # Match TInteger class and C++ kernel expectations
90 type_map.update({'double': 'f64'}) # Will be replaced with TFloat() in Phase 2a
91 type_map.update({'double*': 'memref<?xf64>'})
92 type_map.update({'tarch::la::Vector<2, double>': 'memref<?xi64>'}) # Using memref instead of !llvm.ptr for standard indexing
93 type_map.update({'tarch::la::Vector<3, double>': 'memref<?xi64>'}) # Using memref instead of !llvm.ptr for standard indexing
94 type_map.update({'CellData<double, double>&': '!llvm.ptr'}) #'memref<?xstruct<f64, f64>>'})
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'})
101 # Note: QIn, QOut, cellCentre, cellSize, t, dt, id, numberOfCells, etc. are populated from struct_map below
102
103 # Functor function MLIR signatures for individual function calls
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>) -> ()'})
108
109 # Local variable storage types (for runtime-allocated temporary variables)
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>>'})
119
120 # Lambda/flux arrays (2D double** patterns) - These are ASSIGNED to with nested subscripts
121 # so they need nested memref type mapping
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>>'})
134
135 # Bridge arrays: Used as temporary storage, allocated with outer dimension then inner arrays
136 # Like lambdaLeft, these are semantically "array of arrays" in the C++ kernel
137 # Must use nested memref type to match the allocation pattern
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>>'})
150
151 # MLIR to C++ type mappings for function signature generation
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'})
158
159 # MLIRCellData struct layout with MemRefDescriptor fields
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>)>"
162
163 # Updated patchTypes: 7 descriptors + 3 i64 scalars + 2 ptrs = 12 fields
164 patchTypes = [descriptor2DType, descriptor2DType, descriptor1DType, descriptor1DType,
165 descriptor1DType, descriptor1DType, descriptor1DType, 'i64', 'i64', 'i64', 'ptr', 'ptr']
166
167 struct_map.update({'patchData': [
168 ('QIn' , f"{descriptor2DType}|||!llvm.struct<({', '.join(patchTypes)})>"), # Field 0: 2D descriptor
169 ('QOut' , f"{descriptor2DType}|||!llvm.struct<({', '.join(patchTypes)})>"), # Field 1: 2D descriptor
170 ('cellCentre' , f"{descriptor1DType}|||!llvm.struct<({', '.join(patchTypes)})>"), # Field 2: 1D descriptor
171 ('cellSize' , f"{descriptor1DType}|||!llvm.struct<({', '.join(patchTypes)})>"), # Field 3: 1D descriptor
172 ('t' , f"{descriptor1DType}|||!llvm.struct<({', '.join(patchTypes)})>"), # Field 4: 1D descriptor
173 ('dt' , f"{descriptor1DType}|||!llvm.struct<({', '.join(patchTypes)})>"), # Field 5: 1D descriptor
174 ('id' , f"{descriptor1DType}|||!llvm.struct<({', '.join(patchTypes)})>"), # Field 6: 1D descriptor
175 ('numberOfCells' , f"i64|||!llvm.struct<({', '.join(patchTypes)})>"), # Field 7: i64
176 ('memoryLocation' , f"i64|||!llvm.struct<({', '.join(patchTypes)})>"), # Field 8: i64
177 ('targetDevice' , f"i64|||!llvm.struct<({', '.join(patchTypes)})>"), # Field 9: i64
178 ('QOut_legacy' , f"!llvm.ptr|||!llvm.struct<({', '.join(patchTypes)})>"), # Field 10: function pointer
179 ('maxEigenvalue' , f"memref<?xf64>|||!llvm.struct<({', '.join(patchTypes)})>")]}) # Field 11: function pointer
180
181 # Populate type_map with descriptor field types from struct_map
182 # This maps field names to their converted MLIR memref types
183 for struct_name, fields in struct_map.items():
184 for field_name, field_info in fields:
185 target_type = field_info.split('|||')[0]
186 # Convert descriptor types to their memref equivalents
187 if target_type == descriptor2DType:
188 type_map.update({field_name: 'memref<?x?xf64>'})
189 elif target_type == descriptor1DType:
190 # Distinguish based on field name:
191 # - cellCentre, cellSize, h, x: these are Vector<2,double> fields (stored as i64 pointers)
192 # - t, dt, id: these are scalar f64 fields
193 if field_name in ['cellCentre', 'cellSize', 'h', 'x']:
194 type_map.update({field_name: 'memref<?xi64>'}) # Vector fields (pointers to vectors)
195 else:
196 type_map.update({field_name: 'memref<?xf64>'}) # Scalar fields (f64 values)
197 elif target_type == 'i64':
198 type_map.update({field_name: 'i64'})
199 elif target_type == '!llvm.ptr':
200 # For MLIR pointer types, keep them as-is (they're function pointers, not data pointers)
201 type_map.update({field_name: '!llvm.ptr'})
202 else:
203 # For other types, use the target_type directly
204 type_map.update({field_name: target_type})
205
206 def __init__(self):
207 self._parent = None
208 self._on_left = False
209 self._hoist = False
210 self._index_type = False
211
212 def parent(self, parent = None):
213 if parent is not None:
214 self._parent = parent
215 return self._parent
216
217 def on_left(self, left = None):
218 if left is not None:
219 self._on_left = left
220 return self._on_left
221
222 def hoist(self, flag = None):
223 if flag is not None:
224 self._hoist = flag
225 return self._hoist
226
227 @classmethod
228 def get_mlir_id(cls):
229 current_reg = cls.mlir_id_countmlir_id_count
231 return current_reg
232
233 @classmethod
237 @staticmethod
238 def get_mlir_type(type_string):
239 # Handle string representations by creating temporary type objects
240 if isinstance(type_string, str):
241 if type_string == "int":
242 return TInteger().print_mlir(0)
243 elif type_string == "double":
244 return TFloat().print_mlir(0)
245 elif type_string == "bool":
246 return TBoolean().print_mlir(0)
247 return None
248
249 @staticmethod
250 def generate_pointer_to_2d_element(memref_mlir, row_index, col_index, is_nested, indent_level=0):
251 indent = ' ' * indent_level * Node.spaces_per_tab
252 mlir_code = ""
253
254 if is_nested:
255 # Load inner memref from outer array, then offset within it
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"
258
259 # Extract pointer as index from inner memref
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"
262
263 # Byte offset calculation (col * 8 for f64 elements)
264 element_size_bytes = f"%{Node.get_mlir_id()}"
265 mlir_code += indent + f"{element_size_bytes} = arith.constant 8 : index\n"
266
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"
269
270 # Add offset to base pointer
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"
273
274 # Convert to i64, then to llvm.ptr
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"
277
278 result_ptr = f"%{Node.get_mlir_id()}"
279 mlir_code += indent + f"{result_ptr} = llvm.inttoptr {ptr_i64} : i64 to !llvm.ptr\n"
280 else:
281 # Flat 2D memref: compute (row * width + col) offset
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"
286
287 # Calculate row offset (row * width)
288 row_offset = f"%{Node.get_mlir_id()}"
289 mlir_code += indent + f"{row_offset} = arith.muli {row_index}, {col_width_id} : index\n"
290
291 # Add column to row offset
292 elem_offset = f"%{Node.get_mlir_id()}"
293 mlir_code += indent + f"{elem_offset} = arith.addi {row_offset}, {col_index} : index\n"
294
295 # Convert element offset to bytes
296 element_size_bytes = f"%{Node.get_mlir_id()}"
297 mlir_code += indent + f"{element_size_bytes} = arith.constant 8 : index\n"
298
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"
301
302 # Extract base pointer from memref
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"
305
306 # Add byte offset to base
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"
309
310 # Convert to i64, then to llvm.ptr
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"
313
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"
316
317 return mlir_code, result_ptr
318
319 @classmethod
320 def create_memref_from_ptr(cls, ptr_id, target_type, indent_level=0, size_hint=None, target_var_name=None):
321 indent = ' ' * indent_level * cls.spaces_per_tab
322
323 # Generate unique IDs for all the operations
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()}"
330
331 # Use target variable name if provided, otherwise generate a new one
332 if target_var_name:
333 memref_id = target_var_name
334 else:
335 memref_id = f"%{cls.get_mlir_id()}"
336
337 # Create constants we need
338 c0_i64_id = f"%{cls.get_mlir_id()}"
339 c1_i64_id = f"%{cls.get_mlir_id()}"
340
341 # Use standard LLVM memref descriptor format - this is part of the fix supplied by Nick Brown
342 # For 2D memrefs, use 2D descriptor; for 1D memrefs, use 1D descriptor
343 if target_type.startswith('memref<?x?'):
344 # 2D memref needs 2D descriptor
345 struct_type = '!llvm.struct<(ptr, ptr, i64, array<2 x i64>, array<2 x i64>)>'
346 num_dimensions = 2
347 else:
348 # 1D memref needs 1D descriptor
349 struct_type = '!llvm.struct<(ptr, ptr, i64, array<1 x i64>, array<1 x i64>)>'
350 num_dimensions = 1
351
352 # Determine size based on target type or hint
353 if target_type.startswith('memref<?xmemref') or size_hint == 1:
354 size_val = "1"
355 elif target_type.startswith('memref<?xi64>') and size_hint == 2:
356 size_val = "2"
357 elif target_type.startswith('memref<?xi64>'):
358 size_val = "2" # Default for vector types
359 else:
360 size_val = "1" # Default for scalar-like types
361
362 mlir_str = ""
363
364 # Create the required constants
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"
367
368 # If we need a size constant, create it
369 if size_val != "1":
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
375 else:
376 size_ref = c1_i64_id
377
378 # Create the LLVM memref descriptor
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:
385 # For 2D descriptor, set [3, 1], [4, 0] and [4, 1]
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"
391 else:
392 # For 1D descriptor, just set [4, 0]
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"
396
397 return mlir_str, memref_id
398
399 @classmethod
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):
401 indent = ' ' * indent_level * cls.spaces_per_tab
402
403 # Use target variable name if provided
404 if target_var_name:
405 memref_id = target_var_name
406 else:
407 memref_id = f"%{cls.get_mlir_id()}"
408
409 # Generate unique IDs for all the operations
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()}"
416
417 c1_i64_id = f"%{cls.get_mlir_id()}"
418
419 mlir_str = ""
420
421 # Create constants we need
422 mlir_str += indent + f"{c1_i64_id} = arith.constant 1 : i64\n"
423
424 # Create the LLVM memref descriptor using the extracted fields
425 # We preserve the extracted pointers but hardcode sizes and strides to 1
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"
431
432 # Check if 2D or 1D
433 if is_2d:
434 # 2D descriptor - use all extracted components
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"
440 else:
441 # 1D descriptor - use strides_0_id
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
444
445 mlir_str += indent + f"{memref_id} = builtin.unrealized_conversion_cast {desc_final_id} : {struct_type} to {target_type} {{name = \"{memref_id.strip('%')}\"}}\n"
446
447 return mlir_str, memref_id
448
449 @abstractmethod
450 def print_cpp(self, indent_level = 0):
451 pass
452
453 @abstractmethod
454 def print_mlir(self, indent_level = 0):
455 pass
456
457 @abstractmethod
458 def print_omp(self, indent_level = 0):
459 pass
460
461 @abstractmethod
462 def print_sycl(self, indent_level = 0):
463 pass
464
465 @abstractmethod
466 def print_tree(self, indent_level = 0):
467 pass
468
469class Statement(Node):
470 def __init__(self):
471 super().__init__()
472
473
475 def __init__(self):
476 super().__init__()
477 self._mlir_id = None
478
479 def mlir_id(self, id = None):
480 if id is not None:
481 self._mlir_id = id
482 return self._mlir_id
483
484 def multidimensional_index(self, index_list, start_index=0, dimensions = []):
485 return self
486
487 @abstractmethod
488 def get_type(self):
489 pass
490
491
492class Type(Node):
493
494 def print_tree(self, indent_level=0):
495 return ' '
496
498 return None
499
500
501class TVoid(Type):
502 def print_cpp(self, indent_level=0):
503 return ' ' * indent_level * Node.spaces_per_tab + 'void'
504
505 def print_mlir(self, indent_level=0):
506 return ' ' * indent_level * Node.spaces_per_tab + 'void'
507
508 def print_omp(self, indent_level=0):
509 return ' ' * indent_level * Node.spaces_per_tab + 'void'
510
511 def print_sycl(self, indent_level=0):
512 return ' ' * indent_level * Node.spaces_per_tab + 'void'
513
514 def is_function(self):
515 return False
516
517
518
520 def __init__(self, id, argument_type: Type, namespace = None, is_function = False):
521 super().__init__()
522 self.id = id
523 self._type = argument_type
524 self._type.parent(self)
525 self._namespace = namespace
526 self._is_function = is_function
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:]
531
532 def print_cpp(self, indent_level = 0):
533 return self._type.print_cpp() + " " + self.id
534
535 def print_omp(self, indent_level = 0):
536 return self._type.print_omp() + " " + self.id
537
538 def print_sycl(self, indent_level = 0):
539 return self._type.print_sycl() + " " + self.id
540
541 def print_mlir(self, indent_level = 0):
542 # NOTE: Arguments are inserted in-line into 'func.func' op
543 return f"%{self.id} : {self._type.print_mlir(0)}"
544
545 def print_tree(self, indent_level = 0):
546 return ' ' * indent_level * Node.spaces_per_tab + f"Argument: {self.id}"
547
548
550 _stateless = False
551 _syclDataToCopy = []
552
553 def __init__(self, id, return_type: Type, template = None, namespaces = [], stateless = False, top_level = False):
554 super().__init__()
555 self.id = id
556 self._return_type = return_type
557 self._return_type.parent(self)
558 self._template = template
559 self._arguments = []
560 self._body = []
561 self._namespaces = namespaces
562 FunctionDefinition._stateless = stateless
563 self._top_level = top_level
564
565 def top_level(self, top = None):
566 if top is not None:
567 self._top_level = top
568 return self._top_level
569
570 def add_statement(self, statement: Statement):
571 if statement is not None:
572 self._body.append(statement)
573
574 def add_argument(self, argument: Argument):
575 self._arguments.append(argument)
576 if self._top_level:
577 # NOTE: For now, we assume that the top-level function is always the main entry point
578 # and set global values in Node here - 'namespace' is MLIR type here
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})
582
583 def print_declaration(self, indent_level = 0):
584 namespace_header = f"""namespace {"::".join(self._namespaces)} {{"""
585 namespace_footer = "}"
586 template_string = ""
587 if self._template is not None:
588 template_strings = [argument.print_cpp() for argument in self._template]
589 template_string = "template <" + ",".join(template_strings) + ">\n"
590
591 current_indent = indent_level + 1 # Indentation increases inside namespace
592 argument_prints = [argument.print_cpp() for argument in self._arguments]
593
594 return f"""{namespace_header}
595{template_string}{' ' * current_indent * Node.spaces_per_tab}{self._return_type.print_cpp()} {self.id}({', '.join(argument_prints)});
596{namespace_footer}
597"""
598
599 def print_cpp(self, indent_level=0):
600 namespace_header = f"""namespace {"::".join(self._namespaces)} {{"""
601 namespace_footer = "}"
602 template_string = ""
603 if self._template is not None:
604 template_strings = [argument.print_cpp() for argument in self._template]
605 template_string = "template <" + ",".join(template_strings) + ">\n"
606
607 current_indent = indent_level + 1 # Indentation increases inside namespace
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]
610
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}}}
615{namespace_footer}
616"""
617
618 def print_omp(self, indent_level=0):
619 namespace_header = f"""namespace {"::".join(self._namespaces)} {{"""
620 namespace_footer = "}"
621 template_string = ""
622 if self._template is not None:
623 template_strings = [argument.print_omp() for argument in self._template]
624 template_string = "template <" + ",".join(template_strings) + ">\n"
625
626 current_indent = indent_level + 1 # Indentation increases inside namespace
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]
629
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}}}
634{namespace_footer}
635"""
636
637 def print_sycl(self, indent_level=0):
638 namespace_header = f"""namespace {"::".join(self._namespaces)} {{"""
639 namespace_footer = "}"
640 template_string = ""
641 if self._template is not None:
642 template_strings = [argument.print_sycl() for argument in self._template]
643 template_string = "template <" + ",".join(template_strings) + ">\n"
644
645 current_indent = indent_level + 1 # Indentation increases inside namespace
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;")
650
651 data_copy_prints = []
652 for dataBlockCreation in FunctionDefinition._syclDataToCopy:
653 if len(dataBlockCreation._dataBlock._iteration_range) > 1:
654 step_size = Integer(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])
658
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}}}
664{namespace_footer}
665"""
666
667 def print_mlir(self, indent_level = 0):
668 indent_level += 1 # Increase indentation for MLIR
669
670 if self._top_level:
671 Node.reset_mlir_ids() # Reset the MLIR id count for the top-level function
672
673 # NOTE: We used to generate MLIRbridge.cpp using the dedicated MLIRBridgeGenerator module
674 globals_mlir = []
675
676 # Add external function declarations for bridge functions using dynamic signatures
677 # Extract function signatures directly from self._template (functors only, not constants)
678 if self._template is not None:
679 for argument in self._arguments:
680 mlir_signature = argument._namespace #Node.type_map[cpp_type_str]
681
682 # Only process functors, skip template constants like "NumberOfVolumesPerAxisInPatch"
683 if argument._is_function:
684 # Create the bridge function declaration in MLIR with _bridge suffix
685 # This matches the C++ wrapper function names (flux_bridge, etc.)
686 func_decl = f"func.func private @{argument.id}_bridge{mlir_signature}"
687 globals_mlir.append(' ' * (indent_level) * Node.spaces_per_tab + func_decl)
688
689
690 else:
691 # No template means no MLIR bridge functions are needed
692 raise ValueError(
693 "MLIR code generation requires template parameters to determine bridge function signatures. "
694 "Ensure the function has template parameters that specify the required functors."
695 )
696
697 # Add addressof/loads for globals inside the function
698 function_global_loads = []
699 if self._template is not None:
700 for argument in self._template:
701 # Generate globals for ALL template parameters (including boolean evaluation flags)
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")
705
706 # Get the MLIR type for the argument
707 arg_mlir_type = argument._type.print_mlir(0)
708
709 # Skip global loading for variables with empty types (be more thorough)
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}")
712 else:
713 # For variables with empty types, just create a comment
714 function_global_loads.append(' ' * (indent_level + 1) * Node.spaces_per_tab + f"// Skipping load for {argument.id} (empty type: '{arg_mlir_type}')")
715 else:
716 globals_mlir = []
717
718 Node.context.push_block()
719 argument_prints = [argument.print_mlir(indent_level).strip('') for argument in self._arguments]
720 Node.context.push_block()
721
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
725
726 return_type = self._return_type.print_mlir(0)
727 return_type_str = f" -> ({return_type})" if return_type != 'void' else ''
728
729 return f"""
730builtin.module {{
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)}
735func.return
736 }}
737}}
738"""
739 # Non-top-level case
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
746 return_type = self._return_type.print_mlir(indent_level)
747 return_type_str = f" -> ({return_type})" if return_type != 'void' else ''
748 return f"""
749 func.func @{self.id}({', '.join(argument_prints)}){return_type_str} {{
750{os.linesep.join(statements)}
751func.return
752 }}
753"""
754
755 def print_tree(self, indent_level = 0):
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]
758 return f"""
759FunctionDefinition: {self.id}:
760{os.linesep.join(argument_prints)}
761{os.linesep.join(statement_prints)}"""
762
763 def print_definition_with_timer(self, indent_level = 0):
764 namespace_header = f"""namespace {"::".join(self._namespaces)} {{"""
765 namespace_footer = "}"
766 template_string = ""
767 function_call_string = self.id
768 if self._template is not None:
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]) + ">"
772
773 arguments = copy.deepcopy(self._arguments)
774 arguments.append(Argument("measurement", TCustom("tarch::timing::Measurement&")))
775
776 current_indent = indent_level + 1 # Indentation increases inside namespace
777 argument_prints = [argument.print_cpp() for argument in arguments]
778
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])});
783watch.stop();
784measurement.setValue(watch.getCalendarTime());
785{' ' * current_indent * Node.spaces_per_tab}}}
786{namespace_footer}
787"""
788
789 def print_declaration_with_timer(self, indent_level = 0):
790 namespace_header = f"""namespace {"::".join(self._namespaces)} {{"""
791 namespace_footer = "}"
792 template_string = ""
793 if self._template is not None:
794 template_strings = [argument.print_cpp() for argument in self._template]
795 template_string = "template <" + ",".join(template_strings) + ">\n"
796
797 arguments = copy.deepcopy(self._arguments)
798 arguments.append(Argument("measurement", TCustom("tarch::timing::Measurement&")))
799
800 current_indent = indent_level + 1 # Indentation increases inside namespace
801 argument_prints = [argument.print_cpp() for argument in arguments]
802
803 return f"""{namespace_header}
804{template_string}{' ' * current_indent * Node.spaces_per_tab}{self._return_type.print_cpp()} {self.id}({', '.join(argument_prints)});
805{namespace_footer}
806"""
807
809 def __init__(self, statementToLog, outputStream):
810 self._statement = statementToLog
811 self._output = outputStream
812
813 def print_cpp(self, indent_level=0):
814 return f"""{self._output.print_cpp()} << {self._statement.print_cpp()};"""
815
816 def print_omp(self, indent_level=0):
817 return f"""{self._output.print_omp()} << {self._statement.print_omp()};"""
818
819 def print_sycl(self, indent_level=0):
820 return f"""{self._output.print_sycl()} << {self._statement.print_sycl()};"""
821
822 def print_mlir(self, indent_level=0):
823 pass
824
825 def print_tree(self, indent_level=0):
826 return "LogToFile"
827
828
829# Expressions
831 variables = dict()
832
833 def __init__(self, id, type = None):
834 super().__init__()
835 self.id = id
836 self.type = type
837 if self.type is None:
838 self.type = Name.variables[self.id].get_type()
839 self.type._parent = self
840
841 def multidimensional_index(self, index_list, start_index=0, dimensions = []):
842 if type(Name.variables[self.id].get_type()) is TDataBlock or type(Name.variables[self.id].get_type()) is TDiagonalMatrix:
843 return Name.variables[self.id].multidimensional_index(index_list, start_index, dimensions)
844 else:
845 return self
846
847 def index(self, i):
848 return Name.variables[self.id].index(i)
849
850 def index(self, i, j):
851 return Name.variables[self.id].index(i, j)
852
853 def __add__(self, rhs):
854 return BinaryOperation("+", self, rhs)
855
856 def __sub__(self, rhs):
857 return BinaryOperation("-", self, rhs)
858
859 def __mul__(self, rhs):
860 return BinaryOperation("*", self, rhs)
861
862 def __div__(self, rhs):
863 return BinaryOperation("/", self, rhs)
864
865 def __neg__(self):
866 return UnaryOperation("-", self)
867
868 def __gt__(self, rhs):
869 if type(rhs) is Name:
870 return (Name.variables[self.id] > Name.variables[rhs.id])
871 return (Name.variables[self.id] > rhs)
872
873 def __eq__(self, 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])
877 else:
878 return False
879 else:
880 if self.id in Name.variables:
881 return (Name.variables[self.id] == rhs)
882 else:
883 return False
884
885 def get_type(self):
886 return self.type
887
888 def print_cpp(self, indent_level=0):
889 return ' ' * indent_level * Node.spaces_per_tab + self.id
890
891 def print_omp(self, indent_level=0):
892 return ' ' * indent_level * Node.spaces_per_tab + self.id
893
894 def print_sycl(self, indent_level=0):
895 return ' ' * indent_level * Node.spaces_per_tab + self.id
896
897 def print_mlir(self, indent_level=0):
898 # NOTE: Names are returned and used in-linei
899 return "%" + self.id
900
901 def print_tree(self, indent_level=0):
902 return ' ' * indent_level * Node.spaces_per_tab + "Name:" + self.id
903
904class Index:
905 pass
906
907class Integer(Expression):
908 def __init__(self, value, string = None):
909 super().__init__()
910 self._value = int(value)
911 self._string = string
912
913 def __add__(self, rhs):
914 if self._value == 0:
915 return rhs
916 if type(rhs) is Integer:
917 return Integer(self._value + rhs._value)
918 return BinaryOperation("+", self, rhs)
919
920 def __sub__(self, rhs):
921 if self._value == 0:
922 if type(rhs) is Integer:
923 return Integer(-rhs._value)
924 if type(rhs) is UnaryOperation and rhs._operation == "-":
925 return rhs._operand
926 return -rhs
927 if type(rhs) is Integer:
928 return Integer(self._value - rhs._value)
929 return BinaryOperation("-", self, rhs)
930
931 def __mul__(self, rhs):
932 return BinaryOperation("*", self, rhs)
933
934 def __div__(self, rhs):
935 return BinaryOperation("/", self, rhs)
936
937 def __neg__(self):
938 if self._value == 0:
939 return self
940 else:
941 return UnaryOperation("-", self)
942
943 def __lt__(self, rhs):
944 if type(rhs) is Name:
945 if rhs.id in Name.variables:
946 return (self._value < Name.variables[rhs.id])
947 else:
948 return False
949 return (self._value < rhs)
950
951 def __gt__(self, rhs):
952 if type(rhs) is Name:
953 if rhs.id in Name.variables:
954 return (self._value > Name.variables[rhs.id])
955 else:
956 return False
957 return (self._value > rhs)
958
959 def __eq__(self, rhs):
960 if type(rhs) is Integer:
961 return (self._value == rhs._value)
962
963 def __int__(self):
964 return self._value
965
966 def get_type(self):
967 # Check if this Integer references a struct field in struct_map
968 if self._string is not None and '.' in self._string:
969 struct, field = self._string.split('.')
970 if struct in Node.struct_map:
971 struct_list = Node.struct_map[struct]
972 # Find the field and get its type
973 for item in struct_list:
974 if item[0] == field:
975 # Extract the target type using '|||' separator (struct_map format: "target_type|||outer_struct_type")
976 mlir_type = item[1]
977 target_type = mlir_type.split('|||')[0] if '|||' in mlir_type else mlir_type
978 # Create a type object based on the target type
979 if target_type == 'i64':
980 # Return a wrapper that will generate i64
981 return TLong()
982 elif target_type == 'i32':
983 return TInteger()
984 elif target_type == 'f64':
985 return TFloat()
986 # Default: return regular i32 TInteger
987 return TInteger()
988
989 def print_cpp(self, indent_level=0):
990 if self._string is None:
991 return ' ' * indent_level * Node.spaces_per_tab + str(self._value)
992 else:
993 return ' ' * indent_level * Node.spaces_per_tab + str(self._string)
994
995 def print_omp(self, indent_level=0):
996 if self._string is None:
997 return ' ' * indent_level * Node.spaces_per_tab + str(self._value)
998 else:
999 return ' ' * indent_level * Node.spaces_per_tab + str(self._string)
1000
1001 def print_sycl(self, indent_level=0):
1002 if self._string is None:
1003 return ' ' * indent_level * Node.spaces_per_tab + str(self._value)
1004 else:
1005 return ' ' * indent_level * Node.spaces_per_tab + str(self._string)
1006
1007 def print_mlir(self, indent_level=0):
1008 indent = ' ' * indent_level * Node.spaces_per_tab
1009
1010 if self._string is None:
1011 value = str(self._value)
1012 if self.hoist():
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()}"
1019 return id
1020 return value
1021 else:
1022 # NOTE: We use Strings for the CellData structs (patchData).
1023 # So, we check the name here and we generate the MLIR to
1024 # extract the struct from LLVM pointer
1025 # TODO: We need to raise this to where we make the assignment
1026 match self._string:
1027 case x if '.' in x:
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]
1033 # Extract just the struct type (everything after the separator)
1034 mlir_type_struct = mlir_type.split('|||')[1] if '|||' in mlir_type else mlir_type
1035 # getelementptr returns !llvm.ptr and takes struct type as parameter
1036 mlir_str = f"llvm.getelementptr %{struct}[0, {idx}] : (!llvm.ptr) -> !llvm.ptr, {mlir_type_struct}"
1037 return mlir_str
1038 return self._string
1039
1040 def print_tree(self, indent_level=0):
1041 if self._string is None:
1042 return ' ' * indent_level * Node.spaces_per_tab + "Integer: " + str(self._value)
1043 else:
1044 return ' ' * indent_level * Node.spaces_per_tab + "Integer: " + str(self._string)
1045
1046
1048 def __init__(self, value):
1049 super().__init__()
1050 self._value = value
1051
1052 def __str__(self):
1053 return str(self._value)
1054
1055 def get_type(self):
1056 return TBoolean()
1057
1058 def print_cpp(self, indent_level=0):
1059 return ' ' * indent_level * Node.spaces_per_tab + str(self._value)
1060
1061 def print_omp(self, indent_level=0):
1062 return ' ' * indent_level * Node.spaces_per_tab + str(self._value)
1063
1064 def print_sycl(self, indent_level=0):
1065 return ' ' * indent_level * Node.spaces_per_tab + str(self._value)
1066
1067 def print_mlir(self, indent_level=0):
1068 # NOTE: Literals are returned and used in-line, no hoisting required
1069 return str(self._value)
1070
1071 def print_tree(self, indent_level=0):
1072 return ' ' * indent_level * Node.spaces_per_tab + "Boolean: " + str(self._value)
1073
1075 def __init__(self, value):
1076 super().__init__()
1077 self._value = value
1078
1079 def __str__(self):
1080 return self._value
1081
1082 def get_type(self):
1083 return TString()
1084
1085 def print_cpp(self, indent_level=0):
1086 return ' ' * indent_level * Node.spaces_per_tab + self._value
1087
1088 def print_omp(self, indent_level=0):
1089 return ' ' * indent_level * Node.spaces_per_tab + self._value
1090
1091 def print_sycl(self, indent_level=0):
1092 return ' ' * indent_level * Node.spaces_per_tab + self._value
1093
1094 def print_mlir(self, indent_level=0):
1095 indent = ' ' * indent_level * Node.spaces_per_tab
1096
1097 match self._value:
1098 case x if '.' in x:
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('|||')
1104
1105 # Generate the MLIR operations
1106 struct_val_id = f"%{Node.get_mlir_id()}"
1107
1108 # TODO: Special case: maxEigenvalue is also a function parameter, so we have to hack it here
1109 if field == "maxEigenvalue":
1110 field_id = f"%{field}s"
1111 else:
1112 field_id = f"%{Node.get_mlir_id()}"
1113
1114 # Build the basic MLIR operations
1115 mlir_str = indent + f"{struct_val_id} = llvm.load %{struct} {{name = \"{struct.strip('%')}\"}}: !llvm.ptr -> {mlir_type_struct}\n"
1116
1117 # If the extracted field is a raw pointer (not a descriptor struct),
1118 # extract it directly with a _ptr suffix to distinguish from the final memref
1119 if not mlir_type_target.startswith("!llvm.struct<"):
1120 # For raw pointers (like maxEigenvalue field [11]), use _ptr naming
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
1124 else:
1125 # For descriptor structs, extract normally (will be processed below)
1126 mlir_str += indent + f"{field_id} = llvm.extractvalue {struct_val_id}[{idx}] : {mlir_type_struct}\n"
1127
1128 final_id = field_id
1129
1130 # If the extracted field is a descriptor struct, convert it to appropriate memref
1131 if mlir_type_target.startswith("!llvm.struct<"):
1132 # Use getelementptr[0, field_idx] to properly access struct fields
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"
1137
1138 # Extract all 5 descriptor fields
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()}"
1144
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"
1150
1151 # Determine target memref type based on descriptor dimensions
1152 is_2d = "array<2 x i64>" in mlir_type_target
1153
1154 # For 2D descriptors, also extract sizes[1] and strides[1]
1155 sizes_1_id = None
1156 strides_1_id = None
1157
1158 if is_2d:
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"
1163
1164 # Use type_map to get the correct memref type for this field
1165 # For fields like cellCentre, cellSize, t, dt, we need to use their correct types from type_map
1166 if field in Node.type_map:
1167 type_entry = Node.type_map[field]
1168 # Handle both Type objects and strings
1169 if isinstance(type_entry, str):
1170 target_memref_type = type_entry
1171 else:
1172 # Type object - call print_mlir()
1173 target_memref_type = type_entry.print_mlir(0).strip()
1174 else:
1175 # Fallback to descriptor-based inference
1176 if is_2d:
1177 # 2D descriptor
1178 target_memref_type = "memref<?x?xf64>"
1179 else:
1180 # 1D descriptor (assume array<1 x i64>)
1181 target_memref_type = "memref<?xf64>"
1182
1183 # Create memref descriptor from the extracted fields using the new method
1184 conversion_mlir, memref_id = Node.create_memref_from_extracted_descriptor(
1185 alloc_ptr_id,
1186 aligned_ptr_id,
1187 offset_id,
1188 sizes_0_id,
1189 sizes_1_id,
1190 strides_0_id,
1191 strides_1_id,
1192 target_memref_type,
1193 mlir_type_target,
1194 indent_level,
1195 target_var_name=Node.target_variable_name,
1196 is_2d=is_2d)
1197 mlir_str += conversion_mlir
1198 final_id = memref_id
1199 elif 'memref' in mlir_type_target or mlir_type_target == '!llvm.ptr':
1200 # For explicit memref types or raw pointers (like maxEigenvalue),
1201 # use create_memref_from_ptr with the extracted pointer (_ptr version)
1202 clean_target_type = mlir_type_target.strip('()')
1203 # Use the _ptr field_id for raw pointer types
1204 conversion_mlir, memref_id = Node.create_memref_from_ptr(
1205 field_id, # This is now %maxEigenvalues_ptr for raw pointers
1206 clean_target_type,
1207 indent_level,
1208 target_var_name=Node.target_variable_name)
1209 mlir_str += conversion_mlir
1210 final_id = memref_id
1211
1212 # Always append to context and return variable ID
1213 Node.context.get_block().mlir_append(mlir_str)
1214 return final_id
1215 return self._value
1216
1217 def print_tree(self, indent_level=0):
1218 return ' ' * indent_level * Node.spaces_per_tab + "String: " + self._value
1219
1220
1222 def __init__(self, value, string = None, reference = False):
1223 super().__init__()
1224 self._value = value
1225 self._string = string
1226 self._reference = reference
1227
1228
1229 def __float__(self):
1230 return self._value
1231
1232 def get_type(self):
1233 return TFloat(self._reference)
1234
1235 def print_cpp(self, indent_level=0):
1236 if self._string is None:
1237 return ' ' * indent_level * Node.spaces_per_tab + str(self._value)
1238 else:
1239 return ' ' * indent_level * Node.spaces_per_tab + str(self._string)
1240
1241 def print_omp(self, indent_level=0):
1242 if self._string is None:
1243 return ' ' * indent_level * Node.spaces_per_tab + str(self._value)
1244 else:
1245 return ' ' * indent_level * Node.spaces_per_tab + str(self._string)
1246
1247 def print_sycl(self, indent_level=0):
1248 if self._string is None:
1249 return ' ' * indent_level * Node.spaces_per_tab + str(self._value)
1250 else:
1251 return ' ' * indent_level * Node.spaces_per_tab + str(self._string)
1252
1253 def print_mlir(self, indent_level=0):
1254 indent = ' ' * indent_level * Node.spaces_per_tab
1255
1256 if self._string is None:
1257 value = str(self._value)
1258 if self.hoist():
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()}"
1265 return id
1266 return value
1267 else:
1268 return str(self._string)
1269
1270 def print_tree(self, indent_level=0):
1271 return ' ' * indent_level * Node.spaces_per_tab + "Float: " + str(self._value)
1272
1273
1275 def __init__(self, iteration_range, internal, requires_memory_allocation, id = None, underlying_type = String("double")):
1276 super().__init__()
1277 self._iteration_range = iteration_range
1278 self._internal = internal
1279 self.requires_memory_allocation = requires_memory_allocation
1280 self._underlying_type = underlying_type
1281 self.id = id
1282 self._offset = []
1284 for [x, y] in self._iteration_range:
1285 self._offset.append(Integer(0))
1286 self._memory_range.append([Integer(0), y - x])
1287
1288 def offset(self, offset):
1289 output = DataBlock(copy.deepcopy(self._iteration_range), copy.deepcopy(self._internal), True)
1290 output.id = self.id
1291
1292 for i in range(len(offset)):
1293 if type(offset[i]) is list:
1294 if offset[i][0] is None:
1295 offset[i][0] = self._iteration_range[i][0]
1296 if offset[i][1] is None:
1297 offset[i][1] = self._iteration_range[i][1]
1298
1299 output._offset[i] = offset[i][0] - self._iteration_range[i][0]
1300 output._iteration_range[i][1] = offset[i][1]
1301 output._iteration_range[i][0] = offset[i][0]
1302 else:
1303 raise Exception("Invalid index")
1304
1305 for i in range(len(output._offset), len(output._iteration_range)):
1306 output._offset.append(Integer(0))
1307
1308 return output
1309
1310 def linearise_index(self, indices, dimensions, offset_start_index = 0):
1311 factors = [Integer(1)]
1312 for i in range(0, len(dimensions) - 1):
1313 factors.append(factors[-1] * (self._memory_range[i][1]))
1314 if Node.use_accelerator == False and type(self) is not FaceDataBlock:
1315 index_list = indices[0:-1]
1316 else:
1317 index_list = indices
1318
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]
1321 else:
1322 if type(self._offset[offset_start_index + 0].get_type()) is TDataBlock:
1323 index = index_list[0] + self._offset[offset_start_index + 0].multidimensional_index(index_list, offset_start_index)
1324 else:
1325 index = index_list[0] + self._offset[offset_start_index + 0]
1326
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:
1330 temp = index_list[i] + self._offset[offset_start_index + i].multidimensional_index(indices, offset_start_index)
1331 else:
1332 temp = index_list[i] + self._offset[offset_start_index + i]
1333 index = index + factors[-len(index_list) + i] * temp
1334 return index
1335
1336 def multidimensional_index(self, index_list, start_index = 0, dimensions = []):
1337 if len(self._iteration_range) == 1:
1338 return Subscript(self, index_list[-1])
1339 else:
1340 if Node.use_accelerator == False:
1341 return Subscript(Subscript(self, index_list[-1]), self.linearise_index(index_list, self._iteration_range[0:-1], start_index))
1342 else:
1343 return Subscript(self, self.linearise_index(index_list, self._iteration_range, start_index))
1344
1345 def get_type(self):
1346 db_type = TDataBlock(len(self._iteration_range), self._underlying_type)
1347 return db_type
1348
1349 def print_cpp(self, indent_level=0):
1350 if self.id is None:
1351 return ' ' * indent_level * Node.spaces_per_tab + self._internal.print_cpp()
1352 else:
1353 return ' ' * indent_level * Node.spaces_per_tab + self.id
1354
1355 def print_omp(self, indent_level=0):
1356 if self.id is None:
1357 return ' ' * indent_level * Node.spaces_per_tab + self._internal.print_omp()
1358 else:
1359 return ' ' * indent_level * Node.spaces_per_tab + self.id
1360
1361 def print_sycl(self, indent_level=0):
1362 if self.id is None:
1363 return ' ' * indent_level * Node.spaces_per_tab + self._internal.print_sycl()
1364 else:
1365 return ' ' * indent_level * Node.spaces_per_tab + self.id
1366
1367
1368 def print_mlir(self, indent_level=0):
1369 if self.id is None:
1370 return self._internal.print_mlir(indent_level)
1371 else:
1372 return "%" + self.id
1373
1374 def print_tree(self, indent_level=0):
1375 if self.id is None:
1376 return ' ' * indent_level * Node.spaces_per_tab + f"""DataBlock:
1377{self._internal.print_tree(indent_level + 1)}"""
1378 else:
1379 return ' ' * indent_level * Node.spaces_per_tab + "DataBlock: " + self.id
1380
1381"""
1382FaceDataBlock distinguishes itself from normal DataBlocks by the fact that its internal array is always 1d.
1383"""
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)
1387
1388 def offset(self, offset):
1389 output = FaceDataBlock(copy.deepcopy(self._iteration_range_iteration_range), copy.deepcopy(self._internal), True)
1390 output.id = self.id
1391
1392 for i in range(len(offset)):
1393 if type(offset[i]) is list:
1394 if offset[i][0] is None:
1395 offset[i][0] = self._iteration_range_iteration_range[i][0]
1396 if offset[i][1] is None:
1397 offset[i][1] = self._iteration_range_iteration_range[i][1]
1398
1399 output._offset[i] = offset[i][0] - self._iteration_range_iteration_range[i][0]
1400 output._iteration_range[i][1] = offset[i][1]
1401 output._iteration_range[i][0] = offset[i][0]
1402 else:
1403 raise Exception("Invalid index")
1404
1405 for i in range(len(output._offset), len(output._iteration_range)):
1406 output._offset.append(Integer(0))
1407
1408 return output
1409
1410 def multidimensional_index(self, index_list, start_index = 0, dimensions = []):
1411 return Subscript(self, self.linearise_index(index_list, self._iteration_range_iteration_range, start_index))
1412
1413
1415 def __init__(self, dataBlock, filename):
1416 self._filename = filename
1417
1418 self._dataBlock = dataBlock
1419
1420 loops = []
1421 index = []
1422
1423 loops.append(For([Integer(0), self._dataBlock._iteration_range[-1][1], - self._dataBlock._iteration_range[-1][0]]))
1424 for i in range(len(self._dataBlock._memory_range) - 2, -1, -1):
1425 loops.append(For([Integer(0), self._dataBlock._iteration_range[i][1] - self._dataBlock._iteration_range[i][0]]))
1426
1427 for i in range(len(self._dataBlock._memory_range) - 1, -1, -1):
1428 index.append(loops[i].get_iteration_variable())
1429
1430 loops[-1].add_statement(LogToFile(self._dataBlock.multidimensional_index(index), String("log")))
1431 loops[-1].add_statement(LogToFile(String("\" \""), String("log")))
1432 for i in range(0, len(loops) - 1):
1433 loops[i].add_statement(loops[i + 1])
1434
1435 loops[-2].add_statement(LogToFile(String("\", \""), String("log")))
1436 for i in range(0, len(self._dataBlock._memory_range) - 2):
1437 loops[i].add_statement(LogToFile(String("\"\\n\""), String("log")))
1438
1439 for loop in loops[::-1]:
1440 loop.close_scope()
1441 self._loop = loops[0]
1442
1443 def print_cpp(self, indent_level=0):
1444 return f"""log.open("{self._filename.print_cpp()}");
1445{self._loop.print_cpp()}
1446log.close();
1447"""
1448
1449 def print_omp(self, indent_level=0):
1450 return ""
1451
1452 def print_sycl(self, indent_level=0):
1453 return ""
1454
1455 def print_mlir(self, indent_level=0):
1456 return ""
1457
1458 def print_tree(self, indent_level=0):
1459 return ""
1460
1462 def __init__(self, value: Expression, index: Expression):
1463 super().__init__()
1464 self._value = value
1465 self._index = index
1466
1467 def __neg__(self):
1468 return UnaryOperation("-", self)
1469
1471 return isinstance(self._value, Subscript)
1472
1473 def get_type(self):
1474 return self._value.get_type()
1475
1476 def print_cpp(self, indent_level=0):
1477 return ' ' * indent_level * Node.spaces_per_tab + f"{self._value.print_cpp()}[{self._index.print_cpp()}]"
1478
1479 def print_omp(self, indent_level=0):
1480 return ' ' * indent_level * Node.spaces_per_tab + f"{self._value.print_omp()}[{self._index.print_omp()}]"
1481
1482 def print_sycl(self, indent_level=0):
1483 return ' ' * indent_level * Node.spaces_per_tab + f"{self._value.print_sycl()}[{self._index.print_sycl()}]"
1484
1485 def print_mlir(self, indent_level=0):
1486 indent = ' ' * indent_level * Node.spaces_per_tab
1487 dims = 1
1488 tmp = self._value
1489 while isinstance(tmp, Subscript):
1490 dims += 1
1491 tmp = tmp._value
1492
1493 self._dimensions = dims
1494 value_type = self._value.get_type()
1495 type_str = value_type.print_cpp().strip()
1496
1497 mlir_type = Node.type_map.get(type_str)
1498 if mlir_type is None:
1499 # Fallback: try to get the MLIR type from the type object itself, or use a default
1500 mlir_type = '' + value_type.print_mlir(0)
1501
1502 mlir_str = "".join(Node.context.get_block().pop_mlir())
1503
1504 self._index.hoist(True)
1505 index = self._index.print_mlir(indent_level)
1506 mlir_str += "".join(Node.context.get_block().pop_mlir())
1507
1508 if not getattr(self._index, '_index_type', False):
1509 new_index = f"%{Node.get_mlir_id()}"
1510 # Determine the actual type of the index expression
1511 index_expr_type = self._index.get_type()
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"
1514 self._new_index = new_index
1515 else:
1516 self._new_index = index
1517
1518 # NOTE: special handling: Strings might represent struct field access
1519 if isinstance(self._value, String):
1520 # When indexing a String field (like patchData.input[i]), we need to handle the index
1521 # Get the field type to determine if it's 2D or 1D
1522 base_mlir = self._value.print_mlir(indent_level)
1523 base_type = self._value.get_type()
1524
1525 if isinstance(base_type, TDataBlock):
1526 memref_type = base_type.print_mlir(0).strip()
1527
1528 # Check if it's a 2D flat or 1D nested memref
1529 if "x?x" in memref_type:
1530 # 2D flat memref - but we only have one index, which means we need a second one
1531 # This shouldn't happen for single-level subscripts on 2D arrays
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)
1536 return result_id
1537 else:
1538 # 1D memref - can be nested (memref<?xmemref<?xf64>>) or plain (memref<?xf64>)
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)
1543 return result_id
1544 else:
1545 # Not a memref type, just return the base
1546 return base_mlir
1547
1548 Node.context.push_block()
1549 Node.context.get_block().mlir_append(mlir_str)
1550
1551 self._value.hoist(True)
1552 base_value = self._value.print_mlir(indent_level)
1553
1554 # Handle nested subscripts (like source[i][j]) by generating proper load operations
1555 if self.is_nested_subscript():
1556 result_id = f"%{Node.get_mlir_id()}"
1557
1558 base_expr = tmp # The base expression without subscripts
1559 outer_base = base_expr.print_mlir(indent_level) # The base array
1560
1561 # Get the actual type of the base expression - check mlir_memref_types FIRST
1562 if outer_base in Node.mlir_memref_types:
1563 actual_memref_type = Node.mlir_memref_types[outer_base]
1564 else:
1565 # Fallback to computing from type
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()
1569 else:
1570 actual_memref_type = "memref<?x?xf64>" # Fallback
1571
1572 # Check if base expression has an explicit type in type_map (nested memref override)
1573 base_name = None
1574 if isinstance(base_expr, Name):
1575 base_name = base_expr.id
1576 elif hasattr(base_expr, 'id'):
1577 # DataBlock or similar object with id attribute
1578 base_name = base_expr.id
1579
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]
1582
1583 # Collect both indices from the nested structure
1584 indices = []
1585 current = self
1586 while isinstance(current, Subscript):
1587 indices.insert(0, current._new_index)
1588 current = current._value
1589
1590 # Generate load based on memref type structure
1591 indices_str = ", ".join(indices)
1592 mlir_str = "".join(Node.context.get_block().pop_mlir())
1593
1594 # Check if it's a nested memref (memref<?xmemref<...>>) vs flat multi-D (memref<?x?x...>)
1595 # Handle both Type objects and strings
1596 if isinstance(actual_memref_type, str):
1597 is_nested = actual_memref_type.count("memref<") > 1
1598 actual_memref_type_str = actual_memref_type
1599 else:
1600 # Type object (e.g., TNestedMemref)
1601 is_nested = actual_memref_type.is_nested_memref()
1602 actual_memref_type_str = actual_memref_type.print_mlir(0).strip()
1603
1604 if is_nested and len(indices) >= 2:
1605 # This is a nested memref (memref<?xmemref<?xf64>>)
1606 # Need to load in two steps: first get the inner memref, then load from it
1607 inner_memref_id = f"%{Node.get_mlir_id()}"
1608 # Extract the inner memref type from the outer type
1609 if actual_memref_type_str.startswith("memref<?x") and actual_memref_type_str.endswith(">>"):
1610 inner_type = actual_memref_type_str[9:-1] # Remove "memref<?x" and final ">"
1611 else:
1612 inner_type = "memref<?xf64>" # Fallback
1613
1614 mlir_str += indent + f"{inner_memref_id} = memref.load {outer_base}[{indices[0]}] {{name = \"{outer_base.strip('%')}\"}} : {actual_memref_type_str}\n"
1615 # Now load from the inner memref with remaining indices
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"
1618 else:
1619 # This is a flat multi-D memref (memref<?x?xf64>, memref<?x?x?xf64>, etc.)
1620 mlir_str += indent + f"{result_id} = memref.load {outer_base}[{indices_str}] {{name = \"{outer_base.strip('%')}\"}} : {actual_memref_type_str}\n"
1621
1622 Node.context.push_block()
1623 Node.context.get_block().mlir_append(mlir_str)
1624
1625 return result_id
1626
1627 # Default: just return the base value (may be a memref, pointer, or value)
1628 return base_value
1629
1630 def print_tree(self, indent_level=0):
1631 return f"""Subscript:
1632{self._value.print_tree(indent_level + 1)}
1633{self._index.print_tree(indent_level + 1)}"""
1634
1635
1637 def __init__(self, expression: Expression):
1638 self._expression = expression
1639
1640 def get_type(self):
1641 return self._expression.get_type()
1642
1643 def print_cpp(self, indent_level=0):
1644 return ' ' * indent_level * Node.spaces_per_tab + "&" + self._expression.print_cpp()
1645
1646 def print_omp(self, indent_level=0):
1647 return ' ' * indent_level * Node.spaces_per_tab + "&" + self._expression.print_omp()
1648
1649 def print_sycl(self, indent_level=0):
1650 return ' ' * indent_level * Node.spaces_per_tab + "&" + self._expression.print_sycl()
1651
1652 def print_mlir(self, indent_level=0):
1653 result_id = None
1654 indent = ' ' * indent_level * Node.spaces_per_tab
1655 mlir_str = "".join(Node.context.get_block().pop_mlir())
1656
1657 self._expression.hoist(True)
1658
1659 # Handle references to subscripted expressions (array elements)
1660 if isinstance(self._expression, Subscript):
1661 base_expr = self._expression._value
1662 index_expr = self._expression._index
1663
1664 base_expr.hoist(True)
1665 base_mlir = base_expr.print_mlir(indent_level)
1666
1667 mlir_str += "".join(Node.context.get_block().pop_mlir())
1668
1669 if not hasattr(self._expression, '_new_index'):
1670 # If _new_index doesn't exist, process index and generate it now
1671 index_expr.hoist(True)
1672 index_mlir = index_expr.print_mlir(indent_level)
1673 mlir_str += "".join(Node.context.get_block().pop_mlir())
1674
1675 # Create new index with proper type cast
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"
1679 self._expression._new_index = new_index
1680 else:
1681 new_index = self._expression._new_index
1682
1683 # Determine dimensions and type of the memref
1684 base_type = base_expr.get_type()
1685 if isinstance(base_type, TDataBlock):
1686 result_id = f"%{Node.get_mlir_id()}"
1687
1688 # For nested memrefs (like double**), handle the inner memref properly
1689 if isinstance(base_expr, Subscript):
1690 # base_expr is Subscript[DataBlock, i]
1691 # We need to check if the base_expr itself is a nested subscript or just single-level
1692 if isinstance(base_expr._value, Subscript):
1693 result_id = base_mlir
1694 else:
1695 # Check if this variable has an explicit nested type in `type_map`
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
1700
1701 outer_index = base_expr._new_index
1702
1703 # Use unified method from Node to generate pointer arithmetic
1704 ptr_code, result_id = Node.generate_pointer_to_2d_element(
1705 base_mlir, outer_index, new_index, is_nested_memref, indent_level
1706 )
1707
1708 mlir_str += ptr_code
1709 else:
1710 element_type = "f64" # Default to f64 - modify if needed for other types
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"
1712
1713 # Default behavior for non-subscript references
1714 Node.context.push_block()
1715 Node.context.get_block().mlir_append(mlir_str)
1716
1717 # Process the expression and return its ID
1718 return result_id if result_id is not None else self._expression.print_mlir(indent_level)
1719
1720
1721 def print_tree(self, indent_level=0):
1722 return ' ' * indent_level * Node.spaces_per_tab + f"""Reference:
1723{self._expression.print_tree(indent_level + 1)}"""
1724
1725
1727 def __init__(self, elements, element_type, string = None):
1728 self._elements = elements
1729 self._type = element_type
1730 self._string = string
1731
1732 def print_cpp(self, indent_level=0):
1733 if self._string is 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)}}}"""
1736 return self._string
1737
1738 def print_omp(self, indent_level=0):
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)}}}"""
1741
1742 def print_sycl(self, indent_level=0):
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)}}}"""
1745
1746 def print_mlir(self, indent_level=0):
1747 pass
1748
1749 def print_tree(self, indent_level=0):
1750 return "Vector"
1751
1752 def get_type(self):
1753 if self._string is not None:
1754 return TCustom("auto")
1755 return TCustom(f"tarch::la::Vector<{len(self._elements)}, {self._type.print_cpp()}>")
1756
1757
1759 def __init__(self, width, internal, index_dimensions):
1760 self._internal = internal
1761 self._width = width
1762 self._index_dimensions = index_dimensions
1763 self._id = None
1764
1765 def multidimensional_index(self, index, start_index = 0, dimensions = []):
1766 diagonal_index = index[1]
1767 factor = Integer(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]
1771 return Subscript(self, diagonal_index)
1772
1773 def print_cpp(self, indent_level=0):
1774 if self._id is not None:
1775 return self._id.print_cpp(indent_level)
1776 else:
1777 return self._internal.print_cpp(indent_level)
1778
1779 def print_omp(self, indent_level=0):
1780 if self._id is not None:
1781 return self._id.print_cpp(indent_level)
1782 else:
1783 return self._internal.print_cpp(indent_level)
1784
1785 def print_sycl(self, indent_level=0):
1786 pass
1787
1788 def print_mlir(self, indent_level=0):
1789 pass
1790
1791 def print_tree(self, indent_level=0):
1792 return ' ' * indent_level * Node.spaces_per_tab + "DiagonalMatrix"
1793
1794 def get_type(self):
1795 return TDiagonalMatrix()
1796
1797
1799 def __init__(self, dimensions, internal, index_dimensions):
1800 self._internal = internal
1801 self._dimensions = dimensions
1802 self._index_dimensions = index_dimensions
1803 self._id = None
1804
1805 def index(self, row_index, column_index):
1806 return Subscript(self, row_index * self._dimensions[1] + column_index)
1807
1808 def print_cpp(self, indent_level=0):
1809 if self._id is not None:
1810 return self._id.print_cpp(indent_level)
1811 else:
1812 return self._internal.print_cpp(indent_level)
1813
1814 def print_omp(self, indent_level=0):
1815 if self._id is not None:
1816 return self._id.print_cpp(indent_level)
1817 else:
1818 return self._internal.print_cpp(indent_level)
1819
1820 def print_sycl(self, indent_level=0):
1821 pass
1822
1823 def print_mlir(self, indent_level=0):
1824 pass
1825
1826 def print_tree(self, indent_level=0):
1827 return ' ' * indent_level * Node.spaces_per_tab + "Matrix"
1828
1829 def get_type(self):
1830 return TMatrix()
1831
1832
1834 def __init__(self, id, matrix):
1835 self._id = String(id)
1836 if type(matrix) is Name:
1837 self._matrix = Name.variables[matrix.id]
1838 else:
1839 self._matrix = matrix
1840
1842
1843 lhs_matrix = self._matrix._internal._arguments[0]
1844 if type(lhs_matrix) is Name:
1845 lhs_matrix = Name.variables[lhs_matrix.id]
1846
1847 rhs_matrix = self._matrix._internal._arguments[1]
1848 if type(rhs_matrix) is Name:
1849 rhs_matrix = Name.variables[rhs_matrix.id]
1850
1851 self._assignment = For([Integer(0), lhs_matrix._width])
1852 internal_loop = For([Integer(0), rhs_matrix._width])
1853 index = self._assignment.get_iteration_variable() * rhs_matrix._width + internal_loop.get_iteration_variable()
1854 lhs_index = Subscript(lhs_matrix, self._assignment.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
1860 internal_loop.add_statement(Assignment(Subscript(self._id, index), BinaryOperation("*", lhs_index, rhs_index)))
1861 self._assignment.add_statement(internal_loop)
1862 internal_loop.close_scope()
1863 self._assignment.close_scope()
1864
1865 def print_cpp(self, indent_level=0):
1866 return f"""{self._memory_allocation.print_cpp(indent_level)}
1867{self._assignment.print_cpp(indent_level)}"""
1868
1869 def print_omp(self, indent_level=0):
1870 return f"""{self._memory_allocation.print_omp(indent_level)}
1871{self._assignment.print_omp(indent_level)}"""
1872
1873 def print_sycl(self, indent_level=0):
1874 pass
1875
1876 def print_mlir(self, indent_level=0):
1877 pass
1878
1879 def print_tree(self, indent_level=0):
1880 return ' ' * indent_level * Node.spaces_per_tab + "DiagonalKroneckerProduct"
1881
1882 def get_type(self):
1883 return TDiagonalMatrix()
1884
1885
1887 def __init__(self, id, matrix):
1888 self._id = String(id)
1889 if type(matrix) is Name:
1890 self._matrix = Name.variables[matrix.id]
1891 else:
1892 self._matrix = matrix
1893
1895
1896 self._assignment = For([Integer(0), self._matrix._width])
1897 self._assignment.add_statement(Assignment(Subscript(self._id, self._assignment.get_iteration_variable()), self._matrix._internal))
1898 self._assignment.close_scope()
1899
1900 def print_cpp(self, indent_level=0):
1901 return f"""{self._memory_allocation.print_cpp(indent_level)}
1902{self._assignment.print_cpp(indent_level)}"""
1903
1904 def print_omp(self, indent_level=0):
1905 return f"""{self._memory_allocation.print_omp(indent_level)}
1906{self._assignment.print_omp(indent_level)}"""
1907
1908 def print_sycl(self, indent_level=0):
1909 pass
1910
1911 def print_mlir(self, indent_level=0):
1912 pass
1913
1914 def print_tree(self, indent_level=0):
1915 return ' ' * indent_level * Node.spaces_per_tab + "DiagonalKroneckerProduct"
1916
1917 def get_type(self):
1918 return TDiagonalMatrix()
1919
1920
1922 def __init__(self, id, matrix):
1923 self._id = String(id)
1924 if type(matrix) is Name:
1925 self._matrix = Name.variables[matrix.id]
1926 else:
1927 self._matrix = matrix
1928
1929 self._memory_allocation = MemoryAllocation(self._id, TFloat(), [self._matrix._dimensions[0] * self._matrix._dimensions[1]])
1930
1931 self._initialisation = For([Integer(0), self._matrix._dimensions[0] * self._matrix._dimensions[1]])
1932 self._initialisation.add_statement(Assignment(Subscript(self._id, self._initialisation.get_iteration_variable()), Float(0.0)))
1933 self._initialisation.close_scope()
1934
1935 lhs_matrix = self._matrix._internal._arguments[0]
1936 if type(lhs_matrix) is Name:
1937 lhs_matrix = Name.variables[lhs_matrix.id]
1938
1939 rhs_matrix = self._matrix._internal._arguments[1]
1940 if type(rhs_matrix) is Name:
1941 rhs_matrix = Name.variables[rhs_matrix.id]
1942
1943 if type(lhs_matrix.get_type()) is TMatrix:
1944 lhs_dimensions = lhs_matrix._dimensions
1945 else:
1946 lhs_dimensions = [lhs_matrix._width, lhs_matrix._width]
1947
1948 if type(rhs_matrix.get_type()) is TMatrix:
1949 rhs_dimensions = rhs_matrix._dimensions
1950 else:
1951 rhs_dimensions = [rhs_matrix._width, rhs_matrix._width]
1952
1953 loops = []
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]
1959 else:
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]
1963
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]
1969 else:
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]
1973
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
1980
1981 loops[-1].add_statement(Assignment(Subscript(self._id, index), BinaryOperation("*", lhs_subscript, rhs_subscript)))
1982
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()
1987 self._assignment = loops[0]
1988
1989 def print_cpp(self, indent_level=0):
1990 return f"""{self._memory_allocation.print_cpp(indent_level)}
1991{self._initialisation.print_cpp(indent_level)}
1992{self._assignment.print_cpp(indent_level)}"""
1993
1994 def print_omp(self, indent_level=0):
1995 return f"""{self._memory_allocation.print_omp(indent_level)}
1996{self._initialisation.print_omp(indent_level)}
1997{self._assignment.print_omp(indent_level)}"""
1998
1999 def print_sycl(self, indent_level=0):
2000 pass
2001
2002 def print_mlir(self, indent_level=0):
2003 pass
2004
2005 def print_tree(self, indent_level=0):
2006 return ' ' * indent_level * Node.spaces_per_tab + "KroneckerProduct"
2007
2008 def get_type(self):
2009 return TMatrix()
2010
2011
2013 def __init__(self, value: Expression, index: Expression):
2014 super().__init__()
2015 self._value = value
2016 self._index = index
2017
2018 def get_type(self):
2019 return self._value.get_type()
2020
2021 def print_cpp(self, indent_level=0):
2022 return ' ' * indent_level * Node.spaces_per_tab + f"{self._value.print_cpp()}({self._index.print_cpp()})"
2023
2024 def print_omp(self, indent_level=0):
2025 return ' ' * indent_level * Node.spaces_per_tab + f"{self._value.print_omp()}({self._index.print_omp()})"
2026
2027 def print_sycl(self, indent_level=0):
2028 return ' ' * indent_level * Node.spaces_per_tab + f"{self._value.print_sycl()}({self._index.print_sycl()})"
2029
2030 def print_mlir(self, indent_level=0):
2031 indent = ' ' * indent_level * Node.spaces_per_tab
2032 mlir_str = ""
2033
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)
2039 # Ensure array_index is of type 'index' with correct source type
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"
2045 else:
2046 array_index_cast = array_index_mlir
2047
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"
2050
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"
2053
2054 func_param_mlir = self._index.print_mlir(indent_level)
2055 # Ensure func_param is i32 (emit arith.constant for Integer, cast otherwise)
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"
2062 else:
2063 func_param_i32 = func_param_mlir
2064
2065 result = f"%{Node.get_mlir_id()}"
2066 # Convert i32 index to i64 for pointer arithmetic
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"
2069 # Compute pointer to array element
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"
2072 # Load the value from the computed pointer
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)
2075
2076 return result
2077
2078 def print_tree(self, indent_level=0):
2079 return f"""
2080VectorIndex:
2081{self.id}
2082{self._value.print_tree(indent_level + 1)}
2083{self._index.print_tree(indent_level + 1)}"""
2084
2085
2087 def __init__(self, operation, lhs, rhs):
2088 self._operation = operation
2089 self._lhs = lhs
2090 self._rhs = rhs
2091
2092 def get_type(self):
2093 return self._lhs.get_type()
2094
2095 def print_cpp(self, indent_level = 0):
2096 return ' ' * indent_level * Node.spaces_per_tab + f"({self._lhs.print_cpp()} {self._operation} {self._rhs.print_cpp()})"
2097
2098 def print_omp(self, indent_level = 0):
2099 return ' ' * indent_level * Node.spaces_per_tab + f"({self._lhs.print_omp()} {self._operation} {self._rhs.print_omp()})"
2100
2101 def print_sycl(self, indent_level = 0):
2102 return ' ' * indent_level * Node.spaces_per_tab + f"({self._lhs.print_sycl()} {self._operation} {self._rhs.print_sycl()})"
2103
2104 @staticmethod
2105 def get_mlir_comparison_op(operation, op_type_class):
2106 if op_type_class is TFloat:
2107 op_map = {
2108 "<": "arith.cmpf olt",
2109 ">": "arith.cmpf ogt",
2110 "<=": "arith.cmpf ole",
2111 ">=": "arith.cmpf oge",
2112 "==": "arith.cmpf oeq",
2113 "!=": "arith.cmpf one"
2114 }
2115 else: # TInteger or default
2116 op_map = {
2117 "<": "arith.cmpi slt",
2118 ">": "arith.cmpi sgt",
2119 "<=": "arith.cmpi sle",
2120 ">=": "arith.cmpi sge",
2121 "==": "arith.cmpi eq",
2122 "!=": "arith.cmpi ne"
2123 }
2124 return op_map.get(operation, None)
2125
2126 def print_mlir(self, indent_level = 0):
2127 op_type_class = TDataBlock.get_mlir_type_for_operation(self._lhs, self._rhs)
2128
2129 mlir_operation = Comparison.get_mlir_comparison_op(self._operation, op_type_class)
2130 result_type = "i1" # Comparisons always return boolean
2131
2132 if mlir_operation is None:
2133 raise ValueError(f"Unsupported comparison operation: {self._operation}")
2134
2135 const_zero_id = None
2136
2137 lhs_mlir = self._lhs.print_mlir(indent_level)
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)
2141
2142 rhs_mlir = self._rhs.print_mlir(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)
2146
2147 return ' ' * indent_level * Node.spaces_per_tab + f"""{mlir_operation} {lhs_mlir}, {rhs_mlir} : {result_type}"""
2148
2149 def print_tree(self, indent_level=0):
2150 return f"""
2151Comparison {self._operation}:
2152{self._lhs.print_tree(indent_level + 1)}
2153{self._rhs.print_tree(indent_level + 1)}"""
2154
2156 def __init__(self, operation, lhs, rhs):
2157 super().__init__()
2158 self._operation = operation
2159 self._lhs = lhs
2160 self._rhs = rhs
2161
2162 if type(self._lhs) is BinaryOperation:
2163 if self._lhs._operation == "+":
2164 if self._lhs._rhs == Integer(0):
2165 self._lhs = self._lhs._lhs
2166 elif self._lhs._lhs == Integer(0):
2167 self._lhs = self._lhs._rhs
2168 elif self._lhs._operation == "-":
2169 if self._lhs._rhs == Integer(0):
2170 self._lhs = self._lhs._lhs
2171 elif self._lhs._operation == "*":
2172 if self._lhs._rhs == Integer(1):
2173 self._lhs = self._lhs._lhs
2174 elif self._lhs._lhs == Integer(1):
2175 self._lhs = self._lhs._rhs
2176
2177 if type(self._rhs) is BinaryOperation:
2178 if self._rhs._operation == "+":
2179 if self._rhs._rhs == Integer(0):
2180 self._rhs = self._rhs._lhs
2181 elif self._rhs._lhs == Integer(0):
2182 self._rhs = self._rhs._rhs
2183 elif self._rhs._operation == "-":
2184 if self._rhs._rhs == Integer(0):
2185 self._rhs = self._rhs._lhs
2186 elif self._rhs._operation == "*":
2187 if self._rhs._rhs == Integer(1):
2188 self._rhs = self._rhs._lhs
2189 elif self._rhs._lhs == Integer(1):
2190 self._rhs = self._rhs._rhs
2191
2192 if self._operation == "-" and type(self._rhs) is UnaryOperation and self._rhs._operation == "-":
2193 self._operation = "+"
2194 self._rhs = self._rhs._operand
2195 if self._operation == "+" and type(self._rhs) is UnaryOperation and self._rhs._operation == "-":
2196 self._operation = "-"
2197 self._rhs = self._rhs._operand
2198
2199 def index(self, index):
2200 return BinaryOperation(self._operation, self._lhs.index(index), self._rhs.index(index))
2201
2202 def multidimensional_index(self, index, start_index = 0, dimensions = []):
2203 return BinaryOperation(self._operation, self._lhs.multidimensional_index(index, start_index, dimensions), self._rhs.multidimensional_index(index, start_index, dimensions))
2204
2205
2206 def __add__(self, rhs):
2207 return BinaryOperation("+", self, rhs)
2208
2209 def __sub__(self, rhs):
2210 return BinaryOperation("-", self, rhs)
2211
2212 def __mul__(self, rhs):
2213 return BinaryOperation("*", self, rhs)
2214
2215 def __neg__(self):
2216 return UnaryOperation("-", self)
2217
2218 def get_type(self):
2219 # For arithmetic operations, return the element type, not the pointer type
2220 op_type_class = TDataBlock.get_mlir_type_for_operation(self._lhs, self._rhs)
2221 return op_type_class()
2222
2223 def print_cpp(self, indent_level = 0):
2224 return ' ' * indent_level * Node.spaces_per_tab + f"({self._lhs.print_cpp()} {self._operation} {self._rhs.print_cpp()})"
2225
2226 def print_omp(self, indent_level = 0):
2227 return ' ' * indent_level * Node.spaces_per_tab + f"({self._lhs.print_omp()} {self._operation} {self._rhs.print_omp()})"
2228
2229 def print_sycl(self, indent_level = 0):
2230 return ' ' * indent_level * Node.spaces_per_tab + f"({self._lhs.print_sycl()} {self._operation} {self._rhs.print_sycl()})"
2231
2232 @staticmethod
2233 def get_mlir_arithmetic_op(operation, op_type_class):
2234 if op_type_class is TFloat:
2235 op_map = {
2236 "*": "arith.mulf",
2237 "/": "arith.divf",
2238 "+": "arith.addf",
2239 "-": "arith.subf"
2240 }
2241 else: # TInteger, TLong, or default
2242 op_map = {
2243 "*": "arith.muli",
2244 "/": "arith.divi",
2245 "+": "arith.addi",
2246 "-": "arith.subi"
2247 }
2248 return op_map.get(operation, None)
2249
2250 @staticmethod
2251 def cast_operand(operand_mlir, operand, result_type, indent_level):
2252 indent = ' ' * indent_level * Node.spaces_per_tab
2253 operand_type = operand.get_type()
2254
2255 if type(operand_type) is TDataBlock:
2256 operand_type = operand_type.get_element_type()
2257
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")
2261 return new_op
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")
2265 return new_op
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")
2269 return new_op
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")
2273 return new_op
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")
2277 return new_op
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")
2281 return new_op
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")
2285 return new_op
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")
2289 return new_op
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")
2293 return new_op
2294 return operand_mlir # No casting needed, return original operand
2295
2296
2297 def get_mlir_op(self):
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)
2300
2301 def print_mlir(self, indent_level = 0):
2302 indent = ' ' * indent_level * Node.spaces_per_tab
2303 self._lhs.hoist(True)
2304 self._rhs.hoist(True)
2305
2306 op_type_class = TDataBlock.get_mlir_type_for_operation(self._lhs, self._rhs)
2307 result_type = op_type_class.print_mlir(0)
2308
2309 op_type_class = TDataBlock.get_mlir_type_for_operation(self._lhs, self._rhs)
2310 result_type = op_type_class().print_mlir(0)
2311
2312 const_zero_id = None
2313
2314 self._lhs.hoist(True)
2315 lhs = self._lhs.print_mlir(indent_level).strip(' ')
2316
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)
2320
2321 lhs = BinaryOperation.cast_operand(lhs, self._lhs, op_type_class, indent_level)
2322
2323 self._rhs.hoist(True)
2324 rhs = self._rhs.print_mlir(indent_level).strip(' ')
2325
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)
2329
2330 rhs = BinaryOperation.cast_operand(rhs, self._rhs, op_type_class, indent_level)
2331
2332 mlir_str = ''
2333 for mlir in Node.context.pop_block().get_mlir():
2334 mlir_str += indent + mlir.strip(' ')
2335
2336 self.mlir_id(f"%{Node.get_mlir_id()}")
2337
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)
2341
2342 return self.mlir_id()
2343
2344 def print_tree(self, indent_level=0):
2345 return f"""
2346BinaryOperation {self._operation}:
2347{self._lhs.print_tree(indent_level + 1)}
2348{self._rhs.print_tree(indent_level + 1)}"""
2349
2350
2352 def __init__(self, operation, operand):
2353 super().__init__()
2354 self._operation = operation
2355 self._operand = operand
2356
2357 def __add__(self, rhs):
2358 return BinaryOperation("+", self, rhs)
2359
2360 def __sub__(self, rhs):
2361 return BinaryOperation("-", self, rhs)
2362
2363 def __mul__(self, rhs):
2364 return BinaryOperation("*", self, rhs)
2365
2366 def __neg__(self):
2367 if self._operation == "-":
2368 return self._operand
2369 else:
2370 return UnaryOperation("-", self._operand)
2371
2372 def get_type(self):
2373 return self._operand.get_type()
2374
2375 def print_cpp(self, indent_level = 0):
2376 return f"({self._operation} {self._operand.print_cpp()})"
2377
2378 def print_omp(self, indent_level = 0):
2379 return f"({self._operation} {self._operand.print_omp()})"
2380
2381 def print_sycl(self, indent_level = 0):
2382 return f"({self._operation} {self._operand.print_sycl()})"
2383
2384 def print_mlir(self, indent_level=0):
2385 indent = ' ' * indent_level * Node.spaces_per_tab
2386 self._operand.hoist(True)
2387
2388 operand_type = self._operand.get_type()
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
2392
2393 if isinstance(operand_type, TFloat):
2394 if self._operation == "+":
2395 Node.context.get_block().mlir_append(self._operand.print_mlir(indent_level))
2396 return self.mlir_id()
2397 elif self._operation == "-":
2398 op_mlir = self._operand.print_mlir(indent_level).strip(' ')
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"
2402
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)
2407 return self.mlir_id()
2408 else:
2409 if self._operation == "+":
2410 Node.context.get_block().mlir_append(self._operand.print_mlir(indent_level))
2411 return self.mlir_id()
2412 elif self._operation == "-":
2413 op_mlir = self._operand.print_mlir(indent_level).strip(' ')
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"
2417
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)
2422 return self.mlir_id()
2423
2424 def print_tree(self, indent_level=0):
2425 return f"""
2426UnaryOperation {self._operation}:
2427{self._operand.print_tree(indent_level + 1)}"""
2428
2430 def __init__(self, dataBlock, index):
2431 super().__init__()
2432 if type(dataBlock) is Name:
2433 self._dataBlock = Name.variables[dataBlock.id]
2434 else:
2435 self._dataBlock = dataBlock
2436
2437 self._index = index
2438 self._iteration_range = self._dataBlock._iteration_range
2439 self._memory_range = self._dataBlock._memory_range
2440
2441 def index(self, index):
2442 return VectorIndex(self._dataBlock.index(index), self._index)
2443
2444 def multidimensional_index(self, index, start_index = 0, dimensions = []):
2445 return VectorIndex(self._dataBlock.multidimensional_index(index, start_index, dimensions), self._index)
2446
2447 def get_type(self):
2448 return TDataBlock(self._iteration_range, String("double"))
2449
2450 def print_cpp(self, indent_level = 0):
2451 pass
2452
2453 def print_omp(self, indent_level = 0):
2454 pass
2455
2456 def print_sycl(self, indent_level = 0):
2457 pass
2458
2459 def print_mlir(self, indent_level = 0):
2460 pass
2461
2462 def print_tree(self, indent_level=0):
2463 pass
2464
2465
2466class DataBlockComparison(Expression):
2467 def __init__(self, operation, lhs, rhs):
2468 if type(lhs) is Name and type(lhs.get_type()) is TDataBlock:
2469 self._lhs = Name.variables[lhs.id]
2470 else:
2471 self._lhs = lhs
2472
2473 if type(rhs) is Name and type(rhs.get_type()) is TDataBlock:
2474 self._rhs = Name.variables[rhs.id]
2475 else:
2476 self._rhs = rhs
2477
2478 self._operation = operation
2480 self._offset = None
2481
2482 if type(self._lhs.get_type()) is TDataBlock and type(self._rhs.get_type()) is TDataBlock:
2483 self._iteration_range = copy.deepcopy(self._lhs._iteration_range)
2484 self._memory_range = copy.deepcopy(self._lhs._memory_range)
2485 if 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]:
2486 self._iteration_range[0][1] = self._rhs._iteration_range[0][1]
2487 elif type(self._iteration_range[0][1]) is not BinaryOperation and self._iteration_range[0][1] == Integer(1):
2488 self._iteration_range[0][1] = self._rhs._iteration_range[0][1]
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]:
2490 self._memory_range[0] = self._rhs._memory_range[0]
2491 elif type(self._memory_range[0]) is not BinaryOperation and self._memory_range[0] == Integer(1):
2492 self._memory_range[0] = self._rhs._memory_range[0]
2493
2494 elif type(self._lhs.get_type()) is TDataBlock:
2495 self._iteration_range = copy.deepcopy(self._lhs._iteration_range)
2496 self._memory_range = copy.deepcopy(self._lhs._memory_range)
2497 else:
2498 self._iteration_range = copy.deepcopy(self._rhs._iteration_range)
2499 self._memory_range = copy.deepcopy(self._rhs._memory_range)
2500
2501 def index(self, index):
2502 return Comparison(self._operation, self._lhs.index(index), self._rhs.index(index))
2503
2504 def multidimensional_index(self, index, start_index = 0, dimensions = []):
2505 return Comparison(self._operation, self._lhs.multidimensional_index(index, start_index, dimensions), self._rhs.multidimensional_index(index, start_index, dimensions))
2506
2507 def get_type(self):
2508 return TDataBlock(self._iteration_range, String("int"))
2509
2510 def print_cpp(self, indent_level = 0):
2511 pass
2512
2513 def print_omp(self, indent_level = 0):
2514 pass
2515
2516 def print_sycl(self, indent_level = 0):
2517 pass
2518
2519 def print_mlir(self, indent_level = 0):
2520 pass
2521
2522 def print_tree(self, indent_level=0):
2523 pass
2524
2526 def __init__(self, operation, lhs, rhs, useFunctionSyntax=False):
2527 super().__init__()
2528 if type(lhs) is Name and type(lhs.get_type()) is TDataBlock:
2529 self._lhs = Name.variables[lhs.id]
2530 else:
2531 self._lhs = lhs
2532
2533 if type(rhs) is Name and type(rhs.get_type()) is TDataBlock:
2534 self._rhs = Name.variables[rhs.id]
2535 else:
2536 self._rhs = rhs
2537
2538 self._operation = operation
2540 self._useFunctionSyntax = useFunctionSyntax
2541
2542 if type(self._lhs.get_type()) is TDataBlock and type(self._rhs.get_type()) is TDataBlock:
2543 self._iteration_range = copy.deepcopy(self._lhs._iteration_range)
2544 self._memory_range = copy.deepcopy(self._lhs._memory_range)
2545 if len(self._iteration_range) == 1:
2546 self._iteration_range = copy.deepcopy(self._rhs._iteration_range)
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]:
2548 self._iteration_range[0][1] = self._rhs._iteration_range[0][1]
2549 elif type(self._iteration_range[0][1]) is not BinaryOperation and self._iteration_range[0][1] == Integer(1):
2550 self._iteration_range[0][1] = self._rhs._iteration_range[0][1]
2551
2552 self._memory_range[0] = [Integer(0), self._iteration_range[0][1] - self._iteration_range[0][0]]
2553
2554 elif type(self._lhs.get_type()) is TDataBlock:
2555 self._iteration_range = copy.deepcopy(self._lhs._iteration_range)
2556 self._memory_range = copy.deepcopy(self._lhs._memory_range)
2557 else:
2558 self._iteration_range = copy.deepcopy(self._rhs._iteration_range)
2559 self._memory_range = copy.deepcopy(self._rhs._memory_range)
2560
2561 if (type(self._iteration_range[0][0]) is Integer and type(self._iteration_range[0][1]) is Integer and (self._iteration_range[0][1] - self._iteration_range[0][0])._value == 1):
2562 self._iteration_range[0][0] = Integer(0)
2563 self._iteration_range[0][1] = Integer(1)
2564
2565 self._offset = []
2566 for [x, y] in self._iteration_range:
2567 self._offset.append(x)
2568
2569
2570 def index(self, index):
2571 if self._useFunctionSyntax:
2572 return FunctionCall(self._operation, [self._lhs.index(index), self._rhs.index(index)])
2573 else:
2574 return BinaryOperation(self._operation, self._lhs.index(index), self._rhs.index(index))
2575
2576 def multidimensional_index(self, index, start_index = 0, dimensions = []):
2577 if self._useFunctionSyntax:
2578 return FunctionCall(self._operation, [self._lhs.multidimensional_index(index, start_index, dimensions), self._rhs.multidimensional_index(index, start_index, dimensions)])
2579 else:
2580 return BinaryOperation(self._operation, self._lhs.multidimensional_index(index, start_index, dimensions), self._rhs.multidimensional_index(index, start_index, dimensions))
2581
2582 def get_type(self):
2583 lhs_is_double = False
2584 rhs_is_double = False
2585 if type(self._lhs.get_type()) is TDataBlock:
2586 lhs_is_double = (str(self._lhs.get_type().get_single()) == "double")
2587 else:
2588 lhs_is_double = (str(self._lhs.get_type()) == "double")
2589
2590 if type(self._rhs.get_type()) is TDataBlock:
2591 rhs_is_double = (str(self._rhs.get_type().get_single()) == "double")
2592 else:
2593 rhs_is_double = (str(self._rhs.get_type()) == "double")
2594
2595 if (lhs_is_double or rhs_is_double):
2596 return TDataBlock(self._iteration_range, String("double"))
2597 else:
2598 return TDataBlock(self._iteration_range, String("int"))
2599
2600 def __sub__(self, rhs):
2601 return DataBlockBinaryOperation("-", self, rhs)
2602
2603 def __add__(self, rhs):
2604 return DataBlockBinaryOperation("+", self, rhs)
2605
2606 def print_cpp(self, indent_level = 0):
2607 pass
2608
2609 def print_omp(self, indent_level = 0):
2610 pass
2611
2612 def print_sycl(self, indent_level = 0):
2613 pass
2614
2615 def print_mlir(self, indent_level = 0):
2616 pass
2617
2618 def print_tree(self, indent_level=0):
2619 pass
2620
2621
2623 def __init__(self, operation, dataBlock):
2624 super().__init__()
2625 if type(dataBlock) is Name and type(dataBlock.get_type()) is TDataBlock:
2626 self._dataBlock = Name.variables[dataBlock.id]
2627 else:
2628 self._dataBlock = dataBlock
2629
2630 self._operation = operation
2632 self._offset = None
2633
2634 self._iteration_range = copy.deepcopy(self._dataBlock._iteration_range)
2635 self._memory_range = copy.deepcopy(self._dataBlock._memory_range)
2636
2637 def index(self, index):
2638 return UnaryOperation(self._operation, self._dataBlock.index(index))
2639
2640 def multidimensional_index(self, index, start_index = 0, dimensions = []):
2641 return UnaryOperation(self._operation, self._dataBlock.multidimensional_index(index, start_index, dimensions))
2642
2643 def get_type(self):
2644 return TDataBlock(self._iteration_range, String("double"))
2645
2646 def print_cpp(self, indent_level = 0):
2647 pass
2648
2649 def print_omp(self, indent_level = 0):
2650 pass
2651
2652 def print_sycl(self, indent_level = 0):
2653 pass
2654
2655 def print_mlir(self, indent_level = 0):
2656 pass
2657
2658 def print_tree(self, indent_level=0):
2659 pass
2660
2661
2663 def __init__(self, dataBlock):
2664 if type(dataBlock) is Name:
2665 self._dataBlock = Name.variables[dataBlock.id]
2666 else:
2667 self._dataBlock = dataBlock
2668
2669 def set_output_variable(self, outputVariable):
2670 self._outputVariable = Name.variables[outputVariable.id]
2671 self._loop = For(self._dataBlock._iteration_range[-1])
2672 self._loops = [self._loop]
2673 self._loop.add_statement(Assignment(self._outputVariable.multidimensional_index([self._loop.get_iteration_variable()]), Float(0.0)))
2674 for i in range(len(self._dataBlock._iteration_range) - 2, -1, -1):
2675 self._loops.append(For(self._dataBlock._iteration_range[i]))
2676 self._loops[-2].add_statement(self._loops[-1])
2677 index = []
2678 for loop in self._loops[::-1]:
2679 index.append(loop.get_iteration_variable())
2680 self._loops[-1].add_statement(Assignment(self._outputVariable.multidimensional_index(index), Max(self._outputVariable.multidimensional_index(index), self._dataBlock.multidimensional_index(index))))
2681 for loop in self._loops:
2682 loop.close_scope()
2683
2684 def get_type(self):
2685 return TDataBlock(self._dataBlock._iteration_range, TFloat())
2686
2687 def print_cpp(self, indent_level = 0):
2688 return self._loop.print_cpp(indent_level)
2689
2690 def print_omp(self, indent_level = 0):
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}}}
2697"""
2698 return self._loop.print_omp(indent_level)
2699
2700
2701 def print_sycl(self, indent_level = 0):
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();"""
2712
2713 def print_mlir(self, indent_level = 0):
2714 return self._loop.print_mlir(indent_level)
2715
2716 def print_tree(self, indent_level=0):
2717 return f"""DataBlockUnaryMax:
2718{self._dataBlock.print_tree(indent_level + 1)}"""
2719
2720
2722 def __init__(self, lhs, rhs):
2723 super().__init__("max", lhs, rhs)
2724
2725 def index(self, index):
2726 return Max(self._lhs.index(index), self._rhs.index(index))
2727
2728 def multidimensional_index(self, index, start_index = 0, dimensions = []):
2729 return Max(self._lhs.multidimensional_index(index, start_index, dimensions), self._rhs.multidimensional_index(index, start_index, dimensions))
2730
2731 def print_cpp(self, indent_level = 0):
2732 pass
2733
2734 def print_omp(self, indent_level = 0):
2735 pass
2736
2737 def print_sycl(self, indent_level = 0):
2738 pass
2739
2740 def print_mlir(self, indent_level = 0):
2741 pass
2742
2743 def print_tree(self, indent_level=0):
2744 pass
2745
2746# Types
2747class TCustom(Type):
2748 def __init__(self, type_name):
2749 self._type = type_name
2750
2751 def print_cpp(self, indent_level=0):
2752 return ' ' * indent_level * Node.spaces_per_tab + self._type
2753
2754 def print_omp(self, indent_level=0):
2755 return ' ' * indent_level * Node.spaces_per_tab + self._type
2756
2757 def print_sycl(self, indent_level=0):
2758 return ' ' * indent_level * Node.spaces_per_tab + self._type
2759
2760 def print_mlir(self, indent_level=0):
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
2764 else:
2765 # Fallback to type_map for complex types
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
2769 else:
2770 # Type object (e.g., TVoid, TFloat, etc.)
2771 return ' ' * indent_level * Node.spaces_per_tab + type_entry.print_mlir(0)
2772
2773 def print_tree(self, indent_level=0):
2774 return ' ' * indent_level * Node.spaces_per_tab + "TCustom: " + self._type
2775
2777 try:
2778 return isinstance(self._type, str) and ('xmemref<' in self._type or self._type.count('memref<') > 1)
2779 except Exception:
2780 return False
2781
2782
2784 def print_cpp(self, indent_level = 0):
2785 return ' ' * indent_level * Node.spaces_per_tab + "int"
2786
2787 def print_omp(self, indent_level = 0):
2788 return ' ' * indent_level * Node.spaces_per_tab + "int"
2789
2790 def print_sycl(self, indent_level = 0):
2791 return ' ' * indent_level * Node.spaces_per_tab + "int"
2792
2793 def print_mlir(self, indent_level = 0):
2794 return ' ' * indent_level * Node.spaces_per_tab + "i32"
2795
2796 def print_tree(self, indent_level=0):
2797 return ' ' * indent_level * Node.spaces_per_tab + "TInteger"
2798
2799
2801 """Type for 64-bit integers (long) used in MLIRCellData struct fields"""
2802 def print_cpp(self, indent_level = 0):
2803 return ' ' * indent_level * Node.spaces_per_tab + "long"
2804
2805 def print_omp(self, indent_level = 0):
2806 return ' ' * indent_level * Node.spaces_per_tab + "long"
2807
2808 def print_sycl(self, indent_level = 0):
2809 return ' ' * indent_level * Node.spaces_per_tab + "long"
2810
2811 def print_mlir(self, indent_level = 0):
2812 return ' ' * indent_level * Node.spaces_per_tab + "i64"
2813
2814 def print_tree(self, indent_level=0):
2815 return ' ' * indent_level * Node.spaces_per_tab + "TLong"
2816
2817
2819 def print_cpp(self, indent_level = 0):
2820 return ' ' * indent_level * Node.spaces_per_tab + "bool"
2821
2822 def print_omp(self, indent_level = 0):
2823 return ' ' * indent_level * Node.spaces_per_tab + "bool"
2824
2825 def print_sycl(self, indent_level = 0):
2826 return ' ' * indent_level * Node.spaces_per_tab + "bool"
2827
2828 def print_mlir(self, indent_level = 0):
2829 return ' ' * indent_level * Node.spaces_per_tab + "i1"
2830
2831 def print_tree(self, indent_level=0):
2832 return ' ' * indent_level * Node.spaces_per_tab + "TBoolean"
2833
2834
2836 def __init__(self, reference = False):
2837 self._reference = reference
2838 def print_cpp(self, indent_level = 0):
2839 return ' ' * indent_level * Node.spaces_per_tab + "double" + ("&" if self._reference else "")
2840
2841 def print_omp(self, indent_level = 0):
2842 return ' ' * indent_level * Node.spaces_per_tab + "double" + ("&" if self._reference else "")
2843
2844 def print_sycl(self, indent_level = 0):
2845 return ' ' * indent_level * Node.spaces_per_tab + "double" + ("&" if self._reference else "")
2846
2847 def print_mlir(self, indent_level = 0):
2848 return ' ' * indent_level * Node.spaces_per_tab + "f64"
2849
2850 def print_tree(self, indent_level=0):
2851 return ' ' * indent_level * Node.spaces_per_tab + "TFloat"
2852
2853
2855 def print_cpp(self, indent_level = 0):
2856 return ' ' * indent_level * Node.spaces_per_tab + "const char*"
2857
2858 def print_omp(self, indent_level = 0):
2859 return ' ' * indent_level * Node.spaces_per_tab + "const char*"
2860
2861 def print_sycl(self, indent_level = 0):
2862 return ' ' * indent_level * Node.spaces_per_tab + "const char*"
2863
2864 def print_mlir(self, indent_level=0):
2865 pass
2866
2867 def print_tree(self, indent_level=0):
2868 return ' ' * indent_level * Node.spaces_per_tab + "TString"
2869
2870
2872 def __init__(self, element_type=TFloat()):
2873 super().__init__()
2874 self._element_type = element_type
2875
2876 def print_cpp(self, indent_level=0):
2877 return ' ' * indent_level * Node.spaces_per_tab + f'double**' # C++ representation
2878
2879 def print_mlir(self, indent_level=0):
2880 return ' ' * indent_level * Node.spaces_per_tab + f'memref<?xmemref<?x{self._element_type.print_mlir(0).strip()}>>'
2881
2882 def print_omp(self, indent_level=0):
2883 return ' ' * indent_level * Node.spaces_per_tab + f'double**'
2884
2885 def print_sycl(self, indent_level=0):
2886 return ' ' * indent_level * Node.spaces_per_tab + f'double**'
2887
2889 return True
2890
2892 return self._element_type
2893
2894
2896 def __init__(self, signature_string):
2897 super().__init__()
2898 self._signature = signature_string
2899
2900 def print_cpp(self, indent_level=0):
2901 # C++ function pointers don't have direct equivalent
2902 return ' ' * indent_level * Node.spaces_per_tab + 'void*'
2903
2904 def print_mlir(self, indent_level=0):
2905 return ' ' * indent_level * Node.spaces_per_tab + self._signature
2906
2907 def print_omp(self, indent_level=0):
2908 return ' ' * indent_level * Node.spaces_per_tab + 'void*'
2909
2910 def print_sycl(self, indent_level=0):
2911 return ' ' * indent_level * Node.spaces_per_tab + 'void*'
2912
2913
2915 def __init__(self, element_type=TFloat(), dimensions=1):
2916 super().__init__()
2917 self._element_type = element_type
2918 self._dimensions = dimensions
2919
2920 def print_cpp(self, indent_level=0):
2921 return ' ' * indent_level * Node.spaces_per_tab + 'void*'
2922
2923 def print_mlir(self, indent_level=0):
2924 # Build memref<?x...xf64> with appropriate number of dimensions
2925 dims = 'x'.join(['?'] * self._dimensions)
2926 return ' ' * indent_level * Node.spaces_per_tab + f'memref<{dims}x{self._element_type.print_mlir(0).strip()}>' # Use print_mlir for element type
2927
2928 def print_omp(self, indent_level=0):
2929 return ' ' * indent_level * Node.spaces_per_tab + 'void*'
2930
2931 def print_sycl(self, indent_level=0):
2932 return ' ' * indent_level * Node.spaces_per_tab + 'void*'
2933
2935 return self._element_type
2936
2937
2939 def __init__(self, dimensions, underlying_type):
2940 self._dimensions = dimensions
2941 self._underlying_type = underlying_type
2942
2943 @staticmethod
2944 def needs_memref_load(operand):
2945 if isinstance(operand, VectorIndex):
2946 return False
2947
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):
2952 return True
2953
2954 if isinstance(operand, Subscript) and operand.is_nested_subscript():
2955 return True
2956
2957 return False
2958
2959 @staticmethod
2960 def generate_memref_load(operand, operand_mlir, const_zero_id, indent_level):
2961 indent = ' ' * indent_level * Node.spaces_per_tab
2962 load_id = f"%{Node.get_mlir_id()}"
2963 operand_type = operand.get_type()
2964
2965 if isinstance(operand_type, TDataBlock):
2966 # Handle nested array access (double** style arrays)
2967 if isinstance(operand, Subscript) and operand.is_nested_subscript():
2968 return operand_mlir
2969
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)
2973 else:
2974 element_type = "f64" # Default fallback
2975
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)
2978 return load_id
2979
2980 return operand_mlir
2981
2982 @staticmethod
2983 def ensure_const_zero(const_zero_id, indent_level):
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
2989
2990 def get_single(self):
2991 return self._underlying_type
2992
2994 if isinstance(self._underlying_type, String):
2995 if self._underlying_type._value == "double":
2996 return TFloat
2997 elif self._underlying_type._value == "int":
2998 return TInteger
2999 else:
3000 return None
3001 else:
3002 return type(self._underlying_type)
3003
3004 @staticmethod
3006 lhs_operand_type = lhs.get_type()
3007 if type(lhs_operand_type) is TDataBlock:
3008 lhs_type = lhs_operand_type.get_element_type()
3009 else:
3010 lhs_type = type(lhs_operand_type)
3011
3012 rhs_operand_type = rhs.get_type()
3013 if type(rhs_operand_type) is TDataBlock:
3014 rhs_type = rhs_operand_type.get_element_type()
3015 else:
3016 rhs_type = type(rhs_operand_type)
3017
3018 # Use float if either operand is float
3019 if lhs_type is TFloat or rhs_type is TFloat:
3020 return TFloat
3021 elif lhs_type is TLong or rhs_type is TLong:
3022 return TLong
3023 elif lhs_type is TInteger and rhs_type is TInteger:
3024 return TInteger
3025 elif lhs_type is TBoolean and rhs_type is TBoolean:
3026 return TBoolean
3027
3028 return None
3029
3031 element_type_class = self.get_element_typeget_element_type()
3032
3033 # Check for vector types by examining the underlying type
3034 if isinstance(self._underlying_type, String):
3035 is_vector = "tarch::la::Vector" in self._underlying_type._value
3036 elif hasattr(self._underlying_type, '_type'):
3037 is_vector = "tarch::la::Vector" in self._underlying_type._type
3038 else:
3039 is_vector = "tarch::la::Vector" in self._underlying_type
3040
3041 # Determine MLIR element type based on vector vs scalar
3042 if is_vector:
3043 mlir_element_type = "i64" # Vectors stored as pointers (i64)
3044 else:
3045 mlir_element_type = element_type_class().print_mlir(0)
3046
3047 return element_type_class or TFloat, is_vector, mlir_element_type
3048
3049
3050 def print_cpp(self, indent_level=0):
3051 return ' ' * indent_level * Node.spaces_per_tab + self._underlying_type.print_cpp() + "*"
3052
3053 def print_omp(self, indent_level=0):
3054 return ' ' * indent_level * Node.spaces_per_tab + self._underlying_type.print_omp() + "*"
3055
3056 def print_sycl(self, indent_level=0):
3057 return ' ' * indent_level * Node.spaces_per_tab + self._underlying_type.print_sycl() + "*"
3058
3059 def print_mlir(self, indent_level=0):
3060 dims = []
3061
3062 if isinstance(self._dimensions, int):
3063 dims.append(str(self._dimensions))
3064 else:
3065 for dim in self._dimensions:
3066 match dim:
3067 case int():
3068 dims.append(str(dim))
3069 case Expression():
3070 dims.append(dim.print_mlir(indent_level))
3071 case list():
3072 Node.context.push_block()
3073 for op in dim:
3074 op.print_mlir(indent_level)
3075 dims.append(''.join(Node.context.pop_block().get_mlir()))
3076
3077 # NOTE: The underlying type isn't an MLIR type here but a string representing the C++ type. So...
3078 # We look up the C++ type string in a types map in Node and get the MLIR type
3079 if isinstance(self._underlying_type, Type):
3080 type_name = self._underlying_type.print_mlir(indent_level)
3081 else:
3082 # We have a String type here
3083 type_name = Node.get_mlir_type(self._underlying_type._value)
3084 if type_name is None:
3085 # Fallback to type_map for complex types
3086 type_name = Node.type_map[f"{self._underlying_type}"]
3087
3088 # If type_name is already a memref type (e.g., from Vector types that map to 'memref<?xi64>'),
3089 # return it directly instead of wrapping it again. This prevents creating invalid nested memrefs like 'memref<?xmemref<?xi64>>'.
3090 if isinstance(type_name, str) and type_name.startswith("memref<"):
3091 return ' ' * indent_level * Node.spaces_per_tab + type_name
3092
3093 # Build the memref type with the correct number of dimensions
3094 # NOTE: self._dimensions may include halo dimensions, so we need to infer actual data dimensions
3095 # For now, we use a heuristic: if 1D field, return 1D; if multi-D, assume 2D (common case)
3096 # The type_map provides the correct types for named fields, so this should rarely be used
3097
3098 if isinstance(self._dimensions, int):
3099 if self._dimensions == 1:
3100 actual_dims = 1
3101 else:
3102 # Multi-dimensional iteration ranges with halo are typically 2D or 3D data
3103 # For now, assume 2D as fallback
3104 actual_dims = 2
3105 elif isinstance(self._dimensions, list):
3106 if len(self._dimensions) <= 1:
3107 actual_dims = 1
3108 else:
3109 actual_dims = 2 # Fallback for multi-dim
3110 else:
3111 actual_dims = 1 # Fallback
3112
3113 # Build memref with ? for each actual dimension
3114 dimension_list = ["?"] * actual_dims
3115 memref_dims = "x".join(dimension_list)
3116
3117 return ' ' * indent_level * Node.spaces_per_tab + f"""memref<{memref_dims}x{type_name}>"""
3118
3119 def print_tree(self, indent_level=0):
3120 return ' ' * indent_level * Node.spaces_per_tab + "TDataBlock"
3121
3122
3123def _is_nested_memref(type_map_entry):
3124 if isinstance(type_map_entry, str):
3125 # String: check for 'xmemref<' pattern
3126 return 'xmemref<' in type_map_entry
3127 elif hasattr(type_map_entry, 'is_nested_memref'):
3128 # Type object: call the method
3129 return type_map_entry.is_nested_memref()
3130 else:
3131 return False
3132
3133
3135 def print_cpp(self, indent_level = 0):
3136 pass
3137
3138 def print_omp(self, indent_level = 0):
3139 pass
3140
3141 def print_sycl(self, indent_level = 0):
3142 pass
3143
3144 def print_mlir(self, indent_level=0):
3145 pass
3146
3147 def print_tree(self, indent_level=0):
3148 return ' ' * indent_level * Node.spaces_per_tab + "TMatrix"
3149
3150
3152 def print_cpp(self, indent_level = 0):
3153 pass
3154
3155 def print_omp(self, indent_level = 0):
3156 pass
3157
3158 def print_sycl(self, indent_level = 0):
3159 pass
3160
3161 def print_mlir(self, indent_level=0):
3162 pass
3163
3164 def print_tree(self, indent_level=0):
3165 return ' ' * indent_level * Node.spaces_per_tab + "TDiagonalMatrix"
3166
3167
3169 def print_cpp(self, indent_level = 0):
3170 pass
3171
3172 def print_omp(self, indent_level = 0):
3173 pass
3174
3175 def print_sycl(self, indent_level = 0):
3176 pass
3177
3178 def print_mlir(self, indent_level=0):
3179 pass
3180
3181 def print_tree(self, indent_level=0):
3182 return ' ' * indent_level * Node.spaces_per_tab + "TVector"
3183
3184
3185# Statements
3187 def __init__(self, boolean):
3188 super().__init__()
3189 self._boolean = boolean
3190 self._statements = []
3191
3192 def add_statement(self, statement):
3193 if statement is not None:
3194 self._statements.append(statement)
3195
3196 def print_cpp(self, indent_level=0):
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 + "}"
3201
3202 def print_omp(self, indent_level=0):
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 + "}"
3207
3208 def print_sycl(self, indent_level=0):
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 + "}"
3213
3214 def print_mlir(self, indent_level=0):
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 + "}"
3222
3223
3224 def print_tree(self, indent_level=0):
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)}"""
3228
3230 _inuse_iteration_variables = set()
3231 _iteration_variable_names = ['i', 'j', 'k', 'l', 'n', 'm', 'a', 'b', 'c', 'd']
3232
3233 def __init__(self, iteration_range, iteration_variable_name = None, use_scheduler = False):
3234 super().__init__()
3235 self._iteration_range = iteration_range
3236 self._statements = []
3237 self._use_scheduler = False# use_scheduler
3238
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:
3242 self._iteration_variable = Name(variable_name, TInteger())
3243 self._iteration_variable._index_type = True
3244 For._inuse_iteration_variables.add(variable_name)
3245 break
3246 else:
3247 self._iteration_variable = Name(iteration_variable_name, TInteger())
3248 For._inuse_iteration_variables.update({iteration_variable_name : self._iteration_variable})
3249
3251 return self._iteration_range[1] - self._iteration_range[0]
3252
3254 return self._iteration_variable
3255
3256 def add_statement(self, statement):
3257 self._statements.append(statement)
3258
3259 def close_scope(self):
3260 For._inuse_iteration_variables.remove(self._iteration_variable.id)
3261
3262 def print_cpp(self, indent_level = 0):
3263 statement_prints = [statement.print_cpp(indent_level + 1) for statement in self._statements]
3264 if self._use_scheduler:
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}}}
3268endParallelFor"""
3269 else:
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 + "}"
3273
3274
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 + "}"
3280
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 + "}"
3286
3287 def print_mlir(self, indent_level=0):
3288 indent = ' ' * indent_level * Node.spaces_per_tab
3289 bounds_casting = ''
3290
3291 if isinstance(self._iteration_range[0], list) or isinstance(self._iteration_range[0], tuple):
3292 dims = len(self._iteration_range)
3293 iter_vars = []
3294 lwbs = []
3295 upbs = []
3296 steps = []
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"
3305 lwb = mlir_id
3306 else:
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"
3311 upb = mlir_id
3312 else:
3313 upb = self._iteration_range[d][1].print_mlir(indent_level).strip()
3314 lwbs.append(lwb)
3315 upbs.append(upb)
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)
3319 else:
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]
3330 else:
3331 iter_vars = [self._iteration_variable.print_mlir(indent_level).strip()]
3332
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"
3340 lwbs = [mlir_id]
3341 else:
3342 lwbs += [self._iteration_range[0].print_mlir(indent_level).strip()]
3343
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"
3351 upbs = [mlir_id]
3352 else:
3353 upbs += [self._iteration_range[1].print_mlir(indent_level).strip()]
3354
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"
3358 steps = [mlir_id]
3359
3360 # Collect body
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)
3365
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)
3369
3370 if has_allocation:
3371 # Use sequential scf.for for memory allocation (cannot parallelize)
3372 loop_header = (
3373 f"{indent}scf.for {', '.join(iter_vars)} = "
3374 f"{', '.join(lwbs)} to {', '.join(upbs)} step {', '.join(steps)} {{\n"
3375 )
3376 else:
3377 # Use scf.parallel for computation loops (safe to parallelize)
3378 loop_header = (
3379 f"{indent}scf.parallel ({', '.join(iter_vars)}) = "
3380 f"({', '.join(lwbs)}) to ({', '.join(upbs)}) step ({', '.join(steps)}) {{\n"
3381 )
3382 loop_footer = "\n" + indent + "}\n"
3383
3384 return bounds_casting + loop_header + statements + loop_footer
3385
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)}"""
3390
3391
3392class Comment(Statement):
3393 def __init__(self, comment):
3394 super().__init__()
3395 self._comment = comment
3396
3397 def get_type(self):
3398 return None
3399
3400 def print_cpp(self, indent_level = 0):
3401 return ' ' * indent_level * Node.spaces_per_tab + "//" + self._comment[1:]
3402
3403 def print_omp(self, indent_level = 0):
3404 return ' ' * indent_level * Node.spaces_per_tab + "//" + self._comment[1:]
3405
3406 def print_sycl(self, indent_level = 0):
3407 return ' ' * indent_level * Node.spaces_per_tab + "//" + self._comment[1:]
3408
3409 def print_mlir(self, indent_level = 0):
3410 return ' ' * indent_level * Node.spaces_per_tab + "//" + self._comment[1:]
3411
3412 def print_tree(self, indent_level=0):
3413 return ' ' * indent_level * Node.spaces_per_tab + "Comment"
3414
3415
3416class FunctionCall(Statement):
3417 def __init__(self, id, arguments, is_offloadable = False):
3418 self.id = id
3419 self._arguments = arguments
3420 self._is_offloadable = is_offloadable
3421
3422 def add_argument(self, argument: Expression):
3423 self._arguments.append(argument)
3424
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)});"""
3428
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)});"""
3432
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)});"""
3436
3437 def print_mlir(self, indent_level=0):
3438 indent = ' ' * indent_level * Node.spaces_per_tab
3439 Node.context.push_block()
3440
3441 # Detect if this is a bridge function call
3442 is_bridge_call = self.id in ['sourceTerm', 'flux', 'nonconservativeProduct']
3443
3444 # Process arguments and handle array subscripts
3445 argument_prints = []
3446 for i, argument in enumerate(self._arguments):
3447 argument.hoist(True)
3448
3449 # Check if this is a Reference to a 2D subscript for a bridge call
3450 is_reference_to_2d_subscript = (
3451 is_bridge_call and
3452 isinstance(argument, Reference) and
3453 isinstance(argument._expression, Subscript) and
3454 isinstance(argument._expression._value, Subscript)
3455 )
3456
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
3463
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)
3468
3469 # Get MLIR representations
3470 memref_2d.hoist(True)
3471 memref_2d_mlir = memref_2d.print_mlir(indent_level).strip(' ')
3472
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
3478
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
3482 )
3483 Node.context.get_block().mlir_append(ptr_code)
3484 arg_mlir = final_ptr
3485 else:
3486 # Not a 2D subscript for bridge call - use standard approach
3487 arg_mlir = argument.print_mlir(indent_level).strip(' ')
3488
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()
3495
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}]"
3502 else:
3503 # Single-dimensional DataBlock
3504 element_type_class, is_vector, mlir_element_type = base_type.get_array_element_info()
3505
3506 if is_vector:
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)
3515 arg_mlir = ptr_id
3516 else:
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
3522 else:
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
3530 else:
3531 # For other types, just pass the reference
3532 arg_mlir = f"{base_value}[{index}]"
3533
3534 argument_prints.append(arg_mlir)
3535
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(' '))
3539
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
3546 else:
3547 fn_type = type_entry.print_mlir(0).strip()
3548
3549 bridge_function_name = f"{self.id}_bridge"
3550 call_statement = f"""func.call @{bridge_function_name}({", ".join(argument_prints)}) : {fn_type.strip(' ')}"""
3551
3552 Node.context.get_block().mlir_append(indent + call_statement)
3553 return ''
3554
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)}"""
3559
3560
3561class Max(Statement):
3562 def __init__(self, lhs, rhs):
3563 super().__init__()
3564 self._lhs = lhs
3565 self._rhs = rhs
3566
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()});"""
3569
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()});"""
3572
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()});"""
3575
3576 def print_mlir(self, indent_level=0):
3577 indent = ' ' * indent_level * Node.spaces_per_tab
3578
3579 def unpack_nested_memref_access(expr, indent_level):
3580 indent = ' ' * indent_level * Node.spaces_per_tab
3581
3582 if not (isinstance(expr, Subscript) and isinstance(expr._value, Subscript)):
3583 return expr.print_mlir(indent_level).strip(' ')
3584
3585 outer_subscript = expr._value
3586
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)
3591
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))
3594
3595 inner_memref_id = f"%{Node.get_mlir_id()}"
3596
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(' ')
3599
3600 # Try to determine the outer memref type
3601 # First, check if this variable is registered in type_map
3602 try:
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
3610 else:
3611 # Type object - call print_mlir()
3612 outer_memref_type = type_entry.print_mlir(0).strip()
3613 else:
3614 # Fallback to default nested memref type
3615 outer_memref_type = "memref<?xmemref<?xf64>>"
3616 except:
3617 outer_memref_type = "memref<?xmemref<?xf64>>"
3618
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)
3623
3624 return value_id
3625
3626 Node.context.push_block()
3627
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)
3632
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)
3636 else:
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)
3641
3642 if isinstance(self._rhs, Subscript) and isinstance(self._rhs._value, Subscript):
3643 rhs = unpack_nested_memref_access(self._rhs, indent_level)
3644 else:
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)
3649
3650 argument_block = Node.context.pop_block().get_mlir()
3651 for op in argument_block:
3652 Node.context.get_block().mlir_append(op)
3653
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)
3657
3658 if op_type_class is TInteger:
3659 mlir_op = "arith.maxsi"
3660 else:
3661 mlir_op = "arith.maximumf" # Default to float
3662
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)
3666
3667 return new_id
3668
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)}"""
3673
3674
3675class MemoryAllocation(Statement):
3676 memoryAllocated = []
3677 allocated_types = {} # Map from MLIR variable name to allocated memref type string
3678
3679 def __init__(self, name: Name, object_type: Type, dimensions, specify_type = True, add_to_stack=True):
3680 super().__init__()
3681 self._type = object_type
3682 if type(dimensions[0]) is list:
3683 self._dimensions = [dimension[1] for dimension in dimensions]
3684 else:
3685 self._dimensions = dimensions
3686 self._name = name
3687 self._specify_type = specify_type
3688
3689 if add_to_stack == True:
3690 MemoryAllocation.memoryAllocated[-1].append(self)
3691
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]
3696
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()
3700
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()}];"""
3704 else:
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)}"""
3708 else:
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()}];"""
3710
3711
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()}];"""
3716 else:
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)}"""
3719 else:
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);"""
3723 else:
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);"""
3725 else:
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);"""
3728 else:
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);"""
3730
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()}];"""
3734 else:
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)}"""
3737
3738 def print_mlir(self, indent_level=0):
3739 indent = ' ' * indent_level * Node.spaces_per_tab
3740
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
3747
3748
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()
3754
3755 # Now build the allocation operations
3756 index_id = f"%{Node.get_mlir_id()}"
3757 decl = ''.join(intermediate_ops)
3758
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()}"
3764
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}>"
3773 else:
3774 mlir_type = name_type.print_mlir(0)
3775
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
3779
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
3787
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()
3793
3794 # Add any additional operations
3795 decl += ''.join(additional_ops)
3796
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(' ')
3804 else:
3805 # Type object - call print_mlir()
3806 type_mlir = type_entry.print_mlir(0).strip()
3807 outer_type = f"memref<?xmemref<?x{type_mlir}>>"
3808
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}"
3811
3812 return decl
3813 else:
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
3818
3819 # Use nested memref if in type_map, otherwise check if DataBlock flat
3820 if in_type_map:
3821 # Variable has explicit type specification - use nested memref path
3822 is_datablock = False # Force nested path
3823 actual_dims_calc = -1
3824 else:
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)
3828
3829 # Calculate actual_dims based on DataBlock iteration_range
3830 if is_datablock:
3831 datablock_obj = self._name.get_type()
3832
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
3839 else:
3840 actual_dims_calc = 1
3841 else:
3842 actual_dims_calc = -1 # Not a DataBlock
3843
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)
3848
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
3856
3857 intermediate_decl = ""
3858 dim_ids = []
3859
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()
3867
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"
3870
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
3874
3875 # Re-push a new block for remaining operations
3876 Node.context.push_block()
3877
3878 # Process name and allocate
3879 name_mlir = self._name.print_mlir(indent_level)
3880
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)
3885
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"
3888
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
3893
3894 return intermediate_decl + alloc_line
3895 else:
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]
3901
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()
3906
3907 # Build the allocation operations
3908 decl = ''.join(intermediate_ops)
3909
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"
3915
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()
3919
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]
3924
3925 # Handle both Type objects and strings
3926 if isinstance(type_entry, str):
3927 outer_type = type_entry
3928 else:
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}>>"
3936 else:
3937 # Fallback for non-DataBlock types
3938 name_type_mlir = self._name.get_type().print_mlir(0)
3939
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(' ')
3945 else:
3946 # Type object
3947 name_type_mlir = type_entry.print_mlir(0).strip()
3948 outer_type = f"memref<?x{name_type_mlir}>"
3949
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()
3954
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
3960
3961 # Use the existing loop structure for inner allocations
3962 loop_mlir = self._loop.print_mlir(indent_level)
3963
3964 return decl + loop_mlir
3965
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)}"""
3969
3970
3971class MemoryDeallocation(Statement):
3972 def __init__(self, allocation: MemoryAllocation):
3973 super().__init__()
3974 self._dimensions = allocation._dimensions
3975 self._name = allocation._name
3976 self._allocated_type = None # Will store the MLIR type string when known
3977
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]
3982
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()
3986
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()};"""
3990 else:
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()};"""
3994 else:
3995 return f"""{' ' * indent_level * Node.spaces_per_tab}delete[] {self._name.print_cpp()};"""
3996
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()};"""
4001 else:
4002 return f"""{self._loop.print_omp(indent_level)}
4003{' ' * indent_level * Node.spaces_per_tab}delete[] {self._name.print_omp()};"""
4004 else:
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);"""
4007 else:
4008 return ' ' * indent_level * Node.spaces_per_tab + f"""omp_target_free({self._name.print_omp()}, targetDevice);"""
4009
4010 def print_sycl(self, indent_level=0):
4011 return ' ' * indent_level * Node.spaces_per_tab + f"""::sycl::free({self._name.print_sycl()}, queue);"""
4012 #else:
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);"""
4016
4017 def print_mlir(self, indent_level=0):
4018 indent = ' ' * indent_level * Node.spaces_per_tab
4019
4020 # Get the MLIR variable name
4021 name_mlir = self._name.print_mlir(indent_level)
4022
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
4029 else:
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()
4035
4036 if len(self._dimensions) == 1:
4037 name_type_mlir = f"memref<?x{element_type}>"
4038 else:
4039 # For multidimensional, assume nested memref
4040 name_type_mlir = f"memref<?xmemref<?x{element_type}>>"
4041
4042 return "\n" + indent + f"""memref.dealloc {name_mlir} : {name_type_mlir}\n"""
4043
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)}"""
4047
4048
4049class Construction(Statement):
4050 def __init__(self, name: Name, expression: Expression):
4051 super().__init__()
4052 self._name = name
4053 self._name.on_left(True)
4054 self._expression = expression
4055 Name.variables.update({self._name.id: expression})
4056
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()};"
4059
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()};"
4062
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()};"
4065
4066 def print_mlir(self, indent_level = 0):
4067 indent = ' ' * indent_level * Node.spaces_per_tab
4068
4069 if self._expression._string is None:
4070 expr_mlir = self._expression.print_mlir(indent_level)
4071
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"
4075 else:
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"
4079
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"
4083
4084 Node.context.get_block().mlir_append(mlir_str)
4085 return ''
4086
4087 def print_tree(self, indent_level=0):
4088 return f"""
4089Construction:
4090{self._name.print_tree(indent_level + 1)}
4091{self._expression.print_tree(indent_level + 1)}"""
4092
4093
4094class DataBlockConstructionFromExisting(Statement):
4095 def __init__(self, name: Name, dataBlock: DataBlock):
4096 super().__init__()
4097 self._name = name
4098 self._dataBlock = dataBlock
4099 self._dataBlock.id = name.id
4100 Name.variables.update({self._name.id: self._dataBlock})
4101 FunctionDefinition._syclDataToCopy.append(self)
4102
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:
4106 type_print += "*"
4107 return ' ' * indent_level * Node.spaces_per_tab + f"{type_print} {self._name.print_cpp()} = {self._dataBlock._internal.print_cpp()};"
4108
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:
4112 type_print += "*"
4113 return ' ' * indent_level * Node.spaces_per_tab + f"{type_print} {self._name.print_omp()} = {self._dataBlock._internal.print_omp()};"
4114
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:
4118 type_print += "*"
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}}}
4128#"""
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}}}
4136#"""
4137 return ' ' * indent_level * Node.spaces_per_tab + f"{type_print} {self._name.print_sycl()} = {self._dataBlock._internal.print_sycl()};"
4138
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)
4142
4143 dims = []
4144 for i in range(len(self._dataBlock._iteration_range) - 3, -1, -1):
4145 dims.append(self._dataBlock._iteration_range[i])
4146
4147 if len(self._dataBlock._iteration_range) > 1:
4148 type_print = "!llvm.ptr " #+ f"{self._dataBlock._underlying_type} dims: {dims})"
4149
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
4154
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
4159
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())
4163
4164 # Clear the target variable name - we've finished generating the MLIR for this Node
4165 Node.target_variable_name = None
4166
4167 return mlir_str
4168
4169
4170 def print_tree(self, indent_level=0):
4171 return f"""
4172DataBlockConstructionFromExisting:
4173{self._name.print_tree(indent_level + 1)}
4174{self._dataBlock._internal.print_tree(indent_level + 1)}"""
4175
4176
4177class DataBlockConstructionFromOperation:
4178 def __init__(self, name: Name, dataBlockOperation):
4179 super().__init__()
4180 self._name = name
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
4185 else:
4186 self._operation = dataBlockOperation._internal
4187
4188 self.memoryAllocation = MemoryAllocation(self._name, self._dataBlock.get_type().get_single(), self._dataBlock._memory_range)
4189 self.assignment = DataBlockAssignment(self._name, self._operation)
4190
4191 def print_cpp(self, indent_level = 0):
4192 return f"""{self.memoryAllocation.print_cpp(indent_level)}
4193{self.assignment.print_cpp(indent_level)}"""
4194
4195 def print_omp(self, indent_level = 0):
4196 return f"""{self.memoryAllocation.print_omp(indent_level)}
4197{self.assignment.print_omp(indent_level)}"""
4198
4199 def print_sycl(self, indent_level = 0):
4200 return f"""{self.memoryAllocation.print_sycl(indent_level)}
4201{self.assignment.print_sycl(indent_level)}"""
4202
4203 def print_mlir(self, indent_level = 0):
4204 return f"""{self.memoryAllocation.print_mlir(indent_level)}
4205{self.assignment.print_mlir(indent_level)}"""
4206
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)}"""
4211
4212
4213class DataBlockConstructionFromMatrixVectorProduct:
4214 def __init__(self, name: Name, matrix, vector):
4215 self._name = name
4216 if type(vector) is Name:
4217 self._vector = Name.variables[vector.id]
4218 else:
4219 self._vector = vector
4220
4221 if type(matrix) is Name:
4222 self._matrix = Name.variables[matrix.id]
4223 else:
4224 self._matrix = matrix
4225
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})
4228
4229 self.memoryAllocation = MemoryAllocation(self._name, self._dataBlock.get_type().get_single(), self._dataBlock._memory_range)
4230
4231 element_width = self._vector._iteration_range[0][1] - self._vector._iteration_range[0][0]
4232
4233 loops = []
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]]))
4237
4238 if Node.use_accelerator:
4239 size = Integer(1)
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)))
4244 else:
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)))
4246
4247 loops.append(For([Integer(0), self._matrix._dimensions[1]]))
4248
4249 if Node.use_accelerator:
4250 size = Integer(1)
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())))))
4255 else:
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())))))
4257
4258
4259
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]
4268
4269 def print_cpp(self, indent_level = 0):
4270 return f"""{self.memoryAllocation.print_cpp(indent_level)}
4271{self.assignment.print_cpp(indent_level)}"""
4272
4273 def print_omp(self, indent_level = 0):
4274 return f"""{self.memoryAllocation.print_omp(indent_level)}
4275{self.assignment.print_omp(indent_level)}"""
4276
4277 def print_sycl(self, indent_level = 0):
4278 return f"""{self.memoryAllocation.print_sycl(indent_level)}
4279{self.assignment.print_sycl(indent_level)}"""
4280
4281 def print_mlir(self, indent_level = 0):
4282 return f"""{self.memoryAllocation.print_mlir(indent_level)}
4283{self.assignment.print_mlir(indent_level)}"""
4284
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)}"""
4289
4290
4291class DataBlockConstructionFromDotProduct:
4292 def __init__(self, name: Name, matrix, vector):
4293 self._name = name
4294 if type(vector) is Name:
4295 self._vector = Name.variables[vector.id]
4296 else:
4297 self._vector = vector
4298
4299 if type(matrix) is Name:
4300 self._matrix = Name.variables[matrix.id]
4301 else:
4302 self._matrix = matrix
4303
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})
4308
4309 self.memoryAllocation = MemoryAllocation(self._name, self._dataBlock.get_type().get_single(), self._dataBlock._memory_range)
4310
4311 loops = []
4312 index = []
4313 for [a, b] in self._dataBlock._iteration_range[-1::-1]:
4314 loops.append(For([Integer(0), b - a]))
4315
4316 for loop in loops[::-1]:
4317 index.append(loop.get_iteration_variable())
4318
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])
4322
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])
4328
4329 for loop in loops:
4330 loop.close_scope()
4331 self.assignment = loops[0]
4332
4333 def print_cpp(self, indent_level = 0):
4334 return f"""{self.memoryAllocation.print_cpp(indent_level)}
4335{self.assignment.print_cpp(indent_level)}"""
4336
4337 def print_omp(self, indent_level = 0):
4338 return f"""{self.memoryAllocation.print_omp(indent_level)}
4339{self.assignment.print_omp(indent_level)}"""
4340
4341 def print_sycl(self, indent_level = 0):
4342 return f"""{self.memoryAllocation.print_sycl(indent_level)}
4343{self.assignment.print_sycl(indent_level)}"""
4344
4345 def print_mlir(self, indent_level = 0):
4346 return f"""{self.memoryAllocation.print_mlir(indent_level)}
4347{self.assignment.print_mlir(indent_level)}"""
4348
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)}"""
4353
4354
4355class DataBlockConstructionFromFunction(Statement):
4356 def __init__(self, name: Name, dataBlock: DataBlock):
4357 super().__init__()
4358 self._name = name
4359 self._dataBlock = dataBlock
4360 self._dataBlock.id = name.id
4361 Name.variables.update({self._name.id: self._dataBlock})
4362
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])
4365
4366 inner_loops = []
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)
4374
4375 index = []
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())
4379
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))]
4383 else:
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))]
4385
4386 if type(self._dataBlock._internal) is FunctionCall:
4387 functionCall = FunctionCall(self._dataBlock._internal.id, [])
4388 functionCall._is_offloadable = self._dataBlock._internal._is_offloadable
4389
4390 if type(self._dataBlock._internal._arguments[0]) is Index:
4391 functionCall.add_argument(Vector(index[0:-1], TInteger()))
4392 else:
4393 functionCall.add_argument(Reference(self._dataBlock._internal._arguments[0].multidimensional_index(index, 1)))
4394
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))
4400 else:
4401 functionCall.add_argument(Reference(Name.variables[argument.id].multidimensional_index(index, 1)))
4402 else:
4403 functionCall.add_argument(argument)
4404
4405 functionCall.add_argument(Reference(self._dataBlock.multidimensional_index(index, 1)))
4406 if functionCall._is_offloadable:
4407 functionCall.add_argument(String("Solver::Offloadable::Yes"))
4408 else:
4409 functionCall = Assignment(self._dataBlock.multidimensional_index(index), self._dataBlock._internal)
4410
4411 self.initialisation.add_statement(inner_loops[0])
4412
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)
4416
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)
4422
4423 def print_cpp(self, indent_level = 0):
4424 return f"""{self.memoryAllocation.print_cpp(indent_level)}
4425{self.initialisation.print_cpp(indent_level)}"""
4426
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)"
4430 else:
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)}"""
4435
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();"""
4446
4447
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)
4455
4456 return f"""{mem_alloc}
4457{init}"""
4458
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)}"""
4463
4464class Assignment(Statement):
4465 def __init__(self, lhs: Expression, rhs: Expression):
4466 super().__init__()
4467 self._lhs = lhs
4468 self._lhs.on_left(True)
4469 self._rhs = rhs
4470 self._rhs.on_left(False) # Just to be sure...
4471
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()};"
4474
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()};"
4477
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()};"
4480
4481 def print_mlir(self, indent_level = 0):
4482 indent = ' ' * indent_level * Node.spaces_per_tab
4483 mlir_str = ''
4484 Node.context.push_block()
4485 rhs = self._rhs.print_mlir(indent_level)
4486
4487 rhs_indices = []
4488 rhs_dims = []
4489 rhs_casting_mlir = ''
4490
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()
4498
4499 mlir_str += ''.join(Node.context.pop_block().get_mlir())
4500 mlir_str += rhs_casting_mlir
4501 else:
4502 # Assume scalar here
4503 tmp = ''.join(Node.context.pop_block().get_mlir())
4504 mlir_str += tmp
4505
4506 Node.context.push_block()
4507 lhs = self._lhs.print_mlir(indent_level)
4508
4509 lhs_indices = []
4510 lhs_dims = []
4511 lhs_casting_mlir = ''
4512
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
4521
4522 mlir_str += ''.join(Node.context.pop_block().get_mlir())
4523 mlir_str += lhs_casting_mlir
4524 else:
4525 tmp = ''.join(Node.context.pop_block().get_mlir())
4526 mlir_str += tmp
4527
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 = []
4532
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]
4541
4542 outer_sub = self._lhs
4543 inner_sub = self._lhs._value
4544
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)
4549
4550 if lhs_is_flat_2d:
4551 rhs_value = rhs
4552
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
4558
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)
4562 rhs_value = rhs
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"
4576 rhs_value = rhs_id
4577
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
4582
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]
4588
4589 if isinstance(type_entry, str):
4590 outer_memref_type = type_entry
4591 else:
4592 outer_memref_type = type_entry.print_mlir(0).strip()
4593
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"
4597
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"
4601 else:
4602 # Handle regular flat memref assignment (non-nested)
4603 rhs_id = f"%{Node.get_mlir_id()}"
4604
4605 lhs_index_str = '[' + ','.join(lhs_indices) + ']'
4606 # Ensure at least one dimension for memref type
4607 if lhs_dims:
4608 lhs_type = f"memref<{ 'x'.join(lhs_dims) }xf64>"
4609 else:
4610 lhs_type = "memref<?xf64>"
4611 rhs_index_str = '[' + ','.join(rhs_indices) + ']'
4612 if rhs_dims:
4613 rhs_type = f"memref<{ 'x'.join(rhs_dims) }x f64>"
4614 else:
4615 rhs_type = "memref<?xf64>"
4616
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"
4619 rhs_value = rhs_id
4620 else:
4621 rhs_value = rhs
4622
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
4631
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"
4634 else:
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"
4638
4639 return mlir_str
4640
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)}"""
4645
4646
4647class DataBlockAssignment(Statement):
4648 def __init__(self, lhs: Expression, rhs: Expression):
4649 super().__init__()
4650 if type(lhs) is Name:
4651 self._lhs = Name.variables[lhs.id]
4652 self._lhs._id = lhs.id
4653 else:
4654 self._lhs = lhs
4655
4656 if type(rhs) is Name:
4657 self._rhs = Name.variables[rhs.id]
4658 else:
4659 self._rhs = rhs
4660
4661 self._lhs.on_left(True)
4662 self._rhs.on_left(False) # Just to be sure...
4663
4664 self._loops = []
4665 index = []
4666
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]]))
4670
4671 for i in range(len(self._lhs._memory_range) - 1, -1, -1):
4672 index.append(self._loops[i].get_iteration_variable())
4673
4674 if type(self._rhs.get_type()) is TDataBlock:
4675 rhs_offset = []
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)
4682 else:
4683 rhs_offset = [Integer(0) for i in self._lhs._iteration_range]
4684
4685 lhs_offset = [Integer(0) for i in self._lhs._iteration_range]
4686
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])
4690
4691 for loop in self._loops[::-1]:
4692 loop.close_scope()
4693 self._loop = self._loops[0]
4694 self._numDimensions = len(self._loops)
4695
4696 def print_cpp(self, indent_level = 0):
4697 return self._loop.print_cpp(indent_level)
4698
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)"
4702 else:
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)}"""
4706
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();"""
4718
4719
4720 def print_mlir(self, indent_level = 0):
4721 return self._loop.print_mlir(indent_level)
4722
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)
cast_operand(operand_mlir, operand, result_type, indent_level)
multidimensional_index(self, index, start_index=0, dimensions=[])
get_mlir_arithmetic_op(operation, op_type_class)
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)
get_mlir_comparison_op(operation, op_type_class)
__init__(self, operation, lhs, rhs)
multidimensional_index(self, index, start_index=0, dimensions=[])
__init__(self, operation, lhs, rhs, useFunctionSyntax=False)
multidimensional_index(self, index, start_index=0, dimensions=[])
multidimensional_index(self, index, start_index=0, dimensions=[])
multidimensional_index(self, index, start_index=0, dimensions=[])
multidimensional_index(self, index, start_index=0, dimensions=[])
__init__(self, iteration_range, internal, requires_memory_allocation, id=None, underlying_type=String("double"))
linearise_index(self, indices, dimensions, offset_start_index=0)
multidimensional_index(self, index_list, start_index=0, dimensions=[])
__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)
__init__(self, iteration_range, iteration_variable_name=None, use_scheduler=False)
__init__(self, id, Type return_type, template=None, namespaces=[], stateless=False, top_level=False)
print_omp(self, indent_level=0)
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_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)
__init__(self, statementToLog, outputStream)
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)
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)
__init__(self, Expression expression)
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)
__init__(self, Expression value, Expression index)
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)
generate_memref_load(operand, operand_mlir, const_zero_id, indent_level)
__init__(self, dimensions, underlying_type)
ensure_const_zero(const_zero_id, indent_level)
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)
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)
__init__(self, element_type=TFloat())
__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)
__init__(self, Expression value, Expression index)
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)