From ad534fd7df8644fd3bb8572189145098af2d3da6 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 26 Sep 2025 19:59:29 -0400 Subject: [PATCH 01/10] [DEV] Add interactive visualization features and improve web UI - Add MANIFEST.in and package data configuration for static/template files - Improve README with separate CPU/GPU installation instructions - Add color-by-value mode with Viridis and Mono color schemes - Add interactive cube dragging in 3D visualizations - Add OrbitControls for better camera manipulation - Fix coordinate mapping issues in load/store visualizations - Add support for raw (unmasked) load/store operations - Update CDN imports to use esm.sh for better module resolution - Add auto-launch server on import for easier debugging - Expose operation data to window for debugging - Improve UI controls with better z-index and pointer event handling --- MANIFEST.in | 2 + README.md | 12 +- pyproject.toml | 6 + triton_viz/clients/tracer/tracer.py | 25 +- triton_viz/core/trace.py | 4 + triton_viz/static/gridblock.js | 22 +- triton_viz/static/load.js | 354 ++++++++++++++++++++++++++-- triton_viz/static/load_utils.js | 85 ++++++- triton_viz/static/matmul.js | 11 +- triton_viz/static/store.js | 115 +++++++-- triton_viz/static/visualization.js | 15 +- triton_viz/templates/debug.html | 2 +- triton_viz/templates/index.html | 2 +- triton_viz/visualizer/draw.py | 79 ++++--- triton_viz/visualizer/interface.py | 157 ++++++++++-- 15 files changed, 785 insertions(+), 106 deletions(-) create mode 100644 MANIFEST.in diff --git a/MANIFEST.in b/MANIFEST.in new file mode 100644 index 00000000..9d8e530b --- /dev/null +++ b/MANIFEST.in @@ -0,0 +1,2 @@ +recursive-include triton_viz/templates *.html +recursive-include triton_viz/static * diff --git a/README.md b/README.md index 9f902b23..8b57cf19 100644 --- a/README.md +++ b/README.md @@ -47,13 +47,21 @@ The best part about this tool is that while it does focus on visualizing GPU ope - Python installed (preferably the latest available version). - [Triton](https://github.com/openai/triton/blob/main/README.md) installed. Follow the installation instructions in the linked repository. -Upon successfully installing Triton, install Torch using the following command: +After installing Triton, choose ONE of the following to install PyTorch: + +- CPU-only (no GPU required): + +```sh +pip install --pre torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu +``` + +- GPU (CUDA 12.1 nightly): ```sh pip install --pre torch torchvision torchaudio --index-url https://download.pytorch.org/whl/nightly/cu121 ``` -Upon successful installation of Torch make sure to uninstall `pytorch-triton` using the following command: +If you installed the GPU build, uninstall `pytorch-triton` to avoid conflicts: ```sh pip uninstall pytorch-triton diff --git a/pyproject.toml b/pyproject.toml index 3ed72536..363f76ac 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -34,6 +34,12 @@ dependencies = [ "tqdm", ] +[tool.setuptools] +include-package-data = true + +[tool.setuptools.package-data] +"triton_viz" = ["templates/*.html", "static/*.js", "static/**/*.js", "static/**"] + [project.urls] homepage = "https://github.com/Deep-Learning-Profiling-Tools/triton-viz" diff --git a/triton_viz/clients/tracer/tracer.py b/triton_viz/clients/tracer/tracer.py index 2f5e0a46..4b924059 100644 --- a/triton_viz/clients/tracer/tracer.py +++ b/triton_viz/clients/tracer/tracer.py @@ -1,6 +1,6 @@ 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, RawLoad, RawStore from typing import Callable, Optional, Union import numpy as np @@ -84,6 +84,25 @@ def pre_store_callback(ptr, value, mask, cache_modifier, eviction_policy): Store(tensor.data_ptr(), ptr.data - tensor.data_ptr(), mask.data) ) + # 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) + offsets = ptr.data - tensor.data_ptr() + true_mask = np.ones_like(offsets, dtype=bool) + self.records.append(Load(tensor.data_ptr(), offsets, true_mask)) + + 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) + offsets = ptr.data - tensor.data_ptr() + true_mask = np.ones_like(offsets, dtype=bool) + self.records.append(Store(tensor.data_ptr(), offsets, true_mask)) + def post_reduce_sum_callback(ret, input, axis=None, keep_dims=False): if not self.sample: return @@ -103,6 +122,10 @@ def post_dot_callback(ret, input, other, *args): return OpCallbacks(before_callback=pre_load_callback) elif op_type is Store: return OpCallbacks(before_callback=pre_store_callback) + elif op_type is RawLoad: + return OpCallbacks(before_callback=pre_raw_load_callback) + elif op_type is RawStore: + return OpCallbacks(before_callback=pre_raw_store_callback) elif op_type is ReduceSum: return OpCallbacks(after_callback=post_reduce_sum_callback) elif op_type is Dot: diff --git a/triton_viz/core/trace.py b/triton_viz/core/trace.py index 6cba7815..683c895c 100644 --- a/triton_viz/core/trace.py +++ b/triton_viz/core/trace.py @@ -8,6 +8,10 @@ from .data import Launch from typing import Union +from triton_viz.visualizer.interface import stop_server, launch + +stop_server() +launch(share=False, port=5001) launches: list[Launch] = [] diff --git a/triton_viz/static/gridblock.js b/triton_viz/static/gridblock.js index 6c999ccf..1e58fb30 100644 --- a/triton_viz/static/gridblock.js +++ b/triton_viz/static/gridblock.js @@ -1,6 +1,6 @@ -import { createMatMulVisualization } from './matmul.js'; -import { createLoadVisualization } from './load.js'; -import { createStoreVisualization } from './store.js'; +import { createMatMulVisualization } from './matmul.js?v=3'; +import { createLoadVisualization } from './load.js?v=3'; +import { createStoreVisualization } from './store.js?v=3'; export class GridBlock { constructor(x, y, width, height, gridX, gridY, gridZ, blockData, onClose, containerElement, canvas, drawFunction) { @@ -75,6 +75,13 @@ export class GridBlock { const closeButton = this.createCloseButton(); this.visualizationContainer.appendChild(closeButton); + // Ensure buttons panel sits above the canvas and accepts clicks + const buttonsPanel = this.visualizationContainer.querySelector('div'); + if (buttonsPanel) { + buttonsPanel.style.pointerEvents = 'auto'; + buttonsPanel.style.zIndex = '1002'; + } + this.isDetailedViewVisible = true; this.canvas.style.display = 'none'; this.containerElement.style.display = 'block'; @@ -183,6 +190,15 @@ export class GridBlock { this.contentArea.innerHTML = ''; + // expose current op for debugging + try { + window.last_op = op; + window.last_op_global_shape = op.global_shape; + window.last_global_coords = op.global_coords; + window.last_slice_shape = op.slice_shape; + window.last_slice_coords = op.slice_coords; + } catch (e) {} + switch (op.type) { case 'Dot': this.visualizationCleanupFunction = createMatMulVisualization(this.contentArea, op); diff --git a/triton_viz/static/load.js b/triton_viz/static/load.js index 8574cf67..014fe657 100644 --- a/triton_viz/static/load.js +++ b/triton_viz/static/load.js @@ -1,4 +1,5 @@ -import * as THREE from 'https://cdn.jsdelivr.net/npm/three@0.155.0/build/three.module.js'; +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, @@ -29,7 +30,55 @@ export function createLoadVisualization(containerElement, op) { let isPaused = false; const sideMenu = createSideMenu(containerElement); + // Color map UI + const controlBar = document.createElement('div'); + controlBar.style.position = 'fixed'; + controlBar.style.top = '10px'; + controlBar.style.left = '10px'; + controlBar.style.display = 'flex'; + controlBar.style.gap = '8px'; + controlBar.style.zIndex = '2000'; + controlBar.style.pointerEvents = 'auto'; + const colorizeToggle = document.createElement('button'); + colorizeToggle.textContent = 'Color by Value: OFF'; + controlBar.appendChild(colorizeToggle); + // Color scheme selector + color picker + const schemeSelect = document.createElement('select'); + schemeSelect.style.padding = '4px 6px'; + schemeSelect.style.borderRadius = '4px'; + schemeSelect.style.border = '1px solid #555'; + schemeSelect.style.background = '#2a2a2a'; + schemeSelect.style.color = '#fff'; + schemeSelect.innerHTML = ''; + controlBar.appendChild(schemeSelect); + const colorPicker = document.createElement('input'); + colorPicker.type = 'color'; + colorPicker.value = '#3b82f6'; // default blue + colorPicker.style.width = '36px'; + colorPicker.style.height = '28px'; + colorPicker.style.border = 'none'; + colorPicker.style.outline = 'none'; + colorPicker.title = 'Choose base color for Mono'; + controlBar.appendChild(colorPicker); + const dragToggle = document.createElement('button'); + dragToggle.textContent = 'Drag Cubes: OFF'; + controlBar.appendChild(dragToggle); + containerElement.appendChild(controlBar); + + // expose for debugging + try { + window.last_op_global_shape = op.global_shape; + window.last_global_coords = op.global_coords; + window.last_slice_shape = op.slice_shape; + window.last_slice_coords = op.slice_coords; + } catch (e) {} + + let colorizeOn = false; + let tensorCache = null; // {min, max, shape, dims, values} let hoveredCube = null; + let legendEl = null; + let scheme = 'mono'; + let monoBaseHex = '#3b82f6'; 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) @@ -50,30 +99,80 @@ export function createLoadVisualization(containerElement, op) { scene.add(globalTensor); scene.add(sliceTensor); + // Precompute highlighted coords in Global tensor for quick reset + const highlightedGlobalSet = new Set( + op.global_coords.map(([x, y, z]) => `${x},${y},${z}`) + ); + addLabels(scene, globalTensor, sliceTensor); - setupCamera(scene, camera); + const { center } = setupCamera(scene, camera); + const orbitControls = new OrbitControls(camera, renderer.domElement); + orbitControls.enableDamping = true; + orbitControls.dampingFactor = 0.05; + orbitControls.target.copy(center); + orbitControls.update(); const totalFrames = op.global_coords.length * 2 + 30; const raycaster = new THREE.Raycaster(); const mouse = new THREE.Vector2(); + // Drag state + let dragModeOn = false; + let isDragging = false; + let dragTarget = null; // THREE.Mesh (cube) + const dragPlane = new THREE.Plane(); + const planeIntersect = new THREE.Vector3(); + const worldPosHelper = new THREE.Vector3(); + const dragOffset = new THREE.Vector3(); const onKeyDown = cameraControls(camera, new THREE.Euler(0, 0, 0, 'YXZ')); setupEventListeners(containerElement, camera, renderer, onMouseMove, onKeyDown); + + // Additional pointer events for dragging + containerElement.addEventListener('mousedown', onMouseDown); + containerElement.addEventListener('mouseup', onMouseUp); + containerElement.addEventListener('mouseleave', onMouseUp); animate(); - async function onMouseMove(event) { + function _updateMouseNDC(event) { mouse.x = (event.clientX / containerElement.clientWidth) * 2 - 1; mouse.y = -(event.clientY / containerElement.clientHeight) * 2 + 1; + } + function _raycastAll() { raycaster.setFromCamera(mouse, camera); - const allTensorChildren = [ ...globalTensor.children, ...sliceTensor.children ]; + return raycaster.intersectObjects(allTensorChildren, true); + } + + function _toTopLevelCube(obj) { + let node = obj; + while (node && !(node.userData && node.userData.tensorName)) { + node = node.parent; + } + return node; + } + + async function onMouseMove(event) { + mouse.x = (event.clientX / containerElement.clientWidth) * 2 - 1; + mouse.y = -(event.clientY / containerElement.clientHeight) * 2 + 1; - const intersects = raycaster.intersectObjects(allTensorChildren, true); + raycaster.setFromCamera(mouse, camera); + + // Dragging: update target position on plane + if (isDragging && dragTarget) { + if (raycaster.ray.intersectPlane(dragPlane, planeIntersect)) { + const newWorld = planeIntersect.add(dragOffset); + // Convert world -> parent local + dragTarget.parent.worldToLocal(newWorld); + dragTarget.position.copy(newWorld); + } + } + + const intersects = _raycastAll(); if (hoveredCube) { hoveredCube.getObjectByName('hoverOutline').visible = false; @@ -81,10 +180,7 @@ export function createLoadVisualization(containerElement, op) { } if (intersects.length > 0) { - hoveredCube = intersects[0].object; - while (hoveredCube && !hoveredCube.tensorName) { - hoveredCube = hoveredCube.parent; - } + hoveredCube = _toTopLevelCube(intersects[0].object); if (hoveredCube) { const hoverOutline = hoveredCube.getObjectByName('hoverOutline'); @@ -92,11 +188,12 @@ export function createLoadVisualization(containerElement, op) { hoverOutline.visible = true; } - updateSideMenu(hoveredCube.tensorName, hoveredCube.tensor0, hoveredCube.tensor1, hoveredCube.tensor2, undefined); + const { tensorName, tensor0, tensor1, tensor2 } = hoveredCube.userData; + updateSideMenu(tensorName, tensor0, tensor1, tensor2, undefined); - const res = await getElementValue(hoveredCube.tensorName, hoveredCube.tensor0, hoveredCube.tensor1, hoveredCube.tensor2); + const res = await getElementValue(tensorName, tensor0, tensor1, tensor2); - updateSideMenu(hoveredCube.tensorName, hoveredCube.tensor0, hoveredCube.tensor1, hoveredCube.tensor2, res.value); + updateSideMenu(tensorName, tensor0, tensor1, tensor2, res.value); console.log(`Value: ${res.value}`); } @@ -105,10 +202,171 @@ export function createLoadVisualization(containerElement, op) { } } + // --------- Colormap (Viridis-like) and Legend --------- + function lerp(a, b, t) { return a + (b - a) * t; } + function lerpColor(c1, c2, t) { + const r = Math.round(lerp(c1[0], c2[0], t)) / 255; + const g = Math.round(lerp(c1[1], c2[1], t)) / 255; + const b = Math.round(lerp(c1[2], c2[2], t)) / 255; + return new THREE.Color(r, g, b); + } + // Viridis palette control points (sRGB) + const VIRIDIS = [ + [0.00, [68, 1, 84]], + [0.25, [59, 82, 139]], + [0.50, [33, 145, 140]], + [0.75, [94, 201, 98]], + [1.00, [253, 231, 37]] + ]; + function viridisColor(t) { + if (t <= 0) return new THREE.Color().setRGB(...VIRIDIS[0][1].map(v=>v/255)); + if (t >= 1) return new THREE.Color().setRGB(...VIRIDIS[VIRIDIS.length-1][1].map(v=>v/255)); + for (let i = 0; i < VIRIDIS.length - 1; i++) { + const [p, c] = VIRIDIS[i]; + const [pn, cn] = VIRIDIS[i+1]; + if (t >= p && t <= pn) { + const tt = (t - p) / (pn - p); + return lerpColor(c, cn, tt); + } + } + return new THREE.Color(1,1,1); + } + + function monoColor(t, hex) { + // map to HSL with same hue/sat, varying lightness from 0.9 -> 0.25 + const base = new THREE.Color(hex); + const hsl = {h:0,s:0,l:0}; + base.getHSL(hsl); + const l = lerp(0.9, 0.25, t); + const s = Math.min(1.0, Math.max(0.4, hsl.s)); // ensure enough saturation + return new THREE.Color().setHSL(hsl.h, s, l); + } + + function destroyLegend() { + if (legendEl && legendEl.remove) legendEl.remove(); + legendEl = null; + } + + function createLegend(min, max) { + destroyLegend(); + const wrapper = document.createElement('div'); + wrapper.style.position = 'fixed'; + wrapper.style.left = '10px'; + wrapper.style.top = '50px'; + wrapper.style.padding = '6px 8px'; + wrapper.style.background = 'rgba(0,0,0,0.6)'; + wrapper.style.color = '#fff'; + wrapper.style.font = '12px Arial, sans-serif'; + wrapper.style.borderRadius = '6px'; + wrapper.style.zIndex = '2000'; + + const title = document.createElement('div'); + title.textContent = scheme === 'mono' ? 'Value (Mono)' : 'Value (Viridis)'; + title.style.marginBottom = '4px'; + title.style.opacity = '0.9'; + wrapper.appendChild(title); + + const canvas = document.createElement('canvas'); + canvas.width = 220; canvas.height = 10; + const ctx2 = canvas.getContext('2d'); + for (let x = 0; x < canvas.width; x++) { + const t = x / (canvas.width - 1); + const c = (scheme === 'mono') ? monoColor(t, monoBaseHex) : viridisColor(t); + ctx2.fillStyle = `rgb(${Math.round(c.r*255)},${Math.round(c.g*255)},${Math.round(c.b*255)})`; + ctx2.fillRect(x, 0, 1, canvas.height); + } + wrapper.appendChild(canvas); + + const labels = document.createElement('div'); + labels.style.display = 'flex'; + labels.style.justifyContent = 'space-between'; + labels.style.marginTop = '2px'; + labels.innerHTML = `${min.toFixed(3)}${max.toFixed(3)}`; + wrapper.appendChild(labels); + + containerElement.appendChild(wrapper); + legendEl = wrapper; + } + + function applyColorMapIfNeeded() { + if (!colorizeOn || !tensorCache) return; + const { min, max, dims, values } = tensorCache; + const clamp = (v, lo, hi) => Math.max(lo, Math.min(hi, v)); + const norm = (v) => (max === min ? 0.5 : (v - min) / (max - min)); + const toColor = (t) => { + const u = clamp(norm(t), 0, 1); + return scheme === 'mono' ? monoColor(u, monoBaseHex) : viridisColor(u); + }; + globalTensor.children.forEach((cube) => { + const u = cube.userData; + if (!u) return; + let v = 0.0; + try { + if (dims === 3) v = values[u.tensor0][u.tensor1][u.tensor2]; + else if (dims === 2) v = values[u.tensor0][u.tensor1]; + else if (dims === 1) v = values[u.tensor0]; + } catch (e) { /* ignore bad index */ } + cube.material.color.copy(toColor(v)); + }); + } + + function resetGlobalColors() { + // Restore original colors: Global cubes default to COLOR_GLOBAL, + // highlighted coords (in op.global_coords) are COLOR_SLICE + globalTensor.children.forEach((cube) => { + const u = cube.userData; + if (!u) return; + const key = `${u.tensor0},${u.tensor1},${u.tensor2}`; + const baseColor = highlightedGlobalSet.has(key) ? COLOR_SLICE : COLOR_GLOBAL; + cube.material.color.copy(baseColor); + }); + } + + function resetSliceColors() { + sliceTensor.children.forEach((cube) => { + cube.material.color.copy(COLOR_LEFT_SLICE); + }); + } + + function onMouseDown(event) { + if (!dragModeOn) return; + _updateMouseNDC(event); + const hits = _raycastAll(); + if (hits.length === 0) return; + const cube = _toTopLevelCube(hits[0].object); + if (!cube) return; + // Prepare drag plane using camera forward as normal, passing through cube + const normal = new THREE.Vector3(); + camera.getWorldDirection(normal); + cube.getWorldPosition(worldPosHelper); + dragPlane.setFromNormalAndCoplanarPoint(normal, worldPosHelper); + // Compute offset between intersection and cube world position + raycaster.setFromCamera(mouse, camera); + if (!raycaster.ray.intersectPlane(dragPlane, planeIntersect)) return; + dragOffset.copy(worldPosHelper).sub(planeIntersect); + isDragging = true; + dragTarget = cube; + containerElement.style.cursor = 'grabbing'; + } + + function onMouseUp() { + if (!dragModeOn) return; + isDragging = false; + dragTarget = null; + containerElement.style.cursor = ''; + } + function animate() { requestAnimationFrame(animate); + // If colormap is OFF, ensure colors are reset every frame before animations + if (!colorizeOn) { + resetGlobalColors(); + resetSliceColors(); + } + orbitControls.update(); - if (!isPaused && frame < totalFrames) { + // When Color by Value is OFF, keep scene static (no highlight animation) + if (colorizeOn && !isPaused && frame < totalFrames) { const index = Math.floor(frame / 2); const factor = (frame % 2) / 1.0; @@ -126,6 +384,7 @@ export function createLoadVisualization(containerElement, op) { frame++; } + applyColorMapIfNeeded(); renderer.render(scene, camera); } @@ -134,10 +393,10 @@ export function createLoadVisualization(containerElement, op) { sliceTensor.children.forEach(cube => cube.material.emissive.setHex(0x000000)); const globalCube = globalTensor.children.find(c => - c.tensor0 === globalCoord[0] && c.tensor1 === globalCoord[1] && c.tensor2 === globalCoord[2] + c.userData && c.userData.tensor0 === globalCoord[0] && c.userData.tensor1 === globalCoord[1] && c.userData.tensor2 === globalCoord[2] ); const sliceCube = sliceTensor.children.find(c => - c.tensor0 === sliceCoord[0] && c.tensor1 === sliceCoord[1] && c.tensor2 === sliceCoord[2] + c.userData && c.userData.tensor0 === sliceCoord[0] && c.userData.tensor1 === sliceCoord[1] && c.userData.tensor2 === sliceCoord[2] ); if (globalCube) globalCube.material.emissive.setHex(0x444444); @@ -156,6 +415,71 @@ export function createLoadVisualization(containerElement, op) { return await response.json(); } + async function fetchGlobalTensor() { + try { + const res = await fetch('/api/getLoadTensor', { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ uuid: op.uuid }) + }); + const data = await res.json(); + if (!data || data.error) { + console.warn('getLoadTensor error:', data && data.error); + return null; + } + return data; + } catch (e) { + console.error('getLoadTensor failed', e); + return null; + } + } + + colorizeToggle.addEventListener('click', async () => { + colorizeOn = !colorizeOn; + colorizeToggle.textContent = `Color by Value: ${colorizeOn ? 'ON' : 'OFF'}`; + if (colorizeOn && !tensorCache) { + tensorCache = await fetchGlobalTensor(); + } + if (!colorizeOn) { + // When turning OFF, immediately restore original colors + resetGlobalColors(); + // Also reset slice side to its base color to avoid mixing styles + sliceTensor.children.forEach((cube) => { + cube.material.color.copy(COLOR_LEFT_SLICE); + }); + destroyLegend(); + } else if (tensorCache) { + createLegend(tensorCache.min, tensorCache.max); + } + }); + + schemeSelect.addEventListener('change', () => { + scheme = schemeSelect.value; + if (colorizeOn && tensorCache) { + applyColorMapIfNeeded(); + createLegend(tensorCache.min, tensorCache.max); + } + // color picker visible only for mono + colorPicker.style.display = scheme === 'mono' ? 'block' : 'none'; + }); + + colorPicker.addEventListener('input', (e) => { + monoBaseHex = e.target.value || '#3b82f6'; + if (scheme === 'mono' && colorizeOn && tensorCache) { + applyColorMapIfNeeded(); + createLegend(tensorCache.min, tensorCache.max); + } + }); + + // initialize picker visibility + colorPicker.style.display = 'block'; + + dragToggle.addEventListener('click', () => { + dragModeOn = !dragModeOn; + dragToggle.textContent = `Drag Cubes: ${dragModeOn ? 'ON' : 'OFF'}`; + orbitControls.enabled = !dragModeOn; + }); + 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 2729d04b..52ea6ff9 100644 --- a/triton_viz/static/load_utils.js +++ b/triton_viz/static/load_utils.js @@ -1,4 +1,4 @@ -import * as THREE from 'https://cdn.jsdelivr.net/npm/three@0.155.0/build/three.module.js'; +import * as THREE from 'https://esm.sh/three@0.155.0/build/three.module.js'; export const CUBE_SIZE = 0.2; export const GAP = 0.05; @@ -44,10 +44,10 @@ export function createCube(color, tensorName, x, y, z, cubeGeometry, edgesGeomet hoverOutline.name = 'hoverOutline'; cube.add(hoverOutline); - // Add custom properties to store tensor coordinates - cube.userData.tensor0 = z; + // Add custom properties to store tensor coordinates (x, y, z) + cube.userData.tensor0 = x; cube.userData.tensor1 = y; - cube.userData.tensor2 = x; + cube.userData.tensor2 = z; cube.userData.tensorName = tensorName; cube.name = `${tensorName}_cube_${x}_${y}_${z}`; @@ -77,17 +77,64 @@ export function createTensor(shape, coords, color, tensorName, cubeGeometry, edg } } } + // Build deterministic index from placement order to avoid mismatch + const indexOf = (x, y, z) => z * (width * height) + y * width + x; - console.log(`Highlighting ${coords.length} coordinates in global tensor`); - coords.forEach(([x, y, z]) => { - const cube = tensor.children.find(c => - c.userData.tensor0 === x && c.userData.tensor1 === y && c.userData.tensor2 === z - ); - if (cube) { + // Auto-detect coordinate axis order from incoming coords (try first N samples) + const samples = coords.slice(0, Math.min(256, coords.length)); + const maxIncoming = [0, 0, 0]; + for (const [a, b, c] of samples) { + if (a > maxIncoming[0]) maxIncoming[0] = a; + if (b > maxIncoming[1]) maxIncoming[1] = b; + if (c > maxIncoming[2]) maxIncoming[2] = c; + } + const target = [width - 1, height - 1, depth - 1]; + const perms = [ + [0, 1, 2], + [0, 2, 1], + [1, 0, 2], + [1, 2, 0], + [2, 0, 1], + [2, 1, 0], + ]; + function scorePerm(p) { + // sum of absolute diffs; heavy penalty if any exceeds target + let s = 0; + for (let i = 0; i < 3; i++) { + const diff = Math.abs(maxIncoming[i] - target[p[i]]); + s += diff; + if (maxIncoming[i] > target[p[i]]) s += 1000; // penalize out-of-range + } + return s; + } + let best = perms[0]; + let bestScore = scorePerm(best); + for (let i = 1; i < perms.length; i++) { + const sc = scorePerm(perms[i]); + if (sc < bestScore) { + best = perms[i]; + bestScore = sc; + } + } + + const remap = ([a, b, c]) => { + const arr = [a, b, c]; + return [arr[best.indexOf(0)], arr[best.indexOf(1)], arr[best.indexOf(2)]]; + }; + + console.log(`Highlighting ${coords.length} coordinates in global tensor. perm used: ${best}`); + coords.forEach(([A, B, C]) => { + const [x, y, z] = remap([A, B, C]); + if (x < 0 || x >= width || y < 0 || y >= height || z < 0 || z >= depth) { + console.warn(`Could not find cube at (${A}, ${B}, ${C}) -> mapped to out-of-range (${x}, ${y}, ${z})`); + return; + } + const idx = indexOf(x, y, z); + const cube = tensor.children[idx]; + if (cube && cube.userData && cube.userData.tensor0 === x && cube.userData.tensor1 === y && cube.userData.tensor2 === z) { cube.material.color.set(COLOR_SLICE); - console.log(`Highlighted cube at (${x}, ${y}, ${z})`); } else { - console.warn(`Could not find cube at (${x}, ${y}, ${z})`); + console.warn(`Could not find cube at (${A}, ${B}, ${C}) -> mapped (${x}, ${y}, ${z})`); } }); } else { @@ -122,7 +169,10 @@ export function interpolateColor(color1, color2, factor) { export function updateCubeColor(tensor, coord, startColor, endColor, factor) { const cube = tensor.children.find(c => - c.tensor0 === coord[0] && c.tensor1 === coord[1] && c.tensor2 === coord[2] + c.userData && + c.userData.tensor0 === coord[0] && + c.userData.tensor1 === coord[1] && + c.userData.tensor2 === coord[2] ); if (cube) { cube.material.color.copy(interpolateColor(startColor, endColor, factor)); @@ -152,6 +202,15 @@ export function setupEventListeners(containerElement, camera, renderer, onMouseM }); containerElement.addEventListener('mousemove', onMouseMove); window.addEventListener('keydown', onKeyDown); + + // Mouse wheel zoom + const WHEEL_ZOOM_SPEED = 0.5; + containerElement.addEventListener('wheel', (event) => { + event.preventDefault(); + const direction = event.deltaY > 0 ? 1 : -1; + camera.position.z += direction * WHEEL_ZOOM_SPEED; + camera.updateProjectionMatrix(); + }, { passive: false }); } export function cameraControls(camera, cameraRotation) { diff --git a/triton_viz/static/matmul.js b/triton_viz/static/matmul.js index d87d78bc..c32e9285 100644 --- a/triton_viz/static/matmul.js +++ b/triton_viz/static/matmul.js @@ -1,4 +1,4 @@ -import * as THREE from 'https://cdn.jsdelivr.net/npm/three@0.155.0/build/three.module.js'; +import * as THREE from 'https://esm.sh/three@0.155.0/build/three.module.js'; export function createMatMulVisualization(containerElement, op) { const { input_shape, other_shape, output_shape } = op; @@ -344,6 +344,15 @@ export function createMatMulVisualization(containerElement, op) { window.addEventListener('keydown', onKeyDown); containerElement.addEventListener('mousemove', onMouseMove); + // Mouse wheel zoom for matmul view + const WHEEL_ZOOM_SPEED = 0.5; + containerElement.addEventListener('wheel', (event) => { + event.preventDefault(); + const direction = event.deltaY > 0 ? 1 : -1; + camera.position.z += direction * WHEEL_ZOOM_SPEED; + camera.updateProjectionMatrix(); + }, { passive: false }); + animate(); diff --git a/triton_viz/static/store.js b/triton_viz/static/store.js index a75ca2ee..02cc3273 100644 --- a/triton_viz/static/store.js +++ b/triton_viz/static/store.js @@ -1,4 +1,5 @@ -import * as THREE from 'https://cdn.jsdelivr.net/npm/three@0.155.0/build/three.module.js'; +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, @@ -8,7 +9,7 @@ import { setupCamera, setupEventListeners, cameraControls -} from './load_utils.js'; +} from './load_utils.js?v=3'; export function createStoreVisualization(containerElement, op) { @@ -29,6 +30,20 @@ export function createStoreVisualization(containerElement, op) { let isPaused = false; const sideMenu = createSideMenu(containerElement); + // Controls bar (drag toggle only for Store view) + const controlBar = document.createElement('div'); + controlBar.style.position = 'fixed'; + controlBar.style.top = '10px'; + controlBar.style.left = '10px'; + controlBar.style.display = 'flex'; + controlBar.style.gap = '8px'; + controlBar.style.zIndex = '2000'; + controlBar.style.pointerEvents = 'auto'; + const dragToggle = document.createElement('button'); + dragToggle.textContent = 'Drag Cubes: OFF'; + controlBar.appendChild(dragToggle); + containerElement.appendChild(controlBar); + let dragModeOn = false; let hoveredCube = null; const COLOR_GLOBAL = new THREE.Color(0.2, 0.2, 0.2); // Dark Gray @@ -51,29 +66,72 @@ export function createStoreVisualization(containerElement, op) { scene.add(sliceTensor); addLabels(scene, globalTensor, sliceTensor); - setupCamera(scene, camera); + const { center } = setupCamera(scene, camera); + const orbitControls = new OrbitControls(camera, renderer.domElement); + orbitControls.enableDamping = true; + orbitControls.dampingFactor = 0.05; + orbitControls.target.copy(center); + orbitControls.update(); const totalFrames = op.global_coords.length * 2 + 30; const raycaster = new THREE.Raycaster(); const mouse = new THREE.Vector2(); + // Drag state (optional) + let isDragging = false; + let dragTarget = null; + const dragPlane = new THREE.Plane(); + const planeIntersect = new THREE.Vector3(); + const worldPosHelper = new THREE.Vector3(); + const dragOffset = new THREE.Vector3(); const onKeyDown = cameraControls(camera, new THREE.Euler(0, 0, 0, 'YXZ')); setupEventListeners(containerElement, camera, renderer, onMouseMove, onKeyDown); + containerElement.addEventListener('mousedown', onMouseDown); + containerElement.addEventListener('mouseup', onMouseUp); + containerElement.addEventListener('mouseleave', onMouseUp); + dragToggle.addEventListener('click', () => { + dragModeOn = !dragModeOn; + dragToggle.textContent = `Drag Cubes: ${dragModeOn ? 'ON' : 'OFF'}`; + orbitControls.enabled = !dragModeOn; + }); animate(); - async function onMouseMove(event) { + function _updateMouseNDC(event) { mouse.x = (event.clientX / containerElement.clientWidth) * 2 - 1; mouse.y = -(event.clientY / containerElement.clientHeight) * 2 + 1; + } + function _raycastAll() { raycaster.setFromCamera(mouse, camera); - const allTensorChildren = [ ...globalTensor.children, ...sliceTensor.children ]; + return raycaster.intersectObjects(allTensorChildren, true); + } + + function _toTopLevelCube(obj) { + let node = obj; + while (node && !(node.userData && node.userData.tensorName)) { + node = node.parent; + } + return node; + } - const intersects = raycaster.intersectObjects(allTensorChildren, true); + async function onMouseMove(event) { + _updateMouseNDC(event); + + raycaster.setFromCamera(mouse, camera); + if (isDragging && dragTarget) { + if (raycaster.ray.intersectPlane(dragPlane, planeIntersect)) { + const newWorld = planeIntersect.add(dragOffset); + dragTarget.parent.worldToLocal(newWorld); + dragTarget.position.copy(newWorld); + } + } + + const intersects = _raycastAll(); if (hoveredCube) { hoveredCube.getObjectByName('hoverOutline').visible = false; @@ -81,10 +139,7 @@ export function createStoreVisualization(containerElement, op) { } if (intersects.length > 0) { - hoveredCube = intersects[0].object; - while (hoveredCube && !hoveredCube.tensorName) { - hoveredCube = hoveredCube.parent; - } + hoveredCube = _toTopLevelCube(intersects[0].object); if (hoveredCube) { const hoverOutline = hoveredCube.getObjectByName('hoverOutline'); @@ -92,11 +147,12 @@ export function createStoreVisualization(containerElement, op) { hoverOutline.visible = true; } - updateSideMenu(hoveredCube.tensorName, hoveredCube.tensor0, hoveredCube.tensor1, hoveredCube.tensor2, undefined); + const { tensorName, tensor0, tensor1, tensor2 } = hoveredCube.userData; + updateSideMenu(tensorName, tensor0, tensor1, tensor2, undefined); - const res = await getElementValue(hoveredCube.tensorName, hoveredCube.tensor0, hoveredCube.tensor1, hoveredCube.tensor2); + const res = await getElementValue(tensorName, tensor0, tensor1, tensor2); - updateSideMenu(hoveredCube.tensorName, hoveredCube.tensor0, hoveredCube.tensor1, hoveredCube.tensor2, res.value); + updateSideMenu(tensorName, tensor0, tensor1, tensor2, res.value); console.log(`Value: ${res.value}`); } @@ -105,8 +161,35 @@ export function createStoreVisualization(containerElement, op) { } } + function onMouseDown(event) { + if (!dragModeOn) return; + _updateMouseNDC(event); + const hits = _raycastAll(); + if (hits.length === 0) return; + const cube = _toTopLevelCube(hits[0].object); + if (!cube) return; + const normal = new THREE.Vector3(); + camera.getWorldDirection(normal); + cube.getWorldPosition(worldPosHelper); + dragPlane.setFromNormalAndCoplanarPoint(normal, worldPosHelper); + raycaster.setFromCamera(mouse, camera); + if (!raycaster.ray.intersectPlane(dragPlane, planeIntersect)) return; + dragOffset.copy(worldPosHelper).sub(planeIntersect); + isDragging = true; + dragTarget = cube; + containerElement.style.cursor = 'grabbing'; + } + + function onMouseUp() { + if (!dragModeOn) return; + isDragging = false; + dragTarget = null; + containerElement.style.cursor = ''; + } + function animate() { requestAnimationFrame(animate); + orbitControls.update(); if (!isPaused && frame < totalFrames) { const index = Math.floor(frame / 2); @@ -134,10 +217,10 @@ export function createStoreVisualization(containerElement, op) { sliceTensor.children.forEach(cube => cube.material.emissive.setHex(0x000000)); const globalCube = globalTensor.children.find(c => - c.tensor0 === globalCoord[0] && c.tensor1 === globalCoord[1] && c.tensor2 === globalCoord[2] + c.userData && c.userData.tensor0 === globalCoord[0] && c.userData.tensor1 === globalCoord[1] && c.userData.tensor2 === globalCoord[2] ); const sliceCube = sliceTensor.children.find(c => - c.tensor0 === sliceCoord[0] && c.tensor1 === sliceCoord[1] && c.tensor2 === sliceCoord[2] + c.userData && c.userData.tensor0 === sliceCoord[0] && c.userData.tensor1 === sliceCoord[1] && c.userData.tensor2 === sliceCoord[2] ); if (globalCube) globalCube.material.emissive.setHex(0x444444); @@ -146,7 +229,7 @@ export function createStoreVisualization(containerElement, op) { async function getElementValue(tensorName, x, y, z) { let uuid = op.uuid; - const response = await fetch('/api/getStoreValue', { + const response = await fetch('/api/getLoadValue', { method: 'POST', headers: { 'Content-Type': 'application/json', diff --git a/triton_viz/static/visualization.js b/triton_viz/static/visualization.js index 1a37724c..7c20c56d 100644 --- a/triton_viz/static/visualization.js +++ b/triton_viz/static/visualization.js @@ -1,5 +1,5 @@ -import { GridBlock } from './gridblock.js'; -import { createInfoPopup, showInfoPopup } from './infoPopup.js'; +import { GridBlock } from './gridblock.js?v=3'; +import { createInfoPopup, showInfoPopup } from './infoPopup.js?v=3'; let globalData; let currentView = 'main'; let canvas, ctx; @@ -176,8 +176,9 @@ class KernelGrid { for (let x = 0; x < this.gridSize[0]; x++) { const blockX = this.rect.x + x * (this.blockWidth + 1); const blockY = this.rect.y + y * (this.blockHeight + 1); - const gridKey = `(${x}, ${y}, ${this.currentZ})`; - const blockData = this.visualizationData[gridKey] || []; + const gridKey1 = `(${x}, ${y}, ${this.currentZ})`; + const gridKey2 = `(${x},${y},${this.currentZ})`; + const blockData = this.visualizationData[gridKey1] || this.visualizationData[gridKey2] || []; const block = new GridBlock( blockX, blockY, this.blockWidth, this.blockHeight, x, y, this.currentZ, blockData, @@ -253,8 +254,12 @@ function determineMaxValues(visualizationData) { maxY = 0; maxZ = 0; const keys = Object.keys(visualizationData); + console.log('grid keys:', keys); keys.forEach(key => { - const [x, y, z] = key.replace(/[()]/g, '').split(', ').map(Number); + const [x, y, z] = key + .replace(/[()]/g, '') + .split(',') + .map(s => Number(String(s).trim())); if (x > maxX) maxX = x; if (y > maxY) maxY = y; if (z > maxZ) maxZ = z; diff --git a/triton_viz/templates/debug.html b/triton_viz/templates/debug.html index 9a2017f6..5786fb8d 100644 --- a/triton_viz/templates/debug.html +++ b/triton_viz/templates/debug.html @@ -77,7 +77,7 @@

Triton Viz Debug Page

+ diff --git a/triton_viz/visualizer/draw.py b/triton_viz/visualizer/draw.py index b5f296a1..1dd9037e 100644 --- a/triton_viz/visualizer/draw.py +++ b/triton_viz/visualizer/draw.py @@ -1,34 +1,48 @@ from triton_viz.core.data import ( Tensor, - Grid, - Store, - Load, - Dot, - ExpandDims, ) import numpy as np -from ..core.trace import launches -from ..clients.sanitizer.data import OutOfBoundsRecordBruteForce import sys import torch import uuid sys.setrecursionlimit(100000) + +# helper to tolerate multiple module instances (class identity mismatch) +def _is_type(obj, type_name: str) -> bool: + return getattr(obj, "__class__", type("x", (), {})).__name__ == type_name + + LAST_RECORD_ONLY = True # Generic render helpers def collect_grid(): + from ..core.trace import launches as current_launches + + try: + import sys as _sys + + for name, mod in list(_sys.modules.items()): + if name.endswith("triton_viz.core.trace"): + cand = getattr(mod, "launches", None) + if isinstance(cand, list) and len(cand) > len(current_launches): + current_launches = cand + except Exception: + pass records = [] tensor_tables = [] failures = [] - for launch in launches: + for launch in current_launches: cur_records, cur_tensor_table, cur_failures = collect_launch(launch) records.append(cur_records) tensor_tables.append(cur_tensor_table) failures.append(cur_failures) + if len(records) == 0: + # Gracefully handle when there are no launches yet + return {}, {}, {} assert LAST_RECORD_ONLY, "Only last record is supported for now" return records[-1], tensor_tables[-1], failures[-1] @@ -48,27 +62,30 @@ def collect_launch(launch): i, ) failures = {} - all_grids = {} - last_grid = None - program_records = [] + all_grids: dict[tuple, list] = {} + current_idx: tuple | None = None for r in launch.records: - if isinstance(r, Grid): - if last_grid is not None: - all_grids[last_grid.idx] = program_records - program_records = [] - last_grid = r - program_records.append(r) - if ( - isinstance(r, (OutOfBoundsRecordBruteForce)) - and (r.invalid_access_masks & r.op.masks).any() - ): - failures[last_grid.idx] = True - all_grids[last_grid.idx] = program_records + if _is_type(r, "Grid"): + current_idx = getattr(r, "idx", None) + if current_idx is not None and current_idx not in all_grids: + all_grids[current_idx] = [] + continue + if current_idx is None: + current_idx = (0, 0, 0) + all_grids.setdefault(current_idx, []) + # append non-Grid ops + all_grids[current_idx].append(r) + if _is_type(r, "OutOfBoundsRecordBruteForce"): + try: + if (r.invalid_access_masks & r.op.masks).any(): + failures[current_idx] = True + except Exception: + pass return all_grids, tensor_table, failures def extract_load_coords( - record: Load, global_tensor: Tensor + record, global_tensor: Tensor ) -> tuple[list[tuple[float, float, float]], list[tuple[float, float, float]]]: # Extract coordinates for the global tensor global_shape = make_3d(global_tensor.shape) @@ -149,9 +166,9 @@ def prepare_visualization_data(program_records, tensor_table): for record in program_records: record_uuid = str(uuid.uuid4())[:8] - if isinstance(record, ExpandDims): + if _is_type(record, "ExpandDims"): print(record.input_shape, record.output_shape, record.index) - if isinstance(record, Dot): + if _is_type(record, "Dot"): visualization_data.append( { "type": "Dot", @@ -168,7 +185,7 @@ def prepare_visualization_data(program_records, tensor_table): "intermediate_results": record.intermediate_results, } - elif isinstance(record, Load): + elif _is_type(record, "Load"): global_tensor, slice_tensor = tensor_table[record.ptr] print(global_tensor) global_coords, slice_coords = extract_load_coords(record, global_tensor) @@ -190,7 +207,7 @@ def prepare_visualization_data(program_records, tensor_table): } print(record.masks.shape) - elif isinstance(record, Store): + elif _is_type(record, "Store"): global_tensor, slice_tensor = tensor_table[record.ptr] global_coords, slice_coords = extract_load_coords(record, global_tensor) @@ -214,6 +231,7 @@ def get_visualization_data(): records, tensor_table, failures = collect_grid() visualization_data = {} raw_tensor_data = {} + kernel_src = "" for grid_idx, program_records in records.items(): viz_data, raw_data, kernel_src = prepare_visualization_data( @@ -222,12 +240,13 @@ def get_visualization_data(): visualization_data[str(grid_idx)] = viz_data raw_tensor_data.update(raw_data) - # Get the kernel source code + # Ensure failures dict has JSON-serializable keys + safe_failures = {str(k): v for k, v in failures.items()} return { "visualization_data": visualization_data, "raw_tensor_data": raw_tensor_data, - "failures": failures, + "failures": safe_failures, "kernel_src": kernel_src, } diff --git a/triton_viz/visualizer/interface.py b/triton_viz/visualizer/interface.py index eb32a432..47173ee2 100644 --- a/triton_viz/visualizer/interface.py +++ b/triton_viz/visualizer/interface.py @@ -19,6 +19,8 @@ raw_tensor_data = None precomputed_c_values = {} current_fullscreen_op = None +last_public_url = None +last_local_port = None def precompute_c_values(op_data): @@ -52,6 +54,14 @@ def update_global_data(): # Pass the records to analyze_records analysis_data = analyze_records(all_records) viz_data = get_visualization_data() + try: + keys = list(viz_data.get("visualization_data", {}).keys()) + print(f"[viz] grids: {keys}") + for k in keys: + ops = viz_data["visualization_data"].get(k, []) + print(f"[viz] grid {k} ops: {[op.get('type') for op in ops]}") + except Exception as e: + print("[viz] debug logging failed:", e) global_data = { "ops": { "visualization_data": viz_data["visualization_data"], @@ -93,8 +103,7 @@ def debug_page(): @app.route("/api/data") def get_data(): global global_data - if global_data is None: - update_global_data() + update_global_data() return jsonify(global_data) @@ -191,18 +200,82 @@ def get_load_value(): return jsonify({"error": "Global tensor data not found"}), 200 -def run_flask_with_cloudflared(): - cloudflared_port = 8000 # You can change this port if needed - tunnel_url = _run_cloudflared(cloudflared_port, 8001) # not too important +@app.route("/api/getLoadTensor", methods=["POST"]) +def get_load_tensor(): + """Return entire global tensor for a given Load/Store op, with min/max. + + Response schema: + { + "shape": [d0, d1, d2?], + "dims": 1|2|3, + "min": float, + "max": float, + "values": nested_list # Python list converted from torch tensor + } + """ + global raw_tensor_data + data = request.json + uuid = data.get("uuid") + + if uuid is None or uuid not in raw_tensor_data: + return jsonify({"error": "Operation not found"}), 404 + + op_data = raw_tensor_data[uuid] + 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 + + return jsonify( + { + "shape": list(t.shape), + "dims": len(t.shape), + "min": t_min, + "max": t_max, + "values": t.numpy().tolist(), + } + ) + + +def run_flask_with_cloudflared(port: int = 8000, tunnel_port: int | None = None): + """ + Run the Flask app on a given port and expose it via Cloudflared. + + :param port: Local Flask port to bind to. Defaults to 8000. + :param tunnel_port: Local tunnel control port for cloudflared. Defaults to port + 1. + """ + cloudflared_port = port + if tunnel_port is None: + tunnel_port = cloudflared_port + 1 + global last_public_url, last_local_port + tunnel_url = _run_cloudflared(cloudflared_port, tunnel_port) + last_public_url = tunnel_url + last_local_port = cloudflared_port print(f"Cloudflare tunnel URL: {tunnel_url}") - app.run(port=cloudflared_port) + app.run(host="0.0.0.0", port=cloudflared_port, debug=False, use_reloader=False) + +def launch(share: bool = True, port: int | None = None): + """ + Launch the Triton-Viz Flask server. -def launch(share=True): + """ print("Launching Triton viz tool") + default_port = 8000 if share else 5001 + actual_port = port or int(os.getenv("TRITON_VIZ_PORT", default_port)) + if share: print("--------") - flask_thread = threading.Thread(target=run_flask_with_cloudflared) + flask_thread = threading.Thread( + target=run_flask_with_cloudflared, args=(actual_port, None) + ) flask_thread.start() # Wait for the server to start @@ -210,24 +283,72 @@ def launch(share=True): # Try to get the tunnel URL by making a request to the local server try: - response = requests.get("http://localhost:8000") - public_url = response.url - print("Running on local URL: http://localhost:8000") - print(f"Running on public URL: {public_url}") + local_url = f"http://localhost:{actual_port}" + # touch local server to ensure it's up + _ = requests.get(local_url) + public_url = last_public_url + print(f"Running on local URL: {local_url}") + if public_url: + print(f"Running on public URL: {public_url}") print( "\nThis share link expires in 72 hours. For free permanent hosting and GPU upgrades, check out Spaces: https://huggingface.co/spaces" ) print("--------") + return local_url, public_url except requests.exceptions.RequestException: print("Setting up public URL... Please wait.") + # Even if the readiness check fails, return the intended URLs so callers don't crash + local_url = f"http://localhost:{actual_port}" + public_url = last_public_url + return local_url, public_url else: print("--------") - print("Running on local URL: http://localhost:5001") + local_url = f"http://localhost:{actual_port}" + print(f"Running on local URL: {local_url}") print("--------") - app.run(port=5001) + global last_local_port + last_local_port = actual_port + # Run Flask in a background thread so callers can continue (non-blocking) + def _run_local(): + app.run(host="0.0.0.0", port=actual_port, debug=True, use_reloader=False) -# This function can be called to stop the Flask server if needed -def stop_server(flask_thread): - # Implement a way to stop the Flask server - pass + flask_thread = threading.Thread(target=_run_local, daemon=True) + flask_thread.start() + # Give the server a moment to bind the port + time.sleep(0.5) + return local_url, None + + +def get_last_public_url(): + """Return the last Cloudflare public URL created by launch(share=True).""" + return last_public_url + + +@app.route("/shutdown", methods=["POST", "GET"]) +def _shutdown(): + """Shutdown Flask development server (useful for notebooks).""" + from flask import request as _req + + func = _req.environ.get("werkzeug.server.shutdown") + if func is None: + return jsonify( + {"status": "error", "message": "Not running with the Werkzeug Server"} + ), 400 + func() + return jsonify({"status": "ok", "message": "Server shutting down..."}) + + +def stop_server(port: int | None = None): + """ + Stop the running Flask server by calling the /shutdown endpoint. + If port is None, it will try the last used local port. + """ + target_port = port or last_local_port + if target_port is None: + return False + try: + requests.post(f"http://127.0.0.1:{target_port}/shutdown", timeout=2) + return True + except Exception: + return False From e8d9e675e5d31247e5470b66b369dcd44da202cc Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 26 Sep 2025 20:05:35 -0400 Subject: [PATCH 02/10] restore readme --- README.md | 12 ++---------- 1 file changed, 2 insertions(+), 10 deletions(-) diff --git a/README.md b/README.md index 8b57cf19..9f902b23 100644 --- a/README.md +++ b/README.md @@ -47,21 +47,13 @@ The best part about this tool is that while it does focus on visualizing GPU ope - Python installed (preferably the latest available version). - [Triton](https://github.com/openai/triton/blob/main/README.md) installed. Follow the installation instructions in the linked repository. -After installing Triton, choose ONE of the following to install PyTorch: - -- CPU-only (no GPU required): - -```sh -pip install --pre torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu -``` - -- GPU (CUDA 12.1 nightly): +Upon successfully installing Triton, install Torch using the following command: ```sh pip install --pre torch torchvision torchaudio --index-url https://download.pytorch.org/whl/nightly/cu121 ``` -If you installed the GPU build, uninstall `pytorch-triton` to avoid conflicts: +Upon successful installation of Torch make sure to uninstall `pytorch-triton` using the following command: ```sh pip uninstall pytorch-triton From e674f3db76a457b83f89eccd62805cf7cb23f309 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 26 Sep 2025 20:09:59 -0400 Subject: [PATCH 03/10] remove launch in trace.py --- triton_viz/core/trace.py | 4 ---- 1 file changed, 4 deletions(-) diff --git a/triton_viz/core/trace.py b/triton_viz/core/trace.py index 683c895c..6cba7815 100644 --- a/triton_viz/core/trace.py +++ b/triton_viz/core/trace.py @@ -8,10 +8,6 @@ from .data import Launch from typing import Union -from triton_viz.visualizer.interface import stop_server, launch - -stop_server() -launch(share=False, port=5001) launches: list[Launch] = [] From 4bd778c952c8f8228d65c0d4c26cd3df44cf6bc2 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 26 Sep 2025 20:48:39 -0400 Subject: [PATCH 04/10] remove v3 --- triton_viz/static/gridblock.js | 6 +++--- triton_viz/static/store.js | 2 +- triton_viz/static/visualization.js | 4 ++-- triton_viz/templates/debug.html | 2 +- triton_viz/templates/index.html | 2 +- 5 files changed, 8 insertions(+), 8 deletions(-) diff --git a/triton_viz/static/gridblock.js b/triton_viz/static/gridblock.js index 1e58fb30..bbdda5bb 100644 --- a/triton_viz/static/gridblock.js +++ b/triton_viz/static/gridblock.js @@ -1,6 +1,6 @@ -import { createMatMulVisualization } from './matmul.js?v=3'; -import { createLoadVisualization } from './load.js?v=3'; -import { createStoreVisualization } from './store.js?v=3'; +import { createMatMulVisualization } from './matmul.js'; +import { createLoadVisualization } from './load.js'; +import { createStoreVisualization } from './store.js'; export class GridBlock { constructor(x, y, width, height, gridX, gridY, gridZ, blockData, onClose, containerElement, canvas, drawFunction) { diff --git a/triton_viz/static/store.js b/triton_viz/static/store.js index 02cc3273..9868c16a 100644 --- a/triton_viz/static/store.js +++ b/triton_viz/static/store.js @@ -9,7 +9,7 @@ import { setupCamera, setupEventListeners, cameraControls -} from './load_utils.js?v=3'; +} from './load_utils.js'; export function createStoreVisualization(containerElement, op) { diff --git a/triton_viz/static/visualization.js b/triton_viz/static/visualization.js index 7c20c56d..15738ea5 100644 --- a/triton_viz/static/visualization.js +++ b/triton_viz/static/visualization.js @@ -1,5 +1,5 @@ -import { GridBlock } from './gridblock.js?v=3'; -import { createInfoPopup, showInfoPopup } from './infoPopup.js?v=3'; +import { GridBlock } from './gridblock.js'; +import { createInfoPopup, showInfoPopup } from './infoPopup.js'; let globalData; let currentView = 'main'; let canvas, ctx; diff --git a/triton_viz/templates/debug.html b/triton_viz/templates/debug.html index 5786fb8d..5e26b7c9 100644 --- a/triton_viz/templates/debug.html +++ b/triton_viz/templates/debug.html @@ -77,7 +77,7 @@

Triton Viz Debug Page

+ From 6d0a7395703ee08e20cab81ff4b4ef806efec9de Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 26 Sep 2025 21:03:00 -0400 Subject: [PATCH 05/10] remove _is_type --- triton_viz/visualizer/draw.py | 23 ++++++++++++----------- 1 file changed, 12 insertions(+), 11 deletions(-) diff --git a/triton_viz/visualizer/draw.py b/triton_viz/visualizer/draw.py index 1dd9037e..2e409170 100644 --- a/triton_viz/visualizer/draw.py +++ b/triton_viz/visualizer/draw.py @@ -1,6 +1,12 @@ from triton_viz.core.data import ( Tensor, + Grid, + ExpandDims, + Dot, + Load, + Store, ) +from triton_viz.clients.sanitizer.data import OutOfBoundsRecordBruteForce import numpy as np import sys import torch @@ -9,11 +15,6 @@ sys.setrecursionlimit(100000) -# helper to tolerate multiple module instances (class identity mismatch) -def _is_type(obj, type_name: str) -> bool: - return getattr(obj, "__class__", type("x", (), {})).__name__ == type_name - - LAST_RECORD_ONLY = True # Generic render helpers @@ -65,7 +66,7 @@ def collect_launch(launch): all_grids: dict[tuple, list] = {} current_idx: tuple | None = None for r in launch.records: - if _is_type(r, "Grid"): + if isinstance(r, Grid): current_idx = getattr(r, "idx", None) if current_idx is not None and current_idx not in all_grids: all_grids[current_idx] = [] @@ -75,7 +76,7 @@ def collect_launch(launch): all_grids.setdefault(current_idx, []) # append non-Grid ops all_grids[current_idx].append(r) - if _is_type(r, "OutOfBoundsRecordBruteForce"): + if isinstance(r, OutOfBoundsRecordBruteForce): try: if (r.invalid_access_masks & r.op.masks).any(): failures[current_idx] = True @@ -166,9 +167,9 @@ def prepare_visualization_data(program_records, tensor_table): for record in program_records: record_uuid = str(uuid.uuid4())[:8] - if _is_type(record, "ExpandDims"): + if isinstance(record, ExpandDims): print(record.input_shape, record.output_shape, record.index) - if _is_type(record, "Dot"): + if isinstance(record, Dot): visualization_data.append( { "type": "Dot", @@ -185,7 +186,7 @@ def prepare_visualization_data(program_records, tensor_table): "intermediate_results": record.intermediate_results, } - elif _is_type(record, "Load"): + elif isinstance(record, Load): global_tensor, slice_tensor = tensor_table[record.ptr] print(global_tensor) global_coords, slice_coords = extract_load_coords(record, global_tensor) @@ -207,7 +208,7 @@ def prepare_visualization_data(program_records, tensor_table): } print(record.masks.shape) - elif _is_type(record, "Store"): + elif isinstance(record, Store): global_tensor, slice_tensor = tensor_table[record.ptr] global_coords, slice_coords = extract_load_coords(record, global_tensor) From de5f5414d07f99ad20cb8954b86a5703006e14e8 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 26 Sep 2025 21:08:08 -0400 Subject: [PATCH 06/10] remove MANIFEST.in --- MANIFEST.in | 2 -- 1 file changed, 2 deletions(-) delete mode 100644 MANIFEST.in diff --git a/MANIFEST.in b/MANIFEST.in deleted file mode 100644 index 9d8e530b..00000000 --- a/MANIFEST.in +++ /dev/null @@ -1,2 +0,0 @@ -recursive-include triton_viz/templates *.html -recursive-include triton_viz/static * From fb19ab14be37618740c031cdfe10d8531522b4e8 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Fri, 26 Sep 2025 21:12:32 -0400 Subject: [PATCH 07/10] remove duplicate code in try and except --- triton_viz/visualizer/interface.py | 12 +++++------- 1 file changed, 5 insertions(+), 7 deletions(-) diff --git a/triton_viz/visualizer/interface.py b/triton_viz/visualizer/interface.py index 47173ee2..d01fd8cf 100644 --- a/triton_viz/visualizer/interface.py +++ b/triton_viz/visualizer/interface.py @@ -282,11 +282,12 @@ def launch(share: bool = True, port: int | None = None): time.sleep(5) # Try to get the tunnel URL by making a request to the local server + local_url = f"http://localhost:{actual_port}" + public_url = last_public_url + try: - local_url = f"http://localhost:{actual_port}" # touch local server to ensure it's up _ = requests.get(local_url) - public_url = last_public_url print(f"Running on local URL: {local_url}") if public_url: print(f"Running on public URL: {public_url}") @@ -294,13 +295,10 @@ def launch(share: bool = True, port: int | None = None): "\nThis share link expires in 72 hours. For free permanent hosting and GPU upgrades, check out Spaces: https://huggingface.co/spaces" ) print("--------") - return local_url, public_url except requests.exceptions.RequestException: print("Setting up public URL... Please wait.") - # Even if the readiness check fails, return the intended URLs so callers don't crash - local_url = f"http://localhost:{actual_port}" - public_url = last_public_url - return local_url, public_url + + return local_url, public_url else: print("--------") local_url = f"http://localhost:{actual_port}" From 19a050dee86be43dc33e9d0575c175fa54b7bdc9 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Sun, 28 Sep 2025 19:49:15 -0400 Subject: [PATCH 08/10] fix import --- triton_viz/visualizer/draw.py | 12 ++---------- 1 file changed, 2 insertions(+), 10 deletions(-) diff --git a/triton_viz/visualizer/draw.py b/triton_viz/visualizer/draw.py index 2e409170..23feb087 100644 --- a/triton_viz/visualizer/draw.py +++ b/triton_viz/visualizer/draw.py @@ -21,18 +21,10 @@ def collect_grid(): + # If imported at module level, it may capture an empty launches list before trace.py completes initialization. + # By importing here, we ensure we get the current state of launches with all traced kernel executions. from ..core.trace import launches as current_launches - try: - import sys as _sys - - for name, mod in list(_sys.modules.items()): - if name.endswith("triton_viz.core.trace"): - cand = getattr(mod, "launches", None) - if isinstance(cand, list) and len(cand) > len(current_launches): - current_launches = cand - except Exception: - pass records = [] tensor_tables = [] failures = [] From b97f5f7bd8296cf21a23aabf93feb847987b714f Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Sun, 28 Sep 2025 20:04:07 -0400 Subject: [PATCH 09/10] Enable highlight animation regardless of Color by Value state --- triton_viz/static/load.js | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/triton_viz/static/load.js b/triton_viz/static/load.js index 014fe657..73a86e8f 100644 --- a/triton_viz/static/load.js +++ b/triton_viz/static/load.js @@ -365,8 +365,8 @@ export function createLoadVisualization(containerElement, op) { } orbitControls.update(); - // When Color by Value is OFF, keep scene static (no highlight animation) - if (colorizeOn && !isPaused && frame < totalFrames) { + // Run highlight animation regardless of Color by Value state + if (!isPaused && frame < totalFrames) { const index = Math.floor(frame / 2); const factor = (frame % 2) / 1.0; From 2dec278d8628a5f6760335c8d6329bd33016eb11 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Sun, 28 Sep 2025 20:56:45 -0400 Subject: [PATCH 10/10] [REFACTOR][VISUALIZER] Refactor global variables to use ServerState class Replace global variables last_public_url and last_local_port with a ServerState class to improve code organization and avoid global state management issues. --- triton_viz/visualizer/interface.py | 41 ++++++++++++++++++++++-------- 1 file changed, 31 insertions(+), 10 deletions(-) diff --git a/triton_viz/visualizer/interface.py b/triton_viz/visualizer/interface.py index d01fd8cf..64e1ec28 100644 --- a/triton_viz/visualizer/interface.py +++ b/triton_viz/visualizer/interface.py @@ -19,8 +19,31 @@ raw_tensor_data = None precomputed_c_values = {} current_fullscreen_op = None -last_public_url = None -last_local_port = None + + +# Server state management +class ServerState: + """Encapsulates server state to avoid global variables.""" + + def __init__(self): + self.last_public_url = None + self.last_local_port = None + + def set_public_url(self, url): + self.last_public_url = url + + def set_local_port(self, port): + self.last_local_port = port + + def get_public_url(self): + return self.last_public_url + + def get_local_port(self): + return self.last_local_port + + +# Create a single instance for the module +_server_state = ServerState() def precompute_c_values(op_data): @@ -254,10 +277,9 @@ def run_flask_with_cloudflared(port: int = 8000, tunnel_port: int | None = None) cloudflared_port = port if tunnel_port is None: tunnel_port = cloudflared_port + 1 - global last_public_url, last_local_port tunnel_url = _run_cloudflared(cloudflared_port, tunnel_port) - last_public_url = tunnel_url - last_local_port = cloudflared_port + _server_state.set_public_url(tunnel_url) + _server_state.set_local_port(cloudflared_port) print(f"Cloudflare tunnel URL: {tunnel_url}") app.run(host="0.0.0.0", port=cloudflared_port, debug=False, use_reloader=False) @@ -283,7 +305,7 @@ def launch(share: bool = True, port: int | None = None): # Try to get the tunnel URL by making a request to the local server local_url = f"http://localhost:{actual_port}" - public_url = last_public_url + public_url = _server_state.get_public_url() try: # touch local server to ensure it's up @@ -304,8 +326,7 @@ def launch(share: bool = True, port: int | None = None): local_url = f"http://localhost:{actual_port}" print(f"Running on local URL: {local_url}") print("--------") - global last_local_port - last_local_port = actual_port + _server_state.set_local_port(actual_port) # Run Flask in a background thread so callers can continue (non-blocking) def _run_local(): @@ -320,7 +341,7 @@ def _run_local(): def get_last_public_url(): """Return the last Cloudflare public URL created by launch(share=True).""" - return last_public_url + return _server_state.get_public_url() @app.route("/shutdown", methods=["POST", "GET"]) @@ -342,7 +363,7 @@ def stop_server(port: int | None = None): Stop the running Flask server by calling the /shutdown endpoint. If port is None, it will try the last used local port. """ - target_port = port or last_local_port + target_port = port or _server_state.get_local_port() if target_port is None: return False try: