| def calculate_ds_read_b128_padding( |
| logical_row_words: int, |
| lane_to_fragment_map: callable, |
| max_padding: int = 16 |
| ) -> int: |
| """ |
| Calculate minimal LDS padding (in 32-bit bank words) to eliminate ds_read_b128 conflicts |
| for gfx942 (CDNA 3) hardware. |
| |
| Args: |
| logical_row_words: Logical row width in 32-bit words (W = ceil(K*2/4) for FP16) |
| lane_to_fragment_map: Function(lane_id) -> (row, col) in logical LDS coordinates |
| where col is in FP16 elements (not bank words) |
| max_padding: Maximum padding to search (bank words) |
| |
| Returns: |
| Minimal padding P (bank words) that yields conflict-free ds_read_b128 |
| Returns -1 if no solution found within max_padding |
| |
| Hardware constraints (gfx942): |
| - 32 LDS banks, 4 bytes/bank |
| - ds_read_b128 groups: 8 specific non-contiguous 8-lane groups |
| - Each lane reads 4 consecutive 32-bit words (q=0,1,2,3) |
| - 16-byte alignment required for ds_read_b128 source address |
| """ |
| |
| DS_READ_B128_GROUPS = [ |
| list(range(0, 4)) + list(range(20, 24)), |
| list(range(4, 8)) + list(range(16, 20)), |
| list(range(8, 12)) + list(range(28, 32)), |
| list(range(12, 16)) + list(range(24, 28)), |
| list(range(32, 36)) + list(range(52, 56)), |
| list(range(36, 40)) + list(range(48, 52)), |
| list(range(40, 44)) + list(range(60, 64)), |
| list(range(44, 48)) + list(range(56, 60)) |
| ] |
| |
| def lds_address(lane_id: int, stride_words: int) -> int: |
| """ |
| Calculate LDS byte address for a lane's ds_read_b128 source. |
| Assumes lane_to_fragment_map returns (row, col) in logical FP16 elements. |
| """ |
| row, col_fp16 = lane_to_fragment_map(lane_id) |
| |
| col_bank_word = col_fp16 // 2 |
| |
| return 4 * (row * stride_words + col_bank_word) |
| |
| def is_16byte_aligned(address: int) -> bool: |
| """Check if address is 16-byte aligned (required for ds_read_b128)""" |
| return address % 16 == 0 |
| |
| def has_conflict(stride_words: int) -> bool: |
| """Check if given stride causes any ds_read_b128 bank conflict""" |
| for group in DS_READ_B128_GROUPS: |
| for q in range(4): |
| bank_to_address = {} |
| for lane in group: |
| addr = lds_address(lane, stride_words) |
| if not is_16byte_aligned(addr): |
| return True |
| bank_word = addr // 4 |
| bank = (bank_word + q) % 32 |
| if bank in bank_to_address: |
| |
| if bank_to_address[bank] != addr + 4 * q: |
| return True |
| else: |
| bank_to_address[bank] = addr |
| return False |
| |
| |
| for P in range(max_padding + 1): |
| stride_words = logical_row_words + P |
| if not has_conflict(stride_words): |
| return P |
| return -1 |
|
|
| |
| if __name__ == "__main__": |
| |
| |
| def a_fragment_map(lane_id: int) -> tuple[int, int]: |
| m_in_tile = 2 * (lane_id // 32) + (lane_id % 2) |
| k_in_tile = 4 * (lane_id % 16) |
| |
| |
| return (m_in_tile, k_in_tile) |
| |
| |
| logical_row_words = 64 * 2 // 4 |
| |
| padding = calculate_ds_read_b128_padding( |
| logical_row_words=logical_row_words, |
| lane_to_fragment_map=a_fragment_map, |
| max_padding=16 |
| ) |
| |
| if padding >= 0: |
| print(f"Minimal padding: {padding} bank words") |
| print(f" = {padding * 4} bytes") |
| print(f" = {padding * 2} FP16 elements") |
| print(f"Physical row stride: {logical_row_words + padding} bank words") |
| else: |
| print("No conflict-free padding found within search range") |
| |
| |
| |
| |