Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
106 commits
Select commit Hold shift + click to select a range
e6e0641
Support NKI
Jokeren Jul 28, 2025
c9b2501
Update
Jokeren Jul 28, 2025
6fe57e6
Update
Jokeren Jul 28, 2025
86d1d6f
Update
Jokeren Jul 28, 2025
7aa3d91
Update
Jokeren Jul 28, 2025
4f8f616
Update
Jokeren Jul 28, 2025
9962f53
Update
Jokeren Jul 28, 2025
25c8ceb
Update
Jokeren Jul 28, 2025
cb07831
Update
Jokeren Jul 28, 2025
1810284
Update
Jokeren Jul 28, 2025
e6a6f91
Update
Jokeren Jul 28, 2025
8406f66
Update
Jokeren Jul 28, 2025
b63e74c
Update
Jokeren Jul 28, 2025
241530b
Update
Jokeren Jul 28, 2025
d28a17e
Update
Jokeren Jul 28, 2025
bc0873a
Update
Jokeren Jul 28, 2025
4b6827d
UPdate
Jokeren Jul 28, 2025
62ade2d
Update
Jokeren Jul 28, 2025
d9d7bf7
Update
Jokeren Jul 28, 2025
998a1a5
Update
Jokeren Jul 28, 2025
2a50a04
Update
Jokeren Jul 28, 2025
8e49f9d
Update
Jokeren Jul 28, 2025
45596d3
Update
Jokeren Jul 28, 2025
a54ac94
Update
Jokeren Jul 28, 2025
f1231fe
Update
Jokeren Jul 28, 2025
b1f658a
Update comments
Jokeren Jul 28, 2025
e7e2dd2
stuff
Sep 10, 2025
b3ee8d9
triton-viz should support cpu to allow development on non-trainium node
latentCall145 Sep 10, 2025
454ac73
test kernels to check basic tracing
latentCall145 Sep 26, 2025
65e861d
more flexible advanced indexing (just make arange slices an arange)
latentCall145 Sep 26, 2025
f6b99ae
remove need for SPMD nki kernel call when tracing
latentCall145 Sep 26, 2025
4f863fa
remove
latentCall145 Sep 26, 2025
2932bb9
added calibrate and matmul and show code
gujialiang123 Oct 3, 2025
cb66323
add some ops
latentCall145 Oct 3, 2025
7bd6bea
nki kernels to demo interpreter
latentCall145 Oct 3, 2025
9afefc6
fix misc tests
latentCall145 Oct 3, 2025
95ef0b6
prepare for nki tracing
latentCall145 Oct 3, 2025
2636088
uv
latentCall145 Oct 3, 2025
0644e64
fix simple ndarray indexing (and also add mgrid)
latentCall145 Oct 3, 2025
2037047
masked load stuff (masked store still doesn't work right)
latentCall145 Oct 16, 2025
8f619e8
attempt to connect nki interpreter to triton-viz (still failing)
latentCall145 Oct 16, 2025
5adb5e0
hardcode offsets and callbacks for NKI matmul kernel
latentCall145 Oct 16, 2025
ce4fdd4
Merge remote-tracking branch 'thaihoa/thaihoa/mm' into nki/merge_inte…
mark14wu Oct 16, 2025
f091fd4
Remove uv.lock
mark14wu Oct 16, 2025
9e2d495
feat(flip): add Flip op tracing, 3D flip view, hover value API; fix e…
gujialiang123 Oct 17, 2025
4fa8660
api: align register_op_callback signatures across Client, Sanitizer, …
gujialiang123 Oct 17, 2025
2b2968d
fix(nki): resolve mypy errors (value setter order, unary op calls); p…
gujialiang123 Oct 17, 2025
82f629e
use import for nki test
latentCall145 Oct 24, 2025
86345a1
masked load fast path; use max dtype value instead of 6 for out of bo…
latentCall145 Oct 25, 2025
85972d8
fix masked store
latentCall145 Oct 25, 2025
30d4e5d
clone array so load doesn't write on input array
latentCall145 Oct 25, 2025
f696409
make tests prettier
latentCall145 Oct 25, 2025
9d96499
make all tests pass
latentCall145 Oct 25, 2025
9066d3b
modify undefined values for other dtypes in masked load; lint
latentCall145 Oct 25, 2025
696587d
remove ._value, just use .data; lint
latentCall145 Oct 25, 2025
0fb3dbc
make nki tracing work again
latentCall145 Oct 25, 2025
a225146
don't immediately exit web server when share=false
latentCall145 Oct 26, 2025
15196d0
Merge branch 'thaihoa/local-viz-fix' into nki/merge-interpreter-working
latentCall145 Oct 26, 2025
6571b71
feat(flow): embed memory-flow badges in Load/Store/Dot; add Flow Diag…
gujialiang123 Oct 26, 2025
b959a12
patch nki pt.1
latentCall145 Oct 28, 2025
ed0b309
nki patch pt.2 (new tracer callbacks)
latentCall145 Oct 28, 2025
2e1943e
nki patch pt.3 (nki-specific changes to allow nki tracing)
latentCall145 Oct 28, 2025
301aa66
nki patch pt.4 (update matmul example)
latentCall145 Oct 28, 2025
185464f
Merge branch 'jialiang/work' into nki/merge-interpreter-working
latentCall145 Oct 28, 2025
4826219
Merge branch 'main' into nki/merge-interpreter-working
latentCall145 Oct 28, 2025
e9ecf82
remove NLSlice (was never used)
latentCall145 Oct 28, 2025
b74f471
fix offsets/mask shape mismatch bug in visualizing masked stores
latentCall145 Oct 28, 2025
1910b78
unify DSL patching pt 1. (make nki matmul + load_store examples work)
latentCall145 Oct 29, 2025
a120447
remove some files
latentCall145 Oct 31, 2025
b5797ef
normal args and kwargs
latentCall145 Oct 31, 2025
4e0f02a
allocate instead of array
latentCall145 Oct 31, 2025
a913fda
reorder callbacks
latentCall145 Nov 3, 2025
0c92fbd
more examples
latentCall145 Nov 3, 2025
0ee0896
Merge branch 'main' into nki/patch-dsl
mark14wu Nov 4, 2025
0fc84ec
apply ruff-format; move normalization to draw; stop tracking scripts/
gujialiang123 Nov 5, 2025
aceeae2
Triton/NKI Load/Store unify pt 1. (push to same callback - not working)
latentCall145 Nov 5, 2025
10271ae
Triton/NKI Load/Store unify pt. 2 (add adapters to standardize callba…
latentCall145 Nov 6, 2025
9f37b24
whoops forgot to add allocate adapter
latentCall145 Nov 6, 2025
cd623f5
be gone
latentCall145 Nov 6, 2025
6470e5b
add adapter tests
latentCall145 Nov 6, 2025
187c347
Merge branch 'nki/patch-dsl' of github.com:Deep-Learning-Profiling-To…
latentCall145 Nov 6, 2025
71b8e71
revert unwanted changes from main
latentCall145 Nov 6, 2025
cc9a787
patch backend less bad
latentCall145 Nov 6, 2025
f94d1be
remove unhelpful comments + show expected transformed code
latentCall145 Nov 6, 2025
675abab
debloat nki offsets code
latentCall145 Nov 6, 2025
7c3bec2
fix matmul visualization error
latentCall145 Nov 6, 2025
86e476a
Merge branch 'main' into nki/patch-dsl
latentCall145 Nov 6, 2025
229d4a9
lint
latentCall145 Nov 6, 2025
5bf74c8
Merge branch 'main' into nki/patch-dsl
latentCall145 Nov 7, 2025
096122c
Merge branch 'main' into nki/patch-dsl
mark14wu Nov 7, 2025
54559f1
rename nki_masked_load -> masked_load since it's also used for triton
latentCall145 Nov 13, 2025
f9c72e3
make AWS neuron an optional feature of triton viz (and include CI)
latentCall145 Nov 13, 2025
3a00512
Merge branch 'nki/patch-dsl' of github.com:Deep-Learning-Profiling-To…
latentCall145 Nov 13, 2025
0f6664a
guard nki test import if NKI not installed
latentCall145 Nov 13, 2025
b45cd5e
update import
latentCall145 Nov 13, 2025
6c01d62
[CI] fix aws neuron ci
latentCall145 Nov 13, 2025
fa814b5
python>=3.10 needed (also 3.9 is EOL)
latentCall145 Nov 13, 2025
de93383
[CI] make AWS tests reuse torch + triton installations from build job
latentCall145 Nov 13, 2025
50a0177
make nki set_grid_dim consistent with triton api
latentCall145 Nov 13, 2025
24c9abe
Merge branch 'main' into nki/patch-dsl
latentCall145 Nov 16, 2025
7260fce
Merge branch 'main' into nki/patch-dsl
mark14wu Nov 24, 2025
76f8999
extra nki kernels to viz
latentCall145 Nov 28, 2025
3dc8c02
typo
latentCall145 Dec 5, 2025
0817087
Merge branch 'main' into nki/patch-dsl
latentCall145 Dec 6, 2025
0eb3f22
Merge branch 'main' into nki/patch-dsl
latentCall145 Dec 17, 2025
232e87b
Merge branch 'main' into nki/patch-dsl
mark14wu Dec 23, 2025
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 10 additions & 0 deletions .github/workflows/python-app.yml
Original file line number Diff line number Diff line change
Expand Up @@ -60,3 +60,13 @@ jobs:
run: |
cd triton_viz
python -m pytest tests

- name: Install Triton-Viz with NKI extras
run: |
cd triton_viz
pip install -e .[nki]

- name: Run full (Triton + NKI) pytest suite
run: |
cd triton_viz
python -m pytest tests -m ""
2 changes: 2 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -161,3 +161,5 @@ cython_debug/
# and can be added to the global gitignore or merged into this file. For a more nuclear
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
#.idea/

uv.lock
20 changes: 19 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -44,8 +44,9 @@ The best part about this tool is that while it does focus on visualizing GPU ope

### Prerequisites

- Python installed (preferably the latest available version).
- Python installed (preferably the latest available version), minimum supported version is 3.10.
- [Triton](https://github.com/openai/triton/blob/main/README.md) installed. Follow the installation instructions in the linked repository.
- Note: the below commands must be run in order.

Upon successfully installing Triton, install Torch using the following command:

Expand All @@ -71,6 +72,23 @@ pip install -e .

You're all set!

### Optional: Enable NKI Support

If you want to exercise the Neuron Kernel Interface (NKI) interpreter or run the NKI-specific tests:

1. Follow the [AWS Neuron Torch-NeuronX Ubuntu 22.04 setup guide](https://awsdocs-neuron.readthedocs-hosted.com/en/latest/setup/neuron-setup/pytorch/neuronx/ubuntu/torch-neuronx-ubuntu22.html#setup-torch-neuronx-ubuntu22) to add the Neuron APT repository and install the required system packages (for example `aws-neuronx-tools`, `aws-neuronx-runtime-lib`, `aws-neuronx-collectives`, and their dependencies).
2. Instead of running `pip install -e .` in the above section, install Triton-Viz with the optional NKI extras so the Neuron Python packages (`neuronx-cc`, `libneuronxla`, `torch-neuronx`) are available:

```sh
pip install -e .[nki]
# or pip install triton-viz[nki]
```

### Testing
* To run core Triton-viz tests, run `pytest tests/`.
* (if NKI installed) To run NKI-specific tests, run `pytest tests/ -m nki`.
* To run all tests (Triton + NKI), run `pytest tests/ -m ""`.

## Working with Examples

```sh
Expand Down
81 changes: 81 additions & 0 deletions examples/flip.py
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)
2 changes: 1 addition & 1 deletion examples/load_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,7 @@ def simple_kernel(

# Try to launch visualization
try:
triton_viz.launch()
triton_viz.launch(share=False)
except Exception as e:
print(f"\nError during visualization: {e}")
import traceback
Expand Down
93 changes: 93 additions & 0 deletions examples/matmul_demo.py
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()
116 changes: 116 additions & 0 deletions examples/nki/matmul.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,116 @@
from neuronxcc import nki
Comment thread
latentCall145 marked this conversation as resolved.
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)
Loading