| 384 | |
| 385 | @dataclass |
| 386 | class 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 | |
| 442 | if __name__ == "__main__": |