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
Possible Solution
No response
Logs/Context/Additional Information
No response
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 totop).This kernel is an example (I have never seen it finish compiling successfully on
trn1ortrn2, and in fact ontrn1.2xlargeit runs out of memory and appears to crash the machine):This kernel, on the other hand, does compile in a reasonable amount of time on a trn1.32xlarge:
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:
I have not seen the last kernel here compile on a
trn1instance. It does manage to compile and run on atrn2.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:
Regression Issue
Possible Solution
No response
Logs/Context/Additional Information
No response