Merge branch 'master' into dsp_search

This commit is contained in:
George Hotz
2025-03-26 17:49:15 +08:00
committed by GitHub
3 changed files with 186 additions and 95 deletions
@@ -1,5 +1,5 @@
from tinygrad import dtypes, getenv
from tinygrad.helpers import trange, colored
from tinygrad import dtypes, getenv, Device
from tinygrad.helpers import trange, colored, DEBUG, temp
from tinygrad.nn.datasets import mnist
import torch
from torch import nn, optim
@@ -30,13 +30,15 @@ if __name__ == "__main__":
import tinygrad.frontend.torch
device = torch.device("tiny")
else:
device = torch.device("mps")
device = torch.device({"METAL":"mps","NV":"cuda"}.get(Device.DEFAULT, "cpu"))
if DEBUG >= 1: print(f"using torch backend {device}")
X_train, Y_train, X_test, Y_test = mnist()
X_train = torch.tensor(X_train.float().numpy(), device=device)
Y_train = torch.tensor(Y_train.cast(dtypes.int64).numpy(), device=device)
X_test = torch.tensor(X_test.float().numpy(), device=device)
Y_test = torch.tensor(Y_test.cast(dtypes.int64).numpy(), device=device)
if getenv("TORCHVIZ"): torch.cuda.memory._record_memory_history()
model = Model().to(device)
optimizer = optim.Adam(model.parameters(), 1e-3)
@@ -62,3 +64,6 @@ if __name__ == "__main__":
if target := getenv("TARGET_EVAL_ACC_PCT", 0.0):
if test_acc >= target and test_acc != 100.0: print(colored(f"{test_acc=} >= {target}", "green"))
else: raise ValueError(colored(f"{test_acc=} < {target}", "red"))
if getenv("TORCHVIZ"):
torch.cuda.memory._dump_snapshot(fp:=temp("torchviz.pkl", append_user=True))
print(f"saved torch memory snapshot to {fp}, view in https://pytorch.org/memory_viz")
+124
View File
@@ -0,0 +1,124 @@
import unittest
from tinygrad import dtypes, Device
from tinygrad.device import Buffer
from tinygrad.engine.memory import _internal_memory_planner
global_map = {}
def b(i, base=None, offset=0, pin=False, size=16):
global global_map
if i in global_map: return global_map[i]
global_map[i] = Buffer(Device.DEFAULT, size, dtypes.int8, base=global_map[base] if base is not None else None, offset=offset)
if pin: global_map[i].ref(1)
return global_map[i]
def check_assign(buffers:list[list[Buffer]|tuple[Buffer, ...]]):
assigned = _internal_memory_planner(buffers, noopt_buffers=None)
taken_parts = set()
first_appearance, last_appearance = {}, {}
for i,u in enumerate(buffers):
for buf in u:
if buf.is_allocated() or buf.base.is_allocated() or buf.lb_refcount > 0: continue
if buf.base not in first_appearance: first_appearance[buf.base] = i
last_appearance[buf.base] = i
for i,u in enumerate(buffers):
for buf in u:
if buf.is_allocated() or buf.base.is_allocated() or buf.lb_refcount > 0: continue
cur, base = assigned.get(buf, buf), assigned.get(buf.base, buf.base)
if buf._base is not None:
assert cur.base == base.base and cur.offset == buf.offset + base.offset, f"failed: {buf} {cur} {base} {buf.offset} {base.offset}"
else:
for part in taken_parts:
assert buf.base == part[3] or part[0] != cur.base or part[1] + part[2] <= cur.offset or part[1] >= cur.offset + buf.nbytes
if first_appearance[buf.base] == i: taken_parts.add((cur.base, cur.offset, buf.nbytes, buf.base))
if last_appearance[buf.base] == i: taken_parts.remove((cur.base, cur.offset, buf.nbytes, buf.base))
class TestMemoryPlanner(unittest.TestCase):
def setUp(self):
global global_map
global_map = {}
def test_simple_buffer(self):
bs = [
[b(0), b(1), b(2)],
[b(1), b(2), b(3)],
[b(4), b(3)],
[b(5), b(2)],
]
check_assign(bs)
def test_simple_pinned(self):
bs = [
[b(0, pin=True), b(1), b(2, pin=True)],
[b(1), b(2), b(3)],
[b(4), b(3)],
[b(5), b(2)],
]
check_assign(bs)
def test_all_pinned(self):
bs = [
[b(0, pin=True), b(1, pin=True)],
[b(1), b(2, pin=True)],
[b(4, pin=True), b(3, pin=True)],
]
check_assign(bs)
def test_simple_buffer_offset(self):
bs = [
[b(0, pin=True), b(1, base=0, offset=1, size=8), b(2)],
[b(1), b(2), b(3, base=0, offset=1, size=8)],
[b(4), b(3)],
]
check_assign(bs)
def test_buffer_offset(self):
bs = [
[b(0, pin=True), b(1, base=0, offset=1, size=8), b(2)],
[b(1), b(2), b(3, base=0, offset=1, size=8)],
[b(4), b(3)],
[b(5, base=2, offset=2, size=8), b(3)],
[b(6), b(5), b(0)],
[b(7), b(8, pin=True)],
[b(8), b(9, base=2, offset=2, size=8)],
[b(9), b(3), b(5)],
]
check_assign(bs)
def test_buffer_offset2(self):
bs = [
[b(0, pin=True), b(1), b(2)],
[b(1), b(2), b(3)],
[b(4), b(3)],
[b(5), b(3)],
[b(6), b(5), b(0)],
[b(7), b(8, pin=True)],
[b(8), b(9)],
[b(9), b(3), b(5)],
[b(11), b(0)],
[b(11), b(10), b(5)],
[b(12), b(11), b(0)],
[b(6), b(12), b(7)],
[b(13), b(6), b(11)],
]
check_assign(bs)
def test_all_offsets_of_one(self):
bs = [
[b(0, pin=True), b(1)],
[b(3, base=1, offset=0, size=8), b(2, base=0, offset=0, size=8)],
[b(5, base=1, offset=8, size=8), b(4, base=0, offset=8, size=8)],
[b(7, base=1, offset=4, size=8), b(6, base=0, offset=4, size=8)],
[b(4), b(5), b(2)],
[b(3), b(7)],
[b(10), b(6), b(7)],
[b(11), b(3), b(2)],
[b(12), b(5), b(4), b(3), b(2)],
[b(13), b(6), b(12), b(7)],
]
check_assign(bs)
if __name__ == "__main__":
unittest.main()
+54 -92
View File
@@ -15,11 +15,8 @@
<style>
* {
box-sizing: border-box;
margin-block-start: initial;
margin-block-end: initial;
}
button {
outline: none;
margin: 0;
padding: 0;
}
html, body {
color: #f0f0f5;
@@ -83,39 +80,36 @@
}
.container {
background-color: #0f1018;
padding: 20px;
z-index: 2;
position: relative;
height: 100%;
}
.container > * + *, .rewrite-container > * + * {
.container > * + *, .rewrite-container > * + *, .kernel-list > * + * {
margin-top: 12px;
}
.kernel-list > ul > * + * {
margin-top: 4px;
}
.graph {
position: absolute;
inset: 0;
z-index: 1;
}
.kernel-list-parent {
position: relative;
width: 15%;
padding: 50px 20px 20px 20px;
padding-top: 50px;
border-right: 1px solid #4a4b56;
z-index: 2;
}
.kernel-list {
width: 100%;
height: 100%;
overflow-y: auto;
}
.kernel-list > ul > * + * {
margin-top: 4px;
}
.metadata {
position: relative;
width: 20%;
padding: 20px;
background-color: #0f1018;
border-left: 1px solid #4a4b56;
z-index: 2;
margin-left: auto;
height: 100%;
overflow-y: auto;
}
.resize-handle {
@@ -128,13 +122,6 @@
z-index: 3;
background-color: transparent;
}
#kernel-resize-handle {
right: 0;
}
#metadata-resize-handle {
margin-top: 0;
left: 0;
}
.floating-container {
position: fixed;
top: 10px;
@@ -144,41 +131,26 @@
flex-direction: row;
gap: 8px;
}
.nav-btn {
.btn {
outline: none;
background-color: #1a1b26;
border: 1px solid #4a4b56;
color: #f0f0f5;
height: 32px;
border-radius: 8px;
padding: 6px;
cursor: pointer;
text-decoration: none;
height: 32px;
display: flex;
align-items: center;
padding: 0 6px;
font-weight: bold;
}
.collapse-btn {
width: 32px;
padding: 6px;
}
.btn {
height: 32px;
background-color: #1a1b26;
border: 1px solid #4a4b56;
color: #f0f0f5;
border-radius: 8px;
cursor: pointer;
transition-duration: .5s;
justify-content: center;
text-decoration: none;
}
.btn:hover {
background-color: #2a2b36;
border-color: #5a5b66;
color: #ffffff;
}
.collapsed .kernel-list, .collapsed .metadata {
width: 0;
padding: 0;
overflow: hidden;
.collapsed .container {
display: none;
}
.rewrite-list {
display: flex;
@@ -217,11 +189,11 @@
<div class="main-container">
<div class="floating-container">
<button class="btn collapse-btn">
<svg viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2"><path d="M15 19l-7-7 7-7"/></svg>
<svg viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" width="20"><path d="M15 19l-7-7 7-7"/></svg>
</button>
<a class="btn nav-btn" href="/profiler">Profiler</a>
</div>
<div class="container kernel-list-parent"><div class="container kernel-list"></div></div>
<div class="container kernel-list-parent"><div class="kernel-list"></div></div>
<div class="graph">
<div class="progress-message">Rendering new layout...</div>
<svg id="graph-svg" preserveAspectRatio="xMidYMid meet">
@@ -299,14 +271,12 @@
var expandKernel = true;
const evtSources = [];
async function main() {
const mainContainer = document.querySelector('.main-container');
// ***** LHS kernels list
if (kernels == null) {
kernels = await (await fetch("/kernels")).json();
currentKernel = -1;
}
const kernelListParent = document.querySelector(".container.kernel-list-parent");
const kernelList = document.querySelector(".container.kernel-list");
const kernelList = document.querySelector(".kernel-list");
kernelList.innerHTML = "";
kernels.forEach(([key, items], i) => {
const kernelUl = Object.assign(document.createElement("ul"), { key: `kernel-${i}`, className: i === currentKernel ? "active" : "",
@@ -384,6 +354,7 @@
metadata.innerHTML = "";
metadata.appendChild(vsCodeOpener(kernel.loc.join(":").split("/")));
metadata.appendChild(highlightedCodeBlock(kernel.code_line, "python", true));
appendResizer(metadata, { minWidth: 20, maxWidth: 50 });
// ** code blocks
let code = ret[currentRewrite].uop;
let lang = "python"
@@ -428,48 +399,39 @@
} else {
metadata.appendChild(Object.assign(document.createElement("p"), { textContent: `No rewrites in ${toPath(kernel.loc)}.` }));
}
// ***** collapse/expand
let isCollapsed = false;
const collapseBtn = document.querySelector(".collapse-btn");
collapseBtn.addEventListener("click", () => {
isCollapsed = !isCollapsed;
mainContainer.classList.toggle("collapsed", isCollapsed);
kernelListParent.style.display = isCollapsed ? "none" : "block";
metadata.style.display = isCollapsed ? "none" : "block";
collapseBtn.style.transform = isCollapsed ? "rotate(180deg)" : "rotate(0deg)";
});
// ***** resizer
function createResizer(element, width, type) {
const { minWidth, maxWidth } = width;
const handle = Object.assign(document.createElement("div"), { id: `${type}-resize-handle`, className: "resize-handle" });
element.appendChild(handle);
const resize = (e) => {
const change = e.clientX - element.dataset.startX;
const adjustedChange = type === "kernel" ? change : -change;
const newWidth = ((Number(element.dataset.startWidth) + adjustedChange) / Number(element.dataset.containerWidth)) * 100;
if (newWidth >= minWidth && newWidth <= maxWidth) {
element.style.width = `${newWidth}%`;
}
};
handle.addEventListener("mousedown", (e) => {
e.preventDefault();
element.dataset.startX = e.clientX;
element.dataset.containerWidth = mainContainer.getBoundingClientRect().width;
element.dataset.startWidth = element.getBoundingClientRect().width;
document.documentElement.addEventListener("mousemove", resize, false);
document.documentElement.addEventListener("mouseup", () => {
document.documentElement.removeEventListener("mousemove", resize, false);
element.style.userSelect = "initial";
}, { once: true });
});
}
createResizer(kernelListParent, { minWidth: 15, maxWidth: 50 }, "kernel"); // left resizer
createResizer(metadata, { minWidth: 20, maxWidth: 50 }, "metadata"); // right resizer
}
// **** collapse/expand
let isCollapsed = false;
const mainContainer = document.querySelector('.main-container');
document.querySelector(".collapse-btn").addEventListener("click", (e) => {
isCollapsed = !isCollapsed;
mainContainer.classList.toggle("collapsed", isCollapsed);
e.target.style.transform = isCollapsed ? "rotate(180deg)" : "rotate(0deg)";
});
// **** resizer
function appendResizer(element, { minWidth, maxWidth }, left=false) {
const handle = Object.assign(document.createElement("div"), { className: "resize-handle", style: left ? "right: 0" : "left: 0; margin-top: 0" });
element.appendChild(handle);
const resize = (e) => {
const change = e.clientX - element.dataset.startX;
let newWidth = ((Number(element.dataset.startWidth)+(left ? change : -change))/Number(element.dataset.containerWidth))*100;
element.style.width = `${Math.max(minWidth, Math.min(maxWidth, newWidth))}%`;
};
handle.addEventListener("mousedown", (e) => {
e.preventDefault();
element.dataset.startX = e.clientX;
element.dataset.containerWidth = mainContainer.getBoundingClientRect().width;
element.dataset.startWidth = element.getBoundingClientRect().width;
document.documentElement.addEventListener("mousemove", resize, false);
document.documentElement.addEventListener("mouseup", () => {
document.documentElement.removeEventListener("mousemove", resize, false);
element.style.userSelect = "initial";
}, { once: true });
});
}
appendResizer(document.querySelector(".kernel-list-parent"), { minWidth: 15, maxWidth: 50 }, left=true);
// **** keyboard shortcuts
document.addEventListener("keydown", async function(event) {
// up and down change the UOp or kernel from the list