forked from tinygrad/tinygrad
Merge branch 'master' into dsp_search
This commit is contained in:
@@ -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")
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user