diff --git a/.github/workflows/python-app.yml b/.github/workflows/python-app.yml index bb729408..cb620c29 100644 --- a/.github/workflows/python-app.yml +++ b/.github/workflows/python-app.yml @@ -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 "" diff --git a/.gitignore b/.gitignore index 060367fa..44a40575 100644 --- a/.gitignore +++ b/.gitignore @@ -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 diff --git a/README.md b/README.md index 9f902b23..b29bf707 100644 --- a/README.md +++ b/README.md @@ -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: @@ -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 diff --git a/examples/flip.py b/examples/flip.py new file mode 100644 index 00000000..289250aa --- /dev/null +++ b/examples/flip.py @@ -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) diff --git a/examples/load_store.py b/examples/load_store.py index 24bd55e2..b3289f4d 100644 --- a/examples/load_store.py +++ b/examples/load_store.py @@ -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 diff --git a/examples/matmul_demo.py b/examples/matmul_demo.py new file mode 100644 index 00000000..c6b41db5 --- /dev/null +++ b/examples/matmul_demo.py @@ -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() diff --git a/examples/nki/matmul.py b/examples/nki/matmul.py new file mode 100644 index 00000000..293c1768 --- /dev/null +++ b/examples/nki/matmul.py @@ -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) diff --git a/examples/nki/rmsnorm.py b/examples/nki/rmsnorm.py new file mode 100644 index 00000000..e844a934 --- /dev/null +++ b/examples/nki/rmsnorm.py @@ -0,0 +1,124 @@ +import neuronxcc.nki as nki +import neuronxcc.nki.language as nl +from triton_viz.clients import Tracer +from triton_viz.core.trace import launches +import math +import numpy as np +import torch +import triton_viz + + +def nki_rmsnorm_kernel(a_tensor, g_tensor, result): + # Calculate out_tensor = a_tensor/RMS(a_tensor) * g_tensor + # Where RMS(a_tensor) = sqrt((1/N) * sum(a_tensor * a_tensor)) + # and N = a_tensor.shape[1] + # Reduction (mean) is performed in the free (2nd) dimension + B, D = a_tensor.shape + + # Make sure shapes match + assert D == g_tensor.shape[0] + + # Generate tensor indices to index input tensor + B_TILE = 8 + ix = nl.arange(B_TILE)[:, None] + iw = nl.arange(1)[:, None] + iy = nl.arange(D)[None, :] + + # Load RMSNorm weight once, reused by rows/tiles of a_tensor + g_tile = nl.load(g_tensor.reshape((1, D))[iw, iy], mask=iy < D) + + # Process 2 rows at a time due to 2-partition tile size limitation + # Since we're not reducing across the first dimension + # Tiles can be processed independently + for i in nl.affine_range(math.ceil(B / B_TILE)): + # Load input data from external memory to on-chip memory + mask = (i * B_TILE + ix < B) & (iy < D) + a_tile = nl.load(a_tensor[i * B_TILE + ix, iy], mask=mask) + + # Compute element-wise square of a_tensor + in_square = nl.square(a_tile) + + # Calculate sum of squared elements, along last dimension + # square_sum = nl.sum(in_square, axis=[1], mask=mask) #[:, None] + square_sum = nl.sum(in_square, axis=1, mask=mask)[:, None] + + # Scale and get a reciprocal + mean = square_sum / D + + # Take square root of mean and then reciprocal with + # rsqrt API (one ISA instruction) + rms_reciprocal = nl.rsqrt(mean) + + # Scale the input tensor + out_tile = nl.multiply(a_tile, rms_reciprocal) + + # Broadcast weight along first axis to match tensor shape + # B_active = min(B - i * 2, 2) + g_bcast = g_tile.broadcast_to((B_TILE, D)) + + # Multiply with the RMSNorm weight + out_tile = nl.multiply(out_tile, g_bcast, mask=(i * B_TILE + ix < B)) + + # store the addition results back to external memory (out_tensor) + nl.store(result[i * B_TILE + ix, iy], value=out_tile, mask=mask) + + +# ref +def torch_rmsnorm_kernel(a_tensor, g_tensor): + # Square the tensor (element-wise) + in_square = a_tensor.pow(2) + # Calculate means in the free dimension + mean = in_square.mean(dim=1, keepdim=True) + # Scale by reciprocal of sqrt(mean) + tensor = a_tensor * torch.rsqrt(mean) + + # Scale the output by the weight + return tensor * g_tensor + + +TRITON_VIZ = True +kernel_grid = (1, 1, 1) +B, D = 32, 32 +a_tensor = torch.arange(B * D).float().view(B, D) +g_tensor = torch.arange(D).float() +result = torch.empty_like(a_tensor).numpy() +kernel_args = (a_tensor.numpy(), g_tensor.numpy(), result) + +if TRITON_VIZ: + print("Executing kernel with NKI interpreter...") + traced_kernel = triton_viz.trace(clients=Tracer(), backend="nki")( + nki_rmsnorm_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(nki_rmsnorm_kernel, kernel_return=False) + z2 = nki.simulate_kernel(compiled_kernel[kernel_grid], *kernel_args) + +z2 = result +z1 = torch_rmsnorm_kernel(a_tensor, g_tensor).numpy() +print(np.max(np.abs(z1 - z2))) +assert np.allclose(z1, z2) diff --git a/examples/nki/rope.py b/examples/nki/rope.py new file mode 100644 index 00000000..0ec596ea --- /dev/null +++ b/examples/nki/rope.py @@ -0,0 +1,294 @@ +import os +from typing import Tuple + +import neuronxcc.nki.language as nl +import torch +from neuronxcc import nki +import triton_viz +from triton_viz.clients import Tracer +from triton_viz.core.trace import launches +import numpy as np + +os.environ["NEURON_FRAMEWORK_DEBUG"] = "1" +# Ideally remove disable-dge for production performance, but kept for debugging context +os.environ["NEURON_CC_FLAGS"] = " --disable-dge " + + +def generate_pos_embedding( + head_dim: int, position_ids: torch.Tensor, base: int = 10000 +) -> Tuple[torch.Tensor, torch.Tensor]: + """ + Generate positional embeddings for rotary position encoding (Llama style). + """ + # Core RoPE block + inv_freq = 1.0 / (base ** (torch.arange(0, head_dim, 2) / head_dim)) + + # Expand to [HeadDim/2, 1] + inv_freq_expanded = inv_freq[:, None].float() + # Expand to [1, SeqLen] + position_ids_expanded = position_ids[None, :].float() + + # MatMul -> [HeadDim/2, SeqLen] -> Transpose -> [SeqLen, HeadDim/2] + freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(0, 1) + + # Concatenate to match HeadDim -> [Batch, SeqLen, HeadDim] + emb = torch.cat((freqs, freqs), dim=-1) + cos = emb.cos() + sin = emb.sin() + return cos, sin + + +def div_ceil(n: int, d: int) -> int: + return (n + d - 1) // d + + +def _nki_apply_rotary_embedding_core(q_tile, k_tile, cos_tile, sin_tile, output_tile): + """ + Core NKI implementation of rotary position embedding computation. + + Parameters + ---------- + q_tile : nl.Tensor + Query tensor tile + k_tile : nl.Tensor + Key tensor tile + cos_tile : nl.Tensor + Cosine embedding tile + sin_tile : nl.Tensor + Sine embedding tile + output_tile : nl.Tensor + Output buffer for results + + Notes + ----- + The function applies rotary position embedding to query and key tensors + using the provided cosine and sine embeddings. + """ + + assert q_tile.shape[-1] % 2 == 0, "Sequence length for q_tile must be even!" + assert k_tile.shape[-1] % 2 == 0, "Sequence length for k_tile must be even!" + assert ( + q_tile.shape[-1] == k_tile.shape[-1] + ), "q_tile and k_tile must have the same sequence length" + + seq_len = q_tile.shape[-1] + + # Rotate Q + output_tile[0, :, :] = q_tile * cos_tile + output_tile[0, :, : seq_len // 2] = output_tile[0, :, : seq_len // 2] + ( + -1 * q_tile[:, seq_len // 2 :] * sin_tile[:, : seq_len // 2] + ) + output_tile[0, :, seq_len // 2 :] = output_tile[0, :, seq_len // 2 :] + ( + q_tile[:, : seq_len // 2] * sin_tile[:, seq_len // 2 :] + ) + + # Rotate K + output_tile[1, :, :] = k_tile * cos_tile + output_tile[1, :, : seq_len // 2] = output_tile[1, :, : seq_len // 2] + ( + -1 * k_tile[:, seq_len // 2 :] * sin_tile[:, : seq_len // 2] + ) + output_tile[1, :, seq_len // 2 :] = output_tile[1, :, seq_len // 2 :] + ( + k_tile[:, : seq_len // 2] * sin_tile[:, seq_len // 2 :] + ) + + +def nki_rope_kernel(q, k, cos, sin, output_q, output_k): + """ + NKI implementation of rotary position embedding. + + Parameters + ---------- + q : torch.Tensor + Query tensor of shape [batch_size, num_heads, seq_len, head_dim] + k : torch.Tensor + Key tensor of shape [batch_size, num_heads, seq_len, head_dim] + cos : torch.Tensor + Cosine embeddings + sin : torch.Tensor + Sine embeddings + + Returns + ------- + nl.Tensor + Output tensor containing transformed query and key tensors + + Raises + ------ + AssertionError + If input tensor shapes don't match or head dimension > 128 + """ + assert ( + q.shape == k.shape + ), f"Shape of Q Tensor: {q.shape} doesn't match shape of K Tensor: {k.shape}" + assert ( + cos.shape == sin.shape + ), f"Shape of cos Tensor: {cos.shape} doesn't match shape of sin Tensor: {sin.shape}" + PMAX = 128 + assert ( + q.shape[-1] <= PMAX + ), f"Shape of head dim (last dim) is more than PMAX: {q.shape}" + + head_id = nl.program_id(axis=0) + seq_len = q.shape[1] + num_seq_batches = div_ceil(seq_len, nl.tile_size.pmax) + # output = nl.ndarray([2] + list(q.shape), dtype=q.dtype, buffer=nl.shared_hbm) + i_p, i_f = nl.mgrid[0:PMAX, 0 : q.shape[-1]] + # for seq_batch_id in nl.affine_range(0, num_seq_batches): # TODO + for seq_batch_id in nl.affine_range(num_seq_batches): + # q_hbm_tile = q[batch_id, head_id] + # k_hbm_tile = k[batch_id, head_id] + # cos_hbm_tile = cos[batch_id] + # sin_hbm_tile = sin[batch_id] + + q_tile = nl.load( + # q_hbm_tile[seq_batch_id * nl.tile_size.pmax + i_p, i_f], + q[head_id, seq_batch_id * nl.tile_size.pmax + i_p, i_f], + mask=(seq_batch_id * nl.tile_size.pmax + i_p < seq_len), + ) + k_tile = nl.load( + # k_hbm_tile[seq_batch_id * nl.tile_size.pmax + i_p, i_f], + k[head_id, seq_batch_id * nl.tile_size.pmax + i_p, i_f], + mask=(seq_batch_id * nl.tile_size.pmax + i_p < seq_len), + ) + output_tile = nl.ndarray( + # [2] + [nl.par_dim(k_tile.shape[0]), k_tile.shape[1]], # TODO + [2] + [k_tile.shape[0], k_tile.shape[1]], + dtype=k_tile.dtype, + buffer=nl.sbuf, + ) + cos_tile = nl.load( + # cos_hbm_tile[seq_batch_id * nl.tile_size.pmax + i_p, i_f], + cos[seq_batch_id * nl.tile_size.pmax + i_p, i_f], + mask=(seq_batch_id * nl.tile_size.pmax + i_p < seq_len), + ) + sin_tile = nl.load( + # sin_hbm_tile[seq_batch_id * nl.tile_size.pmax + i_p, i_f], + sin[seq_batch_id * nl.tile_size.pmax + i_p, i_f], + mask=(seq_batch_id * nl.tile_size.pmax + i_p < seq_len), + ) + + _nki_apply_rotary_embedding_core( + q_tile, k_tile, cos_tile, sin_tile, output_tile + ) + + # output_q_hbm_tile = output[0, batch_id, head_id] + # output_k_hbm_tile = output[1, batch_id, head_id] + + nl.store( + # output_q_hbm_tile[seq_batch_id * nl.tile_size.pmax + i_p, i_f], + output_q[head_id, seq_batch_id * nl.tile_size.pmax + i_p, i_f], + output_tile[0, :, :], + mask=(seq_batch_id * nl.tile_size.pmax + i_p < seq_len), + ) + nl.store( + # output_k_hbm_tile[seq_batch_id * nl.tile_size.pmax + i_p, i_f], + output_k[head_id, seq_batch_id * nl.tile_size.pmax + i_p, i_f], + output_tile[1, :, :], + mask=(seq_batch_id * nl.tile_size.pmax + i_p < seq_len), + ) + + # return output + + +# Torch reference implementation +def torch_rope_kernel(q, k, cos, sin): + """Simple torch reference for rotary embedding""" + + def rotate_half(x): + x1 = x[..., : x.shape[-1] // 2] + x2 = x[..., x.shape[-1] // 2 :] + return torch.cat((-x2, x1), dim=-1) + + # + # FIX: Unsqueeze cos/sin to shape [B, 1, S, D] to broadcast over Heads [B, H, S, D] + cos = cos.unsqueeze(0) + sin = sin.unsqueeze(0) + + q_rot = (q * cos) + (rotate_half(q) * sin) + k_rot = (k * cos) + (rotate_half(k) * sin) + return q_rot, k_rot + + +# TRITON_VIZ test section +TRITON_VIZ = True # Set to True to see the trace +# B, H, S, D = 2, 2, 4, 8 +H, S, D = 2, 4, 8 +# kernel_grid = (B, H) +kernel_grid = (H,) + +if TRITON_VIZ: + # Setup test data + # q = torch.randn(B, H, S, D) + # k = torch.randn(B, H, S, D) + q = torch.randn(H, S, D) + k = torch.randn(H, S, D) + position_ids = torch.arange(S) + cos, sin = generate_pos_embedding(D, position_ids) + + # output = torch.empty((2, B, H, S, D)) + output_q = torch.empty((H, S, D)) + output_k = torch.empty((H, S, D)) + # Cast inputs to numpy for NKI tracing + kernel_args = ( + q.numpy(), + k.numpy(), + cos.numpy(), + sin.numpy(), + output_q.numpy(), + output_k.numpy(), + ) + + print("Executing rotary embedding kernel with NKI interpreter...") + traced_kernel = triton_viz.trace(clients=Tracer(), backend="nki")(nki_rope_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)}") + # Optional: Print record details + + try: + triton_viz.launch(share=False) + except Exception as e: + print(f"Visualization error: {e}") + +else: + # Regular execution for comparison + # Ensure inputs are float32 + q = torch.randn(H, S, D, dtype=torch.float32) + k = torch.randn(H, S, D, dtype=torch.float32) + position_ids = torch.arange(S) + cos, sin = generate_pos_embedding(D, position_ids) + + print("Executing NKI JIT-ed rotary embedding kernel...") + compiled_kernel = nki.jit(nki_rope_kernel, kernel_return=False) + output_q_np = np.zeros((H, S, D), dtype=np.float32) + output_k_np = np.zeros((H, S, D), dtype=np.float32) + + nki.simulate_kernel( + compiled_kernel[kernel_grid], + q.numpy(), + k.numpy(), + cos.numpy(), + sin.numpy(), + output_q_np, + output_k_np, + ) + + # Compare with torch reference + expected_q, expected_k = torch_rope_kernel(q, k, cos, sin) + actual_q = torch.from_numpy(output_q_np) + actual_k = torch.from_numpy(output_k_np) + + max_diff_q = torch.max(torch.abs(expected_q - actual_q)) + max_diff_k = torch.max(torch.abs(expected_k - actual_k)) + + print(f"Q max diff: {max_diff_q}") + print(f"K max diff: {max_diff_k}") + + # Relaxed tolerance slightly for float32 accumulation differences + assert torch.allclose(expected_q, actual_q, rtol=1e-3, atol=1e-3), "Q mismatch" + assert torch.allclose(expected_k, actual_k, rtol=1e-3, atol=1e-3), "K mismatch" + print("Results match!") diff --git a/examples/nki/softmax.py b/examples/nki/softmax.py new file mode 100644 index 00000000..99942c0c --- /dev/null +++ b/examples/nki/softmax.py @@ -0,0 +1,73 @@ +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 softmax_kernel(in_tensor, out_tensor): + """NKI softmax_kernel to compute softmax on the last dimension + + Args: + in_tensor: an input tensor of shape [B,D], where B is a multiple of 128 + out_tensor: the resulting output tensor of shape [B,D] + """ + B, D = in_tensor.shape + + # assert nl.tile_size.pmax == 128 + num_tiles = math.ceil(B / 128) + for tile_idx in nl.affine_range(num_tiles): + i_p = tile_idx * 128 + nl.arange(128)[:, None] + i_f = nl.arange(D)[None, :] + mask = (i_p < B) | (i_f < 0) + tile = nl.load(in_tensor[i_p, i_f], mask=mask) + tile_exp = nl.exp(tile) + tile_expsum = nl.sum(tile_exp, -1, keepdims=True, mask=mask) + tile_softmax = tile_exp / tile_expsum + nl.store(out_tensor[i_p, i_f], value=tile_softmax, mask=mask) + + +TRITON_VIZ = True +kernel_grid = (1, 1, 1) +x_small = np.random.rand(16, 32).astype(np.float32) +y_small = np.empty(x_small.shape, dtype=x_small.dtype) +kernel_args = (x_small, y_small) + +if TRITON_VIZ: + print("Executing softmax_kernel with NKI interpreter...") + traced_kernel = triton_viz.trace(clients=Tracer(), backend="nki")(softmax_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 softmax_kernel...") + compiled_kernel = nki.jit(softmax_kernel, kernel_return=False) + nki.simulate_kernel(compiled_kernel[kernel_grid], *kernel_args) + +y_expected = np.exp(x_small) / np.exp(x_small).sum(-1, keepdims=True) +print(np.max(np.abs(y_expected - y_small))) +assert np.allclose(y_expected, y_small) diff --git a/pyproject.toml b/pyproject.toml index 0442001c..dd0cf6e5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -14,7 +14,7 @@ authors = [ ] readme = "README.md" license = {text = "MIT"} -requires-python = ">=3.7" +requires-python = ">=3.10" classifiers = [ "Programming Language :: Python :: 3", "License :: OSI Approved :: MIT License", @@ -46,3 +46,9 @@ homepage = "https://github.com/Deep-Learning-Profiling-Tools/triton-viz" [project.scripts] triton-sanitizer = "triton_viz.wrapper:apply_sanitizer" triton-profiler = "triton_viz.wrapper:apply_profiler" + +[project.optional-dependencies] +nki = [ + "neuronx-cc @ https://pip.repos.neuron.amazonaws.com/neuronx-cc/neuronx_cc-2.21.33363.0%2B82129205-cp310-cp310-linux_x86_64.whl ; python_version == '3.10'", + "neuronx-cc @ https://pip.repos.neuron.amazonaws.com/neuronx-cc/neuronx_cc-2.21.33363.0%2B82129205-cp311-cp311-linux_x86_64.whl ; python_version == '3.11'" +] diff --git a/pytest.ini b/pytest.ini index 0f5ccb8f..1d3f80a6 100644 --- a/pytest.ini +++ b/pytest.ini @@ -2,3 +2,7 @@ python_files = *.py python_classes = *Test python_functions = test_* +addopts = -m "not nki" + +markers = + nki: NKI-specific tests diff --git a/tests/nki/test_nki.py b/tests/nki/test_nki.py new file mode 100644 index 00000000..2ce3985f --- /dev/null +++ b/tests/nki/test_nki.py @@ -0,0 +1,121 @@ +#!/usr/bin/env python3 +""" +Test script to verify NDArray slicing functionality after fixes +""" +import numpy as np +import pytest + +try: + from triton_viz.core.nki import NDArray +except ModuleNotFoundError: + pytest.skip( + "NeuronX dependencies are missing. Install triton-viz[nki] to run these tests.", + allow_module_level=True, + ) + +pytestmark = pytest.mark.nki # only run at "pytest -m nki" + + +def test_ndarray_creation(): + print("Testing NDArray creation...") + + # Test creation with value + data = np.array([[1, 2, 3], [4, 5, 6], [7, 8, 9]]) + nd1 = NDArray(value=data, name="test_array") + print(f"Created NDArray: {nd1}") + print(f"Shape: {nd1.shape}") + print(f"Dtype: {nd1.dtype}") + print(f"Value:\n{nd1.data}") + print() + + # Test creation with shape and dtype + nd2 = NDArray(shape=(2, 3), dtype=np.float32, name="shaped_array") + print(f"Created shaped NDArray: {nd2}") + print(f"Shape: {nd2.shape}") + print(f"Dtype: {nd2.dtype}") + print() + + +def test_slicing(): + print("Testing slicing operations...") + + data = np.array([[1, 2, 3], [4, 5, 6], [7, 8, 9]]) + nd_array = NDArray(value=data, name="test_array") + + # Test [:, :] (all elements) + # slice_all = nd_array[:, :] + slice_all = nd_array[:2, :2] + assert np.allclose(slice_all.data, data[:2, :2]) + print(f"nd_array[:, :] = {slice_all}") + print(f"Value:\n{slice_all.data}") + print() + + # Test [:, 0] (first column) + slice_col = nd_array[:, 0] + assert np.allclose(slice_col.data, data[:, 0]) + print(f"nd_array[:, 0] = {slice_col}") + print(f"Value: {slice_col.data}") + print() + + # Test [0, :] (first row) + slice_row = nd_array[0, :] + assert np.allclose(slice_row.data, data[0, :]) + print(f"nd_array[0, :] = {slice_row}") + print(f"Value: {slice_row.data}") + print() + + # Test advanced indexing + rewritten_slice = ( + np.arange(2)[:, None], + np.arange(3)[None, :], + ) + slice_advanced = nd_array[rewritten_slice] + assert np.allclose(slice_advanced.data, data[rewritten_slice]) + print(f"nd_array[nl.arange(2)[:, None], nl.arange(3)[None, :]] = {slice_advanced}") + print(f"Value:\n{slice_advanced.data}") + print() + + # test nd_array[0, :3, :, 2:4, 2] + data = np.reshape(np.arange(3 * 4 * 5 * 6 * 7), (3, 4, 5, 6, 7)) + nd_array = NDArray(value=data, name="test_array_5d") + rewritten_slice = ( + 0, + np.arange(3)[:, None, None], + np.arange(5)[None, :, None], + np.arange(2, 4)[None, None, :], + 2, + ) # this is what the above slice would be represented as after triton-viz tracing + slice_advanced = nd_array[rewritten_slice] + assert np.allclose(slice_advanced.data, data[rewritten_slice]) + assert np.allclose(slice_advanced.data, data[0, :3, :, 2:4, 2]) + print(f"Value:\n{slice_advanced.data}") + print() + + +def test_arithmetic(): + print("Testing arithmetic operations...") + + data1 = np.array([[1, 2], [3, 4]]) + data2 = np.array([[5, 6], [7, 8]]) + + nd1 = NDArray(value=data1, name="array1") + nd2 = NDArray(value=data2, name="array2") + + # Test addition + result = nd1 + nd2 + print(f"Addition result: {result}") + print(f"Value:\n{result.data}") + print() + + # Test slicing on result + slice_result = result[:, 0] + print(f"Slice of result [:, 0]: {slice_result}") + print(f"Value: {slice_result.data}") + print() + + +if __name__ == "__main__": + test_ndarray_creation() + test_slicing() + test_arithmetic() + print("All tests completed successfully!") diff --git a/tests/test_adapters.py b/tests/test_adapters.py new file mode 100644 index 00000000..b22e24a1 --- /dev/null +++ b/tests/test_adapters.py @@ -0,0 +1,172 @@ +import pytest + +from triton_viz.core.callbacks import OpCallbacks +from triton_viz.core.data import ( + AddPtr, + Load, + ProgramId, + ReduceSum, + Store, +) +from triton_viz.core.patch import ( + AdapterResult, + HAS_NKI, + NKI_ADAPTERS, + PatchOp, + TRITON_ADAPTERS, +) + + +def test_adapter_result_kwargs_copy(): + """Test that AdapterResult SHALLOW-copies kwargs.""" + kwargs = {"named": "value"} + result = AdapterResult(1, **kwargs) + kwargs[ + "named" + ] = "updated" # result.kwargs["named"] is immutable str so it doesn't change from this + assert result.args == (1,) + assert result.kwargs == {"named": "value"} + + mutable_object = ["before"] + kwargs = {"named": mutable_object} + result = AdapterResult(1, **kwargs) + kwargs["named"][ + 0 + ] = "after" # result.kwargs["named"] is mutable so it DOES change from this + assert result.args == (1,) + assert result.kwargs["named"][0] == "after" + + +def test_triton_store_adapter_handles_keys(): + """Adapter mirrors tl.store(ptr, value, mask=..., cache_modifier=..., eviction_policy=..., keys=...).""" + adapter = TRITON_ADAPTERS[Store] + ptr, value, mask, keys = object(), object(), object(), object() + + cache_modifier = "wb" + eviction_policy = "evict_last" + + full = adapter( + ptr, + value, + mask, + cache_modifier=cache_modifier, + eviction_policy=eviction_policy, + keys=keys, + ) + assert isinstance(full, AdapterResult) + assert full.args == (ptr, mask, keys) + assert full.kwargs == {} + + minimal = adapter(ptr, value, mask) + assert minimal.args == (ptr, mask, None) + assert minimal.kwargs == {} + + +def test_triton_load_adapter_keys_passthrough(): + """Adapter mirrors tl.load(ptr, mask=..., other=..., keys=...).""" + adapter = TRITON_ADAPTERS[Load] + ptr, mask, other, keys = object(), object(), object(), object() + + cache_modifier = "wb" + eviction_policy = "evict_last" + is_volatile = False + + full = adapter( + ptr, + mask, + other, + cache_modifier=cache_modifier, + eviction_policy=eviction_policy, + is_volatile=is_volatile, + keys=keys, + ) + assert full.args == (ptr, mask, keys) + assert full.kwargs == {} + + minimal = adapter(ptr, mask, other) + assert minimal.args == (ptr, mask, None) + assert minimal.kwargs == {} + + +def test_triton_reduce_adapter_supports_positional_and_keyword_axis(): + """Adapter retains axis/keep_dims regardless of positional or keyword usage.""" + adapter = TRITON_ADAPTERS[ReduceSum] + tensor = object() + axis = 1 + keep_dims = True + + positional = adapter(tensor, axis, keep_dims) + assert positional.args == (tensor, axis, keep_dims) + assert positional.kwargs == {} + + keyword = adapter(tensor, axis=2, keep_dims=False) + assert keyword.args == (tensor, 2, False) + assert keyword.kwargs == {} + + +def test_triton_addptr_adapter_orders_arguments(): + """Adapter maps tl.addptr(ptr, offset) to (ptr, offset).""" + adapter = TRITON_ADAPTERS[AddPtr] + ptr = object() + offset = object() + result = adapter(ptr, offset) + assert result.args == (ptr, offset) + assert result.kwargs == {} + + +def test_program_id_adapter_returns_axis_only(): + """Adapter maps tl.program_id(axis) to (axis,).""" + adapter = TRITON_ADAPTERS[ProgramId] + program_id = 0 + result = adapter(program_id) + assert result.args == (program_id,) + assert result.kwargs == {} + + +@pytest.mark.skipif(not HAS_NKI, reason="NKI extras not installed") +def test_nki_load_store_adapters_align_with_clients(): + """NKI adapters normalize masked load/store to (tensor, mask, keys).""" + src = object() + dst = object() + keys = object() + mask = object() + value = object() + + load_result = NKI_ADAPTERS[Load](src, keys, mask=mask) + assert load_result.args == (src, mask, keys) + assert load_result.kwargs == {} + + store_result = NKI_ADAPTERS[Store](dst, keys, value, mask=mask) + assert store_result.args == (dst, mask, keys) + assert store_result.kwargs == {} + + +def test_patchop_uses_adapter_for_callbacks(): + """PatchOp must feed the adapted arguments into before/after callbacks.""" + before_log: list[tuple[str, tuple, dict]] = [] + after_log: list[tuple[str, tuple, dict]] = [] + + def before_callback(*args, **kwargs): + before_log.append(("before", args, kwargs)) + + def after_callback(ret, *args, **kwargs): + after_log.append(("after", args, kwargs)) + + adapter = TRITON_ADAPTERS[Store] + callbacks = OpCallbacks( + before_callback=before_callback, after_callback=after_callback + ) + patch_op = PatchOp( + op=lambda *_a, **_k: "return-value", + op_type=Store, + callbacks=callbacks, + adapter=adapter, + ) + + ptr = object() + mask = object() + patch_op(ptr, "value", mask) + + expected_args = (ptr, mask, None) + assert before_log == [("before", expected_args, {})] + assert after_log == [("after", expected_args, {})] diff --git a/tests/test_core.py b/tests/test_core.py index d0bf3eca..9d733f3b 100644 --- a/tests/test_core.py +++ b/tests/test_core.py @@ -40,9 +40,9 @@ def my_kernel(x_ptr, y_ptr, out_ptr, BLOCK_SIZE: tl.constexpr): tl.store(out_ptr + offs, tl.load(x_ptr + offs) + tl.load(y_ptr + offs)) # Should be wrapped as a Trace object. - from triton_viz.core.trace import Trace + from triton_viz.core.trace import TritonTrace - assert isinstance(my_kernel, Trace) + assert isinstance(my_kernel, TritonTrace) # Verify client de-duplication and addition logic clients = my_kernel.client_manager.clients diff --git a/tests/test_masked_load.py b/tests/test_masked_load.py new file mode 100644 index 00000000..f966d808 --- /dev/null +++ b/tests/test_masked_load.py @@ -0,0 +1,770 @@ +import numpy as np +import pytest +from triton_viz.core.masked_load import masked_load, masked_store + + +def print_op_details( + test_name, + operation_type, + input_data, + keys, + values=None, + mask=None, + output=None, + error_expected=None, +): + """Print detailed information about a masked load/store operation.""" + print(f"{test_name}:") + print(f"Operation type: {operation_type}") + print("---------------------------------") + print(f"Input:\n{input_data}") + print("---------------------------------") + print(f"Values:\n{values}") + print("---------------------------------") + print(f"Keys:\n{keys}") + print("---------------------------------") + print(f"Mask:\n{mask}") + print("---------------------------------") + print(f"Result {operation_type}:\n{output}") + if error_expected: + print(f"Expecting {error_expected}...") + print() + + +def test_masked_load_none_case(): + """Test masked_load with mask=None (direct indexing).""" + arr = np.array([1, 2, 3, 4, 5]) + result = masked_load(arr, (slice(1, 4),), mask=None) + expected = arr[1:4] + + print_op_details( + "Test 1: mask=None case", "load", arr, (slice(1, 4),), mask=None, output=result + ) + + assert np.array_equal(result, expected) + + +def test_masked_load_in_bounds(): + """Test masked_load with in-bounds indexing and mask.""" + arr = np.array([[1, 2, 3], [4, 5, 6], [7, 8, 9]]) + mask = np.array([[True, True, False], [True, False, True], [False, True, True]]) + result = masked_load(arr, (slice(0, 3), slice(0, 3)), mask=mask) + result[ + [0, 1, 2], [2, 1, 0] + ] = 0 # manually set masked out positions to 0 for testing + expected = np.array([[1, 2, 0], [4, 0, 6], [0, 8, 9]]) + + print_op_details( + "Test 2: In-bounds indexing with mask", + "load", + arr, + (slice(0, 3), slice(0, 3)), + mask=mask, + output=result, + ) + + assert np.array_equal(result, expected) + + +def test_masked_load_out_of_bounds(): + """Test masked_load with out-of-bounds indexing and mask.""" + arr = np.array([[1, 2], [3, 4]]) + mask = np.array( + [ + [True, False, False], + [False, True, False], + [False, False, False], + [False, False, False], + [False, False, False], + ] + ) # note that arr.shape != mask.shape, this is intended + + # normally this arr[slice] would OOB but the mask at the OOB idxs are false so they're not loaded + result = masked_load(arr, (slice(0, 5), slice(0, 3)), mask=mask) + + print_op_details( + "Test 3: Out-of-bounds indexing with mask", + "load", + arr, + (slice(0, 5), slice(0, 3)), + mask=mask, + output=result, + ) + + assert result.shape == (5, 3) + + +def test_masked_load_integer_oob(): + """Test masked_load with integer out-of-bounds indexing.""" + arr = np.array([10, 20, 30]) + mask = np.array([False]) + result = masked_load(arr, ([5],), mask=mask) # idx 5 is OOB + + print_op_details( + "Test 4: Integer indexing out of bounds", + "load", + arr, + ([5],), + mask=mask, + output=result, + ) + + assert isinstance(result, np.ndarray) + + +def test_masked_load_array_indexing(): + """Test masked_load with array indexing.""" + arr = np.array([100, 200, 300, 400, 500]) + indices = np.array([0, 2, 4]) + mask = np.array([True, True, True]) + result = masked_load(arr, (indices,), mask=mask) + expected = arr[indices] + + print_op_details( + "Test 5: Array indexing", "load", arr, (indices,), mask=mask, output=result + ) + + assert np.array_equal(result, expected) + + +def test_masked_load_mixed_indexing(): + """Test masked_load with mixed indexing types.""" + arr = np.array([[1, 2, 3, 4], [5, 6, 7, 8], [9, 10, 11, 12]]) + mask = np.array([True, False, True, False]) + result = masked_load(arr, (slice(0, 4), 2), mask=mask) # set col 2 + + print_op_details( + "Test 6: Mixed indexing types", + "load", + arr, + (slice(0, 4), 2), + mask=mask, + output=result, + ) + + assert result.shape == (4,) + + +def test_masked_load_complex_indexing(): + """Test masked_load with complex multi-dimensional indexing.""" + arr = np.array( + [ + [ + [1, 2, 3], + [4, 5, 6], + [7, 8, 9], + ], + [ + [11, 12, 13], + [14, 15, 16], + [17, 18, 19], + ], + [ + [21, 22, 23], + [24, 25, 26], + [27, 28, 29], + ], + ] + ) + arr_slice = ( + slice(None, None, None), + None, + np.arange(20)[:, None], + np.arange(3)[None, :], + None, + None, + ) + mask = np.mgrid[:3, :1, :20, :3, :1, :1] + mask = (mask[0] > 0) & (mask[2] < 3) & (mask[3] < 2) + result = masked_load(arr, arr_slice, mask=mask) + UD = np.iinfo(arr.dtype).max # UD = undefined (out of bounds) + expected = np.array( + [ + [ + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + ], + [ + [11, 12, UD], + [14, 15, UD], + [17, 18, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + ], + [ + [21, 22, UD], + [24, 25, UD], + [27, 28, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + ], + ] + )[:, None, :, :, None, None] + + print( + "Load Test 7: Complex multi-dimensional indexing (omitted details, check code to see inputs/outputs)" + ) + + assert np.allclose(result, expected) + + +def test_masked_load_incomplete_slices(): + """Test masked_load with incomplete slices.""" + arr = np.array( + [ + [ + [1, 2, 3], + [4, 5, 6], + [7, 8, 9], + ], + [ + [11, 12, 13], + [14, 15, 16], + [17, 18, 19], + ], + [ + [21, 22, 23], + [24, 25, 26], + [27, 28, 29], + ], + ] + ) + arr_slice = (slice(None, None, None),) + mask = np.mgrid[:3, :3, :3] + mask = (0 < mask[0]) & (mask[0] < 2) & (mask[1] < 2) & (mask[2] != 1) + result = masked_load(arr, arr_slice, mask=mask) + UD = np.iinfo(arr.dtype).max # UD = undefined (out of bounds) + expected = np.array( + [ + [ + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + ], + [ + [11, UD, 13], + [14, UD, 16], + [UD, UD, UD], + ], + [ + [UD, UD, UD], + [UD, UD, UD], + [UD, UD, UD], + ], + ] + ) + + print_op_details( + "Test 8: Incomplete slices", "load", arr, arr_slice, mask=mask, output=result + ) + + assert np.allclose(result, expected) + + +def test_masked_load_index_error_with_true_mask(): + """Test that IndexError is raised when mask=True at out-of-bounds location.""" + arr = np.array( + [ + [ + [1, 2, 3], + [4, 5, 6], + [7, 8, 9], + ], + [ + [11, 12, 13], + [14, 15, 16], + [17, 18, 19], + ], + [ + [21, 22, 23], + [24, 25, 26], + [27, 28, 29], + ], + ] + ) + arr_slice = (slice(None, None, None), slice(0, 4)) + mask = np.mgrid[:3, :4, :3] + mask = (mask[0] < 3) & (mask[1] < 3) & (mask[2] < 3) + mask[-1, -1, -1] = True # index error + + print_op_details( + "Test 9: Index error should still happen when mask=True at OOB location", + "load", + arr, + arr_slice, + mask=mask, + error_expected="IndexError", + ) + + with pytest.raises(IndexError): + masked_load(arr, arr_slice, mask=mask) + print("Success: Correctly raised IndexError\n") + + +def test_masked_load_mask_shape_mismatch(): + """Test that AssertionError is raised when mask shape doesn't match array slice.""" + arr = np.array( + [ + [ + [1, 2, 3], + [4, 5, 6], + [7, 8, 9], + ], + [ + [11, 12, 13], + [14, 15, 16], + [17, 18, 19], + ], + [ + [21, 22, 23], + [24, 25, 26], + [27, 28, 29], + ], + ] + ) + arr_slice = (slice(None, None, None),) + mask = np.mgrid[:3, :4, :3] + mask = (mask[0] < 3) & (mask[1] < 3) & (mask[2] < 3) + + print_op_details( + "Test 10: Error if mask shape wrong", + "load", + arr, + arr_slice, + mask=mask, + error_expected="AssertionError", + ) + + with pytest.raises(AssertionError): + masked_load(arr, arr_slice, mask=mask) + print("Success: Correctly raised AssertionError\n") + + +def test_masked_load_correct_mask_shape(): + """Test masked_load with correctly shaped mask.""" + arr = np.array( + [ + [ + [1, 2, 3], + [4, 5, 6], + [7, 8, 9], + ], + [ + [11, 12, 13], + [14, 15, 16], + [17, 18, 19], + ], + [ + [21, 22, 23], + [24, 25, 26], + [27, 28, 29], + ], + ] + ) + arr_slice = (slice(None, None, None),) + mask = np.arange(27).reshape(3, 3, 3) % 2 == 0 + result = masked_load(arr, arr_slice, mask=mask) + + print_op_details( + "Test 11: Correct mask shape", "load", arr, arr_slice, mask=mask, output=result + ) + + assert result.shape == mask.shape + + +def test_masked_store_none_case(): + """Test masked_store with mask=None (direct indexing).""" + arr = np.array([1, 2, 3, 4, 5]) + values = np.array([10, 20, 30]) + arr_copy = arr.copy() + masked_store(arr_copy, (slice(1, 4),), values, mask=None) + expected = arr.copy() + expected[1:4] = values + + print_op_details( + "Store Test 1: mask=None case", + "store", + arr, + (slice(1, 4),), + values=values, + mask=None, + output=arr_copy, + ) + + assert np.array_equal(arr_copy, expected) + + +def test_masked_store_in_bounds(): + """Test masked_store with in-bounds indexing and mask.""" + arr = np.array([[1, 2, 3], [4, 5, 6], [7, 8, 9]]) + values = np.array([[100, 200, 0], [400, 0, 600], [0, 800, 900]]) + mask = np.array([[True, True, False], [True, False, True], [False, True, True]]) + arr_copy = arr.copy() + masked_store(arr_copy, (slice(0, 3), slice(0, 3)), values, mask=mask) + expected = np.array([[100, 200, 3], [400, 5, 600], [7, 800, 900]]) + + print_op_details( + "Store Test 2: In-bounds indexing with mask", + "store", + arr, + (slice(0, 3), slice(0, 3)), + values=values, + mask=mask, + output=arr_copy, + ) + + assert np.array_equal(arr_copy, expected) + + +def test_masked_store_out_of_bounds(): + """Test masked_store with out-of-bounds indexing and mask.""" + arr = np.array([[1, 2], [3, 4]]) + values = np.array( + [ + [100, 0, 0], + [0, 200, 0], + [0, 0, 0], + [0, 0, 0], + [0, 0, 0], + ] + ) + mask = np.array( + [ + [True, False, False], + [False, True, False], + [False, False, False], + [False, False, False], + [False, False, False], + ] + ) + arr_copy = arr.copy() + masked_store(arr_copy, (slice(0, 5), slice(0, 3)), values, mask=mask) + expected = np.array([[100, 2], [3, 200]]) + + print_op_details( + "Store Test 3: Out-of-bounds indexing with mask", + "store", + arr, + (slice(0, 5), slice(0, 3)), + values=values, + mask=mask, + output=arr_copy, + ) + + assert np.array_equal(arr_copy, expected) + + +def test_masked_store_integer_oob(): + """Test masked_store with integer out-of-bounds indexing.""" + arr = np.array([10, 20, 30]) + values = np.array([999]) + mask = np.array([False]) # mask=False at OOB index + arr_copy = arr.copy() + masked_store(arr_copy, ([5],), values, mask=mask) # Index 5 is OOB + + print_op_details( + "Store Test 4: Integer indexing out of bounds", + "store", + arr, + ([5],), + values=values, + mask=mask, + output=arr_copy, + ) + + assert np.array_equal(arr_copy, arr) + + +def test_masked_store_array_indexing(): + """Test masked_store with array indexing.""" + arr = np.array([100, 200, 300, 400, 500]) + indices = np.array([0, 2, 4]) + values = np.array([111, 333, 555]) + mask = np.array([True, True, True]) + arr_copy = arr.copy() + masked_store(arr_copy, (indices,), values, mask=mask) + expected = np.array([111, 200, 333, 400, 555]) + + print_op_details( + "Store Test 5: Array indexing", + "store", + arr, + (indices,), + values=values, + mask=mask, + output=arr_copy, + ) + + assert np.array_equal(arr_copy, expected) + + +def test_masked_store_mixed_indexing(): + """Test masked_store with mixed indexing types.""" + arr = np.array([[1, 2, 3, 4], [5, 6, 7, 8], [9, 10, 11, 12]]) + values = np.array([100, 200, 300, 400]) + mask = np.array([True, False, True, False]) + arr_copy = arr.copy() + masked_store(arr_copy, (slice(0, 4), 2), values, mask=mask) # set col 2 + expected = np.array([[1, 2, 100, 4], [5, 6, 7, 8], [9, 10, 300, 12]]) + + print_op_details( + "Store Test 6: Mixed indexing types", + "store", + arr, + (slice(0, 4), 2), + values=values, + mask=mask, + output=arr_copy, + ) + + assert np.array_equal(arr_copy, expected) + + +def test_masked_store_complex_indexing(): + """Test masked_store with complex multi-dimensional indexing.""" + arr = np.array( + [ + [ + [1, 2, 3], + [4, 5, 6], + [7, 8, 9], + ], + [ + [11, 12, 13], + [14, 15, 16], + [17, 18, 19], + ], + [ + [21, 22, 23], + [24, 25, 26], + [27, 28, 29], + ], + ] + ) + arr_slice = ( + slice(None, None, None), + None, + np.arange(20)[:, None], + np.arange(3)[None, :], + None, + None, + ) + mask = np.mgrid[:3, :1, :20, :3, :1, :1] + mask = (mask[0] > 0) & (mask[2] < 3) & (mask[3] < 2) + values = np.array( + [ + [ + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + ], + [ + [31, 32, 99], + [34, 35, 99], + [37, 38, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + ], + [ + [41, 42, 99], + [44, 45, 99], + [47, 48, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + [99, 99, 99], + ], + ] + )[:, None, :, :, None, None] + arr_copy = arr.copy() + masked_store(arr_copy, arr_slice, values, mask=mask) + expected_modified = np.array( + [ + [ + [1, 2, 3], + [4, 5, 6], + [7, 8, 9], + ], + [ + [31, 32, 13], + [34, 35, 16], + [37, 38, 19], + ], + [ + [41, 42, 23], + [44, 45, 26], + [47, 48, 29], + ], + ] + ) + + print( + "Store Test 7: Complex multi-dimensional indexing (omitted details, check code to see inputs/outputs)" + ) + + assert np.array_equal(arr_copy, expected_modified) + + +def test_masked_store_oob_with_true_mask(): + """Test that IndexError is raised when mask=True at out-of-bounds location.""" + arr = np.array([[1, 2], [3, 4]]) + values = np.array([[100, 200], [300, 400], [500, 600]]) + mask = np.array( + [ + [True, True], + [True, True], + [True, False], # This position is OOB and mask=True + ] + ) + arr_copy = arr.copy() + + print_op_details( + "Store Test 8: OOB with mask=True should raise IndexError", + "store", + arr, + (slice(0, 3), slice(0, 2)), + values=values, + mask=mask, + error_expected="IndexError", + ) + + with pytest.raises(IndexError): + masked_store(arr_copy, (slice(0, 3), slice(0, 2)), values, mask=mask) + print("Success: Correctly raised IndexError\n") + + +def test_masked_store_values_mask_shape_mismatch(): + """Test that AssertionError is raised when values and mask shapes don't match.""" + arr = np.array([1, 2, 3]) + values = np.array([10, 20, 30]) # Shape (3,) + mask = np.array([True, False]) # Shape (2,) - different from values + arr_copy = arr.copy() + + print_op_details( + "Store Test 9: Values and mask shape mismatch should raise AssertionError", + "store", + arr, + (slice(0, 3),), + values=values, + mask=mask, + error_expected="AssertionError", + ) + + with pytest.raises(AssertionError): + masked_store(arr_copy, (slice(0, 3),), values, mask=mask) + print("Success: Correctly raised AssertionError\n") + + +def test_masked_store_values_shape_mismatch(): + """Test that IndexError is raised when values shape is incompatible.""" + arr = np.array([1, 2, 3]) + values = np.array([10, 20, 30, 40]) # Wrong shape + mask = np.array([True, False, True, False]) + arr_copy = arr.copy() + + print_op_details( + "Store Test 10: Values shape mismatch", + "store", + arr, + (slice(0, 3),), + values=values, + mask=mask, + error_expected="IndexError", + ) + + with pytest.raises(IndexError): + masked_store(arr_copy, (slice(0, 3),), values, mask=mask) + print("Success: Correctly raised IndexError\n") diff --git a/triton_viz/clients/profiler/profiler.py b/triton_viz/clients/profiler/profiler.py index 4822c1dc..4e8ac101 100644 --- a/triton_viz/clients/profiler/profiler.py +++ b/triton_viz/clients/profiler/profiler.py @@ -198,9 +198,7 @@ def _get_mask_stats(mask: TensorHandle) -> tuple[int, int]: false_count = np.count_nonzero(np.logical_not(mask.data)) return total_count, false_count - def pre_load_callback( - ptr, mask, other, cache_modifier, eviction_policy, is_volatile - ): + def pre_load_callback(ptr, mask, keys): self._report_load_store_bytes("load", ptr, mask) if not self.disable_load_mask_percentage_check: total_count, false_count = _get_mask_stats(mask) @@ -229,7 +227,7 @@ def load_overrider( dtype_np = _get_np_dtype(dtype_tt) return TensorHandle(np.zeros_like(ptr.data, dtype=dtype_np), dtype_tt) - def pre_store_callback(ptr, value, mask, cache_modifier, eviction_policy): + def pre_store_callback(ptr, mask, keys): self._report_load_store_bytes("store", ptr, mask) if not self.disable_load_mask_percentage_check: total_count, false_count = _get_mask_stats(mask) diff --git a/triton_viz/clients/sanitizer/sanitizer.py b/triton_viz/clients/sanitizer/sanitizer.py index 1ca9f380..4d80a015 100644 --- a/triton_viz/clients/sanitizer/sanitizer.py +++ b/triton_viz/clients/sanitizer/sanitizer.py @@ -572,9 +572,7 @@ def grid_callback(self, grid: tuple[int, ...]) -> None: self.tensors = sorted(self.tensors, key=lambda x: x.data_ptr()) def register_op_callback(self, op_type: type[Op]) -> OpCallbacks: - def pre_load_callback( - ptr, mask, other, cache_modifier, eviction_policy, is_volatile - ): + def pre_load_callback(ptr, mask, _keys): first_loc = np.unravel_index(np.argmax(mask, axis=None), mask.data.shape) first_ptr = ptr.data[first_loc] tensor = _get_tensor(self.tensors, first_ptr) @@ -582,14 +580,12 @@ def pre_load_callback( self._report(op_type, oob) ptr.data = tensor.data_ptr() + oob["corrected_offsets"] - def pre_store_callback(ptr, value, mask, cache_modifier, eviction_policy): + def pre_store_callback(ptr, mask, _keys): first_loc = np.unravel_index(np.argmax(mask, axis=None), mask.data.shape) first_ptr = ptr.data[first_loc] tensor = _get_tensor(self.tensors, first_ptr) oob = check_out_of_bounds_access(ptr.data, mask.data, tensor) - self._report( - op_type, check_out_of_bounds_access(ptr.data, mask.data, tensor) - ) + self._report(op_type, oob) ptr.data = tensor.data_ptr() + oob["corrected_offsets"] if op_type is Load: @@ -1535,8 +1531,9 @@ def concretize(self_or_cls, obj=None): obj.lhs.concretize(), obj.rhs.concretize(), obj.binary_numpy_op ) elif obj.op == "load": - from ...core.patch import original_ops + from ...core.patch import OPERATION_REGISTRY + original_ops = OPERATION_REGISTRY["triton"]["original_ops"] ptr_concrete = obj.ptr.concretize() # concretize mask # create an all-True mask if mask is None @@ -1564,9 +1561,10 @@ def concretize(self_or_cls, obj=None): else: # Special handling for cast_impl and bitcast operations if obj.op in ("cast_impl", "bitcast"): - from ...core.patch import original_ops + from ...core.patch import OPERATION_REGISTRY from ...core.data import CastImpl, Bitcast + original_ops = OPERATION_REGISTRY["triton"]["original_ops"] src_concrete = obj.src.concretize() # dst_type is stored as a SymbolicExpr const node, need to extract the value dst_type_value = ( diff --git a/triton_viz/clients/tracer/tracer.py b/triton_viz/clients/tracer/tracer.py index c9841838..48c3d513 100644 --- a/triton_viz/clients/tracer/tracer.py +++ b/triton_viz/clients/tracer/tracer.py @@ -1,8 +1,19 @@ from ...core.client import Client from ...core.callbacks import OpCallbacks, ForLoopCallbacks -from ...core.data import Op, Load, Store, ReduceSum, Dot, Grid +from ...core.data import ( + Op, + Load, + Store, + ReduceSum, + Dot, + Grid, + Allocate, + Flip, +) +from triton_viz.core.masked_load import masked_load from typing import Callable, Optional, Union import numpy as np +import traceback def _convert_grid_idx(grid_idx) -> Optional[tuple[int, int, int]]: @@ -71,25 +82,110 @@ def grid_callback(self, grid: tuple[int, ...]): self.tensors = sorted(self.tensors, key=lambda x: x.data_ptr()) def register_op_callback(self, op_type: type[Op]) -> OpCallbacks: - def pre_load_callback( - ptr, mask, other, cache_modifier, eviction_policy, is_volatile - ): + def _extract_user_frames() -> list[traceback.FrameSummary]: + stack: list[traceback.FrameSummary] = list(traceback.extract_stack()) + # drop current frames (this function and callers) + stack = stack[:-2] + cleaned: list[traceback.FrameSummary] = [] + for f in stack: + fn = f.filename.replace("\\", "/") + if any( + s in fn + for s in [ + "triton_viz/core/", + "triton_viz/clients/", + "triton/runtime/", + "triton/language/", + "site-packages/triton/", + "runpy.py", + "IPython", + ] + ): + continue + cleaned.append(f) + if cleaned: + return cleaned + # fallback to last non "<...>" frame + for f in reversed(stack): + if not f.filename.startswith("<"): + return [f] + return stack[-1:] + + def post_allocate_callback(ret): + assert hasattr(ret, "data") + self.tensors.append(ret) + + def _convert_keys_to_numpy(keys): + """Convert any NDArrays in keys to numpy arrays.""" + if isinstance(keys, (tuple, list)): + return tuple(_convert_keys_to_numpy(k) for k in keys) + elif hasattr(keys, "data"): + return keys.data + else: + return keys + + def pre_load_callback(ptr, mask, keys): + if not self.sample: + return + + if keys is None: # i.e. for triton, ptr = pointer + offsets + first_ptr = np.reshape(ptr.data, (-1))[0] + tensor = self._get_tensor(first_ptr) + offsets = ptr.data - tensor.data_ptr() + else: + keys = _convert_keys_to_numpy(keys) + offsets = masked_load(ptr.get_offsets().data, keys, mask=mask.data) + tensor = ptr + + rec = Load(tensor.data_ptr(), offsets, mask.data) + rec.call_path = _extract_user_frames() + self.records.append(rec) + + def pre_store_callback(ptr, mask, keys): + if not self.sample: + return + + if keys is None: # i.e. for triton, ptr = pointer + offsets, so keys=None + first_ptr = np.reshape(ptr.data, (-1))[0] + tensor = self._get_tensor(first_ptr) + offsets = ptr.data - tensor.data_ptr() + mask_data = mask.data + else: + keys = _convert_keys_to_numpy(keys) + if mask is None: + offsets = masked_load(ptr.get_offsets().data, keys) + mask_data = np.ones_like(offsets).astype(bool) + else: + mask_data = mask.data + offsets = masked_load(ptr.get_offsets().data, keys, mask=mask_data) + tensor = ptr + + rec = Store(tensor.data_ptr(), offsets, mask_data) + rec.call_path = _extract_user_frames() + self.records.append(rec) + + # Raw (unmasked) ops: synthesize a full True mask based on ptr shape + def pre_raw_load_callback(ptr): if not self.sample: return first_ptr = np.reshape(ptr.data, (-1))[0] tensor = self._get_tensor(first_ptr) - self.records.append( - Load(tensor.data_ptr(), ptr.data - tensor.data_ptr(), mask.data) - ) + offsets = ptr.data - tensor.data_ptr() + true_mask = np.ones_like(offsets, dtype=bool) + rec = Load(tensor.data_ptr(), offsets, true_mask) + rec.call_path = _extract_user_frames() + self.records.append(rec) - def pre_store_callback(ptr, value, mask, cache_modifier, eviction_policy): + def pre_raw_store_callback(ptr, value): if not self.sample: return first_ptr = np.reshape(ptr.data, (-1))[0] tensor = self._get_tensor(first_ptr) - self.records.append( - Store(tensor.data_ptr(), ptr.data - tensor.data_ptr(), mask.data) - ) + offsets = ptr.data - tensor.data_ptr() + true_mask = np.ones_like(offsets, dtype=bool) + rec = Store(tensor.data_ptr(), offsets, true_mask) + rec.call_path = _extract_user_frames() + self.records.append(rec) def post_reduce_sum_callback(ret, input, axis=None, keep_dims=False): if not self.sample: @@ -98,17 +194,41 @@ def post_reduce_sum_callback(ret, input, axis=None, keep_dims=False): output_shape = ret.handle.data.shape self.records.append(ReduceSum(input_shape, axis, keep_dims, output_shape)) - def post_dot_callback(ret, input, other, *args): + def post_dot_callback(ret, input, other): if not self.sample: return input_shape = input.data.shape other_shape = other.data.shape ret_shape = ret.data.shape - self.records.append( - Dot(input_shape, other_shape, ret_shape, input.data, other.data) - ) + # Pass input/other raw arrays so draw.py can render MatMul + rec = Dot(input_shape, other_shape, ret_shape, input.data, other.data) + rec.call_path = _extract_user_frames() + self.records.append(rec) - if op_type is Load: + def post_flip_callback(ret, x, *args, **kwargs): + if not self.sample: + return + # Try to capture dim argument + dim = None + if args: + dim = args[0] + if "dim" in kwargs: + dim = kwargs.get("dim") + try: + in_shape = tuple(x.data.shape) + out_shape = tuple(ret.data.shape) + except Exception: + in_shape = getattr(getattr(x, "handle", None), "data", None) + out_shape = getattr(getattr(ret, "handle", None), "data", None) + in_shape = tuple(getattr(in_shape, "shape", []) or []) + out_shape = tuple(getattr(out_shape, "shape", []) or []) + rec = Flip(in_shape, out_shape, int(dim) if dim is not None else 0) + rec.call_path = _extract_user_frames() + self.records.append(rec) + + if op_type is Allocate: + return OpCallbacks(after_callback=post_allocate_callback) + elif op_type is Load: return OpCallbacks(before_callback=pre_load_callback) elif op_type is Store: return OpCallbacks(before_callback=pre_store_callback) @@ -116,6 +236,8 @@ def post_dot_callback(ret, input, other, *args): return OpCallbacks(after_callback=post_reduce_sum_callback) elif op_type is Dot: return OpCallbacks(after_callback=post_dot_callback) + # Flip is wrapped at tl.flip; we don't have an interpreter op to hook here. + # The wrapper in patch_lang will append Flip records directly to tracer. return OpCallbacks() diff --git a/triton_viz/core/__init__.py b/triton_viz/core/__init__.py index e6a5e16b..73512549 100644 --- a/triton_viz/core/__init__.py +++ b/triton_viz/core/__init__.py @@ -26,6 +26,8 @@ Rsqrt, CastImpl, ) +from .masked_load import masked_load, masked_store +from .nki_extract_slice import StoreCallTransformer, transform_code __all__ = [ "trace", @@ -55,4 +57,8 @@ "Idiv", "Rsqrt", "CastImpl", + "masked_load", + "masked_store", + "StoreCallTransformer", + "transform_code", ] diff --git a/triton_viz/core/client.py b/triton_viz/core/client.py index 60fb8e1e..c5d9c2ed 100644 --- a/triton_viz/core/client.py +++ b/triton_viz/core/client.py @@ -10,12 +10,11 @@ unpatch_op, patch_for_loop, unpatch_for_loop, - op_list, patch_calls, ) from functools import wraps from .callbacks import OpCallbacks, ForLoopCallbacks -from .patch import patch_lang, unpatch_lang +from .patch import patch_lang, unpatch_lang, OPERATION_REGISTRY class Client(ABC): @@ -54,7 +53,9 @@ def grid_idx_callback(self, grid_idx: tuple[int, ...]): ... @abstractmethod - def register_op_callback(self, op_type: type[Op]) -> OpCallbacks: + def register_op_callback( + self, op_type: type[Op], *args: Any, **kwargs: Any + ) -> OpCallbacks: ... @abstractmethod @@ -125,27 +126,31 @@ def wrapped(*args, **kwargs): jit_fn.warmup = jit_fn.warmup.__wrapped__ @contextmanager - def patch_run(self, fn): - with patch_calls(): + def patch_run(self, fn, backend: str): + with patch_calls(backend): for client in self.clients.values(): - for op in op_list: - # patch ops + # get operations for the specified backend + backend_ops: list[type[Op]] = OPERATION_REGISTRY[backend]["op_list"] + for op in backend_ops: + # patch ops callbacks = client.register_op_callback(op) - patch_op(op, callbacks) + patch_op(op, callbacks, backend=backend) # patch for loops loop_callbacks = client.register_for_loop_callback() patch_for_loop(loop_callbacks) # Remaps core language functions to interpreted ones - patch_lang(fn) + patch_lang(fn, backend) try: yield finally: - for op in op_list: - unpatch_op(op) + backend_ops = OPERATION_REGISTRY[backend]["op_list"] + + for op in backend_ops: + unpatch_op(op, backend) unpatch_for_loop() - unpatch_lang() + unpatch_lang(backend) def pre_run_callback(self, fn: Callable) -> bool: rets = [] diff --git a/triton_viz/core/data.py b/triton_viz/core/data.py index 0c19e1e4..4f66c0f6 100644 --- a/triton_viz/core/data.py +++ b/triton_viz/core/data.py @@ -16,17 +16,26 @@ @dataclass class Op: name: ClassVar[str] = "op" - call_path: list[traceback.StackSummary] = field(init=False, default_factory=list) + call_path: list[traceback.FrameSummary] = field(init=False, default_factory=list) def __post_init__(self): - self.call_path = traceback.extract_stack()[:-2] + full_stack = traceback.extract_stack()[:-2] + # keep original for fallback + self.call_path = full_stack clean_call_path = [] - for frame in self.call_path: + for frame in full_stack: if not any( triton_frame in frame.filename for triton_frame in TRITON_FRAMES ): clean_call_path.append(frame) - self.call_path = clean_call_path + # if filtering removed all frames, fallback to last meaningful frame(s) + if clean_call_path: + self.call_path = clean_call_path + else: + for frame in reversed(full_stack): + if not str(frame.filename).startswith("<"): + self.call_path = [frame] + break @dataclass @@ -34,6 +43,12 @@ class ProgramId(Op): name: ClassVar[str] = "program_id" +@dataclass +class Allocate(Op): + name: ClassVar[str] = "allocate" + ptr: int + + @dataclass class RawStore(Op): name: ClassVar[str] = "raw_store" @@ -45,6 +60,11 @@ class Store(Op): ptr: int offsets: npt.NDArray[np.int_] masks: npt.NDArray[np.bool_] + mem_src: str = "SBUF" + mem_dst: str = "HBM" + backend: str = "nki" + bytes: int = 0 + time_idx: int = 0 @dataclass @@ -58,6 +78,12 @@ class Load(Op): ptr: int offsets: npt.NDArray[np.int_] masks: npt.NDArray[np.bool_] + # buffer: str + mem_src: str = "HBM" + mem_dst: str = "SBUF" + backend: str = "nki" + bytes: int = 0 + time_idx: int = 0 @dataclass @@ -95,6 +121,17 @@ def update_intermediate(self, row: int, col: int, result: float): self.intermediate_results[(row, col)] = result +@dataclass +class Flip(Op): + name: ClassVar[str] = "flip" + input_shape: tuple + output_shape: tuple + dim: int + # Optional payloads to help frontend render actual values when available + input_data: list | None = None + output_data: list | None = None + + @dataclass class MakeRange(Op): name: ClassVar[str] = "make_range" diff --git a/triton_viz/core/masked_load.py b/triton_viz/core/masked_load.py new file mode 100644 index 00000000..f13ff19a --- /dev/null +++ b/triton_viz/core/masked_load.py @@ -0,0 +1,185 @@ +import numpy as np + +get = lambda x, default: default if x is None else x + + +def normalize_slice(ndarray: np.ndarray, keys: tuple) -> tuple[tuple, tuple]: + """Separate singleton dims (None) and add bounds to slices like :2, 2:, or :""" + singleton_dims = [] + dim_idx = 0 + new_keys = [] + for key_dim, k in enumerate(keys): + if k is None: + singleton_dims.append(key_dim) + continue + + arr_dim = ndarray.shape[dim_idx] + if isinstance(k, slice): + # add bounds to [:, :N, N:]-type slices + start = get(k.start, 0) + stop = get(k.stop, arr_dim) + step = get(k.step, 1) + if start < 0: + start = max(0, arr_dim + start) + if stop < 0: + stop = max(0, arr_dim + stop) + k = slice(start, stop, step) + + new_keys.append(k) + dim_idx += 1 + return tuple(new_keys), tuple(singleton_dims) + + +def _calculate_target_shape( + keys: tuple, original_shape: tuple[int, ...] +) -> tuple[int, ...]: + """Calculate the shape needed for the padded array to handle indexing.""" + target_shape = list(original_shape) + for i, (arr_dim, key) in enumerate(zip(original_shape, keys)): + if isinstance(key, slice): + target_shape[i] = max(arr_dim, key.stop) + elif isinstance(key, (int, np.integer)): + target_shape[i] = max(arr_dim, int(key) + 1) + elif isinstance(key, (list, np.ndarray)): + target_shape[i] = max(arr_dim, np.max(key) + 1) + return tuple(target_shape) + + +def _get_valid_indices( + keys: tuple, original_shape: tuple[int, ...], result_shape: tuple[int, ...] +) -> np.ndarray: + """Get a boolean mask indicating which result indices are within original array bounds.""" + coords = np.mgrid[[slice(0, s) for s in result_shape]] + valid_mask = np.ones(result_shape, dtype=bool) + + for key_idx, (arr_dim, key) in enumerate(zip(original_shape, keys)): + if isinstance(key, slice): + start = get(key.start, 0) + step = get(key.step, 1) + original_coords = coords[key_idx] * step + start + valid_mask &= (0 <= original_coords) & (original_coords < arr_dim) + elif isinstance(key, (list, np.ndarray)): + valid_indices = np.array(key) + valid_mask &= (0 <= valid_indices) & (valid_indices < arr_dim) + elif isinstance(key, (int, np.integer)): + valid_mask &= 0 <= int(key) < arr_dim + + return valid_mask + + +def masked_load( + ndarray: np.ndarray, keys: tuple, mask: np.ndarray = None +) -> np.ndarray: + """ + Load array elements with masking for out-of-bounds errors. + + Args: + ndarray: Input numpy array + keys: Indexing keys (tuple of slices, integers, arrays, etc.) + mask: Boolean mask array. If None, returns direct indexing. + Where mask is False at OOB indices, garbage values are kept. + Mask shape should match the shape of the result, not the input array. + + Returns: + Indexed array with masked error handling + """ + try: # fast path case - if keys aren't OOB, just go with that + out = ndarray[keys].copy() + if mask is None: + return out + out[~mask] = np.iinfo(ndarray.dtype).max + return out + except Exception: + pass + + # Convert keys to tuple if it's not already + if not isinstance(keys, tuple): + keys = (keys,) + + # normalize (remove Nones in slices) + keys, singleton_dims = normalize_slice(ndarray, keys) + target_shape = _calculate_target_shape(keys, ndarray.shape) + padded_array = np.empty(target_shape, dtype=ndarray.dtype) + + # Calculate the region where we can safely copy the original array + copy_slices = [] + for i in range(len(ndarray.shape)): + copy_slices.append(slice(0, ndarray.shape[i])) + padded_array[tuple(copy_slices)] = ndarray[tuple(copy_slices)] + + result = padded_array[keys] + assert np.expand_dims(result, singleton_dims).shape == mask.shape + + mask = np.squeeze(mask, singleton_dims) + + # Determine which indices are actually within original array bounds + in_bounds_mask = _get_valid_indices(keys, ndarray.shape, result.shape) + + # Check if there are any OOB indices where mask=True + oob_with_true_mask = (~in_bounds_mask) & mask + oob_coords = np.where(oob_with_true_mask) + if len(oob_coords[0]) > 0: + # Get the first OOB coordinate + oob_idx = tuple(coord[0] for coord in oob_coords) + raise IndexError( + f"index {oob_idx} is out of bounds for array of size {ndarray.shape}" + ) + + valid_mask = mask & in_bounds_mask + # Use appropriate max value based on dtype + if np.issubdtype(ndarray.dtype, np.integer): + result[~valid_mask] = np.iinfo(ndarray.dtype).max + elif np.issubdtype(ndarray.dtype, np.floating): + result[~valid_mask] = np.finfo(ndarray.dtype).max + else: + result[~valid_mask] = 0 # false + + return np.expand_dims(result, singleton_dims) + + +def masked_store( + ndarray: np.ndarray, keys: tuple, value: np.ndarray, mask: np.ndarray = None +) -> None: + """ + Store array elements with masking for out-of-bounds errors. + General idea of this procedure is to get all the indices that the slice + mask needs for a single advanced indexing assign. + - to handle OOB accesses we infer the min possible array size where array[keys] wouldn't have any OOBs + - we can get the indices for each dim by selecting from an mgrid + + Args: + ndarray: Input numpy array to store values into + keys: Indexing keys (tuple of slices, integers, arrays, etc.) + value: Values to store + mask: Boolean mask array. If None, performs direct indexing. + Where mask is False at OOB indices, values are not stored. + Mask shape should match the shape of the result, not the input array. + + Returns: + None (modifies ndarray in-place) + """ + # Handle mask=None case + if mask is None: + ndarray[keys] = value + return + + assert value.shape == mask.shape + + # Convert keys to tuple if it's not already + if not isinstance(keys, tuple): + keys = (keys,) + + flat_mask = mask.ravel() # can only index with bool tensors if they're 1d + + # normalize (remove Nones in slices) + keys, singleton_dims = normalize_slice(ndarray, keys) + target_shape = _calculate_target_shape(keys, ndarray.shape) + mgrid = np.mgrid[tuple(slice(0, dim, 1) for dim in target_shape)] + idxs = [] + for mgrid_dim in mgrid: + # get the section that would be sliced if the array was big enough + sliced_section = mgrid_dim[keys] + + # get just the section where mask=True + dim_idxs = sliced_section.ravel()[flat_mask] + idxs.append(dim_idxs) + ndarray[tuple(idxs)] = value.ravel()[flat_mask] diff --git a/triton_viz/core/nki.py b/triton_viz/core/nki.py new file mode 100644 index 00000000..4517a251 --- /dev/null +++ b/triton_viz/core/nki.py @@ -0,0 +1,549 @@ +import numpy as np + +try: + import neuronxcc.nki.language as nl +except ( + ModuleNotFoundError +) as exc: # pragma: no cover - only hit when optional deps missing + raise ModuleNotFoundError( + "NeuronX dependencies are missing. Install triton-viz[nki] to enable the NKI interpreter." + ) from exc +import inspect +from .nki_extract_slice import transform_code +from .masked_load import masked_load, masked_store + + +class NDArray: + def __init__(self, buffer=None, name="", **kwargs): + self.buffer = buffer + self.name = name + self.kwargs = kwargs + val = None + if "shape" in kwargs and "dtype" in kwargs: + shape = kwargs.pop("shape") + dtype = kwargs.pop("dtype") + val = np.ndarray(shape, dtype=dtype) + if "value" in kwargs: + assert val is None or val.shape == kwargs["value"].shape + val = kwargs["value"] + self._data_ptr = None + self.data = val + + @property + def shape(self): + return self.data.shape if self.data is not None else None + + @property + def dtype(self): + return self.data.dtype if self.data is not None else None + + def data_ptr(self): + if self._data_ptr is None: + self._data_ptr = self.data.ctypes.data + return self._data_ptr + + def stride(self): + return self.data.strides + + def element_size(self): + return self.dtype.itemsize + + def cpu(self): + return self + + def detach(self): + return self + + def numpy(self): + return self.data + + def get_offsets(self): + """ + Generate offset arrays for each dimension based on shape and stride. + Given array with shape (A, ..., Z) and strides (a, ..., z), return offsets: + a * arange(A)[:, None, ..., None] + ... + z * arange(Z)[None, None, ..., :] + """ + offsets = 0 + for dim_size, stride in zip(self.shape, self.stride()): + offsets = np.expand_dims(offsets, -1) + np.arange(dim_size) * stride + return NDArray(value=offsets, name=self.name) + + def __repr__(self): + return f"NDArray(shape={self.shape}, dtype={self.dtype}, name={self.name})" + + def __getitem__(self, keys): + """Implement slicing operations for NDArray""" + if self.data is None: + raise AttributeError("NDArray has no value to slice") + if not isinstance(keys, tuple): + keys = (keys,) + + # Apply the slicing to the underlying numpy array + new_keys = [k.data if isinstance(k, NDArray) else k for k in keys] + sliced_value = self.data[tuple(new_keys)] + + # Create a new NDArray with the sliced data + return NDArray(value=sliced_value, name=f"{self.name}_slice") + + def __setitem__(self, keys, value): + if not isinstance(keys, tuple): + keys = (keys,) + + # Apply the slicing to the underlying numpy array + new_keys = [k.data if isinstance(k, NDArray) else k for k in keys] + self.data[tuple(new_keys)] = value.data + + return self + + def _binary_op(self, other, op_func, op_name, op_symbol): + if isinstance(other, NDArray): + return NDArray( + value=op_func(self.data, other.data), + name=f"{self.name}_{op_name}_{other.name}", + ) + elif np.isscalar(other): + return NDArray( + value=op_func(self.data, other), name=f"{self.name}_{op_name}_scalar" + ) + raise TypeError( + f"Unsupported operand type(s) for {op_symbol}: 'NDArray' and '{type(other).__name__}'" + ) + + def _rbinary_op(self, other, op_func, op_name, op_symbol): + if isinstance(other, NDArray): + return NDArray( + value=op_func(other.data, self.data), + name=f"{other.name}_{op_name}_{self.name}", + ) + elif np.isscalar(other): + return NDArray( + value=op_func(other, self.data), name=f"scalar_{op_name}_{self.name}" + ) + raise TypeError( + f"Unsupported operand type(s) for {op_symbol}: '{type(other).__name__}' and 'NDArray'" + ) + + # Define operator +/-/*// + def __add__(self, other): + return self._binary_op(other, lambda a, b: a + b, "add", "+") + + def __radd__(self, other): + return self._rbinary_op(other, lambda a, b: a + b, "add", "+") + + def __sub__(self, other): + return self._binary_op(other, lambda a, b: a - b, "sub", "-") + + def __rsub__(self, other): + return self._rbinary_op(other, lambda a, b: a - b, "sub", "-") + + def __mul__(self, other): + return self._binary_op(other, lambda a, b: a * b, "mul", "*") + + def __rmul__(self, other): + return self._rbinary_op(other, lambda a, b: a * b, "mul", "*") + + def __truediv__(self, other): + return self._binary_op(other, lambda a, b: a / b, "div", "/") + + def __rtruediv__(self, other): + return self._rbinary_op(other, lambda a, b: a / b, "div", "/") + + def __lt__(self, other): + return self._binary_op(other, lambda a, b: a < b, "lt", "<") + + def __gt__(self, other): + return self._binary_op(other, lambda a, b: a > b, "gt", ">") + + def __le__(self, other): + return self._binary_op(other, lambda a, b: a <= b, "le", "<=") + + def __ge__(self, other): + return self._binary_op(other, lambda a, b: a >= b, "ge", ">=") + + def __and__(self, other): + return self._binary_op(other, lambda a, b: a & b, "and", "&") + + def __or__(self, other): + return self._binary_op(other, lambda a, b: a | b, "or", "|") + + def reshape(self, *args, **kwargs): + return NDArray( + value=self.data.reshape(*args), name=f"{self.name}_reshape", **kwargs + ) + + def broadcast_to(self, *args, **kwargs): + return NDArray( + value=np.broadcast_to(self.data, *args), + name=f"{self.name}_broadcast_to", + **kwargs, + ) + + +class Builder: + def __init__(self, grid_dims=None): + # TODO: infinite grid dims for NKI + self.grid_dims = grid_dims if grid_dims is not None else (1, 1, 1) + self.grid_x = None + self.grid_y = None + self.grid_z = None + self.fn = None + self.shared_hbm_arrays = {} + + def set_grid_dim(self, *grid_dims): + self.grid_dims = grid_dims + + def set_grid_idx(self, x, y, z): + self.grid_x = x + self.grid_y = y + self.grid_z = z + + def ndarray(self, shape, dtype, *, buffer=None, name=None, **kwargs): + if buffer == nl.shared_hbm: + if name is None: + # file name + function name + line number + frame = inspect.currentframe().f_back + file_name = frame.f_code.co_filename + function_name = frame.f_code.co_name + line_number = frame.f_lineno + name = f"{file_name}_{function_name}_{line_number}" + if name in self.shared_hbm_arrays: + # Return the existing shared HBM array + ret = self.shared_hbm_arrays[name] + else: + # Create a new shared HBM array and store it + ret = NDArray( + buffer=buffer, name=name, shape=shape, dtype=dtype, **kwargs + ) + self.shared_hbm_arrays[name] = ret + else: + ret = NDArray(buffer=buffer, name=name, shape=shape, dtype=dtype, **kwargs) + return ret + + def zeros(self, shape, dtype, *, buffer=None, name=None, **kwargs): + value = np.zeros(shape, dtype=dtype) + return self.ndarray( + shape, dtype, buffer=buffer, name=name, value=value, **kwargs + ) + + def arange(self, *args): + return NDArray(value=np.arange(*args)) + + def program_id(self, axis: int): + if axis == 0: + return self.grid_x + elif axis == 1: + return self.grid_y + elif axis == 2: + return self.grid_z + else: + raise ValueError(f"Invalid axis: {axis}. Must be 0, 1, or 2.") + + def load(self, src: NDArray, *, mask=None, dtype=None, **kwargs): + value = src.data + if isinstance(mask, NDArray): + value = value[mask.data] + elif mask is not None: + value = value[mask] + if dtype is not None: + value = value.astype(dtype) + + mask_value = getattr(mask, "data", np.ones_like(src)) + new_shape = [] + for i, _v in enumerate(mask_value.shape): + assert np.unique(mask_value.sum(i)).size <= 2 + new_dim = mask_value.sum(i).flatten()[0] + new_shape.append(new_dim) + + value = value.reshape(new_shape) + return NDArray(value=value, name=src.name, **kwargs) + + def load_transpose2d(self, src: NDArray, *, mask=None, dtype=None, **kwargs): + # THTODO + value = src.data + return self.load( + NDArray(value=value.T, name=src.name), mask=mask, dtype=dtype, **kwargs + ) + + def store(self, dst: NDArray, value: NDArray, *, mask=None, **kwargs): + dst.data[mask.data] = value.data.ravel() + return dst + + def _convert_keys_to_numpy(self, keys): + """Convert any NDArrays in keys to numpy arrays.""" + if isinstance(keys, (tuple, list)): + return tuple(self._convert_keys_to_numpy(k) for k in keys) + elif isinstance(keys, NDArray): + return keys.data + else: + return keys + + def masked_load(self, src: NDArray, keys, *, mask=None, **kwargs): + """Load array elements with masking for out-of-bounds errors.""" + # Convert NDArray to numpy array + ndarray = src.data + mask_value = getattr(mask, "data", mask) if mask is not None else None + + # Convert any NDArrays in keys to numpy arrays + numpy_keys = self._convert_keys_to_numpy(keys) + + # Call the actual masked_load function + result = masked_load(ndarray, numpy_keys, mask=mask_value) + + # Convert result back to NDArray + return NDArray(value=result, name=f"{src.name}_masked_load", **kwargs) + + def masked_store(self, dst: NDArray, keys, value: NDArray, *, mask=None, **kwargs): + """Store array elements with masking for out-of-bounds errors.""" + # Convert NDArrays to numpy arrays + ndarray = dst.data + value_array = value.data + mask_value = getattr(mask, "data", mask) if mask is not None else None + + # Convert any NDArrays in keys to numpy arrays + numpy_keys = self._convert_keys_to_numpy(keys) + + # Call the actual masked_store function + masked_store(ndarray, numpy_keys, value_array, mask=mask_value) + + return dst + + def _unary_op(self, x: NDArray, np_func, op_name, **kwargs): + return NDArray(value=np_func(x.data), name=f"{x.name}_{op_name}", **kwargs) + + # Elementwise operator implementations + def exp(self, x: NDArray, **kwargs): + return self._unary_op(x, np.exp, "exp", **kwargs) + + def relu(self, x: NDArray, **kwargs): + return self._unary_op(x, lambda v: np.maximum(v, 0), "relu", **kwargs) + + def sigmoid(self, x: NDArray, **kwargs): + return self._unary_op(x, lambda v: 1 / (1 + np.exp(-v)), "sigmoid", **kwargs) + + def tanh(self, x: NDArray, **kwargs): + return self._unary_op(x, np.tanh, "tanh", **kwargs) + + def silu(self, x: NDArray, **kwargs): + # SiLU(x) = x * sigmoid(x) + sigmoid_x = 1 / (1 + np.exp(-x.data)) + return NDArray(value=x.data * sigmoid_x, name=f"{x.name}_silu", **kwargs) + + def gelu(self, x: NDArray, **kwargs): + # GELU(x) = 0.5 * x * (1 + tanh(sqrt(2/pi) * (x + 0.044715 * x^3))) + sqrt_2_pi = np.sqrt(2 / np.pi) + inner = sqrt_2_pi * (x.data + 0.044715 * np.power(x.data, 3)) + return NDArray( + value=0.5 * x.data * (1 + np.tanh(inner)), name=f"{x.name}_gelu", **kwargs + ) + + def sqrt(self, x: NDArray, **kwargs): + return self._unary_op(x, np.sqrt, "sqrt", **kwargs) + + def abs(self, x: NDArray, **kwargs): + return self._unary_op(x, np.abs, "abs", **kwargs) + + def log(self, x: NDArray, **kwargs): + return self._unary_op(x, np.log, "log", **kwargs) + + def pow(self, x: NDArray, exponent, **kwargs): + if isinstance(exponent, NDArray): + return NDArray( + value=np.power(x.data, exponent.data), + name=f"{x.name}_pow_{exponent.name}", + **kwargs, + ) + elif np.isscalar(exponent): + return NDArray( + value=np.power(x.data, exponent), + name=f"{x.name}_pow_{exponent}", + **kwargs, + ) + else: + raise TypeError(f"Unsupported exponent type: {type(exponent)}") + + def reciprocal(self, x: NDArray, **kwargs): + return self._unary_op(x, lambda v: 1 / v, "reciprocal", **kwargs) + + def matmul(self, x: NDArray, y: NDArray, transpose_x=False, mask=None, **kwargs): + x_value = x.data + if transpose_x: + x_value = x_value.T + y_value = y.data + return NDArray( + value=(x_value @ y_value), name=f"{x.name}_{y.name}_matmul", **kwargs + ) + + def copy(self, x: NDArray, **kwargs): + return self._unary_op(x, np.copy, "copy", **kwargs) + + def sum(self, x: NDArray, *args, mask=None, **kwargs): + if mask is not None: + kwargs["where"] = mask.data + return NDArray( + value=x.data.sum(*args, **kwargs), name=f"{x.name}_sum", **kwargs + ) + + def square(self, x: NDArray, **kwargs): + return self._unary_op(x, np.square, "square", **kwargs) + + def rsqrt(self, x: NDArray, **kwargs): + return self._unary_op(x, lambda v: 1 / np.sqrt(v), "rsqrt", **kwargs) + + def multiply(self, x: NDArray, y: NDArray, **kwargs): + if isinstance(y, NDArray): + return NDArray( + value=np.multiply(x.data, y.data), + name=f"{x.name}_multiply_{y.name}", + **kwargs, + ) + elif np.isscalar(y): + return NDArray( + value=np.multiply(x.data, y), + name=f"{x.name}_multiply_scalar", + **kwargs, + ) + else: + raise TypeError(f"Unsupported type for multiply: {type(y)}") + + def range(self, stop): + return range(stop) + + +nki_builder = Builder() + + +def nki_patch_lang(): + nl.ndarray = nki_builder.ndarray + nl.program_id = nki_builder.program_id + nl.arange = nki_builder.arange + nl.load = nki_builder.load + nl.store = nki_builder.store + + # Also expose masked_load and masked_store functions + nl.masked_load = nki_builder.masked_load + nl.masked_store = nki_builder.masked_store + # see https://awsdocs-neuron.readthedocs-hosted.com/en/latest/general/nki/api/nki.language.html + + # TODO: implement + # matmul-specific + # nl.shared_hbm + # nl.psum + nl.affine_range = nki_builder.range + nl.par_dim + nl.zeros = nki_builder.zeros + nl.mgrid = NDArray(value=np.mgrid, buffer=nl.sbuf, name="mgrid") + nl.matmul = nki_builder.matmul + nl.copy = nki_builder.copy + nl.sum = nki_builder.sum + nl.square = nki_builder.square + nl.rsqrt = nki_builder.rsqrt + nl.multiply = nki_builder.multiply + + # attention-specific + nl.load_transpose2d = nki_builder.load_transpose2d + # nisa.affine_select + # nl.tensor_reduce + # nisa.activation + nl.broadcast_to + # nisa.nc_transpose + + # Elementwise operators + nl.exp = nki_builder.exp + nl.relu = nki_builder.relu + nl.sigmoid = nki_builder.sigmoid + nl.tanh = nki_builder.tanh + nl.silu = nki_builder.silu + nl.gelu = nki_builder.gelu + nl.sqrt = nki_builder.sqrt + nl.abs = nki_builder.abs + nl.log = nki_builder.log + nl.pow = nki_builder.pow + nl.reciprocal = nki_builder.reciprocal + + nl.device_print = print + + +def nki_unpatch_lang(): + # reload the original functions + import importlib + + importlib.reload(nl) + + +class NKIInterpretedFunction: + def __init__(self, fn): + self.fn = fn + + def run(self, *args, **kwargs): + grid_dims = kwargs.pop( + "grid", (1, 1, 1) + ) # Remove grid from kwargs to avoid passing it to the function + # make it 3d if not + if len(grid_dims) == 1: + grid_dims = (grid_dims[0], 1, 1) + elif len(grid_dims) == 2: + grid_dims = (grid_dims[0], grid_dims[1], 1) + elif len(grid_dims) != 3: + raise ValueError( + f"Grid must be 1, 2, or 3 dimensions, got {len(grid_dims)}" + ) + nki_builder.set_grid_dim(*grid_dims) + nki_builder.shared_hbm_arrays = {} + nki_builder.fn = self.fn + + kwargs.pop("warmup", None) # Remove warmup from kwargs if it exists + client_manager = kwargs.pop( + "client_manager", None + ) # Remove client_manager from kwargs if it exists + + # Call grid_callback once before grid execution (similar to Triton) + if client_manager is not None: + client_manager.grid_callback(grid_dims) + + # Apply AST transformer to convert nl.load/nl.store calls to nl.masked_load/nl.masked_store + if hasattr(self.fn, "__code__"): + # Get the source code of the function + source_code = inspect.getsource(self.fn) + # Transform the source code using the AST transformer + transformed_code = transform_code(source_code) + # Create a new function from the transformed code + exec_globals = self.fn.__globals__.copy() + import random + import string + import os + + rand_str = "".join( + random.choices(string.ascii_letters + string.digits, k=16) + ) + os.makedirs("/tmp/triton-viz", exist_ok=True) + filename = f"/tmp/triton-viz/{rand_str}.py" + with open(filename, "w") as f: + f.write(transformed_code) + code_obj = compile(transformed_code, filename=filename, mode="exec") + exec(code_obj, exec_globals) + self.fn = exec_globals[self.fn.__name__] + + # convert args to NDArray if they are not already + args = [arg if isinstance(arg, NDArray) else NDArray(value=arg) for arg in args] + + name_args = inspect.getcallargs(self.fn, *args) + call_args = {} + for name, arg in name_args.items(): + call_args[name] = arg + ret = arg + client_manager.arg_callback(name, arg, ret) + + for x in range(grid_dims[0]): + for y in range(grid_dims[1]): + for z in range(grid_dims[2]): + nki_builder.set_grid_idx(x, y, z) + + # Call grid_idx_callback for each grid iteration (similar to Triton) + if client_manager is not None: + client_manager.grid_idx_callback((x, y, z)) + + if not client_manager.pre_run_callback(self.fn): + return + self.fn(*args, **kwargs) + if not client_manager.post_run_callback(self.fn): + return diff --git a/triton_viz/core/nki_extract_slice.py b/triton_viz/core/nki_extract_slice.py new file mode 100644 index 00000000..5b7b9188 --- /dev/null +++ b/triton_viz/core/nki_extract_slice.py @@ -0,0 +1,162 @@ +import ast + + +class StoreCallTransformer(ast.NodeTransformer): + """ + A targeted AST transformer to rewrite `nl.store(x[...], ...)` calls + into `masked_store(x, slice_obj, ...)` and `nl.load(x[...])` calls + into `masked_load(x, slice_obj)` with a preceding assignment. + """ + + def visit_Expr(self, node: ast.Expr) -> ast.AST | list[ast.AST]: + """ + Intercept and transform expression statements. + """ + # We only care about expressions that are function calls. + if not isinstance(node.value, ast.Call): + return self.generic_visit(node) + + call_node = node.value + + if not ( + isinstance(call_node.func, ast.Attribute) + and isinstance(call_node.func.value, ast.Name) + and call_node.func.value.id == "nl" + and call_node.func.attr in ("store", "load") + ): + return self.generic_visit(node) + + if not call_node.args or not isinstance(call_node.args[0], ast.Subscript): + return self.generic_visit(node) + + subscript_node = call_node.args[0] + sliced_object = subscript_node.value + slice_content = subscript_node.slice + remaining_args = call_node.args[1:] + + # Convert to nl.masked_load or nl.masked_store + func_name = "masked_" + call_node.func.attr + + new_call = ast.Call( + func=ast.Attribute( + value=ast.Name(id="nl", ctx=ast.Load()), attr=func_name, ctx=ast.Load() + ), + args=[ + sliced_object, + self._create_slice_value_node(slice_content), + *remaining_args, + ], + keywords=call_node.keywords, + ) + + return ast.Expr(value=new_call) + + def visit_Assign(self, node: ast.Assign) -> ast.AST: + """ + Handle assignment statements where the value is an nl.load or nl.store call. + """ + if not isinstance(node.value, ast.Call): + return self.generic_visit(node) + + call_node = node.value + + if not ( + isinstance(call_node.func, ast.Attribute) + and isinstance(call_node.func.value, ast.Name) + and call_node.func.value.id == "nl" + and call_node.func.attr in ("store", "load") + ): + return self.generic_visit(node) + + if not call_node.args or not isinstance(call_node.args[0], ast.Subscript): + return self.generic_visit(node) + + subscript_node = call_node.args[0] + sliced_object = subscript_node.value + slice_content = subscript_node.slice + remaining_args = call_node.args[1:] + + # Convert to nl.masked_load or nl.masked_store + func_name = "masked_" + call_node.func.attr + + new_call = ast.Call( + func=ast.Attribute( + value=ast.Name(id="nl", ctx=ast.Load()), attr=func_name, ctx=ast.Load() + ), + args=[ + sliced_object, + self._create_slice_value_node(slice_content), + *remaining_args, + ], + keywords=call_node.keywords, + ) + + return ast.Assign( + targets=node.targets, + value=new_call, + type_comment=getattr(node, "type_comment", None), + ) + + def _create_slice_value_node(self, node: ast.AST) -> ast.AST: + """ + (This helper function is unchanged and correct) + Recursively transforms a slice's AST content into a constructible object. + """ + match node: + case ast.Slice(lower, upper, step): + return ast.Call( + func=ast.Name(id="slice", ctx=ast.Load()), + args=[ + lower or ast.Constant(value=None), + upper or ast.Constant(value=None), + step or ast.Constant(value=None), + ], + keywords=[], + ) + case ast.Tuple(elts): + return ast.Tuple( + elts=[self._create_slice_value_node(e) for e in elts], + ctx=ast.Load(), + ) + case _: + return node + + +def transform_code(source_code: str) -> str: + """ + Applies the StoreCallTransformer to a string of Python code. + """ + tree = ast.parse(source_code) + transformer = StoreCallTransformer() + new_tree = transformer.visit(tree) + ast.fix_missing_locations(new_tree) + return ast.unparse(new_tree) + + +source_code = """ +import numpy as np + +# This load call should also be transformed +sbuf_value = nl.load(x[2:4, None, ..., nl.arange(128)[None, :]], mask=mask) + +# This is the call we want to transform +nl.store(x[2:4, None, ..., nl.arange(128)[None, :]], sbuf_value, mask=mask) + + +# This call should be ignored +other_lib.store(y[10:]) + +# This call should also be ignored +nl.some_other_func(z[:5]) +""" + +expected_transformed_code = """\ +import numpy as np +sbuf_value = nl.masked_load(x, (slice(2, 4, None), None, ..., nl.arange(128)[None, :]), mask=mask) +nl.masked_store(x, (slice(2, 4, None), None, ..., nl.arange(128)[None, :]), sbuf_value, mask=mask) +other_lib.store(y[10:]) +nl.some_other_func(z[:5])\ +""" + +if __name__ == "__main__": + assert transform_code(source_code) == expected_transformed_code diff --git a/triton_viz/core/patch.py b/triton_viz/core/patch.py index 63d31c84..afc5b329 100644 --- a/triton_viz/core/patch.py +++ b/triton_viz/core/patch.py @@ -1,15 +1,17 @@ import triton.language as tl from contextlib import contextmanager from collections.abc import Callable +from dataclasses import dataclass from typing import Any, Optional +from functools import partialmethod from tqdm import tqdm from .config import config as cfg -from dataclasses import dataclass from .callbacks import OpCallbacks, ForLoopCallbacks from .data import ( Op, + Allocate, RawLoad, Load, RawStore, @@ -46,6 +48,7 @@ AtomicCas, AtomicRMW, ) +from .data import Flip # separate import to avoid reordering noise import inspect import ast from triton.runtime.interpreter import ( @@ -59,7 +62,101 @@ from triton.tools.tensor_descriptor import TensorDescriptor from triton.runtime import JITFunction -op_list = [ +HAS_NKI = False +nki_builder = None +try: + from triton_viz.core.nki import nki_builder # type: ignore + + HAS_NKI = True +except ModuleNotFoundError: + pass + + +@dataclass +class AdapterResult: + """ + For each backend, ops may have slightly different function signatures + which we run through (backend, function)-specific adapters to return + standardized args/kwargs for client callbacks. + """ + + args: tuple[Any, ...] + kwargs: dict[str, Any] + + def __init__(self, *args: Any, **kwargs: Any) -> None: + self.args = args + self.kwargs = kwargs + + +def passthrough_adapter(*args: Any, **kwargs: Any) -> AdapterResult: + """Return arguments unchanged for clients that expect the original signature.""" + return AdapterResult(*args, **kwargs) + + +def _program_id_adapter(axis: Any, *_args: Any, **_kwargs: Any) -> AdapterResult: + return AdapterResult(axis) + + +def _triton_raw_store_adapter( + ptr: Any, value: Any, *_args: Any, **_kwargs: Any +) -> AdapterResult: + return AdapterResult(ptr, value) + + +def _triton_store_adapter( + ptr: Any, _value: Any, mask: Any, *_args: Any, **kwargs: Any +) -> AdapterResult: + keys = kwargs.get("keys") + return AdapterResult(ptr, mask, keys) + + +def _triton_raw_load_adapter(ptr: Any, *_args: Any, **_kwargs: Any) -> AdapterResult: + return AdapterResult(ptr) + + +def _triton_load_adapter( + ptr: Any, mask: Any, _other: Any, *_args: Any, **kwargs: Any +) -> AdapterResult: + return AdapterResult(ptr, mask, kwargs.get("keys")) + + +def _triton_dot_adapter(a: Any, b: Any, *_args: Any, **_kwargs: Any) -> AdapterResult: + return AdapterResult(a, b) + + +def _triton_reduce_sum_adapter( + input_tensor, axis=None, keep_dims=False, *_args, **_kwargs +) -> AdapterResult: + return AdapterResult(input_tensor, axis, keep_dims) + + +def _triton_addptr_adapter( + ptr: Any, offset: Any, *_args: Any, **_kwargs: Any +) -> AdapterResult: + return AdapterResult(ptr, offset) + + +def _nki_allocate_adapter(*_args: Any, **_kwargs: Any) -> AdapterResult: + return AdapterResult() + + +def _nki_load_adapter( + src: Any, keys: Any, *, mask: Optional[Any] = None, **_kwargs: Any +) -> AdapterResult: + return AdapterResult(src, mask, keys) + + +def _nki_store_adapter( + dst: Any, keys: Any, value: Any, *, mask: Optional[Any] = None, **_kwargs: Any +) -> AdapterResult: + return AdapterResult(dst, mask, keys) + + +def _nki_dot_adapter(x: Any, y: Any, *_args: Any, **_kwargs: Any) -> AdapterResult: + return AdapterResult(x, y) + + +TRITON_OP_LIST = [ ProgramId, RawStore, Store, @@ -97,8 +194,40 @@ AtomicRMW, ] -# Hardcoded operation attribute names to avoid issues with lambda functions -_OP_ATTR_NAMES = { +TRITON_ORIGINAL_OPS = { + ProgramId: interpreter_builder.create_get_program_id, + RawStore: interpreter_builder.create_store, + Store: interpreter_builder.create_masked_store, + RawLoad: interpreter_builder.create_load, + Load: interpreter_builder.create_masked_load, + Dot: interpreter_builder.create_dot, + UnaryOp: interpreter_builder.unary_op, + BinaryOp: interpreter_builder.binary_op, + TernaryOp: interpreter_builder.ternary_op, + MakeRange: interpreter_builder.create_make_range, + AddPtr: interpreter_builder.create_addptr, + ExpandDims: interpreter_builder.create_expand_dims, + Broadcast: interpreter_builder.create_broadcast, + Splat: interpreter_builder.create_splat, + MakeBlockPointer: interpreter_builder.create_make_block_ptr, + TensorPointerLoad: interpreter_builder.create_tensor_pointer_load, + TensorPointerStore: interpreter_builder.create_tensor_pointer_store, + Idiv: interpreter_builder.create_idiv, + Rsqrt: interpreter_builder.create_rsqrt, + CastImpl: interpreter_builder.cast_impl, + Reshape: interpreter_builder.create_reshape, + Join: interpreter_builder.create_join, + Fabs: interpreter_builder.create_fabs, + Ashr: interpreter_builder.create_ashr, + Advance: interpreter_builder.create_advance, + FpToFp: interpreter_builder.create_fp_to_fp, + Umulhi: interpreter_builder.create_umulhi, + Bitcast: interpreter_builder.create_bitcast, + AtomicCas: interpreter_builder.create_atomic_cas, + AtomicRMW: interpreter_builder.create_atomic_rmw, +} + +TRITON_OP_ATTR_NAMES = { ProgramId: "create_get_program_id", RawStore: "create_store", Store: "create_masked_store", @@ -131,38 +260,87 @@ AtomicRMW: "create_atomic_rmw", } -original_ops = { - ProgramId: interpreter_builder.create_get_program_id, - RawStore: interpreter_builder.create_store, - Store: interpreter_builder.create_masked_store, - RawLoad: interpreter_builder.create_load, - Load: interpreter_builder.create_masked_load, - Dot: interpreter_builder.create_dot, - UnaryOp: interpreter_builder.unary_op, - BinaryOp: interpreter_builder.binary_op, - TernaryOp: interpreter_builder.ternary_op, - MakeRange: interpreter_builder.create_make_range, - AddPtr: interpreter_builder.create_addptr, - ExpandDims: interpreter_builder.create_expand_dims, - Broadcast: interpreter_builder.create_broadcast, - Splat: interpreter_builder.create_splat, - MakeBlockPointer: interpreter_builder.create_make_block_ptr, - TensorPointerLoad: interpreter_builder.create_tensor_pointer_load, - TensorPointerStore: interpreter_builder.create_tensor_pointer_store, - Idiv: interpreter_builder.create_idiv, - Rsqrt: interpreter_builder.create_rsqrt, - CastImpl: interpreter_builder.cast_impl, - Reshape: interpreter_builder.create_reshape, - Join: interpreter_builder.create_join, - Fabs: interpreter_builder.create_fabs, - Ashr: interpreter_builder.create_ashr, - Advance: interpreter_builder.create_advance, - FpToFp: interpreter_builder.create_fp_to_fp, - Umulhi: interpreter_builder.create_umulhi, - Bitcast: interpreter_builder.create_bitcast, - AtomicCas: interpreter_builder.create_atomic_cas, - AtomicRMW: interpreter_builder.create_atomic_rmw, +TRITON_ADAPTERS: dict[type[Op], Callable[..., AdapterResult]] = { + ProgramId: _program_id_adapter, + RawStore: _triton_raw_store_adapter, + Store: _triton_store_adapter, + RawLoad: _triton_raw_load_adapter, + Load: _triton_load_adapter, + Dot: _triton_dot_adapter, + ReduceSum: _triton_reduce_sum_adapter, + AddPtr: _triton_addptr_adapter, +} + +for op_type in TRITON_OP_LIST: + TRITON_ADAPTERS.setdefault(op_type, passthrough_adapter) + +NKI_OP_LIST: list[type[Op]] = [] +NKI_ORIGINAL_OPS: dict[type[Op], Callable] = {} +NKI_OP_ATTR_NAMES: dict[type[Op], str] = {} +NKI_ADAPTERS: dict[type[Op], Callable[..., AdapterResult]] = {} +if HAS_NKI: + assert nki_builder is not None + + NKI_OP_LIST = [ + Allocate, + ProgramId, + Load, + Store, + Dot, + UnaryOp, + MakeRange, + ] + + NKI_ORIGINAL_OPS = { + ProgramId: nki_builder.program_id, + Allocate: nki_builder.ndarray, + Load: nki_builder.masked_load, + Store: nki_builder.masked_store, + Dot: nki_builder.matmul, + UnaryOp: nki_builder._unary_op, + MakeRange: nki_builder.arange, + } + + NKI_OP_ATTR_NAMES = { + ProgramId: "program_id", + Allocate: "ndarray", + Load: "masked_load", + Store: "masked_store", + Dot: "matmul", + UnaryOp: "_unary_op", + MakeRange: "arange", + } + + NKI_ADAPTERS = { + ProgramId: _program_id_adapter, + Allocate: _nki_allocate_adapter, + Load: _nki_load_adapter, + Store: _nki_store_adapter, + Dot: _nki_dot_adapter, + } + + for op_type in NKI_OP_LIST: + NKI_ADAPTERS.setdefault(op_type, passthrough_adapter) + + +OPERATION_REGISTRY: dict[str, dict[str, Any]] = { + "triton": { + "builder": interpreter_builder, + "op_list": TRITON_OP_LIST, + "original_ops": TRITON_ORIGINAL_OPS, + "op_attr_names": TRITON_OP_ATTR_NAMES, + "adapters": TRITON_ADAPTERS, + }, + "nki": { + "builder": nki_builder, + "op_list": NKI_OP_LIST, + "original_ops": NKI_ORIGINAL_OPS, + "op_attr_names": NKI_OP_ATTR_NAMES, + "adapters": NKI_ADAPTERS, + }, } + + reduce_map: dict[type[Op], Callable] = { ReduceMax: tl.max, ReduceMin: tl.min, @@ -185,20 +363,19 @@ def __init__( op: Callable, op_type: type[Op], callbacks: OpCallbacks, + adapter: Callable[..., AdapterResult], ): self.op = op self.op_type = op_type self.callbacks = callbacks + self.adapter = adapter def __call__(self, *args, **kwargs): if self.callbacks.before_callback: - self.callbacks.before_callback(*args, **kwargs) + before_args = self.adapter(*args, **kwargs) + self.callbacks.before_callback(*before_args.args, **before_args.kwargs) if self.callbacks.op_overrider: - if ( - self.op_type in reduce_map - or self.op_type in scan_map - or self.op_type in reshape_map - ): + if self.op_type in {**reduce_map, **scan_map, **reshape_map}: # see triton.runtime.interpreter:ReduceOps.sum # First, convert input from tl.tensor to TensorHandle. Here, input tensor is args[0] # Then, convert return value from TensorHandle to tl.tensor @@ -218,28 +395,34 @@ def __call__(self, *args, **kwargs): ret = self.op(*args, **kwargs) if self.callbacks.after_callback: # Pass ret so that we don't have to derive output shape from args - self.callbacks.after_callback(ret, *args, **kwargs) + after_args = self.adapter(*args, **kwargs) + self.callbacks.after_callback(ret, *after_args.args, **after_args.kwargs) return ret -def patch_op(op_type: type[Op], callbacks: OpCallbacks): +def patch_op(op_type: type[Op], callbacks: OpCallbacks, backend: str): """ Register a callback to be called before and after an operator is executed. :param op_type: The type of the operator to register the callback for. :param callbacks: The OpCallbacks object containing before_callback, after_callback, and op_overrider. + :param backend: The backend to use ('triton', 'nki', or None for current backend). """ - if op_type in original_ops: - # create a new function that calls the before_callback, the original op and the after_callback - op_name = _OP_ATTR_NAMES[op_type] - original_op = original_ops[op_type] - patched_op = PatchOp(original_op, op_type, callbacks) - setattr( - interpreter_builder, - op_name, - lambda *args, **kwargs: patched_op(*args, **kwargs), - ) - elif op_type in reduce_map or op_type in scan_map or op_type in reshape_map: + if backend not in OPERATION_REGISTRY: + raise ValueError(f"Unknown backend: {backend}") + + backend_ops = OPERATION_REGISTRY[backend]["original_ops"] + backend_attr_names = OPERATION_REGISTRY[backend]["op_attr_names"] + backend_adapters = OPERATION_REGISTRY[backend]["adapters"] + backend_builder = OPERATION_REGISTRY[backend]["builder"] + + if op_type in backend_ops: + op_name = backend_attr_names[op_type] + original_op = backend_ops[op_type] + adapter = backend_adapters[op_type] + patched_op = PatchOp(original_op, op_type, callbacks, adapter) + setattr(backend_builder, op_name, patched_op) + elif backend == "triton" and op_type in {**reduce_map, **scan_map, **reshape_map}: if op_type in reduce_map: op_name = reduce_map[op_type].__name__ elif op_type in scan_map: @@ -247,28 +430,34 @@ def patch_op(op_type: type[Op], callbacks: OpCallbacks): elif op_type in reshape_map: op_name = reshape_map[op_type].__name__ original_op = getattr(tl, op_name) - patched_op = PatchOp(original_op, op_type, callbacks) - setattr(tl, op_name, lambda *args, **kwargs: patched_op(*args, **kwargs)) - elif op_type in math_map: + adapter = backend_adapters[op_type] + patched_op = PatchOp(original_op, op_type, callbacks, adapter) + setattr(tl, op_name, patched_op) + elif backend == "triton" and op_type in math_map: op_name = math_map[op_type].__name__ original_op = getattr(tl.math, op_name) - patched_op = PatchOp(original_op, op_type, callbacks) - setattr(tl.math, op_name, lambda *args, **kwargs: patched_op(*args, **kwargs)) + adapter = backend_adapters[op_type] + patched_op = PatchOp(original_op, op_type, callbacks, adapter) + setattr(tl.math, op_name, patched_op) else: raise ValueError(f"Patching operator {op_type} not supported") -def unpatch_op(op_type: type[Op]): +def unpatch_op(op_type: type[Op], backend: str): """ Unregister a callback for an operator. :param op_type: The type of the operator to unregister the callback for. """ - if op_type in original_ops: - original_op = original_ops[op_type] + backend_ops = OPERATION_REGISTRY[backend]["original_ops"] + backend_attr_names = OPERATION_REGISTRY[backend]["op_attr_names"] + backend_builder = OPERATION_REGISTRY[backend]["builder"] + + if op_type in backend_ops: + original_op = backend_ops[op_type] # type: ignore # Use hardcoded name from _OP_ATTR_NAMES - op_name = _OP_ATTR_NAMES[op_type] - setattr(interpreter_builder, op_name, original_op) + op_name = backend_attr_names[op_type] # type: ignore + setattr(backend_builder, op_name, original_op) class _LoopIter: @@ -526,12 +715,111 @@ def unpatch_for_loop(): _loop_patcher.unpatch() -def patch_lang(fn): - triton_patch_lang(fn) - fn.__globals__["_triton_viz_loop_patcher"] = _loop_patcher +def patch_lang(fn, backend): + if backend == "triton": + triton_patch_lang(fn) + elif backend == "nki": + from triton_viz.core.nki import nki_patch_lang + nki_patch_lang() + else: + raise ValueError( + f"Unsupported backend {backend} received. Triton-viz only supports one of ('triton', 'nki')." + ) -def unpatch_lang(): + fn.__globals__["_triton_viz_loop_patcher"] = _loop_patcher + # Wrap tl.flip to emit a Flip record after computing result + try: + _orig_flip = getattr(tl, "flip", None) + if _orig_flip is not None and not getattr( + _orig_flip, "__triton_viz_wrapped__", False + ): + + def _viz_flip(x, *args, **kwargs): + # Call original flip implementation + ret = _orig_flip(x, *args, **kwargs) + # Best-effort extract dim + dim = None + if args: + dim = args[0] + if "dim" in kwargs: + dim = kwargs.get("dim") + # Best-effort shapes + in_shape = None + out_shape = None + x_arr = None + r_arr = None + try: + # interpreter tensors may expose .data or .handle.data + x_data = getattr(x, "data", None) + if x_data is None and hasattr(x, "handle"): + x_data = getattr(x.handle, "data", None) + if x_data is not None: + in_shape = tuple(x_data.shape) + x_arr = x_data + except Exception: + pass + try: + r_data = getattr(ret, "data", None) + if r_data is None and hasattr(ret, "handle"): + r_data = getattr(ret.handle, "data", None) + if r_data is not None: + out_shape = tuple(r_data.shape) + r_arr = r_data + except Exception: + pass + + # Emit a Flip record to the active tracer, if available + try: + global _current_client_manager + cm = _current_client_manager + if cm is not None and hasattr(cm, "clients"): + tracer = cm.get_client("tracer") + if tracer is not None: + input_payload = None + output_payload = None + try: + # Avoid huge payloads: cap to 64k elements + def _maybe_list(arr): + import numpy as _np + + if arr is None: + return None + try: + if arr.size <= 65536: + return _np.asarray(arr).tolist() + except Exception: + pass + return None + + input_payload = _maybe_list(x_arr) + output_payload = _maybe_list(r_arr) + except Exception: + pass + + rec = Flip( + input_shape=in_shape or tuple(), + output_shape=out_shape or (in_shape or tuple()), + dim=int(dim) if dim is not None else 0, + input_data=input_payload, + output_data=output_payload, + ) + # attach call path already handled by Flip.__post_init__ + tracer.records.append(rec) + except Exception: + # Never fail kernel execution due to viz + pass + return ret + + # mark wrapper to avoid double-wrapping on subsequent patch_lang calls + setattr(_viz_flip, "__triton_viz_wrapped__", True) + tl.flip = _viz_flip # type: ignore[assignment] + except Exception: + # If wrapping fails, continue without Flip records + pass + + +def unpatch_lang(backend): # TODO: once this (https://github.com/triton-lang/triton/pull/8735) # gets into a stable release, we can simplify this unpatching logic by upgrading Triton. # This PR is ugly to implement in triton-viz directly because it piggybacks off @@ -540,13 +828,18 @@ def unpatch_lang(): import importlib import sys - for name in ("core", "math", "extra"): - mod = getattr(tl, name, None) - if mod is not None and mod.__name__ in sys.modules: - importlib.reload(mod) + if backend == "triton": + for name in ("core", "math", "extra"): + mod = getattr(tl, name, None) + if mod is not None and mod.__name__ in sys.modules: + importlib.reload(mod) - if tl.__name__ in sys.modules: - importlib.reload(tl) + if tl.__name__ in sys.modules: + importlib.reload(tl) + elif backend == "nki": + from triton_viz.core.nki import nki_unpatch_lang + + nki_unpatch_lang() from triton.language import semantic as tl_semantic from triton.compiler import code_generator as codegen @@ -621,10 +914,13 @@ def _to_cpu(arg): return args_hst, kwargs_hst -def _grid_executor_call(self, *args_dev, **kwargs): +def _grid_executor_call(self, *args_dev, backend=None, **kwargs): + assert backend is not None if kwargs.pop("warmup", False): return + builder = OPERATION_REGISTRY[backend]["builder"] + def run_grid_loops(grid): for x in tqdm( range(grid[0]), @@ -644,7 +940,7 @@ def run_grid_loops(grid): leave=False, disable=not (cfg.report_grid_execution_progress and grid[2] > 1), ): - interpreter_builder.set_grid_idx(x, y, z) + builder.set_grid_idx(x, y, z) client_manager.grid_idx_callback((x, y, z)) if not client_manager.pre_run_callback(self.fn): continue # Skip this block @@ -661,11 +957,16 @@ def run_grid_loops(grid): k: v for k, v in kwargs.items() if k in argspec.args or k in triton_viz_args } client_manager = kwargs.pop("client_manager") + + # Expose client_manager to tl.flip wrapper via a module-global + global _current_client_manager + _current_client_manager = client_manager kwargs.pop("jit_fn") if cfg.virtual_memory: args_hst, kwargs_hst = _init_args_hst(args_dev, kwargs) else: args_hst, kwargs_hst = self._init_args_hst(args_dev, kwargs) + # Prepare call arguments args = inspect.getcallargs(self.fn, *args_hst, **kwargs_hst) call_args = {} @@ -682,7 +983,8 @@ def run_grid_loops(grid): grid = self.grid(call_args) if callable(self.grid) else self.grid assert len(grid) <= 3 grid = grid + (1,) * (3 - len(grid)) - interpreter_builder.set_grid_dim(*grid) + + builder.set_grid_dim(*grid) client_manager.grid_callback(grid) if cfg.enable_timing: import time @@ -699,17 +1001,20 @@ def run_grid_loops(grid): self._restore_args_dev(args_dev, args_hst, kwargs, kwargs_hst) -def _jit_function_call(self, *args, **kwargs): - patch_lang(self.fn) +def _jit_function_call( + self, *args, backend=None, **kwargs +): # NOTE: is this ever called? + assert backend is not None + patch_lang(self.fn, backend) return self.fn(*args, **kwargs) @contextmanager -def patch_calls(): +def patch_calls(backend): old_grid_executor_call = GridExecutor.__call__ old_jit_function_call = JITFunction.__call__ - GridExecutor.__call__ = _grid_executor_call - JITFunction.__call__ = _jit_function_call + GridExecutor.__call__ = partialmethod(_grid_executor_call, backend=backend) + JITFunction.__call__ = partialmethod(_jit_function_call, backend=backend) try: yield finally: diff --git a/triton_viz/core/trace.py b/triton_viz/core/trace.py index 3ffe3bba..03f013a9 100644 --- a/triton_viz/core/trace.py +++ b/triton_viz/core/trace.py @@ -8,18 +8,24 @@ from ..clients import Sanitizer, Profiler, Tracer from .client import ClientManager, Client from .data import Launch -from typing import Callable, Optional, Union +from typing import Callable, Optional, Union, TypeVar launches: list[Launch] = [] +T = TypeVar("T") + def dummy_benchmarker(fn, quantiles): fn() return (1.0, 1.0, 1.0) -class Trace(KernelInterface): +class TraceInterface: + def __init__(self, client: Union[str, Client]) -> None: + self.client_manager = ClientManager() + self.add_client(client) + @staticmethod def _normalize_client(client: Union[str, Client]) -> Client: if isinstance(client, str): @@ -39,19 +45,27 @@ def _normalize_client(client: Union[str, Client]) -> Client: def add_client(self, new_client: Union[str, Client]) -> None: self.client_manager.add_clients([self._normalize_client(new_client)]) + def finalize(self): + self.client_manager.finalize() + launches.append(self.client_manager.launch) + + +class TritonTrace(KernelInterface, TraceInterface): def __init__( self, runner: Union[JITFunction, InterpretedFunction, Autotuner, Heuristics], client: Union[str, Client], ) -> None: - self.fn = runner + self.jit_fn: Optional[JITFunction] = None + self.base_fn: Optional[Callable] = None + self.interpreted_fn: Optional[InterpretedFunction] = None def unpack_kernel( - source: Union["Trace", JITFunction, InterpretedFunction, Heuristics], + source: Union["TritonTrace", JITFunction, InterpretedFunction, Heuristics], ) -> tuple[ Optional[JITFunction], Optional[Callable], Optional[InterpretedFunction] ]: - if isinstance(source, Trace): + if isinstance(source, TritonTrace): return source.jit_fn, source.base_fn, source.interpreted_fn if isinstance(source, JITFunction): base_fn = source.fn @@ -89,9 +103,12 @@ def unpack_kernel( self.jit_fn, self.base_fn, self.interpreted_fn = unpack_kernel(runner) self.runner = self.interpreted_fn self.warmup_runner = self.jit_fn + self.arg_names = runner.arg_names - self.client_manager = ClientManager() - self.add_client(client) + + self.fn = runner + + TraceInterface.__init__(self, client) # Preserve common function attributes for compatibility # with code that expects to access these attributes on the kernel @@ -128,7 +145,7 @@ def run(self, *args, **kwargs): if self.warmup_runner: self.warmup_runner.warmup(*args, **kwargs) - with self.client_manager.patch_run(self.base_fn): + with self.client_manager.patch_run(self.base_fn, backend="triton"): kwargs.update({"client_manager": self.client_manager}) kwargs.update({"jit_fn": self.jit_fn}) ret = self.runner.run(*args, **kwargs) @@ -145,9 +162,24 @@ def warmup(self, *args, **kwargs): if self.warmup_runner: self.warmup_runner.warmup(*args, **kwargs) - def finalize(self): - self.client_manager.finalize() - launches.append(self.client_manager.launch) + +class NKITrace(KernelInterface, TraceInterface): + def __init__(self, kernel, client: str | Client) -> None: + from neuronxcc.nki.compile import GenericKernel + from .nki import NKIInterpretedFunction + + if isinstance(kernel, GenericKernel): + # This is wrong + self.interpreter_fn = NKIInterpretedFunction(kernel.func) + self.func = kernel.func + elif isinstance(kernel, NKIInterpretedFunction): + self.interpreter_fn = kernel + self.func = kernel.fn + else: + self.interpreter_fn = NKIInterpretedFunction(kernel) + self.func = kernel + + TraceInterface.__init__(self, client) def __getattr__(self, name): # Forward any missing attributes to the underlying runner @@ -178,8 +210,21 @@ def __getattr__(self, name): f"'{type(self).__name__}' object has no attribute '{name}'" ) + def __getitem__(self, *grid): + return KernelInterface.__getitem__(self, tuple(*grid)) + + def __call__(self, *args, **kwargs): + return self[(1, 1, 1)](*args, **kwargs) + + def run(self, *args, **kwargs): + with self.client_manager.patch_run(self.func, backend="nki"): + kwargs.update({"client_manager": self.client_manager}) + ret = self.interpreter_fn.run(*args, **kwargs) + self.finalize() + return ret -def trace(clients: Union[str, Client, None] = None): + +def trace(clients: Union[str, Client, None] = None, backend: str = "triton"): """ Create a trace object that can be used to run a kernel with instrumentation clients. @@ -192,19 +237,37 @@ def trace(clients: Union[str, Client, None] = None): if not isinstance(clients, (str, Client)): raise TypeError(f"Expected str or Client, got {type(clients)}") - def decorator(kernel) -> Trace: + def decorator(kernel) -> TraceInterface: # When sanitizer is disabled, skip tracing and return the original kernel unchanged if cfg.disable_sanitizer: return kernel # First-time wrapping + # Triton backend need JIT/Interpreter/Autotuner; + # NKI allow Python function( NKIInterpretedFunction) + if backend == "nki": + return NKITrace(kernel, clients) if isinstance( kernel, (JITFunction, InterpretedFunction, Autotuner, Heuristics) ): - return Trace(kernel, clients) - - # If the object is already a Trace, just append the new client(s) - if isinstance(kernel, Trace): + if backend == "triton": + return TritonTrace(kernel, clients) + else: + raise ValueError(f"Unknown backend: {backend}") + + # Handle NKI functions specifically + if backend == "nki": + from .nki import NKIInterpretedFunction + + if isinstance(kernel, NKIInterpretedFunction): + return NKITrace(kernel, clients) + else: + # Wrap plain functions as NKIInterpretedFunction + interpreted_fn = NKIInterpretedFunction(kernel) + return NKITrace(interpreted_fn, clients) + + # If the object is already initialized as a TraceInterface, just append the new client(s) + if isinstance(kernel, TraceInterface): trace = kernel trace.add_client(clients) return trace diff --git a/triton_viz/static/flip.js b/triton_viz/static/flip.js new file mode 100644 index 00000000..73f5758c --- /dev/null +++ b/triton_viz/static/flip.js @@ -0,0 +1,175 @@ +import * as THREE from 'https://esm.sh/three@0.155.0/build/three.module.js'; +import { OrbitControls } from 'https://esm.sh/three@0.155.0/examples/jsm/controls/OrbitControls.js'; +import { setupScene, setupGeometries, createCube, CUBE_SIZE, GAP } from './load_utils.js'; +import { createFlip3D } from './flip_3d.js'; + +export function createFlipVisualization(containerElement, op) { + const API_BASE = window.__TRITON_VIZ_API__ || ''; + const overlay = document.createElement('div'); + Object.assign(overlay.style, { + position: 'absolute', left: '0', top: '0', right: '0', bottom: '0', + zIndex: 3000, pointerEvents: 'auto' + }); + containerElement.appendChild(overlay); + + const { scene, camera, renderer } = setupScene(overlay, 0x15151b); + const controls = new OrbitControls(camera, renderer.domElement); + controls.enableDamping = true; controls.dampingFactor = 0.06; + const { cubeGeometry, edgesGeometry, lineMaterial } = setupGeometries(); + + // Helper color mapping + function hsv(h,s,v){ + const c=v*s; const x=c*(1-Math.abs((h/60)%2-1)); const m=v-c; let r,g,b; + if(h<60){[r,g,b]=[c,x,0]} else if(h<120){[r,g,b]=[x,c,0]} else if(h<180){[r,g,b]=[0,c,x]} + else if(h<240){[r,g,b]=[0,x,c]} else if(h<300){[r,g,b]=[x,0,c]} else {[r,g,b]=[c,0,x]} + return new THREE.Color(r+m, g+m, b+m); + } + const shape = op.output_shape || op.input_shape || []; + const dim = op.dim || 0; + const is2D = (shape.length === 2); + const H = is2D ? (shape[0]||8) : 1; + const W = is2D ? (shape[1]||16) : ((shape[0]||16)); + + // Build input and output layers + const root = new THREE.Group(); scene.add(root); + + // Input row + const inputGroup = new THREE.Group(); + for (let r=0;r { + if (!is2D) return { rr: 0, cc: (W-1-c) }; + if (dim === 0) return { rr: (H-1-r), cc: c }; + return { rr: r, cc: (W-1-c) }; + }; + for (let r=0;r{ + if (stepsCleanup){ try{ stepsCleanup(); }catch(e){} stepsCleanup=null; stepsBtn.textContent='Flip Steps 3D'; return; } + // choose length along flipped dimension if 2D; otherwise use total W + const n = is2D ? (dim===0? H: W) : W; + stepsCleanup = createFlip3D(overlay, { length: Math.min(128, n) }); + stepsBtn.textContent = 'Close Steps'; + }); + + async function getFlipValue(which, r, c){ + const body = { uuid: op.uuid, which, x: c, y: r }; + const res = await fetch(`${API_BASE}/api/getFlipValue`, { method:'POST', headers:{'Content-Type':'application/json'}, body: JSON.stringify(body)}); + return await res.json(); + } + + const raycaster = new THREE.Raycaster(); + const mouse = new THREE.Vector2(); + let hoveredCube = null; + function clearHoverOutlines(){ + // reset outlines on all cubes to avoid residual highlights + const groups = [inputGroup, outputGroup]; + for (const g of groups){ + if (!g) continue; + for (const child of g.children){ + const ho = child && child.getObjectByName && child.getObjectByName('hoverOutline'); + if (ho) ho.visible = false; + } + } + } + function syncHoverHighlight(){ + // periodically enforce: only current hovered cube is highlighted + clearHoverOutlines(); + if (hoveredCube){ + const ho = hoveredCube.getObjectByName && hoveredCube.getObjectByName('hoverOutline'); + if (ho) ho.visible = true; + } + } + // run every ~60ms (about 16 FPS) to avoid leaving residual highlights + const _hoverTimer = setInterval(syncHoverHighlight, 60); + function _updateMouseNDC(event) { + const rect = renderer.domElement.getBoundingClientRect(); + const dpr = (window.devicePixelRatio || 1); + const px = (event.clientX - rect.left) * dpr; + const py = (event.clientY - rect.top ) * dpr; + const w = rect.width * dpr, h = rect.height * dpr; + mouse.x = (px / w) * 2 - 1; mouse.y = -(py / h) * 2 + 1; + } + function findTopLevel(obj){ let n=obj; while(n && !(n.userData && n.userData.tensorName)) n=n.parent; return n; } + function updatePanel(which, r, c, value){ + if (!which){ sideMenu.innerHTML=''; return; } + sideMenu.innerHTML = `

${which} Tensor

Row: ${r+1}

Col: ${c+1}

Value: ${value !== undefined ? value : '...'}

`; + } + + async function onMouseMove(event){ + _updateMouseNDC(event); raycaster.setFromCamera(mouse, camera); + const objs = [...inputGroup.children, ...outputGroup.children]; + const hits = raycaster.intersectObjects(objs, true); + // only update target; periodic timer will enforce single highlight + hoveredCube = null; + if (hits.length === 0){ updatePanel('',0,0,undefined); return; } + let cube = findTopLevel(hits[0].object); if (!cube) { updatePanel('',0,0,undefined); return; } + hoveredCube = cube; + const which = cube.userData.tensorName === 'Input' ? 'input' : 'output'; + const r = cube.userData.tensor1 || 0; const c = cube.userData.tensor0 || 0; + try { const res = await getFlipValue(which, r, c); updatePanel(which, r, c, res.value); } catch(e){ updatePanel(which, r, c, undefined); } + // do not toggle here; timer will handle visibility to avoid residuals + } + overlay.addEventListener('mousemove', onMouseMove); + overlay.addEventListener('mouseleave', ()=>{ hoveredCube=null; updatePanel('',0,0,undefined); syncHoverHighlight(); }); + + // Camera framing + const box = new THREE.Box3().setFromObject(root); const center = box.getCenter(new THREE.Vector3()); const size = box.getSize(new THREE.Vector3()); + const maxDim = Math.max(size.x, size.y, size.z); const fov = camera.fov * (Math.PI/180); let cameraZ = Math.abs(maxDim/2/Math.tan(fov/2)); cameraZ *= 1.6; + camera.position.set(center.x, center.y, center.z + cameraZ); camera.lookAt(center); + + function animate(){ requestAnimationFrame(animate); controls.update(); renderer.render(scene, camera); } + animate(); + + const closeBtn = document.createElement('button'); closeBtn.textContent = 'Close'; Object.assign(closeBtn.style, { position:'absolute', top:'50px', left:'10px', zIndex:'2001' }); overlay.appendChild(closeBtn); + function cleanup(){ clearInterval(_hoverTimer); overlay.removeEventListener('mousemove', onMouseMove); overlay.removeEventListener('mouseleave', ()=>{}); if (stepsCleanup){ try{ stepsCleanup(); }catch(e){} stepsCleanup=null; } if (overlay && overlay.remove) overlay.remove(); } + closeBtn.addEventListener('click', cleanup); + try { window.current_op_uuid = op.uuid; } catch (e) {} + return cleanup; +} diff --git a/triton_viz/static/flip_3d.js b/triton_viz/static/flip_3d.js new file mode 100644 index 00000000..f87f7fee --- /dev/null +++ b/triton_viz/static/flip_3d.js @@ -0,0 +1,189 @@ +import * as THREE from 'https://esm.sh/three@0.155.0/build/three.module.js'; +import { OrbitControls } from 'https://esm.sh/three@0.155.0/examples/jsm/controls/OrbitControls.js'; +import { setupScene, setupGeometries, createCube, CUBE_SIZE, GAP } from './load_utils.js'; + +// 3D Flip visualization with layered rows (initial + step_i view + step_i swap) +// API: createFlip3D(container, { length, steps }) => cleanup() +export function createFlip3D(containerElement, options) { + // overlay wrapper to guarantee visibility above existing content + const overlay = document.createElement('div'); + Object.assign(overlay.style, { + position: 'absolute', left: '0', top: '0', right: '0', bottom: '0', + zIndex: 3000, pointerEvents: 'auto' + }); + containerElement.appendChild(overlay); + const length = Math.max(2, (options && options.length) || 32); + const steps = (options && options.steps) || computeSteps(length); + + const { scene, camera, renderer } = setupScene(overlay, 0x15151b); + const controls = new OrbitControls(camera, renderer.domElement); + controls.enableDamping = true; controls.dampingFactor = 0.06; + + const { cubeGeometry, edgesGeometry, lineMaterial } = setupGeometries(); + + // Utilities + function hsv(h,s,v){ + const c=v*s; const x=c*(1-Math.abs((h/60)%2-1)); const m=v-c; let r,g,b; + if(h<60){[r,g,b]=[c,x,0]} else if(h<120){[r,g,b]=[x,c,0]} else if(h<180){[r,g,b]=[0,c,x]} + else if(h<240){[r,g,b]=[0,x,c]} else if(h<300){[r,g,b]=[x,0,c]} else {[r,g,b]=[c,0,x]} + return new THREE.Color(r+m, g+m, b+m); + } + const hueAt = (val)=> 360 * (val/(length-1||1)); + + function createTextSprite(text){ + const canvas = document.createElement('canvas'); + const ctx = canvas.getContext('2d'); + ctx.font = 'Bold 28px Arial'; + const metrics = ctx.measureText(text); + canvas.width = metrics.width + 8; canvas.height = 40; + ctx.font = 'Bold 28px Arial'; + ctx.fillStyle = 'white'; ctx.fillText(text, 0, 28); + const tex = new THREE.CanvasTexture(canvas); + const mat = new THREE.SpriteMaterial({ map: tex, transparent: true }); + const sp = new THREE.Sprite(mat); sp.scale.set(4, 1.2, 1); + return sp; + } + + // Build sequences with side-by-side panels per step + function computeSteps(n){ const arr=[]; let s=Math.floor(n/2); while(s>=1){ arr.push(s); s=Math.floor(s/2);} return arr; } + function swapBlocksArray(array, seg){ + const res = array.slice(); const group = Math.max(1, 2*seg); + for(let g=0; g= seg) ? 1 : 0; + const col = p % seg; + const color = hsv(hueAt(arrVals[i]), 1, 1); + const cube = createCube(color, 'Flip', col, row, g, cubeGeometry, edgesGeometry, lineMaterial); + const x = xOffset + g*(seg*per + groupGap) + col*per; + const yPos = y - row*per; + cube.position.set(x, yPos, 0); + panel.add(cube); + } + const label = createTextSprite(labelText); + label.position.set(xOffset - (CUBE_SIZE*3.2), y + CUBE_SIZE*0.5, 0); + panel.add(label); + return panel; + } + + // Build all layers + const layers = []; + const root = new THREE.Group(); scene.add(root); + + const rowGap = CUBE_SIZE * 4.2; // vertical distance per step row + const columnGap = CUBE_SIZE * 6.0; // horizontal gap between view and swap panels + + let current_y = 0; + // initial 1D row (single panel) + const initPanel = new THREE.Group(); + for (let i=0;ii); + for (let s=0; s g.visible = (i===0)); + + // Center camera to see all layers + const totalW = Math.max(layoutWidth(steps[0]||Math.floor(length/2), length) + columnGap + layoutWidth(steps[0]||Math.floor(length/2), length), length*(CUBE_SIZE+GAP)); + const totalH = (layers.length) * rowGap; + camera.position.set(totalW*0.6, -totalH*0.35, Math.max(6, totalW*0.9)); + camera.lookAt(new THREE.Vector3(totalW*0.5, -totalH*0.4, 0)); + + // UI controls + const ui = document.createElement('div'); + Object.assign(ui.style, { position: 'absolute', top: '10px', left: '10px', display: 'flex', gap: '8px', zIndex: 2000 }); + const playBtn = document.createElement('button'); playBtn.textContent = 'Play'; + const stepLabel = document.createElement('span'); stepLabel.textContent = 'Layer: 1/' + layers.length; stepLabel.style.color = '#fff'; + const speedSel = document.createElement('select'); ['0.5x','1x','2x','4x'].forEach(s=>{ const o=document.createElement('option'); o.value=s; o.textContent=s; if(s==='1x') o.selected=true; speedSel.appendChild(o); }); + ui.appendChild(playBtn); ui.appendChild(stepLabel); ui.appendChild(speedSel); overlay.appendChild(ui); + const closeBtn = document.createElement('button'); closeBtn.textContent = 'Close'; + ui.appendChild(closeBtn); + + // Animation loop: reveal rows sequentially + let playing=false; let lastTs=0; let idxVisible=0; let raf=0; + function rate(){ return speedSel.value==='0.5x'?800: speedSel.value==='1x'?400: speedSel.value==='2x'?200:100; } + function loop(){ + controls.update(); + if (playing && (lastTs===0 || performance.now()-lastTs>=rate())){ + if (idxVisible < layers.length-1){ + idxVisible += 1; layers[idxVisible].visible = true; stepLabel.textContent = `Layer: ${idxVisible+1}/${layers.length}`; lastTs = performance.now(); + } else { playing=false; playBtn.textContent='Replay'; } + } + renderer.render(scene, camera); + raf = requestAnimationFrame(loop); + } + raf = requestAnimationFrame(loop); + + playBtn.addEventListener('click', ()=>{ + if (!playing && idxVisible===layers.length-1){ + // reset visibility + layers.forEach((r, i)=> r.visible = (i===0)); idxVisible = 0; stepLabel.textContent = `Layer: 1/${layers.length}`; + } + playing = !playing; playBtn.textContent = playing? 'Pause' : (idxVisible===layers.length-1? 'Replay' : 'Play'); + if (playing) lastTs = 0; + }); + + closeBtn.addEventListener('click', ()=> cleanup()); + + function cleanup(){ cancelAnimationFrame(raf); controls.dispose(); renderer.dispose?.(); if (overlay && overlay.remove) overlay.remove(); } + return cleanup; +} diff --git a/triton_viz/static/flip_demo.js b/triton_viz/static/flip_demo.js new file mode 100644 index 00000000..34ecc5d9 --- /dev/null +++ b/triton_viz/static/flip_demo.js @@ -0,0 +1,178 @@ +// Lightweight 2D canvas demo to visualize stepwise flip via view/reshape + swap +// API: createFlipDemo(containerElement, options) -> cleanup() +// options: { length: number, steps?: number[] } where steps are segment sizes halving each step + +export function createFlipDemo(containerElement, options) { + const length = Math.max(1, (options && options.length) || 16); + const steps = (options && options.steps) || computeSteps(length); + + const wrapper = document.createElement('div'); + Object.assign(wrapper.style, { + position: 'absolute', left: '0', top: '0', right: '0', bottom: '0', + display: 'flex', flexDirection: 'column', gap: '10px', padding: '10px' + }); + + const toolbar = document.createElement('div'); + Object.assign(toolbar.style, { display: 'flex', gap: '8px', alignItems: 'center' }); + const playBtn = document.createElement('button'); playBtn.textContent = 'Play'; + const stepLabel = document.createElement('span'); stepLabel.textContent = 'Step: 0'; + const speedSel = document.createElement('select'); + ;['0.5x','1x','2x','4x'].forEach(s=>{ const o=document.createElement('option'); o.value=s; o.textContent=s; if(s==='1x') o.selected=true; speedSel.appendChild(o); }); + toolbar.appendChild(playBtn); toolbar.appendChild(stepLabel); toolbar.appendChild(speedSel); + + const canvas = document.createElement('canvas'); + canvas.width = Math.min(1200, containerElement.clientWidth - 20); + canvas.height = Math.min(500, containerElement.clientHeight - 20); + const ctx = canvas.getContext('2d'); + + wrapper.appendChild(toolbar); + wrapper.appendChild(canvas); + containerElement.appendChild(wrapper); + + // data: 0..length-1 + let data = new Array(length).fill(0).map((_,i)=>i); + let currentStep = 0; + let playing = false; + let rafId = 0; + + function computeSteps(n){ + // e.g. for 16 -> [8,4,2,1] + const arr=[]; let s = Math.floor(n/2); + while(s>=1){ arr.push(s); s=Math.floor(s/2);} return arr; + } + + function lerp(a,b,t){ return a+(b-a)*t; } + + function hsv(h,s,v){ + let c=v*s; let x=c*(1-Math.abs((h/60)%2-1)); let m=v-c; let r,g,b; + if(h<60){[r,g,b]=[c,x,0]} else if(h<120){[r,g,b]=[x,c,0]} else if(h<180){[r,g,b]=[0,c,x]} + else if(h<240){[r,g,b]=[0,x,c]} else if(h<300){[r,g,b]=[x,0,c]} else {[r,g,b]=[c,0,x]} + return `rgb(${Math.round((r+m)*255)},${Math.round((g+m)*255)},${Math.round((b+m)*255)})`; + } + + function drawArray(arr, y, boxW, boxH, highlightPairs){ + for(let i=0;i{ + const x1 = 20 + i*boxW + (boxW-2)/2; + const x2 = 20 + j*boxW + (boxW-2)/2; + const yy = y + boxH + 6; + ctx.beginPath(); ctx.moveTo(x1, yy); ctx.lineTo(x2, yy); ctx.stroke(); + }); + } + } + + function calcPairs(seg){ + // pair indices across two adjacent blocks of size `seg` inside each group of size `2*seg` + const pairs=[]; const n=length; const group = Math.max(1, 2*seg); + for(let g=0; gi), 40, boxW, boxH); + + let y = 100; + for(let s=0;s= rate){ + if (currentStep < steps.length){ + data = performSwap(data, steps[currentStep]); + currentStep += 1; + stepLabel.textContent = `Step: ${currentStep}`; + _lastStepTime = Date.now(); + } else { + playing = false; playBtn.textContent = 'Replay'; + } + } + rafId = requestAnimationFrame(tick); + } + let _lastStepTime = 0; + + playBtn.addEventListener('click', ()=>{ + if (!playing && currentStep===steps.length){ + // reset + data = new Array(length).fill(0).map((_,i)=>i); + currentStep = 0; stepLabel.textContent='Step: 0'; + } + playing = !playing; playBtn.textContent = playing ? 'Pause' : (currentStep===steps.length ? 'Replay' : 'Play'); + if (playing){ _lastStepTime = 0; cancelAnimationFrame(rafId); rafId = requestAnimationFrame(tick);} + }); + + // initial paint + drawFrame(0); + + return function cleanup(){ cancelAnimationFrame(rafId); if (wrapper && wrapper.remove) wrapper.remove(); }; +} diff --git a/triton_viz/static/gridblock.js b/triton_viz/static/gridblock.js index bbdda5bb..8e70b823 100644 --- a/triton_viz/static/gridblock.js +++ b/triton_viz/static/gridblock.js @@ -1,6 +1,8 @@ import { createMatMulVisualization } from './matmul.js'; +import { createFlipVisualization } from './flip.js'; import { createLoadVisualization } from './load.js'; import { createStoreVisualization } from './store.js'; +import { createFlowDiagram } from './nki.js'; export class GridBlock { constructor(x, y, width, height, gridX, gridY, gridZ, blockData, onClose, containerElement, canvas, drawFunction) { @@ -74,6 +76,9 @@ export class GridBlock { const closeButton = this.createCloseButton(); this.visualizationContainer.appendChild(closeButton); + // Add an explicit Back button (top-left) to return to main canvas view + const backButton = this.createBackButton(); + this.visualizationContainer.appendChild(backButton); // Ensure buttons panel sits above the canvas and accepts clicks const buttonsPanel = this.visualizationContainer.querySelector('div'); @@ -90,6 +95,71 @@ export class GridBlock { if (this.blockData.length > 0) { this.displayOpVisualization(this.blockData[0]); } + + // Add a global "Show Code" toggle once (applies to any op type) + const codeBtnId = 'global-code-toggle-btn'; + if (!document.getElementById(codeBtnId)) { + const btn = document.createElement('button'); + btn.id = codeBtnId; + btn.textContent = 'Show Code: OFF'; + Object.assign(btn.style, { + position: 'fixed', + right: '10px', + top: '10px', + zIndex: '2001' + }); + document.body.appendChild(btn); + + let panel = null; + const destroyPanel = () => { if (panel && panel.remove) panel.remove(); panel = null; }; + const createPanel = async (uuid, frameIdx = 0, context = 8) => { + destroyPanel(); + const wrapper = document.createElement('div'); + Object.assign(wrapper.style, { + position: 'fixed', right: '10px', top: '50px', width: '520px', maxHeight: '60vh', overflow: 'auto', + padding: '8px 10px', background: 'rgba(0,0,0,0.65)', color: '#fff', font: '12px Menlo, Consolas, monospace', + borderRadius: '6px', zIndex: '2000' + }); + const header = document.createElement('div'); + header.textContent = 'Operation Code & Context'; + header.style.marginBottom = '6px'; + header.style.opacity = '0.9'; + wrapper.appendChild(header); + try { + const API_BASE = window.__TRITON_VIZ_API__ || ''; + const res = await fetch(`${API_BASE}/api/op_code`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ uuid, frame_idx: frameIdx, context }) }); + const data = await res.json(); + const meta = document.createElement('div'); + meta.style.marginBottom = '4px'; + meta.textContent = `${data.filename || ''}:${data.lineno || ''}`; + wrapper.appendChild(meta); + const pre = document.createElement('pre'); + pre.style.margin = '0'; + pre.style.whiteSpace = 'pre'; + const lines = (data.lines || []).map(l => { + const mark = (data.highlight === l.no) ? '▶ ' : ' '; + return `${mark}${String(l.no).padStart(6,' ')} | ${l.text||''}`; + }).join('\n'); + pre.textContent = lines || '(no code available)'; + wrapper.appendChild(pre); + } catch (e) { + const err = document.createElement('div'); + err.textContent = 'Failed to load code context.'; + wrapper.appendChild(err); + } + document.body.appendChild(wrapper); + panel = wrapper; + }; + + btn.addEventListener('click', async () => { + const turnOn = btn.textContent.endsWith('OFF'); + btn.textContent = `Show Code: ${turnOn ? 'ON' : 'OFF'}`; + if (!this.blockData || this.blockData.length === 0) { if (!turnOn) destroyPanel(); return; } + // Prefer latest clicked/active uuid; fallback to first in current grid block + const activeUuid = window.current_op_uuid || (this.blockData[0] && this.blockData[0].uuid); + if (turnOn && activeUuid) await createPanel(activeUuid, 0, 8); else destroyPanel(); + }); + } } createVisualizationContainer() { @@ -128,12 +198,23 @@ export class GridBlock { }); let currentSelectedTab = null; + // Tabs for each op this.blockData.forEach((op, index) => { const opTab = this.createOperationTab(op, index === 0); opTab.addEventListener('click', () => this.handleTabClick(opTab, op, currentSelectedTab)); headerBar.appendChild(opTab); if (index === 0) currentSelectedTab = opTab; }); + // Extra: NKI view tab (aggregates all ops) + const nkiTab = document.createElement('button'); + nkiTab.textContent = 'Flow'; + Object.assign(nkiTab.style, { flex:'0 0 auto', marginRight:'5px', background:'#333', color:'#fff', border:'none', padding:'10px', cursor:'pointer' }); + nkiTab.addEventListener('click', () => { + if (currentSelectedTab) currentSelectedTab.style.backgroundColor = '#333'; + nkiTab.style.backgroundColor = '#555'; + this.displayFlowDiagram(); + }); + headerBar.appendChild(nkiTab); return headerBar; } @@ -209,6 +290,9 @@ export class GridBlock { case 'Store': this.visualizationCleanupFunction = createStoreVisualization(this.contentArea, op); break; + case 'Flip': + this.visualizationCleanupFunction = createFlipVisualization(this.contentArea, op); + break; default: const unsupportedMsg = document.createElement('p'); unsupportedMsg.textContent = `Visualization not supported for ${op.type} operation`; @@ -217,6 +301,14 @@ export class GridBlock { } } + displayFlowDiagram() { + if (!this.contentArea) return; + if (this.visualizationCleanupFunction) { this.visualizationCleanupFunction(); this.visualizationCleanupFunction = null; } + this.contentArea.innerHTML = ''; + // Pass the entire block data (ops in this program) to Flow view + this.visualizationCleanupFunction = createFlowDiagram(this.contentArea, this.blockData || []); + } + createCloseButton() { const closeButton = document.createElement('button'); @@ -231,6 +323,26 @@ export class GridBlock { return closeButton; } + createBackButton() { + const backButton = document.createElement('button'); + backButton.textContent = 'Back'; + Object.assign(backButton.style, { + position: 'fixed', + left: '10px', + bottom: '10px', + zIndex: '2002', + background: 'rgba(0,0,0,0.65)', + color: '#fff', + border: '1px solid #666', + padding: '6px 10px', + borderRadius: '6px', + cursor: 'pointer', + boxShadow: '0 2px 8px rgba(0,0,0,0.5)' + }); + backButton.addEventListener('click', () => this.hideDetailedView()); + return backButton; + } + hideDetailedView() { if (!this.isDetailedViewVisible) return; diff --git a/triton_viz/static/load.js b/triton_viz/static/load.js index 73a86e8f..ba84edb5 100644 --- a/triton_viz/static/load.js +++ b/triton_viz/static/load.js @@ -14,7 +14,8 @@ import { export function createLoadVisualization(containerElement, op) { console.log(op.uuid); - fetch('/api/setop', { + const API_BASE = window.__TRITON_VIZ_API__ || ''; + fetch(`${API_BASE}/api/setop`, { method: 'POST', headers: { 'Content-Type': 'application/json', @@ -25,6 +26,9 @@ export function createLoadVisualization(containerElement, op) { .then(data => console.log('Set current op:', data)) .catch((error) => console.error('Error:', error)); + // expose current op uuid globally for generic code panel in gridblock + try { window.current_op_uuid = op.uuid; } catch(e){} + let currentStep = 0; let frame = 0; let isPaused = false; @@ -63,6 +67,28 @@ export function createLoadVisualization(containerElement, op) { const dragToggle = document.createElement('button'); dragToggle.textContent = 'Drag Cubes: OFF'; controlBar.appendChild(dragToggle); + const codeToggle = document.createElement('button'); + codeToggle.textContent = 'Show Code: OFF'; + controlBar.appendChild(codeToggle); + // Mouse pick calibration (dx/dy in pixels) + const calibWrap = document.createElement('div'); + calibWrap.style.display = 'flex'; + calibWrap.style.alignItems = 'center'; + calibWrap.style.gap = '4px'; + const calibLabel = document.createElement('span'); + calibLabel.textContent = 'Calib:'; + calibLabel.style.opacity = '0.8'; + const dxMinus = document.createElement('button'); dxMinus.textContent = '−X'; + const dxPlus = document.createElement('button'); dxPlus.textContent = '+X'; + const dyMinus = document.createElement('button'); dyMinus.textContent = '−Y'; + const dyPlus = document.createElement('button'); dyPlus.textContent = '+Y'; + const dxdyInfo = document.createElement('span'); dxdyInfo.style.minWidth = '70px'; dxdyInfo.style.textAlign = 'center'; + const dxdyReset = document.createElement('button'); dxdyReset.textContent = 'Reset'; + calibWrap.appendChild(calibLabel); + calibWrap.appendChild(dxMinus); calibWrap.appendChild(dxPlus); + calibWrap.appendChild(dyMinus); calibWrap.appendChild(dyPlus); + calibWrap.appendChild(dxdyInfo); calibWrap.appendChild(dxdyReset); + controlBar.appendChild(calibWrap); containerElement.appendChild(controlBar); // expose for debugging @@ -79,6 +105,7 @@ export function createLoadVisualization(containerElement, op) { let legendEl = null; let scheme = 'mono'; let monoBaseHex = '#3b82f6'; + let codePanel = null; const COLOR_GLOBAL = new THREE.Color(0.2, 0.2, 0.2); // Dark Gray const COLOR_SLICE = new THREE.Color(0.0, 0.7, 1.0); // Cyan (starting color for global slice) @@ -105,6 +132,27 @@ export function createLoadVisualization(containerElement, op) { ); addLabels(scene, globalTensor, sliceTensor); + + // Overlay memory flow badges if available (NKI only) + try { + const badge = document.createElement('div'); + badge.style.position = 'fixed'; + badge.style.right = '10px'; + badge.style.top = '60px'; + badge.style.zIndex = '2500'; + badge.style.background = 'rgba(0,0,0,0.65)'; + badge.style.color = '#fff'; + badge.style.padding = '6px 8px'; + badge.style.borderRadius = '6px'; + badge.style.font = '12px Arial'; + const ms = (op.mem_src||'').toUpperCase(); + const md = (op.mem_dst||'').toUpperCase(); + const by = Number(op.bytes||0); + if (ms && md) { + badge.innerHTML = `Memory Flow
${ms} → ${md}${by?`
${by} B`:''}`; + containerElement.appendChild(badge); + } + } catch(e){} const { center } = setupCamera(scene, camera); const orbitControls = new OrbitControls(camera, renderer.domElement); orbitControls.enableDamping = true; @@ -116,6 +164,11 @@ export function createLoadVisualization(containerElement, op) { const raycaster = new THREE.Raycaster(); const mouse = new THREE.Vector2(); + // persistent calibration offsets + let mouseDx = Number(localStorage.getItem('viz_mouse_dx') || 0); + let mouseDy = Number(localStorage.getItem('viz_mouse_dy') || 0); + function updateDxDyLabel(){ dxdyInfo.textContent = `dx=${mouseDx}, dy=${mouseDy}`; } + updateDxDyLabel(); // Drag state let dragModeOn = false; let isDragging = false; @@ -135,8 +188,13 @@ export function createLoadVisualization(containerElement, op) { animate(); function _updateMouseNDC(event) { - mouse.x = (event.clientX / containerElement.clientWidth) * 2 - 1; - mouse.y = -(event.clientY / containerElement.clientHeight) * 2 + 1; + const rect = renderer.domElement.getBoundingClientRect(); + const dpr = (window.devicePixelRatio || 1); + const px = (event.clientX - rect.left + mouseDx) * dpr; + const py = (event.clientY - rect.top + mouseDy) * dpr; + const w = rect.width * dpr, h = rect.height * dpr; + mouse.x = (px / w) * 2 - 1; + mouse.y = -(py / h) * 2 + 1; } function _raycastAll() { @@ -157,8 +215,7 @@ export function createLoadVisualization(containerElement, op) { } async function onMouseMove(event) { - mouse.x = (event.clientX / containerElement.clientWidth) * 2 - 1; - mouse.y = -(event.clientY / containerElement.clientHeight) * 2 + 1; + _updateMouseNDC(event); raycaster.setFromCamera(mouse, camera); @@ -288,6 +345,63 @@ export function createLoadVisualization(containerElement, op) { legendEl = wrapper; } + function destroyCodePanel() { + if (codePanel && codePanel.remove) codePanel.remove(); + codePanel = null; + } + + async function createCodePanel(frameIdx = 0, context = 8) { + destroyCodePanel(); + const wrapper = document.createElement('div'); + wrapper.style.position = 'fixed'; + wrapper.style.right = '10px'; + wrapper.style.top = '50px'; + wrapper.style.width = '520px'; + wrapper.style.maxHeight = '60vh'; + wrapper.style.overflow = 'auto'; + wrapper.style.padding = '8px 10px'; + wrapper.style.background = 'rgba(0,0,0,0.65)'; + wrapper.style.color = '#fff'; + wrapper.style.font = '12px Menlo, Consolas, monospace'; + wrapper.style.borderRadius = '6px'; + wrapper.style.zIndex = '2000'; + + const header = document.createElement('div'); + header.textContent = 'Operation Code & Context'; + header.style.marginBottom = '6px'; + header.style.opacity = '0.9'; + wrapper.appendChild(header); + + try { + const res = await fetch(`${API_BASE}/api/op_code`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ uuid: op.uuid, frame_idx: frameIdx, context }) + }); + const data = await res.json(); + const meta = document.createElement('div'); + meta.style.marginBottom = '4px'; + meta.textContent = `${data.filename || ''}:${data.lineno || ''}`; + wrapper.appendChild(meta); + const pre = document.createElement('pre'); + pre.style.margin = '0'; + pre.style.whiteSpace = 'pre'; + const lines = (data.lines || []).map(l => { + const mark = (data.highlight === l.no) ? '▶ ' : ' '; + return `${mark}${String(l.no).padStart(6,' ')} | ${l.text||''}`; + }).join('\n'); + pre.textContent = lines || '(no code available)'; + wrapper.appendChild(pre); + } catch (e) { + const err = document.createElement('div'); + err.textContent = 'Failed to load code context.'; + wrapper.appendChild(err); + } + + containerElement.appendChild(wrapper); + codePanel = wrapper; + } + function applyColorMapIfNeeded() { if (!colorizeOn || !tensorCache) return; const { min, max, dims, values } = tensorCache; @@ -405,7 +519,7 @@ export function createLoadVisualization(containerElement, op) { async function getElementValue(tensorName, x, y, z) { let uuid = op.uuid; - const response = await fetch('/api/getLoadValue', { + const response = await fetch(`${API_BASE}/api/getLoadValue`, { method: 'POST', headers: { 'Content-Type': 'application/json', @@ -417,7 +531,7 @@ export function createLoadVisualization(containerElement, op) { async function fetchGlobalTensor() { try { - const res = await fetch('/api/getLoadTensor', { + const res = await fetch(`${API_BASE}/api/getLoadTensor`, { method: 'POST', headers: { 'Content-Type': 'application/json' }, body: JSON.stringify({ uuid: op.uuid }) @@ -480,6 +594,23 @@ export function createLoadVisualization(containerElement, op) { orbitControls.enabled = !dragModeOn; }); + codeToggle.addEventListener('click', async () => { + const on = codeToggle.textContent.endsWith('OFF'); + codeToggle.textContent = `Show Code: ${on ? 'ON' : 'OFF'}`; + if (on) { + await createCodePanel(0, 8); + } else { + destroyCodePanel(); + } + }); + + // calibration handlers + dxMinus.addEventListener('click', ()=>{ mouseDx -= 1; localStorage.setItem('viz_mouse_dx', String(mouseDx)); updateDxDyLabel(); }); + dxPlus.addEventListener('click', ()=>{ mouseDx += 1; localStorage.setItem('viz_mouse_dx', String(mouseDx)); updateDxDyLabel(); }); + dyMinus.addEventListener('click', ()=>{ mouseDy -= 1; localStorage.setItem('viz_mouse_dy', String(mouseDy)); updateDxDyLabel(); }); + dyPlus.addEventListener('click', ()=>{ mouseDy += 1; localStorage.setItem('viz_mouse_dy', String(mouseDy)); updateDxDyLabel(); }); + dxdyReset.addEventListener('click', ()=>{ mouseDx = 0; mouseDy = 0; localStorage.setItem('viz_mouse_dx','0'); localStorage.setItem('viz_mouse_dy','0'); updateDxDyLabel(); }); + function updateSideMenu(tensorName, x, y, z, value) { if (!tensorName) { sideMenu.innerHTML = ''; diff --git a/triton_viz/static/load_utils.js b/triton_viz/static/load_utils.js index 52ea6ff9..80e468c1 100644 --- a/triton_viz/static/load_utils.js +++ b/triton_viz/static/load_utils.js @@ -12,6 +12,9 @@ export function setupScene(container, backgroundColor = 0x000000) { scene.background = new THREE.Color(backgroundColor); const camera = new THREE.PerspectiveCamera(45, container.clientWidth / container.clientHeight, 0.1, 1000); const renderer = new THREE.WebGLRenderer({ antialias: true }); + // Honour device pixel ratio to align raycaster with drawn pixels + const dpr = (window.devicePixelRatio || 1); + renderer.setPixelRatio(dpr); renderer.setSize(container.clientWidth, container.clientHeight); container.appendChild(renderer.domElement); @@ -58,9 +61,21 @@ export function createCube(color, tensorName, x, y, z, cubeGeometry, edgesGeomet export function createTensor(shape, coords, color, tensorName, cubeGeometry, edgesGeometry, lineMaterial) { console.log(`Creating ${tensorName} tensor:`, shape, coords); const tensor = new THREE.Group(); - let [width, height, depth] = shape; - depth = depth || 1; - height = height || 1; + // Normalize shape to width (X), height (Y), depth (Z) + let width, height, depth; + if (shape.length === 1) { + width = shape[0]; + height = 1; + depth = 1; + } else if (shape.length === 2) { + // Backend provides (H, W) for 2D tensors; interpret as width=W, height=H + height = shape[0]; + width = shape[1]; + depth = 1; + } else { + // Assume incoming order already matches [width, height, depth] + [width, height, depth] = shape; + } if (tensorName === 'Global') { console.log(`Creating global tensor with dimensions: ${width}x${height}x${depth}`); @@ -155,7 +170,20 @@ export function createTensor(shape, coords, color, tensorName, cubeGeometry, edg } export function calculateTensorSize(shape) { - const [width, height, depth] = shape; + // Normalize shape for size calculation consistent with createTensor + let width, height, depth; + if (shape.length === 1) { + width = shape[0]; + height = 1; + depth = 1; + } else if (shape.length === 2) { + // (H, W) -> width=W, height=H + height = shape[0]; + width = shape[1]; + depth = 1; + } else { + [width, height, depth] = shape; + } return new THREE.Vector3( width * (CUBE_SIZE + GAP), height * (CUBE_SIZE + GAP), diff --git a/triton_viz/static/matmul.js b/triton_viz/static/matmul.js index c32e9285..d4d89aba 100644 --- a/triton_viz/static/matmul.js +++ b/triton_viz/static/matmul.js @@ -1,9 +1,11 @@ import * as THREE from 'https://esm.sh/three@0.155.0/build/three.module.js'; +import { OrbitControls } from 'https://esm.sh/three@0.155.0/examples/jsm/controls/OrbitControls.js'; export function createMatMulVisualization(containerElement, op) { + const API_BASE = window.__TRITON_VIZ_API__ || ''; const { input_shape, other_shape, output_shape } = op; console.log(op.uuid) - fetch('/api/setop', { + fetch(`${API_BASE}/api/setop`, { method: 'POST', headers: { 'Content-Type': 'application/json', @@ -13,13 +15,12 @@ export function createMatMulVisualization(containerElement, op) { .then(response => response.json()) .then(data => console.log('Set current op:', data)) .catch((error) => console.error('Error:', error)); - let currentStep = 0; - const totalSteps = input_shape[1]; - let frame = 0; + let currentStep = 0; // kept for compatibility with getValue; not used for animation const sideMenu = document.createElement('div'); - sideMenu.style.position = 'absolute'; + // Use fixed position and high z-index to ensure it's above WebGL canvas + sideMenu.style.position = 'fixed'; sideMenu.style.top = '10px'; sideMenu.style.right = '10px'; sideMenu.style.width = '200px'; @@ -29,7 +30,31 @@ export function createMatMulVisualization(containerElement, op) { sideMenu.style.fontFamily = 'Arial, sans-serif'; sideMenu.style.fontSize = '14px'; sideMenu.style.borderRadius = '5px'; + sideMenu.style.zIndex = '3000'; + sideMenu.style.pointerEvents = 'auto'; containerElement.appendChild(sideMenu); + + // Memory flow badge for Dot (SBUF→PSUM) + try { + const badge = document.createElement('div'); + badge.style.position = 'fixed'; + badge.style.right = '10px'; + badge.style.top = '60px'; + badge.style.zIndex = '2500'; + badge.style.background = 'rgba(0,0,0,0.65)'; + badge.style.color = '#fff'; + badge.style.padding = '6px 8px'; + badge.style.borderRadius = '6px'; + badge.style.font = '12px Arial'; + const ms = (op.mem_src||'').toUpperCase(); + const md = (op.mem_dst||'').toUpperCase(); + if (ms && md) { + const tiles = Array.isArray(op.tile_shape)? `${op.tile_shape[0]}x${op.tile_shape[1]}` : ''; + const k = Number(op.k||0); + badge.innerHTML = `Memory Flow
${ms} → ${md}${tiles?`
Tile: ${tiles}`:''}${k?`
K: ${k}`:''}`; + containerElement.appendChild(badge); + } + } catch(e){} let hoveredCube = null; @@ -37,7 +62,7 @@ export function createMatMulVisualization(containerElement, op) { async function getElementValue( matrixName, row, col) { let uuid = op.uuid; - const response = await fetch('/api/getValue', { + const response = await fetch(`${API_BASE}/api/getValue`, { method: 'POST', headers: { 'Content-Type': 'application/json', @@ -48,43 +73,70 @@ export function createMatMulVisualization(containerElement, op) { } - function updateSideMenu(matrix, x, y) { - if (!matrix) { + function updateSideMenu(matrixOrName, x, y, vectors, value) { + if (!matrixOrName) { sideMenu.innerHTML = ''; return; } let matrixName; let dims; - if (matrix === matrixA) { - matrixName = 'A'; - dims = input_shape; - } else if (matrix === matrixB) { - matrixName = 'B'; - dims = other_shape; - } else if (matrix === matrixC) { - matrixName = 'C'; - dims = output_shape; + // Accept both string name ('A'|'B'|'C') and matrix group object + if (typeof matrixOrName === 'string') { + const name = matrixOrName.toUpperCase(); + if (name === 'A') { matrixName = 'A'; dims = input_shape; } + else if (name === 'B') { matrixName = 'B'; dims = other_shape; } + else if (name === 'C') { matrixName = 'C'; dims = output_shape; } + else { sideMenu.innerHTML = ''; return; } } else { - sideMenu.innerHTML = ''; - return; + if (matrixOrName === matrixA) { matrixName = 'A'; dims = input_shape; } + else if (matrixOrName === matrixB) { matrixName = 'B'; dims = other_shape; } + else if (matrixOrName === matrixC) { matrixName = 'C'; dims = output_shape; } + else { sideMenu.innerHTML = ''; return; } } console.log(matrixName, "x:", (x + 1), "y:", (y + 1)); + let extra = ''; + if (matrixName === 'C' && vectors && !vectors.error) { + const aRow = vectors.a_row || []; + const bCol = vectors.b_col || []; + const k = vectors.k || 0; + extra = ` +
+
From A row: [${aRow.slice(0,8).join(', ')}${aRow.length>8?' …':''}]
+
From B col: [${bCol.slice(0,8).join(', ')}${bCol.length>8?' …':''}]
+
k: ${k}
+ `; + } + sideMenu.innerHTML = `

Matrix ${matrixName}

Row: ${y + 1}

Column: ${x + 1}

Dimensions: ${dims[0]} x ${dims[1]}

+

Value: ${value !== undefined ? value : 'Loading...'}

+ ${extra} `; } const raycaster = new THREE.Raycaster(); const mouse = new THREE.Vector2(); - + // persistent calibration offsets shared with Load view for consistency + let mouseDx = Number(localStorage.getItem('viz_mouse_dx') || 0); + let mouseDy = Number(localStorage.getItem('viz_mouse_dy') || 0); + + + function _updateMouseNDC(event) { + const rect = renderer.domElement.getBoundingClientRect(); + const dpr = (window.devicePixelRatio || 1); + const px = (event.clientX - rect.left + mouseDx) * dpr; + const py = (event.clientY - rect.top + mouseDy) * dpr; + const w = rect.width * dpr, h = rect.height * dpr; + mouse.x = (px / w) * 2 - 1; + mouse.y = -(py / h) * 2 + 1; + } async function onMouseMove(event) { - mouse.x = (event.clientX / containerElement.clientWidth) * 2 - 1; - mouse.y = -(event.clientY / containerElement.clientHeight) * 2 + 1; + _updateMouseNDC(event); raycaster.setFromCamera(mouse, camera); @@ -122,10 +174,47 @@ export function createMatMulVisualization(containerElement, op) { `Value: ${res.value}` ); - updateSideMenu(hoveredCube.matrixName, hoveredCube.matrixRow, hoveredCube.matrixCol); + // If hovering C, fetch vectors and update panel + let vectors = null; + let valueForPanel = res && res.value !== undefined ? res.value : undefined; + if (hoveredCube.matrixName === 'C') { + try { + const resp = await fetch(`${API_BASE}/api/getMatmulVectors`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ uuid: op.uuid, row: hoveredCube.matrixRow, col: hoveredCube.matrixCol }) + }); + vectors = await resp.json(); + // Compute C[row,col] directly from vectors if available + if (vectors && !vectors.error && Array.isArray(vectors.a_row) && Array.isArray(vectors.b_col)) { + const aRow = vectors.a_row; + const bCol = vectors.b_col; + const len = Math.min(aRow.length, bCol.length); + let sum = 0.0; + for (let i = 0; i < len; i++) sum += aRow[i] * bCol[i]; + valueForPanel = sum; + } + } catch (e) { vectors = { error: String(e) }; } + + // highlight A[row,:], B[:,col], C[row,col] + resetColors(); + const row = hoveredCube.matrixRow; + const col = hoveredCube.matrixCol; + const highlightA = Array.from({ length: input_shape[1] }, (_, i) => [row, i]); + const highlightB = Array.from({ length: other_shape[0] }, (_, i) => [i, col]); + const highlightC = [[row, col]]; + highlightCubes(matrixA, highlightA, COLOR_HIGHLIGHT); + highlightCubes(matrixB, highlightB, COLOR_HIGHLIGHT); + highlightCubes(matrixC, highlightC, COLOR_FILLED); + } else { + // hovering A/B or others -> just reset + resetColors(); + } + updateSideMenu(hoveredCube.matrixName, hoveredCube.matrixRow, hoveredCube.matrixCol, vectors, valueForPanel); } } else { updateSideMenu(null); + resetColors(); } } @@ -150,8 +239,14 @@ export function createMatMulVisualization(containerElement, op) { const camera = new THREE.PerspectiveCamera(45, containerElement.clientWidth / containerElement.clientHeight, 0.1, 1000); const renderer = new THREE.WebGLRenderer({ antialias: true }); + const dpr = (window.devicePixelRatio || 1); + renderer.setPixelRatio(dpr); renderer.setSize(containerElement.clientWidth, containerElement.clientHeight); containerElement.appendChild(renderer.domElement); + // Ensure canvas is under overlays + renderer.domElement.style.position = 'relative'; + renderer.domElement.style.zIndex = '1'; + try { containerElement.style.pointerEvents = 'auto'; } catch(e){} const ambientLight = new THREE.AmbientLight(0xffffff, 0.5); scene.add(ambientLight); @@ -221,10 +316,13 @@ export function createMatMulVisualization(containerElement, op) { camera.position.set(center.x, center.y, center.z + cameraZ); camera.lookAt(center); + const controls = new OrbitControls(camera, renderer.domElement); + controls.enableDamping = true; + controls.dampingFactor = 0.05; + controls.target.copy(center); + controls.update(); - let isPaused = false; - - const totalFrames = input_shape[0] * other_shape[1]; + // remove auto animation; highlight is driven by mouse hover only function highlightCubes(matrix, indices, highlightColor) { indices.forEach(([i, j]) => { @@ -241,29 +339,12 @@ export function createMatMulVisualization(containerElement, op) { function resetColors() { matrixA.children.forEach(cube => cube.material.color.copy(COLOR_A)); matrixB.children.forEach(cube => cube.material.color.copy(COLOR_B)); + matrixC.children.forEach(cube => cube.material.color.copy(COLOR_C)); } function animate() { requestAnimationFrame(animate); - - if (!isPaused && frame < totalFrames) { - resetColors(); - - const row = Math.floor(frame / other_shape[1]); - const col = frame % other_shape[1]; - currentStep = frame % totalSteps + 1; - - const highlightA = Array.from({ length: input_shape[1] }, (_, i) => [row, i]); - const highlightB = Array.from({ length: other_shape[0] }, (_, i) => [i, col]); - const highlightC = [[row, col]]; - - highlightCubes(matrixA, highlightA, COLOR_HIGHLIGHT); - highlightCubes(matrixB, highlightB, COLOR_HIGHLIGHT); - highlightCubes(matrixC, highlightC, COLOR_FILLED); - - frame++; - } - + controls.update(); renderer.render(scene, camera); } @@ -274,27 +355,46 @@ export function createMatMulVisualization(containerElement, op) { } const controlPanel = document.createElement('div'); - controlPanel.style.position = 'absolute'; + controlPanel.style.position = 'fixed'; controlPanel.style.bottom = '10px'; controlPanel.style.left = '10px'; controlPanel.style.display = 'flex'; controlPanel.style.gap = '10px'; - - const playPauseButton = document.createElement('button'); - playPauseButton.textContent = 'Play/Pause'; - playPauseButton.addEventListener('click', () => { - isPaused = !isPaused; - }); - - const resetButton = document.createElement('button'); - resetButton.textContent = 'Reset'; - resetButton.addEventListener('click', () => { - frame = 0; - resetColors(); - }); - - controlPanel.appendChild(playPauseButton); - controlPanel.appendChild(resetButton); + controlPanel.style.zIndex = '3000'; + controlPanel.style.pointerEvents = 'auto'; + + // Removed animation controls; keep panel for color-by-value toggle only + + // Color by Value (for C matrix) toggle + legend (mono blue) + const colorToggle = document.createElement('button'); + colorToggle.textContent = 'Color by Value: OFF'; + controlPanel.appendChild(colorToggle); + let colorOn = false; + let legendEl = null; + function destroyLegend(){ if(legendEl && legendEl.remove) legendEl.remove(); legendEl=null; } + function createLegend(min,max){ + destroyLegend(); + const w = document.createElement('div'); + Object.assign(w.style,{position:'absolute', left:'10px', bottom:'60px', background:'rgba(0,0,0,0.6)', color:'#fff', padding:'6px 8px', borderRadius:'6px', zIndex:'2000', pointerEvents:'auto'}); + const c = document.createElement('canvas'); c.width=220; c.height=10; const ctx=c.getContext('2d'); + for(let x=0;x${min.toFixed(3)}${max.toFixed(3)}`; + const ttl=document.createElement('div'); ttl.textContent='Value (C)'; ttl.style.marginBottom='4px'; ttl.style.opacity='0.9'; + w.appendChild(ttl); w.appendChild(c); w.appendChild(lab); containerElement.appendChild(w); legendEl=w; + } + async function fetchCValues(){ + try{ const res=await fetch(`${API_BASE}/api/getMatmulC`,{method:'POST',headers:{'Content-Type':'application/json'},body:JSON.stringify({uuid: op.uuid})}); + return await res.json(); }catch(e){ return {error:String(e)} } + } + function applyColorCWithData(data){ if(!colorOn||!data||data.error) return; try{ + const M=(data.shape||[])[0]||0, N=(data.shape||[])[1]||0; const vals2d=data.values||[]; const vals=[]; const cells=[]; + matrixC.children.forEach((cube,i)=>{ const row=Math.floor(i/N), col=i%N; cells.push({cube,row,col}); vals.push((row{ const v=vals[idx]; const u=(mx===mn)?0.5:(v-mn)/(mx-mn); const r=u,g=0.2,b=1-u; e.cube.material.color.setRGB(r,g,b); }); + createLegend(mn,mx); + }catch(e){} + } + colorToggle.addEventListener('click', async ()=>{ colorOn=!colorOn; colorToggle.textContent=`Color by Value: ${colorOn?'ON':'OFF'}`; if(!colorOn){ destroyLegend(); resetColors(); } else { const data=await fetchCValues(); applyColorCWithData(data); } }); containerElement.appendChild(controlPanel); @@ -343,6 +443,7 @@ export function createMatMulVisualization(containerElement, op) { window.addEventListener('resize', onResize); window.addEventListener('keydown', onKeyDown); containerElement.addEventListener('mousemove', onMouseMove); + try { window.current_op_uuid = op.uuid; } catch (e) {} // Mouse wheel zoom for matmul view const WHEEL_ZOOM_SPEED = 0.5; diff --git a/triton_viz/static/nki.js b/triton_viz/static/nki.js new file mode 100644 index 00000000..fd155bc8 --- /dev/null +++ b/triton_viz/static/nki.js @@ -0,0 +1,107 @@ +export function createFlowDiagram(containerElement, opsByProgram) { +// Minimal NKI flow view: three lanes (HBM, SBUF, PSUM) and arrows per op time_idx +// opsByProgram: array of op objects for a grid block (Load/Store/Dot/Copy) with mem_* fields + const laneNames = ["HBM", "SBUF", "PSUM"]; + const laneY = { HBM: 60, SBUF: 160, PSUM: 260 }; + const width = containerElement.clientWidth || 1200; + const height = 360; + + const canvas = document.createElement('canvas'); + canvas.width = width; + canvas.height = height; + canvas.style.width = '100%'; + canvas.style.height = '360px'; + canvas.style.background = '#111'; + canvas.style.border = '1px solid #333'; + containerElement.appendChild(canvas); + const ctx = canvas.getContext('2d'); + + // Draw lanes + ctx.font = '14px Arial'; + ctx.fillStyle = '#ddd'; + ctx.strokeStyle = '#444'; + laneNames.forEach(name => { + const y = laneY[name]; + ctx.beginPath(); + ctx.moveTo(80, y); + ctx.lineTo(width - 20, y); + ctx.stroke(); + ctx.fillText(name, 20, y + 5); + }); + + // Normalize time to X range + const events = opsByProgram + .filter(op => op.time_idx !== undefined && op.time_idx >= 0) + .map(op => ({ + t: op.time_idx, + src: (op.mem_src || '').toUpperCase(), + dst: (op.mem_dst || '').toUpperCase(), + bytes: Number(op.bytes || 0), + type: op.type, + uuid: op.uuid + })); + if (events.length === 0) { + ctx.fillStyle = '#aaa'; + ctx.fillText('No NKI metadata found. Run NKI example or add mem_src/mem_dst/bytes/time_idx to records.', 80, height/2); + return () => { containerElement.innerHTML = ''; }; + } + events.sort((a,b)=>a.t-b.t); + const tMin = events[0].t; + const tMax = events[events.length-1].t || (tMin+1); + const toX = t => 80 + (width-120) * (t - tMin) / Math.max(1, (tMax - tMin)); + + // Color mapping per type + const colorFor = (type) => { + switch ((type||'').toLowerCase()){ + case 'load': return '#00bcd4'; + case 'store': return '#ff9800'; + case 'dot': return '#8bc34a'; + default: return '#9e9e9e'; + } + }; + + // Draw arrows for each event + events.forEach(e => { + const y1 = laneY[e.src] ?? laneY.SBUF; + const y2 = laneY[e.dst] ?? laneY.SBUF; + const x = toX(e.t); + const col = colorFor(e.type); + // Arrow thickness by bytes (log scaled) + const w = Math.max(1, Math.log10((e.bytes||1))); + ctx.strokeStyle = col; + ctx.lineWidth = w; + ctx.beginPath(); + ctx.moveTo(x, y1); + ctx.lineTo(x, y2); + ctx.stroke(); + // Arrow head + ctx.beginPath(); + const dir = Math.sign(y2 - y1) || 1; + const hx = x; const hy = y2; + ctx.moveTo(hx, hy); + ctx.lineTo(hx - 6, hy - 6*dir); + ctx.lineTo(hx + 6, hy - 6*dir); + ctx.closePath(); + ctx.fillStyle = col; + ctx.fill(); + // Label + ctx.fillStyle = '#ccc'; + ctx.font = '11px Arial'; + const label = `${e.type} ${e.bytes||0}B`; + ctx.fillText(label, x + 6, (y1+y2)/2 - 6); + }); + + // Legend + const legend = document.createElement('div'); + legend.style.position = 'relative'; + legend.style.marginTop = '6px'; + legend.style.color = '#ddd'; + legend.style.font = '12px Arial'; + legend.innerHTML = `Flow Diagram   ■ Load  ■ Dot(PSUM)  ■ Store`; + containerElement.appendChild(legend); + + return () => { containerElement.innerHTML = ''; }; +} + +// Backward compatibility alias +export const createNKIFlow = createFlowDiagram; diff --git a/triton_viz/static/store-utils.js b/triton_viz/static/store-utils.js index d87d78bc..12b9df11 100644 --- a/triton_viz/static/store-utils.js +++ b/triton_viz/static/store-utils.js @@ -3,7 +3,8 @@ import * as THREE from 'https://cdn.jsdelivr.net/npm/three@0.155.0/build/three.m export function createMatMulVisualization(containerElement, op) { const { input_shape, other_shape, output_shape } = op; console.log(op.uuid) - fetch('/api/setop', { + const API_BASE = window.__TRITON_VIZ_API__ || ''; + fetch(`${API_BASE}/api/setop`, { method: 'POST', headers: { 'Content-Type': 'application/json', @@ -37,7 +38,7 @@ export function createMatMulVisualization(containerElement, op) { async function getElementValue( matrixName, row, col) { let uuid = op.uuid; - const response = await fetch('/api/getValue', { + const response = await fetch(`${API_BASE}/api/getValue`, { method: 'POST', headers: { 'Content-Type': 'application/json', diff --git a/triton_viz/static/store.js b/triton_viz/static/store.js index 9868c16a..b276c083 100644 --- a/triton_viz/static/store.js +++ b/triton_viz/static/store.js @@ -10,11 +10,14 @@ import { setupEventListeners, cameraControls } from './load_utils.js'; +import { createFlipDemo } from './flip_demo.js'; +import { createFlip3D } from './flip_3d.js'; export function createStoreVisualization(containerElement, op) { console.log(op.uuid); - fetch('/api/setop', { + const API_BASE = window.__TRITON_VIZ_API__ || ''; + fetch(`${API_BASE}/api/setop`, { method: 'POST', headers: { 'Content-Type': 'application/json', @@ -45,6 +48,7 @@ export function createStoreVisualization(containerElement, op) { containerElement.appendChild(controlBar); let dragModeOn = false; let hoveredCube = null; + let flipCleanup = null; const COLOR_GLOBAL = new THREE.Color(0.2, 0.2, 0.2); // Dark Gray const COLOR_SLICE = new THREE.Color(0.0, 0.7, 1.0); // Cyan (starting color for global slice) @@ -66,6 +70,27 @@ export function createStoreVisualization(containerElement, op) { scene.add(sliceTensor); addLabels(scene, globalTensor, sliceTensor); + + // Overlay memory flow badges if available (NKI only) + try { + const badge = document.createElement('div'); + badge.style.position = 'fixed'; + badge.style.right = '10px'; + badge.style.top = '60px'; + badge.style.zIndex = '2500'; + badge.style.background = 'rgba(0,0,0,0.65)'; + badge.style.color = '#fff'; + badge.style.padding = '6px 8px'; + badge.style.borderRadius = '6px'; + badge.style.font = '12px Arial'; + const ms = (op.mem_src||'').toUpperCase(); + const md = (op.mem_dst||'').toUpperCase(); + const by = Number(op.bytes||0); + if (ms && md) { + badge.innerHTML = `Memory Flow
${ms} → ${md}${by?`
${by} B`:''}`; + containerElement.appendChild(badge); + } + } catch(e){} const { center } = setupCamera(scene, camera); const orbitControls = new OrbitControls(camera, renderer.domElement); orbitControls.enableDamping = true; @@ -95,6 +120,7 @@ export function createStoreVisualization(containerElement, op) { dragToggle.textContent = `Drag Cubes: ${dragModeOn ? 'ON' : 'OFF'}`; orbitControls.enabled = !dragModeOn; }); + // Removed Flip demo button from Store view; Flip visualization is available under Flip op. animate(); function _updateMouseNDC(event) { @@ -229,7 +255,7 @@ export function createStoreVisualization(containerElement, op) { async function getElementValue(tensorName, x, y, z) { let uuid = op.uuid; - const response = await fetch('/api/getLoadValue', { + const response = await fetch(`${API_BASE}/api/getLoadValue`, { method: 'POST', headers: { 'Content-Type': 'application/json', diff --git a/triton_viz/static/visualization.js b/triton_viz/static/visualization.js index 15738ea5..bbfe8deb 100644 --- a/triton_viz/static/visualization.js +++ b/triton_viz/static/visualization.js @@ -379,7 +379,8 @@ function draw() { async function fetchData() { try { - const response = await fetch('/api/data'); + const API_BASE = window.__TRITON_VIZ_API__ || ''; + const response = await fetch(`${API_BASE}/api/data`); globalData = await response.json(); console.log(globalData); diff --git a/triton_viz/templates/calibrate.html b/triton_viz/templates/calibrate.html new file mode 100644 index 00000000..3e3d20da --- /dev/null +++ b/triton_viz/templates/calibrate.html @@ -0,0 +1,96 @@ + + + + + + Mouse Calibration + + + +
+
点击每个目标中心点进行采样(至少 4 次)。
+
样本: 0 | dx=0, dy=0
+ + +
+ + + + diff --git a/triton_viz/visualizer/draw.py b/triton_viz/visualizer/draw.py index 23feb087..be3bbd5b 100644 --- a/triton_viz/visualizer/draw.py +++ b/triton_viz/visualizer/draw.py @@ -5,6 +5,7 @@ Dot, Load, Store, + Flip, ) from triton_viz.clients.sanitizer.data import OutOfBoundsRecordBruteForce import numpy as np @@ -89,9 +90,10 @@ def extract_load_coords( record.masks, ) + # Report (x, y, z); delinearized returns (z, y, x) global_coords = [ (float(xi), float(yi), float(zi)) - for xi, yi, zi in zip(global_z, global_y, global_x) + for xi, yi, zi in zip(global_x, global_y, global_z) if xi != -1 and yi != -1 and zi != -1 ] @@ -169,13 +171,91 @@ def prepare_visualization_data(program_records, tensor_table): "other_shape": record.other_shape, "output_shape": record.output_shape, "uuid": record_uuid, + # provide C values for color-by-value + "c_shape": record.output_shape, + # NKI meta (accumulate to PSUM) + "mem_src": getattr(record, "mem_src", None), + "mem_dst": getattr(record, "mem_dst", None), + "tile_shape": getattr(record, "tile_shape", None), + "k": int(getattr(record, "k", 0)), + "time_idx": int(getattr(record, "time_idx", -1)), } ) + # Normalize Dot operands to NumPy arrays for downstream endpoints + try: + import numpy as _np + + def _to_numpy_cpu(x): + if isinstance(x, _np.ndarray): + return x + if hasattr(x, "cpu"): + try: + return x.detach().cpu().numpy() + except Exception: + return _np.asarray(x) + if hasattr(x, "data"): + return _np.asarray(getattr(x, "data")) + if hasattr(x, "_value"): + return _np.asarray(getattr(x, "_value")) + return _np.asarray(x) + + a_np = _to_numpy_cpu(record.input_data) + b_np = _to_numpy_cpu(record.other_data) + except Exception: + import numpy as _np + + a_np = _np.asarray(record.input_data) + b_np = _np.asarray(record.other_data) + raw_tensor_data[record_uuid] = { - "input_data": torch.tensor(record.input_data), - "other_data": torch.tensor(record.other_data), + "input_data": a_np, + "other_data": b_np, "intermediate_results": record.intermediate_results, + "tracebacks": [ + { + "filename": f.filename, + "lineno": f.lineno, + "line": f.line, + "name": f.name, + } + for f in getattr(record, "call_path", []) + ], + # prepare C values after kernel (if available from intermediate or recompute best-effort) + } + + elif isinstance(record, Flip): + visualization_data.append( + { + "type": "Flip", + "input_shape": record.input_shape, + "output_shape": record.output_shape, + "dim": int(getattr(record, "dim", 0)), + "uuid": record_uuid, + } + ) + + raw_tensor_data[record_uuid] = { + "tracebacks": [ + { + "filename": f.filename, + "lineno": f.lineno, + "line": f.line, + "name": f.name, + } + for f in getattr(record, "call_path", []) + ], + # best-effort payload for potential future value viz + "input_shape": list(record.input_shape), + "output_shape": list(record.output_shape), + "dim": int(getattr(record, "dim", 0)), + # optionally include data for hover value queries + "input_data": None + if getattr(record, "input_data", None) is None + else torch.tensor(record.input_data), + "output_data": None + if getattr(record, "output_data", None) is None + else torch.tensor(record.output_data), } elif isinstance(record, Load): @@ -191,12 +271,53 @@ def prepare_visualization_data(program_records, tensor_table): "global_coords": global_coords, "slice_coords": slice_coords, "uuid": record_uuid, + # NKI flow meta (optional) + "mem_src": getattr(record, "mem_src", None), + "mem_dst": getattr(record, "mem_dst", None), + "bytes": int(getattr(record, "bytes", 0)), + "time_idx": int(getattr(record, "time_idx", -1)), } ) + # Normalize to NumPy array for downstream APIs, and cache basic stats + try: + import numpy as _np + + gt = global_tensor.data + if hasattr(gt, "cpu") and callable(getattr(gt, "cpu", None)): + try: + arr = gt.detach().cpu().numpy() + except Exception: + arr = _np.asarray(gt) + elif hasattr(gt, "data"): + arr = _np.asarray(getattr(gt, "data")) + elif hasattr(gt, "_value"): + arr = _np.asarray(getattr(gt, "_value")) + else: + arr = _np.asarray(gt) + except Exception: + import numpy as _np + + arr = _np.asarray([]) + + t_min = float(np.min(arr)) if arr.size else 0.0 + t_max = float(np.max(arr)) if arr.size else 0.0 + raw_tensor_data[record_uuid] = { - "global_tensor": global_tensor.data.cpu(), # Ensure it's on CPU - "dims": len(global_tensor.data.cpu().shape), + "global_tensor": arr, + "dims": int(arr.ndim), + "shape": list(arr.shape), + "min": t_min, + "max": t_max, + "tracebacks": [ + { + "filename": f.filename, + "lineno": f.lineno, + "line": f.line, + "name": f.name, + } + for f in getattr(record, "call_path", []) + ], } print(record.masks.shape) @@ -213,9 +334,25 @@ def prepare_visualization_data(program_records, tensor_table): "global_coords": global_coords, "slice_coords": slice_coords, "uuid": record_uuid, + "mem_src": getattr(record, "mem_src", None), + "mem_dst": getattr(record, "mem_dst", None), + "bytes": int(getattr(record, "bytes", 0)), + "time_idx": int(getattr(record, "time_idx", -1)), } ) + raw_tensor_data[record_uuid] = { + "tracebacks": [ + { + "filename": f.filename, + "lineno": f.lineno, + "line": f.line, + "name": f.name, + } + for f in getattr(record, "call_path", []) + ], + } + return visualization_data, raw_tensor_data, "" diff --git a/triton_viz/visualizer/interface.py b/triton_viz/visualizer/interface.py index c815cd22..00c9947d 100644 --- a/triton_viz/visualizer/interface.py +++ b/triton_viz/visualizer/interface.py @@ -3,7 +3,6 @@ from .analysis import analyze_records from .draw import get_visualization_data import os -import torch from flask_cloudflared import _run_cloudflared import requests import time @@ -59,9 +58,7 @@ def precompute_c_values(op_data): for j in range(cols): precomputed[(i, j)] = [0] * (inner_dim + 1) for k in range(1, inner_dim + 1): - precomputed[(i, j)][k] = torch.dot( - input_data[i, :k], other_data[:k, j] - ).item() + precomputed[(i, j)][k] = input_data[i, :k] @ other_data[:k, j] return precomputed @@ -113,6 +110,35 @@ def update_global_data(): global_data["analysis"] = df_dict +def _safe_read_file_segment(filename: str, lineno: int, context: int = 8): + try: + # only allow files under current working dir for safety + cwd = os.path.realpath(os.getcwd()) + path = os.path.realpath(filename) + if not path.startswith(cwd): + return None + start = max(1, lineno - context) + end = lineno + context + lines = [] + with open(path, "r", encoding="utf-8", errors="ignore") as f: + for i, line in enumerate(f, start=1): + if i < start: + continue + if i > end: + break + lines.append({"no": i, "text": line.rstrip("\n")}) + return { + "filename": path, + "lineno": lineno, + "start": start, + "end": end, + "highlight": lineno, + "lines": lines, + } + except Exception: + return None + + @app.route("/") def index(): update_global_data() @@ -125,6 +151,12 @@ def debug_page(): return render_template("debug.html") +@app.route("/calibrate") +def calibrate_page(): + # A minimal page to calibrate mouse picking dx/dy + return render_template("calibrate.html") + + @app.route("/api/data") def get_data(): global global_data @@ -148,6 +180,161 @@ def set_current_op(): ) +@app.route("/api/op_code", methods=["POST"]) +def get_op_code(): + global raw_tensor_data, current_fullscreen_op + data = request.json or {} + # ensure data prepared + if raw_tensor_data is None: + update_global_data() + uuid = data.get("uuid") or current_fullscreen_op + frame_idx = int(data.get("frame_idx", 0)) + context = int(data.get("context", 8)) + if not uuid or raw_tensor_data is None or uuid not in raw_tensor_data: + return jsonify({"error": "Operation not found"}), 404 + tb_list = raw_tensor_data[uuid].get("tracebacks") or [] + if not tb_list: + return jsonify({"error": "Traceback not available"}), 200 + # Heuristic: pick the BEST user frame under CWD (closest to op) + cwd = os.path.realpath(os.getcwd()) + + def _score(tb: dict) -> int: + fn = tb.get("filename") or "" + line = (tb.get("line") or "").strip() + name = tb.get("name") or "" + p = os.path.realpath(fn) + score = 0 + + # 强优先:明确的 Triton 语言内核操作符 + if "tl.load" in line or "tl.store" in line: + score += 100 + elif "tl." in line: + score += 50 + + # 位置相关:项目代码优先,三方/框架降权 + if p.startswith(cwd): + score += 5 + if any( + s in p + for s in ["site-packages", "triton_viz/", "triton/", "runpy.py", "IPython"] + ): + score -= 10 + + # 语义相关:函数名看起来像 kernel 的加分 + if name.endswith("_kernel") or "kernel" in name: + score += 3 + + # 体验相关:examples 目录小幅加分 + if "examples" in p: + score += 1 + + return score + + # Prefer frames from tail (closest to current) and highest score + best = None + best_score = -(10**9) + for tb in reversed(tb_list): + sc = _score(tb) + if sc > best_score: + best, best_score = tb, sc + chosen = best if best is not None else None + if chosen is None: + # fallback: use requested frame or last non- frame + frame_idx = max(0, min(frame_idx, len(tb_list) - 1)) + chosen = tb_list[frame_idx] + for tb in reversed(tb_list): + fn = tb.get("filename") or "" + if not fn.startswith("<"): + chosen = tb + break + tb = chosen + filename = tb.get("filename") + lineno = int(tb.get("lineno", 0)) + line_of_code = tb.get("line") + seg = _safe_read_file_segment(filename, lineno, context) + if seg is None: + # fallback with single line + seg = { + "filename": filename, + "lineno": lineno, + "start": lineno, + "end": lineno, + "highlight": lineno, + "lines": [{"no": lineno, "text": line_of_code or ""}], + } + return jsonify(seg) + + +@app.route("/api/getMatmulC", methods=["POST"]) +def get_matmul_c(): + global raw_tensor_data + data = request.json or {} + uuid = data.get("uuid") + if not uuid or uuid not in raw_tensor_data: + return jsonify({"error": "Operation not found"}), 404 + op = raw_tensor_data[uuid] + a = op.get("input_data") + b = op.get("other_data") + if a is None or b is None: + return jsonify({"error": "MatMul tensors not available"}), 200 + try: + import numpy as _np + + a_np = _np.asarray(a) + b_np = _np.asarray(b) + c_np = a_np @ b_np + cmin = float(_np.min(c_np)) if c_np.size else 0.0 + cmax = float(_np.max(c_np)) if c_np.size else 0.0 + return jsonify( + { + "shape": list(c_np.shape), + "min": cmin, + "max": cmax, + "values": c_np.tolist(), + } + ) + except Exception as e: + return jsonify({"error": f"MatMul compute failed: {e}"}), 200 + + +@app.route("/api/getMatmulVectors", methods=["POST"]) +def get_matmul_vectors(): + """Return A[row, :] and B[:, col] for a given Dot op. + + Request JSON: { uuid, row, col } + Response JSON: { + "row": int, + "col": int, + "a_row": [float, ...], + "b_col": [float, ...], + "k": int + } + """ + global raw_tensor_data + data = request.json or {} + uuid = data.get("uuid") + row = int(data.get("row", 0)) + col = int(data.get("col", 0)) + if not uuid or uuid not in raw_tensor_data: + return jsonify({"error": "Operation not found"}), 404 + op = raw_tensor_data[uuid] + a = op.get("input_data") + b = op.get("other_data") + if a is None or b is None: + return jsonify({"error": "MatMul tensors not available"}), 200 + try: + import numpy as _np + + a_np = _np.asarray(a) + b_np = _np.asarray(b) + a_row = a_np[row, :].tolist() + b_col = b_np[:, col].tolist() + k = len(a_row) + return jsonify({"row": row, "col": col, "a_row": a_row, "b_col": b_col, "k": k}) + except Exception as e: + return jsonify({"error": f"MatMul vectors failed: {e}"}), 200 + + @app.route("/api/getValue", methods=["POST"]) def get_value(): global raw_tensor_data, precomputed_c_values, current_fullscreen_op @@ -210,14 +397,18 @@ def get_load_value(): x is not None and y is not None and z is not None ): try: - value = 0.0 - if op_data["dims"] == 3: - value = op_data["global_tensor"][x, y, z].item() - elif op_data["dims"] == 2: - value = op_data["global_tensor"][x, y].item() - elif op_data["dims"] == 1: - value = op_data["global_tensor"][x].item() - + import numpy as _np + + arr = _np.asarray(op_data["global_tensor"]) # already NumPy in draw.py + yy, xx, zz = int(y), int(x), int(z) + if arr.ndim >= 3: + value = float(arr[yy, xx, zz]) + elif arr.ndim == 2: + value = float(arr[yy, xx]) + elif arr.ndim == 1: + value = float(arr[xx]) + else: + value = 0.0 return jsonify({"value": value}) except IndexError: return jsonify({"error": "Coordinates out of bounds"}), 200 @@ -225,6 +416,43 @@ def get_load_value(): return jsonify({"error": "Global tensor data not found"}), 200 +@app.route("/api/getFlipValue", methods=["POST"]) +def get_flip_value(): + """Return a value from input or output of a Flip op by linear index or 2D coords. + + Request JSON: { uuid, which: "input"|"output", x, y } + - If tensor is 1D, use x only; if 2D, use x (col), y (row) + """ + global raw_tensor_data + data = request.json or {} + uuid = data.get("uuid") + which = (data.get("which") or "input").lower() + x = data.get("x") + y = data.get("y") + if not uuid or uuid not in raw_tensor_data: + return jsonify({"error": "Operation not found"}), 404 + op = raw_tensor_data[uuid] + t = op.get("input_data") if which == "input" else op.get("output_data") + shape = op.get("input_shape") if which == "input" else op.get("output_shape") + if t is None or shape is None: + return jsonify({"error": "Flip tensor data unavailable"}), 200 + try: + if len(shape) <= 1: + idx = int(x or 0) + return jsonify({"value": t[idx].item()}) + elif len(shape) == 2: + rr = int(y or 0) + cc = int(x or 0) + return jsonify({"value": t[rr, cc].item()}) + else: + # higher dims not directly supported; fallback to flat index + idx = int(x or 0) + flat = t.view(-1) + return jsonify({"value": flat[idx].item()}) + except Exception as e: + return jsonify({"error": f"Flip get value failed: {e}"}), 200 + + @app.route("/api/getLoadTensor", methods=["POST"]) def get_load_tensor(): """Return entire global tensor for a given Load/Store op, with min/max. @@ -249,24 +477,29 @@ def get_load_tensor(): if "global_tensor" not in op_data: return jsonify({"error": "Global tensor data not found"}), 200 - t = op_data["global_tensor"].cpu() - try: - t_min = float(t.min().item()) - t_max = float(t.max().item()) - except Exception: - # In case of empty tensor - t_min = 0.0 - t_max = 0.0 + import numpy as _np - return jsonify( - { - "shape": list(t.shape), - "dims": len(t.shape), - "min": t_min, - "max": t_max, - "values": t.numpy().tolist(), - } - ) + try: + arr = _np.asarray(op_data["global_tensor"]) # already NumPy in draw.py + t_shape = op_data.get("shape") or list(arr.shape) + t_dims = op_data.get("dims") or int(arr.ndim) + t_min = op_data.get("min") + t_max = op_data.get("max") + if t_min is None: + t_min = float(_np.min(arr)) if arr.size else 0.0 + if t_max is None: + t_max = float(_np.max(arr)) if arr.size else 0.0 + return jsonify( + { + "shape": t_shape, + "dims": int(t_dims), + "min": float(t_min), + "max": float(t_max), + "values": arr.tolist(), + } + ) + except Exception as e: + return jsonify({"error": f"getLoadTensor failed: {e}"}), 200 def run_flask_with_cloudflared(port: int = 8000, tunnel_port: int | None = None):