Split A into (inter, intra) for the target scope kind.
(A: ActiveSet, scope_kind: str)
| 347 | |
| 348 | |
| 349 | def scope_switch(A: ActiveSet, scope_kind: str) -> Split: |
| 350 | """Split A into (inter, intra) for the target scope kind.""" |
| 351 | if scope_kind == THREAD: |
| 352 | return Split(inter={"laneid": A.laneid, "warpid": A.warpid, "cta_id": A.cta_id}, intra={}) |
| 353 | if scope_kind == WARP: |
| 354 | return Split(inter={"warpid": A.warpid, "cta_id": A.cta_id}, intra={"laneid": A.laneid}) |
| 355 | if scope_kind == CTA: |
| 356 | return Split(inter={"cta_id": A.cta_id}, intra={"laneid": A.laneid, "warpid": A.warpid}) |
| 357 | if scope_kind == CLUSTER: |
| 358 | return Split(inter={}, intra={"laneid": A.laneid, "warpid": A.warpid, "cta_id": A.cta_id}) |
| 359 | if scope_kind == WARPGROUP: |
| 360 | factored = _factor_warpid(A.warpid) |
| 361 | if factored is None: |
| 362 | raise ExecContextError( |
| 363 | "scope_switch(warpgroup) failed: warpid axis" |
| 364 | f" (extent={A.warpid.extent}, offset={A.warpid.offset})" |
| 365 | " crosses warpgroup boundary and is not aligned" |
| 366 | ) |
| 367 | wid_in_wg, wgid = factored |
| 368 | return Split( |
| 369 | inter={"wgid": wgid, "cta_id": A.cta_id}, |
| 370 | intra={"laneid": A.laneid, "wid_in_wg": wid_in_wg}, |
| 371 | ) |
| 372 | if scope_kind == KERNEL: |
| 373 | return Split(inter={"laneid": A.laneid, "warpid": A.warpid, "cta_id": A.cta_id}, intra={}) |
| 374 | raise ValueError(f"unknown scope kind: {scope_kind!r}") |
| 375 | |
| 376 | |
| 377 | @dataclass(frozen=True) |
searching dependent graphs…