MCPcopy Create free account
hub / github.com/apache/tvm / BufferReplacer

Class BufferReplacer

python/tvm/tirx/transform/common.py:40–172  ·  view source on GitHub ↗

Replace buffer with another buffer. Also replace the data of the buffer with another var.

Source from the content-addressed store, hash-verified

38
39# FIXME: this pass does not replace var in the shape/layout of a buffer
40class 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,

Calls

no outgoing calls

Tested by 1

Used in the wild real call sites across dependent graphs

searching dependent graphs…