MCPcopy Create free account
hub / github.com/KnowingNothing/MatmulTutorial / SimpleTileScheduler

Class SimpleTileScheduler

cutlass.py/tile_scheduler.py:386–439  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

384
385@dataclass
386class SimpleTileScheduler:
387 dev_coord: DeviceCoord
388 cluster_m: int
389 cluster_n: int
390 block_m: int
391 block_n: int
392 linear_idx: int = 0
393 m_blocks: int = 0
394 n_blocks: int = 0
395
396 @dataclass
397 class WorkInfo:
398 m_idx: int
399 n_idx: int
400 valid: bool
401
402 def init(self, M, N):
403 self.linear_idx = (
404 dev_coord.blockIdx.x + dev_coord.blockIdx.y * dev_coord.gridDim.x
405 )
406 self.get_blocks_m_n(M, N)
407
408 def get_current_work_info(self):
409 m_idx, n_idx = self.get_current_m_n_idx()
410 return SimpleTileScheduler.WorkInfo(
411 m_idx, n_idx, self.linear_idx < self.m_blocks * self.n_blocks
412 )
413
414 def advance(self, number=1):
415 self.linear_idx += number * self.dev_coord.gridDim.x * self.dev_coord.gridDim.y
416
417 def get_current_m_n_idx(self):
418 div_cluster_x = self.linear_idx // self.cluster_m
419 mod_cluster_x = self.linear_idx % self.cluster_m
420 div_cluster_xy = div_cluster_x // self.cluster_n
421 mod_cluster_xy = div_cluster_x % self.cluster_n
422 cluster_per_row = self.n_blocks // self.cluster_n
423 cluster_row = div_cluster_xy // cluster_per_row
424 cluster_col = div_cluster_xy % cluster_per_row
425 m_idx = cluster_row * self.cluster_m + mod_cluster_x
426 n_idx = cluster_col * self.cluster_n + mod_cluster_xy
427 return (m_idx, n_idx)
428
429 def get_blocks_m_n(self, M, N):
430 self.m_blocks = (
431 ((M + self.block_m - 1) // self.block_m + self.cluster_m - 1)
432 // self.cluster_m
433 * self.cluster_m
434 )
435 self.n_blocks = (
436 ((N + self.block_n - 1) // self.block_n + self.cluster_n - 1)
437 // self.cluster_n
438 * self.cluster_n
439 )
440
441
442if __name__ == "__main__":

Callers 1

tile_scheduler.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected