-
Notifications
You must be signed in to change notification settings - Fork 33
[FEATURE] NKI interpreter #206
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Merged
Merged
Changes from all commits
Commits
Show all changes
106 commits
Select commit
Hold shift + click to select a range
e6e0641
Support NKI
Jokeren c9b2501
Update
Jokeren 6fe57e6
Update
Jokeren 86d1d6f
Update
Jokeren 7aa3d91
Update
Jokeren 4f8f616
Update
Jokeren 9962f53
Update
Jokeren 25c8ceb
Update
Jokeren cb07831
Update
Jokeren 1810284
Update
Jokeren e6a6f91
Update
Jokeren 8406f66
Update
Jokeren b63e74c
Update
Jokeren 241530b
Update
Jokeren d28a17e
Update
Jokeren bc0873a
Update
Jokeren 4b6827d
UPdate
Jokeren 62ade2d
Update
Jokeren d9d7bf7
Update
Jokeren 998a1a5
Update
Jokeren 2a50a04
Update
Jokeren 8e49f9d
Update
Jokeren 45596d3
Update
Jokeren a54ac94
Update
Jokeren f1231fe
Update
Jokeren b1f658a
Update comments
Jokeren e7e2dd2
stuff
b3ee8d9
triton-viz should support cpu to allow development on non-trainium node
latentCall145 454ac73
test kernels to check basic tracing
latentCall145 65e861d
more flexible advanced indexing (just make arange slices an arange)
latentCall145 f6b99ae
remove need for SPMD nki kernel call when tracing
latentCall145 4f863fa
remove
latentCall145 2932bb9
added calibrate and matmul and show code
gujialiang123 cb66323
add some ops
latentCall145 7bd6bea
nki kernels to demo interpreter
latentCall145 9afefc6
fix misc tests
latentCall145 95ef0b6
prepare for nki tracing
latentCall145 2636088
uv
latentCall145 0644e64
fix simple ndarray indexing (and also add mgrid)
latentCall145 2037047
masked load stuff (masked store still doesn't work right)
latentCall145 8f619e8
attempt to connect nki interpreter to triton-viz (still failing)
latentCall145 5adb5e0
hardcode offsets and callbacks for NKI matmul kernel
latentCall145 ce4fdd4
Merge remote-tracking branch 'thaihoa/thaihoa/mm' into nki/merge_inte…
mark14wu f091fd4
Remove uv.lock
mark14wu 9e2d495
feat(flip): add Flip op tracing, 3D flip view, hover value API; fix e…
gujialiang123 4fa8660
api: align register_op_callback signatures across Client, Sanitizer, …
gujialiang123 2b2968d
fix(nki): resolve mypy errors (value setter order, unary op calls); p…
gujialiang123 82f629e
use import for nki test
latentCall145 86345a1
masked load fast path; use max dtype value instead of 6 for out of bo…
latentCall145 85972d8
fix masked store
latentCall145 30d4e5d
clone array so load doesn't write on input array
latentCall145 f696409
make tests prettier
latentCall145 9d96499
make all tests pass
latentCall145 9066d3b
modify undefined values for other dtypes in masked load; lint
latentCall145 696587d
remove ._value, just use .data; lint
latentCall145 0fb3dbc
make nki tracing work again
latentCall145 a225146
don't immediately exit web server when share=false
latentCall145 15196d0
Merge branch 'thaihoa/local-viz-fix' into nki/merge-interpreter-working
latentCall145 6571b71
feat(flow): embed memory-flow badges in Load/Store/Dot; add Flow Diag…
gujialiang123 b959a12
patch nki pt.1
latentCall145 ed0b309
nki patch pt.2 (new tracer callbacks)
latentCall145 2e1943e
nki patch pt.3 (nki-specific changes to allow nki tracing)
latentCall145 301aa66
nki patch pt.4 (update matmul example)
latentCall145 185464f
Merge branch 'jialiang/work' into nki/merge-interpreter-working
latentCall145 4826219
Merge branch 'main' into nki/merge-interpreter-working
latentCall145 e9ecf82
remove NLSlice (was never used)
latentCall145 b74f471
fix offsets/mask shape mismatch bug in visualizing masked stores
latentCall145 1910b78
unify DSL patching pt 1. (make nki matmul + load_store examples work)
latentCall145 a120447
remove some files
latentCall145 b5797ef
normal args and kwargs
latentCall145 4e0f02a
allocate instead of array
latentCall145 a913fda
reorder callbacks
latentCall145 0c92fbd
more examples
latentCall145 0ee0896
Merge branch 'main' into nki/patch-dsl
mark14wu 0fc84ec
apply ruff-format; move normalization to draw; stop tracking scripts/
gujialiang123 aceeae2
Triton/NKI Load/Store unify pt 1. (push to same callback - not working)
latentCall145 10271ae
Triton/NKI Load/Store unify pt. 2 (add adapters to standardize callba…
latentCall145 9f37b24
whoops forgot to add allocate adapter
latentCall145 cd623f5
be gone
latentCall145 6470e5b
add adapter tests
latentCall145 187c347
Merge branch 'nki/patch-dsl' of github.com:Deep-Learning-Profiling-To…
latentCall145 71b8e71
revert unwanted changes from main
latentCall145 cc9a787
patch backend less bad
latentCall145 f94d1be
remove unhelpful comments + show expected transformed code
latentCall145 675abab
debloat nki offsets code
latentCall145 7c3bec2
fix matmul visualization error
latentCall145 86e476a
Merge branch 'main' into nki/patch-dsl
latentCall145 229d4a9
lint
latentCall145 5bf74c8
Merge branch 'main' into nki/patch-dsl
latentCall145 096122c
Merge branch 'main' into nki/patch-dsl
mark14wu 54559f1
rename nki_masked_load -> masked_load since it's also used for triton
latentCall145 f9c72e3
make AWS neuron an optional feature of triton viz (and include CI)
latentCall145 3a00512
Merge branch 'nki/patch-dsl' of github.com:Deep-Learning-Profiling-To…
latentCall145 0f6664a
guard nki test import if NKI not installed
latentCall145 b45cd5e
update import
latentCall145 6c01d62
[CI] fix aws neuron ci
latentCall145 fa814b5
python>=3.10 needed (also 3.9 is EOL)
latentCall145 de93383
[CI] make AWS tests reuse torch + triton installations from build job
latentCall145 50a0177
make nki set_grid_dim consistent with triton api
latentCall145 24c9abe
Merge branch 'main' into nki/patch-dsl
latentCall145 7260fce
Merge branch 'main' into nki/patch-dsl
mark14wu 76f8999
extra nki kernels to viz
latentCall145 3dc8c02
typo
latentCall145 0817087
Merge branch 'main' into nki/patch-dsl
latentCall145 0eb3f22
Merge branch 'main' into nki/patch-dsl
latentCall145 232e87b
Merge branch 'main' into nki/patch-dsl
mark14wu File filter
Filter by extension
Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
There are no files selected for viewing
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,81 @@ | ||
| import torch | ||
| import triton | ||
| import triton.language as tl | ||
|
|
||
| import triton_viz | ||
| from triton_viz.clients import Tracer | ||
|
|
||
|
|
||
| @triton_viz.trace(clients=Tracer()) | ||
| @triton.jit | ||
| def flip_1d_kernel(x_ptr, y_ptr, n, BLOCK: tl.constexpr): | ||
| pid = tl.program_id(0) | ||
| offs = pid * BLOCK + tl.arange(0, BLOCK) | ||
| mask = offs < n | ||
| x = tl.load(x_ptr + offs, mask=mask) | ||
| _ = tl.flip(x, dim=0) | ||
| rev_offs = (n - 1) - offs | ||
| tl.store(y_ptr + rev_offs, x, mask=mask) | ||
|
|
||
|
|
||
| @triton_viz.trace(clients=Tracer()) | ||
| @triton.jit | ||
| def flip_2d_kernel( | ||
| x_ptr, | ||
| y_ptr, | ||
| H, | ||
| W, | ||
| stride_h, | ||
| stride_w, | ||
| dim: tl.constexpr, | ||
| BLOCK_M: tl.constexpr, | ||
| BLOCK_N: tl.constexpr, | ||
| ): | ||
| pid_m = tl.program_id(0) | ||
| pid_n = tl.program_id(1) | ||
|
|
||
| offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) | ||
| offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) | ||
|
|
||
| mm = offs_m[:, None] | ||
| nn = offs_n[None, :] | ||
| mask = (mm < H) & (nn < W) | ||
|
|
||
| ptrs = x_ptr + mm * stride_h + nn * stride_w | ||
| vals = tl.load(ptrs, mask=mask, other=0) | ||
|
|
||
| _ = tl.flip(vals, dim=(0 if dim == 0 else 1)) | ||
|
|
||
| if dim == 0: | ||
| rm = (H - 1) - mm | ||
| rn = nn | ||
| else: | ||
| rm = mm | ||
| rn = (W - 1) - nn | ||
|
|
||
| out_ptrs = y_ptr + rm * stride_h + rn * stride_w | ||
| tl.store(out_ptrs, vals, mask=mask) | ||
|
|
||
|
|
||
| def run_1d(n: int = 256): | ||
| x = torch.arange(n, dtype=torch.int32) | ||
| y = torch.empty_like(x) | ||
| grid = (triton.cdiv(n, 128),) | ||
| flip_1d_kernel[grid](x, y, n, BLOCK=128) | ||
| assert torch.equal(y, torch.flip(x, dims=[0])) | ||
|
|
||
|
|
||
| def run_2d(h: int = 16, w: int = 32, dim: int = 1): | ||
| x = torch.arange(h * w, dtype=torch.int32).reshape(h, w) | ||
| y = torch.empty_like(x) | ||
| grid = (triton.cdiv(h, 16), triton.cdiv(w, 16)) | ||
| flip_2d_kernel[grid]( | ||
| x, y, h, w, x.stride(0), x.stride(1), dim, BLOCK_M=16, BLOCK_N=16 | ||
| ) | ||
| assert torch.equal(y, torch.flip(x, dims=[dim])) | ||
|
|
||
|
|
||
| if __name__ == "__main__": | ||
| run_1d(256) | ||
| run_2d(16, 32, dim=1) | ||
| triton_viz.launch(share=True, port=8003) |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,93 @@ | ||
| import torch | ||
| import triton | ||
| import triton.language as tl | ||
|
|
||
| import triton_viz | ||
| from triton_viz.clients import Tracer | ||
|
|
||
|
|
||
| # Simple matmul kernel producing a C = A @ B (fp32, small sizes for demo) | ||
| @triton_viz.trace(clients=Tracer()) | ||
| @triton.jit | ||
| def matmul_kernel( | ||
| a_ptr, | ||
| b_ptr, | ||
| c_ptr, | ||
| M: tl.constexpr, | ||
| N: tl.constexpr, | ||
| K: tl.constexpr, | ||
| stride_am, | ||
| stride_ak, | ||
| stride_bk, | ||
| stride_bn, | ||
| stride_cm, | ||
| stride_cn, | ||
| BLOCK_M: tl.constexpr, | ||
| BLOCK_N: tl.constexpr, | ||
| BLOCK_K: tl.constexpr, | ||
| ): | ||
| pid_m = tl.program_id(0) | ||
| pid_n = tl.program_id(1) | ||
|
|
||
| offs_m = pid_m * BLOCK_M + tl.arange(0, BLOCK_M) | ||
| offs_n = pid_n * BLOCK_N + tl.arange(0, BLOCK_N) | ||
| offs_k = tl.arange(0, BLOCK_K) | ||
|
|
||
| a_ptrs = a_ptr + offs_m[:, None] * stride_am + offs_k[None, :] * stride_ak | ||
| b_ptrs = b_ptr + offs_k[:, None] * stride_bk + offs_n[None, :] * stride_bn | ||
|
|
||
| acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) | ||
|
|
||
| for k in range(0, K, BLOCK_K): | ||
| a = tl.load( | ||
| a_ptrs, mask=(offs_m[:, None] < M) & (k + offs_k[None, :] < K), other=0.0 | ||
| ) | ||
| b = tl.load( | ||
| b_ptrs, mask=(k + offs_k[:, None] < K) & (offs_n[None, :] < N), other=0.0 | ||
| ) | ||
| acc += tl.dot(a, b) | ||
| a_ptrs += BLOCK_K * stride_ak | ||
| b_ptrs += BLOCK_K * stride_bk | ||
|
|
||
| c_ptrs = c_ptr + offs_m[:, None] * stride_cm + offs_n[None, :] * stride_cn | ||
| tl.store(c_ptrs, acc, mask=(offs_m[:, None] < M) & (offs_n[None, :] < N)) | ||
|
|
||
|
|
||
| def run_demo(): | ||
| torch.manual_seed(0) | ||
| M, N, K = 64, 64, 64 | ||
| a = torch.randn((M, K), dtype=torch.float32) | ||
| b = torch.randn((K, N), dtype=torch.float32) | ||
| c = torch.empty((M, N), dtype=torch.float32) | ||
|
|
||
| BLOCK_M, BLOCK_N, BLOCK_K = 32, 32, 32 | ||
| grid = (triton.cdiv(M, BLOCK_M), triton.cdiv(N, BLOCK_N)) | ||
|
|
||
| matmul_kernel[grid]( | ||
| a, | ||
| b, | ||
| c, | ||
| M, | ||
| N, | ||
| K, | ||
| a.stride(0), | ||
| a.stride(1), | ||
| b.stride(0), | ||
| b.stride(1), | ||
| c.stride(0), | ||
| c.stride(1), | ||
| BLOCK_M=BLOCK_M, | ||
| BLOCK_N=BLOCK_N, | ||
| BLOCK_K=BLOCK_K, | ||
| ) | ||
|
|
||
| # Verify correctness | ||
| ref = a @ b | ||
| assert torch.allclose(c, ref, atol=1e-3), "matmul result mismatch" | ||
|
|
||
| # Launch viz UI in blocking (share=True) mode | ||
| triton_viz.launch(share=True) | ||
|
|
||
|
|
||
| if __name__ == "__main__": | ||
| run_demo() |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -0,0 +1,116 @@ | ||
| from neuronxcc import nki | ||
| import neuronxcc.nki.language as nl | ||
|
|
||
| import triton_viz | ||
| from triton_viz.clients import Tracer | ||
| from triton_viz.core.trace import launches | ||
| import numpy as np | ||
| import math | ||
|
|
||
|
|
||
| def matmul_kernel(lhs, rhs, result): | ||
| """NKI matmul_kernel to compute a matrix multiplication operation in a tiled manner | ||
|
|
||
| Args: | ||
| lhs: an input tensor of shape [K,M], where both K and M are multiples for | ||
| 128. It is the left-hand-side argument of the matrix multiplication, | ||
| delivered transposed for optimal performance. | ||
| rhs: an input tensor of shape [K,N], where K is a multiple of 128, and N | ||
| is a multiple of 512. It is the right-hand-side argument of the | ||
| matrix multiplication. | ||
| Returns: | ||
| result: the resulting output tensor of shape [M,N] | ||
| """ | ||
|
|
||
| M, K = lhs.shape | ||
| K_, N = rhs.shape | ||
| assert K == K_, "lhs and rhs must have the same contraction dimension" | ||
|
|
||
| TILE_M = 2 | ||
| TILE_K = 2 | ||
| TILE_N = 4 | ||
|
|
||
| # Use affine_range to loop over tiles | ||
| for m in nl.affine_range(math.ceil(M / TILE_M)): | ||
| for n in nl.affine_range(math.ceil(N / TILE_N)): | ||
| # Allocate a tensor in PSUM | ||
| res_psum = nl.zeros((TILE_M, TILE_N), nl.int32, buffer=nl.psum) | ||
|
|
||
| for k in nl.affine_range(math.ceil(K / TILE_K)): | ||
| # Declare the tiles on SBUF | ||
| lhs_tile = nl.ndarray((TILE_K, TILE_M), dtype=lhs.dtype, buffer=nl.sbuf) | ||
| rhs_tile = nl.ndarray((TILE_K, TILE_N), dtype=rhs.dtype, buffer=nl.sbuf) | ||
|
|
||
| # Load tiles from lhs and rhs | ||
| lhs_p = nl.arange(TILE_M)[:, None] + m * TILE_M | ||
| lhs_f = nl.arange(TILE_K)[None, :] + k * TILE_K | ||
| lhs_mask = (lhs_p < M) & (lhs_f < K) | ||
| lhs_tile = nl.load(lhs[lhs_p, lhs_f], mask=lhs_mask) | ||
|
|
||
| rhs_p = nl.arange(TILE_K)[:, None] + k * TILE_K | ||
| rhs_f = nl.arange(TILE_N)[None, :] + n * TILE_N | ||
| rhs_mask = (rhs_p < K) & (rhs_f < N) | ||
| rhs_tile = nl.load(rhs[rhs_p, rhs_f], mask=rhs_mask) | ||
|
|
||
| # Accumulate partial-sums into PSUM | ||
| x = nl.matmul(lhs_tile[...], rhs_tile[...], transpose_x=False) | ||
| res_psum += x | ||
|
|
||
| # Copy the result from PSUM back to SBUF, and cast to expected output data-type | ||
| res_sb = nl.copy(res_psum, dtype=result.dtype) | ||
|
|
||
| out_p = nl.arange(TILE_M)[:, None] + m * TILE_M | ||
| out_f = nl.arange(TILE_N)[None, :] + n * TILE_N | ||
| out_mask = (out_p < M) & (out_f < N) | ||
| nl.store( | ||
| result[m * TILE_M : (m + 1) * TILE_M, n * TILE_N : (n + 1) * TILE_N], | ||
| value=res_sb, | ||
| mask=out_mask, | ||
| ) | ||
|
|
||
|
|
||
| TRITON_VIZ = True | ||
| kernel_grid = (1, 1, 1) | ||
| lhs_small = np.arange(16).astype(np.float32).reshape(4, 4) | ||
| rhs_small = np.arange(32).astype(np.float32).reshape(4, 8) | ||
| # lhs_small = np.arange(9).astype(np.float32).reshape(3, 3) | ||
| # rhs_small = np.arange(18).astype(np.float32).reshape(3, 6) | ||
| result = np.empty((lhs_small.shape[0], rhs_small.shape[1]), dtype=lhs_small.dtype) | ||
| kernel_args = (lhs_small, rhs_small, result) | ||
|
|
||
| if TRITON_VIZ: | ||
| print("Executing matmul_kernel with NKI interpreter...") | ||
| traced_kernel = triton_viz.trace(clients=Tracer(), backend="nki")(matmul_kernel) | ||
| kernel_instance = traced_kernel[kernel_grid] | ||
| kernel_instance(*kernel_args) | ||
|
|
||
| print(f"Number of launches: {len(launches)}") | ||
| if launches: | ||
| launch = launches[-1] | ||
| print(f"Number of records: {len(launch.records)}") | ||
| for i, record in enumerate(launch.records): | ||
| print(f"Record {i}: {type(record).__name__}") | ||
| if hasattr(record, "ptr"): | ||
| print(f" ptr: {record.ptr}") | ||
| if hasattr(record, "offsets"): | ||
| print(f" offsets shape: {record.offsets.shape}") | ||
| if hasattr(record, "masks"): | ||
| print(f" masks shape: {record.masks.shape}") | ||
|
|
||
| # Try to launch visualization | ||
| try: | ||
| triton_viz.launch(share=False) | ||
| except Exception as e: | ||
| print(f"\nError during visualization: {e}") | ||
| import traceback | ||
|
|
||
| traceback.print_exc() | ||
| else: | ||
| print("Executing NKI JIT-ed matmul_kernel...") | ||
| compiled_kernel = nki.jit(matmul_kernel, kernel_return=False) | ||
| z2 = nki.simulate_kernel(compiled_kernel[kernel_grid], *kernel_args) | ||
|
|
||
| z2 = result | ||
| z1 = lhs_small @ rhs_small | ||
| print(np.max(np.abs(z1 - z2))) | ||
| assert np.allclose(z1, z2) | ||
Oops, something went wrong.
Add this suggestion to a batch that can be applied as a single commit.
This suggestion is invalid because no changes were made to the code.
Suggestions cannot be applied while the pull request is closed.
Suggestions cannot be applied while viewing a subset of changes.
Only one suggestion per line can be applied in a batch.
Add this suggestion to a batch that can be applied as a single commit.
Applying suggestions on deleted lines is not supported.
You must change the existing code in this line in order to create a valid suggestion.
Outdated suggestions cannot be applied.
This suggestion has been applied or marked resolved.
Suggestions cannot be applied from pending reviews.
Suggestions cannot be applied on multi-line comments.
Suggestions cannot be applied while the pull request is queued to merge.
Suggestion cannot be applied right now. Please check back later.
Uh oh!
There was an error while loading. Please reload this page.