custom
code
sovereign-compute
File size: 5,079 Bytes
e92f76f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
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
    """
    # gfx942 ds_read_b128 lane groups (from AMD documentation)
    DS_READ_B128_GROUPS = [
        list(range(0, 4)) + list(range(20, 24)), # G0: 0-3 + 20-23
        list(range(4, 8)) + list(range(16, 20)), # G1: 4-7 + 16-19
        list(range(8, 12)) + list(range(28, 32)), # G2: 8-11 + 28-31
        list(range(12, 16)) + list(range(24, 28)), # G3: 12-15 + 24-27
        list(range(32, 36)) + list(range(52, 56)), # G4: 32-35 + 52-55
        list(range(36, 40)) + list(range(48, 52)), # G5: 36-39 + 48-51
        list(range(40, 44)) + list(range(60, 64)), # G6: 40-43 + 60-63
        list(range(44, 48)) + list(range(56, 60)) # G7: 44-47 + 56-59
    ]
    
    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)
        # Convert FP16 column to bank-word column (2 FP16 = 1 bank word)
        col_bank_word = col_fp16 // 2
        # Physical address in bytes: 4 * (row * stride_words + col_bank_word)
        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): # q = 0,1,2,3 for the 4 dwords in b128
                bank_to_address = {} # Maps bank -> first address seen at this bank/q
                for lane in group:
                    addr = lds_address(lane, stride_words)
                    if not is_16byte_aligned(addr):
                        return True # Alignment violation
                    bank_word = addr // 4 # Convert byte address to bank-word index
                    bank = (bank_word + q) % 32 # Bank for this dword phase
                    if bank in bank_to_address:
                        # Conflict: different addresses mapping to same bank in same phase
                        if bank_to_address[bank] != addr + 4 * q:
                            return True
                    else:
                        bank_to_address[bank] = addr
        return False
    
    # Search for minimal padding
    for P in range(max_padding + 1):
        stride_words = logical_row_words + P
        if not has_conflict(stride_words):
            return P
    return -1 # No solution found

# EXAMPLE USAGE FOR gfx942 v_mfma_f32_16x16x16f16:
if __name__ == "__main__":
    # Lane-to-fragment map for A operand in v_mfma_f32_16x16x16f16
    # (From previous fragment: 8 FP16 elements as [2 rows × 4 columns])
    def a_fragment_map(lane_id: int) -> tuple[int, int]:
        m_in_tile = 2 * (lane_id // 32) + (lane_id % 2) # Row start [0,14] step 2
        k_in_tile = 4 * (lane_id % 16) # Column start [0,60] step 4
        # For ds_read_b128, we read 8 consecutive FP16 elements (4 bank words)
        # Starting at (m_in_tile, k_in_tile)
        return (m_in_tile, k_in_tile) # Returns logical (row, col) in FP16 elements
    
    # For FP16 row with 64 elements (typical MFMA K dimension)
    logical_row_words = 64 * 2 // 4 # 32 bank words
    
    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")
        
    # To verify, plug padding into your kernel's LDS layout:
    # .align 256
    # .lgs A_tile: .skip ((64 + padding*2) * 16 * 2) ; 64 rows, (64+2P) cols, FP16