viz: fix newlines breaking label colors (#11030)

* viz: fix newlines breaking label colors

* TestViz.test_colored_label

* TestWordWrap
This commit is contained in:
qazal
2025-06-30 13:39:44 +03:00
committed by GitHub
parent 5911b71404
commit 2ea4737930
5 changed files with 33 additions and 8 deletions
+14 -1
View File
@@ -1,6 +1,6 @@
import ctypes, gzip, unittest
from tinygrad import Variable
from tinygrad.helpers import Context, ContextVar, argfix
from tinygrad.helpers import Context, ContextVar, argfix, colored, word_wrap
from tinygrad.helpers import merge_dicts, strip_parens, prod, round_up, fetch, fully_flatten, from_mv, to_mv, polyN, time_to_str, cdiv, cmod, getbits
from tinygrad.tensor import get_shape
from tinygrad.shape.view import get_contraction, get_contraction_with_reduce
@@ -363,5 +363,18 @@ class TestArgFix(unittest.TestCase):
def test_list(self):
self.assertEqual(argfix([True, False]), (True, False))
class TestWordWrap(unittest.TestCase):
def test_wrap_simple(self):
wrap = 10
st = "x"*wrap*2
st2 = word_wrap(st, wrap)
self.assertEqual(len(st2.splitlines()), 2)
def test_wrap_colored(self):
wrap = 10
st = colored("x"*wrap*2, "red")
st2 = word_wrap(st, wrap=wrap)
self.assertEqual(len(st2.splitlines()), 2)
if __name__ == '__main__':
unittest.main()
+12 -2
View File
@@ -1,10 +1,11 @@
import unittest, decimal, json
from dataclasses import dataclass
from tinygrad.uop.ops import UOp, UPat, Ops, PatternMatcher, TrackedPatternMatcher
from tinygrad.uop.ops import graph_rewrite, track_rewrites, TRACK_MATCH_STATS
from tinygrad.uop.symbolic import sym
from tinygrad.dtype import dtypes
from tinygrad.helpers import PROFILE
from tinygrad.helpers import PROFILE, colored, ansistrip
from tinygrad.device import Buffer
@track_rewrites(name=True)
@@ -14,7 +15,7 @@ def exec_rewrite(sink:UOp, pm_lst:list[PatternMatcher], names:None|list[str]=Non
return sink
# real VIZ=1 pickles these tracked values
from tinygrad.viz.serve import get_metadata
from tinygrad.viz.serve import get_metadata, uop_to_json
from tinygrad.uop.ops import tracked_keys, tracked_ctxs, active_rewrites, _name_cnt
def get_viz_list(): return get_metadata(tracked_keys, tracked_ctxs)
@@ -107,6 +108,15 @@ class TestViz(unittest.TestCase):
lst = get_viz_list()
self.assertEqual(lst[0]["name"], "(a+1) n1")
def test_colored_label(self):
# NOTE: dataclass repr prints literal escape codes instead of unicode chars
@dataclass(frozen=True)
class TestStruct:
colored_field: str
a = UOp(Ops.CUSTOM, arg=TestStruct(colored("xyz", "magenta")+colored("12345", "blue")))
a2 = uop_to_json(a)[id(a)]
self.assertEqual(ansistrip(a2["label"]), f"CUSTOM\n{TestStruct.__qualname__}(colored_field='xyz12345')")
# VIZ displays nested graph_rewrites in a tree view
def leaf_rewrite(x:UOp): return x.rtag(1) if x.tag is None else None
+5 -1
View File
@@ -73,7 +73,11 @@ def get_child(obj, key):
elif isinstance(obj, dict): obj = obj[k]
else: obj = getattr(obj, k)
return obj
def word_wrap(x, wrap=80): return x if len(x) <= wrap or '\n' in x[0:wrap] else (x[0:wrap] + "\n" + word_wrap(x[wrap:], wrap))
def word_wrap(x, wrap=80):
if len(ansistrip(x)) <= wrap: return x
i = 0
while len(ansistrip(x[:i])) < wrap and i < len(x): i += 1
return x[:i] + "\n" + word_wrap(x[i:], wrap)
def pluralize(st:str, cnt:int): return f"{cnt} {st}"+('' if cnt == 1 else 's')
class LazySeq(Generic[T]): # NOTE: Mapping requires __iter__ and __len__, Sequence requires supporting __len__ and slicing in __getitem__
-2
View File
@@ -11,8 +11,6 @@ onmessage = (e) => {
if (additions.length !== 0) g.setNode("addition", {label:"", style:"fill: rgba(26, 27, 38, 0.5);", padding:0});
for (let [k, {label, src, ref, ...rest }] of Object.entries(graph)) {
const idx = ref ? ctxs.findIndex(k => k.ref === ref) : -1;
// replace JSON.parse string literal with real ESC
label = label.replace(/\\x1b\r?\n*\[/g, "\u001B[");
if (idx != -1) label += `\ncodegen@${ctxs[idx].function_name}`;
// adjust node dims by label size (excluding escape codes) + add padding
let [width, height] = [0, 0];
+2 -2
View File
@@ -1,5 +1,5 @@
#!/usr/bin/env python3
import multiprocessing, pickle, difflib, os, threading, json, time, sys, webbrowser, socket, argparse, socketserver, functools, decimal
import multiprocessing, pickle, difflib, os, threading, json, time, sys, webbrowser, socket, argparse, socketserver, functools, decimal, codecs
from http.server import BaseHTTPRequestHandler
from urllib.parse import parse_qs, urlparse
from typing import Any, TypedDict, Generator
@@ -53,7 +53,7 @@ def uop_to_json(x:UOp) -> dict[int, dict]:
excluded.update(u.src)
for u in toposort:
if u in excluded: continue
argst = str(u.arg)
argst = codecs.decode(str(u.arg), "unicode_escape")
if u.op is Ops.VIEW:
argst = ("\n".join([f"{shape_to_str(v.shape)} / {shape_to_str(v.strides)}"+("" if v.offset == 0 else f" / {srender(v.offset)}")+
(f"\nMASK {mask_to_str(v.mask)}" if v.mask is not None else "") for v in unwrap(u.st).views]))