diff --git a/tinygrad/viz/index.html b/tinygrad/viz/index.html
index 8fc8cdbeec..4f16b317b3 100644
--- a/tinygrad/viz/index.html
+++ b/tinygrad/viz/index.html
@@ -39,12 +39,12 @@
}
ul {
padding: 0;
- color: #7c7d85;
+ opacity: 0.6;
white-space: nowrap;
cursor: pointer;
}
ul.active {
- color: #f0f0f5;
+ opacity: 1;
}
ul.disabled {
opacity: 0.6;
@@ -283,6 +283,12 @@
hljs.highlightElement(codeEl);
return pre;
};
+ const coloredToHTML = (str) => {
+ const colors = ['gray','red','green','yellow','blue','magenta','cyan','white'];
+ return str.replace(/\u001b\[(\d+)m(.*?)\u001b\[0m/g, (_, code, st) => {
+ return `${st}`;
+ })
+ }
// **** main loop
var ret = [];
@@ -308,7 +314,7 @@
if (i === currentKernel) {
requestAnimationFrame(() => kernelUl.scrollIntoView({ behavior: "auto", block: "nearest" }));
}
- const p = Object.assign(document.createElement("p"), { id: `kernel-${key}`, innerText: key, style: "cursor: pointer;"});
+ const p = Object.assign(document.createElement("p"), { id: `kernel-${key}`, innerHTML: coloredToHTML(key), style: "cursor: pointer;"});
kernelUl.appendChild(p)
items.forEach((u, j) => {
const rwUl = Object.assign(document.createElement("ul"), { innerText: `${toPath(u.loc)} - ${u.match_count}`, key: `uop-rewrite-${j}`,
diff --git a/tinygrad/viz/serve.py b/tinygrad/viz/serve.py
index 291107258c..e0013e5b88 100755
--- a/tinygrad/viz/serve.py
+++ b/tinygrad/viz/serve.py
@@ -3,7 +3,7 @@ import multiprocessing, pickle, functools, difflib, os, threading, json, time, s
from http.server import HTTPServer, BaseHTTPRequestHandler
from urllib.parse import parse_qs, urlparse
from typing import Any, Callable, TypedDict, Generator
-from tinygrad.helpers import colored, getenv, to_function_name, tqdm, unwrap, word_wrap
+from tinygrad.helpers import colored, getenv, tqdm, unwrap, word_wrap
from tinygrad.ops import TrackedGraphRewrite, UOp, Ops, lines, GroupOp
from tinygrad.codegen.kernel import Kernel
from tinygrad.device import ProfileEvent, ProfileDeviceEvent, ProfileRangeEvent, ProfileGraphEvent
@@ -37,7 +37,7 @@ def to_metadata(k:Any, v:TrackedGraphRewrite) -> GraphRewriteMetadata:
return {"loc":v.loc, "match_count":len(v.matches), "code_line":lines(v.loc[0])[v.loc[1]-1].strip(),
"kernel_code":pcall(_prg, k) if isinstance(k, Kernel) else None}
def get_metadata(keys:list[Any], contexts:list[list[TrackedGraphRewrite]]) -> list[tuple[str, list[GraphRewriteMetadata]]]:
- return [(to_function_name(k.name) if isinstance(k, Kernel) else str(k), [to_metadata(k, v) for v in vals]) for k,vals in zip(keys, contexts)]
+ return [(k.name if isinstance(k, Kernel) else str(k), [to_metadata(k, v) for v in vals]) for k,vals in zip(keys, contexts)]
# ** Complete rewrite details for a graph_rewrite call