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

Method _normalize_buffer_arg

python/tvm/s_tir/schedule/schedule.py:3244–3310  ·  view source on GitHub ↗
(
        self,
        block: SBlockRV,
        buffer: tuple[str, int] | int | str | Buffer,
        required_buffer_type=None,
    )

Source from the content-addressed store, hash-verified

3242 return block
3243
3244 def _normalize_buffer_arg(
3245 self,
3246 block: SBlockRV,
3247 buffer: tuple[str, int] | int | str | Buffer,
3248 required_buffer_type=None,
3249 ) -> tuple[str, int, Buffer]:
3250 block_obj: SBlock = self.get(block)
3251 block_name = block_obj.name_hint
3252
3253 def iter_buffers():
3254 for i, read in enumerate(block_obj.reads):
3255 yield "read", i, read.buffer
3256 for i, write in enumerate(block_obj.writes):
3257 yield "write", i, write.buffer
3258
3259 if isinstance(buffer, int):
3260 buffer = (required_buffer_type, buffer)
3261
3262 if isinstance(buffer, str):
3263 possible_buffers = {}
3264 # String lookup requires ensuring that the name is unique
3265 for buffer_index_type, buffer_index, buf in iter_buffers():
3266 if buf.name == buffer:
3267 possible_buffers[buf] = (buffer_index_type, buffer_index)
3268
3269 assert possible_buffers, f"Could not find buffer '{buffer}' in block '{block_name}'"
3270 assert len(possible_buffers) == 1, (
3271 f"Multiple buffers named '{buffer}' in block '{block_name}'"
3272 )
3273 buffer_obj, (buffer_index_type, buffer_index) = next(iter(possible_buffers.items()))
3274
3275 elif isinstance(buffer, Buffer):
3276 # Buffer lookup has unique id, can break out early
3277 found = False
3278 for buffer_index_type, buffer_index, buffer_obj in iter_buffers():
3279 if buffer_obj.same_as(buffer):
3280 found = True
3281 break
3282
3283 assert found, f"Could not find buffer '{buffer.name}' in block '{block_name}'"
3284
3285 elif isinstance(buffer, tuple):
3286 buffer_index_type, buffer_index = buffer
3287 assert buffer_index_type in ["read", "write"], (
3288 f"Invalid buffer_index_type. "
3289 f"Expected 'read' or 'write', "
3290 f"but received {buffer_index_type}"
3291 )
3292 buffer_list = block_obj.reads if buffer_index_type == "read" else block_obj.writes
3293 assert 0 <= buffer_index < len(buffer_list), (
3294 f"Invalid buffer_index {buffer_index}. "
3295 f"Block {block_name} has only "
3296 f"{len(buffer_list)} {buffer_index_type} buffers."
3297 )
3298 buffer_obj = buffer_list[buffer_index].buffer
3299
3300 else:
3301 raise TypeError(f"Invalid type for argument 'buffer': {type(buffer)}")

Callers 7

cache_readMethod · 0.95
cache_writeMethod · 0.95
cache_inplaceMethod · 0.95
reindexMethod · 0.95
set_scopeMethod · 0.95
transform_layoutMethod · 0.95
set_axis_separatorMethod · 0.95

Calls 3

getMethod · 0.95
itemsMethod · 0.45
same_asMethod · 0.45

Tested by

no test coverage detected