Skip to content

Extremely long compile times and high compiler memory usage for some large NKI kernels #1374

Description

@natetyoung

Describe the bug

I am working on some kernels for large fused matrix multiplications, and have been encountering long compile times (~hours) and high memory usage (40GB+), specifically in walrus_driver (according to top).
This kernel is an example (I have never seen it finish compiling successfully on trn1 or trn2, and in fact on trn1.2xlarge it runs out of memory and appears to crash the machine):

@nki.jit
def mm2(At, B, D, m_factor=1):
    E = nl.ndarray((min(1311 * (25 // m_factor), 32 * 1024), 4 * 1024), dtype=nl.float32, buffer=nl.shared_hbm, name="E")
    for m0 in nl.static_range(25 // m_factor):
        E_write_1 = nl.zeros((120, 4096, 11), dtype=nl.float32, buffer=nl.sbuf, name="E_write_1") # time 0
        for n0 in nl.static_range(8):
            Ct_write_0 = nl.zeros((128, 1311, 16), dtype=nl.float32, buffer=nl.sbuf, name="Ct_write_0") # Fused root time 1
            for k0 in nl.static_range(32):
                actual_m_start = m0 * 1311
                actual_m_end = min(actual_m_start + 1311, 32 * 1024)
                m_elems = actual_m_end - actual_m_start
                At_read_0 = nl.load(At[k0 * 128:(k0 + 1) * 128, actual_m_start:actual_m_end]) # time 2
                for n1 in nl.static_range(16):
                    actual_n = n0 * 2048 + n1 * 128
                    B_read_0 = nl.load(B[k0 * 128:(k0 + 1) * 128, actual_n:actual_n + 128]) # time 3
                    for m1 in nl.static_range(3):
                        ct_m_idx_start = m1 * 512
                        ct_m_idx_end = min(ct_m_idx_start + 512, 1311)
                        ct_m_elems = ct_m_idx_end - ct_m_idx_start
                        Ct_write_0_psum = nl.ndarray((128, ct_m_elems), dtype=nl.float32, buffer=nl.psum)
                        nisa.tensor_copy(Ct_write_0_psum, Ct_write_0[:, ct_m_idx_start:ct_m_idx_end, n1])
                        nisa.nc_matmul(Ct_write_0_psum, B_read_0, At_read_0[:, ct_m_idx_start:ct_m_idx_end])
                        nisa.tensor_copy(Ct_write_0[:, ct_m_idx_start:ct_m_idx_end, n1], Ct_write_0_psum)
                    # free(B_read_0) # time 4
                # free(At_read_0) # time 5
            Ct_read_1 = Ct_write_0 # Fused alias time 7
            for n1 in nl.static_range(16):
                actual_n = n0 * 2048 + n1 * 128
                for l0 in nl.static_range(32):
                    D_read_1 = nl.load(D[actual_n:actual_n + 128, l0 * 128:(l0 + 1) * 128]) # time 8
                    for m1 in nl.static_range(11):
                        ct_m_idx_start = m1 * 120
                        ct_m_idx_end = min(ct_m_idx_start + 120, 1311)
                        ct_m_elems = ct_m_idx_end - ct_m_idx_start
                        E_write_1_psum = nl.ndarray((ct_m_elems, 128), dtype=nl.float32, buffer=nl.psum)
                        nisa.tensor_copy(E_write_1_psum, E_write_1[:ct_m_elems, l0 * 128:(l0 + 1) * 128, m1])
                        nisa.nc_matmul(E_write_1_psum, Ct_read_1[:, ct_m_idx_start:ct_m_idx_end, n1], D_read_1)
                        nisa.tensor_copy(E_write_1[:ct_m_elems, l0 * 128:(l0 + 1) * 128, m1], E_write_1_psum)
                    # free(D_read_1) # time 9
            # free(C_read_1) # time 10
        for m1 in nl.static_range(11):
            actual_m_start = m0 * 1311 + m1 * 120
            actual_m_end = min(actual_m_start + 120, E.shape[0])
            m_elems = actual_m_end - actual_m_start
            nl.store(E[actual_m_start:actual_m_end, :], E_write_1[:m_elems, :, m1])
        # free(E_write_1) # time 11
    return E

This kernel, on the other hand, does compile in a reasonable amount of time on a trn1.32xlarge:

M0_TOTAL = 256
K0, P0, L0 = 32, 128, 8

SHARD = int(sys.argv[1]) if len(sys.argv) > 1 else 0
NSHARDS = int(sys.argv[2]) if len(sys.argv) > 2 else 1
M0_PER_SHARD = M0_TOTAL // NSHARDS


@nki.jit
def compute_einsum(A, B, C, D):
    E = nl.ndarray((M0_PER_SHARD * 128, 4 * 1024), dtype = nl.float32, buffer = nl.shared_hbm)
    for m0 in nl.affine_range(M0_PER_SHARD):
        E_write_chunks = []
        for _ in range(L0):
            E_write_chunks.append(nl.zeros((128, 512), dtype = nl.float32, buffer = nl.psum))
        A_chunks = []
        for k0 in range(K0):
            A_chunks.append(nl.load_transpose2d(A[m0 * 128 : (m0 + 1) * 128, k0 * 128 : (k0 + 1) * 128]))
        for p0 in nl.affine_range(P0):
            C_write_0 = nl.zeros((128, 128), dtype = nl.float32, buffer = nl.psum)
            for k0 in nl.affine_range(K0):
                B_read_0 = nl.load(B[k0 * 128 : (k0 + 1) * 128, p0 * 128 : (p0 + 1) * 128])
                nisa.nc_matmul(C_write_0, B_read_0, A_chunks[k0]) # 128-wide contraction, 128-wide middle-dim chunk
            C_read_1 = nl.zeros(C_write_0.shape, dtype = nl.float32)
            nisa.tensor_copy(C_read_1, C_write_0)
            for l0 in nl.affine_range(L0):
                D_subchunk = nl.load(D[p0 * 128 : (p0 + 1) * 128, l0 * 512 : (l0 + 1) * 512])
                nisa.nc_matmul(E_write_chunks[l0], C_read_1, D_subchunk)
        E_sbuf = nl.zeros((128, 4 * 1024), dtype = nl.float32)
        for l0 in range(L0):
            nisa.tensor_copy(E_sbuf[:, l0 * 512 : (l0 + 1) * 512], E_write_chunks[l0])
        nl.store(E[m0 * 128 : (m0 + 1) * 128, : ], E_sbuf)
    return E

It is smaller when unrolled since it is parallel, so I figured that was the difference, but unfortunately this kernel, similar to the first one but made smaller using dynamic loops, does not compile in a reasonable amount of time:

@nki.jit
def mm2(At, B, D):
    E = nl.ndarray((32 * 1024, 4 * 1024), dtype=nl.float32, buffer=nl.shared_hbm, name="E")
    m_start_idx = nl.zeros((1, 1), dtype=nl.int32, buffer=nl.sbuf, name="m_start_idx")
    for m_ in nl.dynamic_range(24): # 24 * 1311 + 1304
        E_write_1 = nl.zeros((120, 4096, 11), dtype=nl.float32, buffer=nl.sbuf, name="E_write_1") # time 0
        n_start_idx = nl.zeros((1, 1), dtype=nl.int32, buffer=nl.sbuf, name="n_start_idx")
        for n0 in nl.dynamic_range(8):
            Ct_write_0 = nl.zeros((128, 1311, 16), dtype=nl.float32, buffer=nl.sbuf, name="Ct_write_0") # Fused root time 1
            for k0 in nl.static_range(32):
                At_read_access_pattern = At.ap(
                    pattern=[[32*1024, 128], [1, 1311]],
                    offset = k0 * 128 * 32 * 1024,
                    scalar_offset=m_start_idx,
                    indirect_dim=1
                )
                At_read_0 = nl.load(At_read_access_pattern)
                for n1 in nl.static_range(16):
                    B_ap = B.ap(
                        pattern=[[16*1024, 128], [1, 128]],
                        offset = n1 * 128,
                        scalar_offset=n_start_idx,
                        indirect_dim=1
                    )
                    B_read_0 = nl.load(B_ap) # time 3
                    for m1 in nl.static_range(3):
                        ct_m_idx_start = m1 * 512
                        ct_m_idx_end = min(ct_m_idx_start + 512, 1311)
                        ct_m_elems = ct_m_idx_end - ct_m_idx_start
                        Ct_write_0_psum = nl.ndarray((128, ct_m_elems), dtype=nl.float32, buffer=nl.psum, name=f"Ct_write_0_psum_{m1}_{n1}_{k0}")
                        nisa.tensor_copy(Ct_write_0_psum, Ct_write_0[:, ct_m_idx_start:ct_m_idx_end, n1])
                        nisa.nc_matmul(Ct_write_0_psum, B_read_0, At_read_0[:, ct_m_idx_start:ct_m_idx_end])
                        nisa.tensor_copy(Ct_write_0[:, ct_m_idx_start:ct_m_idx_end, n1], Ct_write_0_psum)
                    # free(B_read_0) # time 4
                # free(At_read_0) # time 5
            Ct_read_1 = Ct_write_0 # Fused alias time 7
            for n1 in nl.static_range(16):
                for l0 in nl.static_range(32):
                    D_ap = D.ap(
                        pattern=[[4*1024, 128], [1, 128]],
                        offset = l0 * 128,
                        scalar_offset=n_start_idx,
                        indirect_dim=0
                    )
                    D_read_1 = nl.load(D_ap) # time 8
                    for m1 in nl.static_range(11):
                        ct_m_idx_start = m1 * 120
                        ct_m_idx_end = min(ct_m_idx_start + 120, 1311)
                        ct_m_elems = ct_m_idx_end - ct_m_idx_start
                        E_write_1_psum = nl.ndarray((ct_m_elems, 128), dtype=nl.float32, buffer=nl.psum, name=f"E_write_1_psum_{m1}_{n1}_{l0}")
                        nisa.tensor_copy(E_write_1_psum, E_write_1[:ct_m_elems, l0 * 128:(l0 + 1) * 128, m1])
                        nisa.nc_matmul(E_write_1_psum, Ct_read_1[:, ct_m_idx_start:ct_m_idx_end, n1], D_read_1)
                        nisa.tensor_copy(E_write_1[:ct_m_elems, l0 * 128:(l0 + 1) * 128, m1], E_write_1_psum)
                    # free(D_read_1) # time 9
            # free(C_read_1) # time 10
            nisa.tensor_scalar(dst=n_start_idx, data=n_start_idx, op0=nl.add, operand0=2048)

        for m1 in nl.static_range(11):
            E_ap = E.ap(
                pattern=[[4*1024, 120], [1, 4096]],
                offset = m1 * 120 * 4 * 1024,
                scalar_offset=m_start_idx,
                indirect_dim=0
            )
            nl.store(E_ap, E_write_1[:, :, m1])
        # free(E_write_1) # time 11
        nisa.tensor_scalar(dst=m_start_idx, data=m_start_idx, op0=nl.add, operand0=1311)
    m0 = 24
    E_write_1 = nl.zeros((120, 4096, 11), dtype=nl.float32, buffer=nl.sbuf, name="E_write_1_ragged") # time 0
    n_start_idx = nl.zeros((1, 1), dtype=nl.int32, buffer=nl.sbuf, name="n_start_idx_ragged")
    for n0 in nl.dynamic_range(8):
        Ct_write_0 = nl.zeros((128, 1311, 16), dtype=nl.float32, buffer=nl.sbuf, name="Ct_write_0_ragged") # Fused root time 1
        for k0 in nl.static_range(32):
            actual_m_start = m0 * 1311
            actual_m_end = min(actual_m_start + 1311, 32 * 1024)
            m_elems = actual_m_end - actual_m_start
            At_read_0 = nl.load(At[k0 * 128:(k0 + 1) * 128, actual_m_start:actual_m_end]) # time 2
            for n1 in nl.static_range(16):
                B_ap = B.ap(
                    pattern=[[16*1024, 128], [1, 128]],
                    offset = n1 * 128,
                    scalar_offset=n_start_idx,
                    indirect_dim=1
                )
                B_read_0 = nl.load(B_ap) # time 3
                for m1 in nl.static_range(3):
                    ct_m_idx_start = m1 * 512
                    ct_m_idx_end = min(ct_m_idx_start + 512, m_elems)
                    ct_m_elems = ct_m_idx_end - ct_m_idx_start
                    Ct_write_0_psum_r = nl.ndarray((128, ct_m_elems), dtype=nl.float32, buffer=nl.psum, name=f"Ct_write_0_psum_ragged_{m1}_{n1}_{k0}")
                    nisa.tensor_copy(Ct_write_0_psum_r, Ct_write_0[:, ct_m_idx_start:ct_m_idx_end, n1])
                    nisa.nc_matmul(Ct_write_0_psum_r, B_read_0, At_read_0[:, ct_m_idx_start:ct_m_idx_end])
                    nisa.tensor_copy(Ct_write_0[:, ct_m_idx_start:ct_m_idx_end, n1], Ct_write_0_psum_r)
                # free(B_read_0) # time 4
            # free(At_read_0) # time 5
        Ct_read_1 = Ct_write_0 # Fused alias time 7
        for n1 in nl.static_range(16):
            for l0 in nl.static_range(32):
                D_ap = D.ap(
                    pattern=[[4*1024, 128], [1, 128]],
                    offset = l0 * 128,
                    scalar_offset=n_start_idx,
                    indirect_dim=0
                )
                D_read_1 = nl.load(D_ap) # time 8
                for m1 in nl.static_range(11):
                    ct_m_idx_start = m1 * 120
                    ct_m_idx_end = min(ct_m_idx_start + 120, 1311)
                    ct_m_elems = ct_m_idx_end - ct_m_idx_start
                    E_write_1_psum = nl.ndarray((ct_m_elems, 128), dtype=nl.float32, buffer=nl.psum, name=f"E_write_1_psum_ragged_{m1}_{n1}_{l0}")
                    nisa.tensor_copy(E_write_1_psum, E_write_1[:ct_m_elems, l0 * 128:(l0 + 1) * 128, m1])
                    nisa.nc_matmul(E_write_1_psum, Ct_read_1[:, ct_m_idx_start:ct_m_idx_end, n1], D_read_1)
                    nisa.tensor_copy(E_write_1[:ct_m_elems, l0 * 128:(l0 + 1) * 128, m1], E_write_1_psum)
                # free(D_read_1) # time 9
        # free(C_read_1) # time 10
        nisa.tensor_scalar(dst=n_start_idx, data=n_start_idx, op0=nl.add, operand0=2048)
    for m1 in nl.static_range(11):
        actual_m_start = m0 * 1311 + m1 * 120
        actual_m_end = min(actual_m_start + 120, E.shape[0])
        m_elems = actual_m_end - actual_m_start
        nl.store(E[actual_m_start:actual_m_end, :], E_write_1[:m_elems, :, m1])
    # free(E_write_1) # time 11
    return E

I have not seen the last kernel here compile on a trn1 instance. It does manage to compile and run on a trn2.3xlarge, but it takes quite a while (30min+ IIRC) and 30GB+ to do so.

Model Name

N/A

Describe the workload type

matmuls

Instance Type

trn1.32xlarge, trn2.3xlarge

Release version

aws-neuronx-collectives/unknown,now 2.33.10.0-068180c7a amd64 [installed]
aws-neuronx-dkms/unknown,now 2.29.0.0 all [installed]
aws-neuronx-oci-hook/unknown,now 2.17.30.0 amd64 [installed]
aws-neuronx-runtime-lib/unknown,now 2.33.10.0-3dcef56f0 amd64 [installed]
aws-neuronx-tools/unknown,now 2.31.13.0-a9e473f33 amd64 [installed]
libneuronxla 2.2.17544.0+fb9962bf
neuron-agentic-development 1.2
neuronx-cc 2.26.6360.0+6f180f47
neuronx-distributed 0.19.28492+435aae2b
torch 2.9.1
torch-neuronx 2.9.0.2.15.32035+de43f57c
torch-xla 2.9.0
torchvision 0.24.1

Reproduction Steps

Large kernels above can be tested using this testbench:

from mm2_kernel import mm2
import numpy as np

# Generate input tensors.
At = np.ones((4 * 1024, 32 * 1024), dtype=np.float32)
B = np.ones((4 * 1024, 16 * 1024), dtype=np.float32)
D = np.ones((16 * 1024, 4 * 1024), dtype=np.float32)

E = np.zeros((32 * 1024, 4 * 1024), dtype=np.float32)

E = mm2(At, B, D)

E_baseline = np.dot(np.dot(At.T, B), D)

# Print the result.
print(E.shape)
print(E_baseline.shape)

print("Difference norm:")
print(np.linalg.norm(E - E_baseline))

print("Max difference:")
print(np.max(np.abs(E - E_baseline)))

Regression Issue

  • Select this option if this issue appears to be a regression.

Possible Solution

No response

Logs/Context/Additional Information

No response

Metadata

Metadata

Assignees

No one assigned

    Labels

    Type

    No type

    Projects

    No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions