Replace buffer with another buffer. Also replace the data of the buffer with another var.
| 38 | |
| 39 | # FIXME: this pass does not replace var in the shape/layout of a buffer |
| 40 | class BufferReplacer(StmtExprMutator): |
| 41 | """ |
| 42 | Replace buffer with another buffer. |
| 43 | Also replace the data of the buffer with another var. |
| 44 | """ |
| 45 | |
| 46 | def __init__( |
| 47 | self, buffer_map: dict[Buffer, Buffer] | None = None, var_map: dict[Var, Var] | None = None |
| 48 | ): |
| 49 | super().__init__() |
| 50 | self.buffer_map = buffer_map if buffer_map is not None else {} |
| 51 | self.var_map = var_map if var_map is not None else {} |
| 52 | self.buffer_attr_var_mutated = False |
| 53 | for old_buffer, new_buffer in self.buffer_map.items(): |
| 54 | self.var_map[old_buffer.data] = new_buffer.data |
| 55 | |
| 56 | def mutate_buffer(self, buffer: Buffer): |
| 57 | if buffer in self.buffer_map: |
| 58 | return self.buffer_map[buffer] |
| 59 | |
| 60 | # Track mutations for this specific buffer only. Without this reset, |
| 61 | # unrelated buffers can be spuriously cloned and introduce alias buffers. |
| 62 | prev_mutated = self.buffer_attr_var_mutated |
| 63 | self.buffer_attr_var_mutated = False |
| 64 | new_data = self.visit_expr(buffer.data) |
| 65 | new_shape = [self.visit_expr(expr) for expr in buffer.shape] |
| 66 | if isinstance(buffer.layout, TileLayout): |
| 67 | new_shard = [] |
| 68 | new_replicate = [] |
| 69 | for iter in buffer.layout.shard: |
| 70 | new_iter = Iter( |
| 71 | self.visit_expr(iter.extent), self.visit_expr(iter.stride), iter.axis |
| 72 | ) |
| 73 | new_shard.append(new_iter) |
| 74 | for iter in buffer.layout.replica: |
| 75 | new_iter = Iter( |
| 76 | self.visit_expr(iter.extent), self.visit_expr(iter.stride), iter.axis |
| 77 | ) |
| 78 | new_replicate.append(new_iter) |
| 79 | new_layout = TileLayout.from_iters( |
| 80 | new_shard, new_replicate, offset=buffer.layout.offset |
| 81 | ) |
| 82 | else: |
| 83 | new_layout = buffer.layout |
| 84 | buffer_attr_mutated = self.buffer_attr_var_mutated |
| 85 | self.buffer_attr_var_mutated = prev_mutated or buffer_attr_mutated |
| 86 | if not buffer_attr_mutated: |
| 87 | return None |
| 88 | new_buffer = decl_buffer( |
| 89 | new_shape, |
| 90 | buffer.dtype, |
| 91 | buffer.name, |
| 92 | new_data, |
| 93 | buffer.strides, |
| 94 | buffer.elem_offset, |
| 95 | buffer.scope(), |
| 96 | buffer.data_alignment, |
| 97 | buffer.offset_factor, |
no outgoing calls
searching dependent graphs…