This commit is contained in:
2025-12-21 12:34:00 -04:00
parent 59c02dd87f
commit 439c4319ec
3 changed files with 216 additions and 21 deletions
+170
View File
@@ -0,0 +1,170 @@
from __future__ import annotations
from typing import Any
from tinygrad.helpers import DEBUG, GlobalCounters, all_same, dedup, colored, ansilen, PROFILE, ProfilePointEvent, cpu_events, time_to_str, TRACEMETA
from tinygrad.uop.ops import UOp, Ops, sym_infer
from tinygrad.device import Device, Buffer
# **************** ExecutionUnit ****************
class ExecutionUnit:
"""
A bound, ready-to-execute unit. Replaces CapturedJit.
Takes ExecItems and binds them to real device resources on execution.
"""
def __init__(self, items: list):
"""
Create an ExecutionUnit from ExecItems.
Args:
items: ExecItems (with bufs as Buffers or UOps, lib optionally set)
"""
from tinygrad.engine.realize import ExecItem
self.items: list[ExecItem] = items
self.buffer_map: dict[UOp, Buffer] = {}
self.var_vals: dict[str, int] = {}
# Create bound items with runners - lazy, done on first call
self._bound_items: list[tuple[Any, list[Buffer], tuple, dict[str, int]]]|None = None
self._graphs: list|None = None
self._first_run = True
def _bind(self):
"""Create runners from lib and bind buffers."""
from tinygrad.engine.realize import CompiledRunner, BufferCopy, BufferXfer, ViewOp, EncDec, get_runner, get_program
self._bound_items = []
for item in self.items:
# Get buffers - either from buffer_map (for UOps) or directly (for already-bound Buffers)
bufs: list[Buffer] = []
for b in item.bufs:
if b is None:
continue
if isinstance(b, UOp):
bufs.append(self.buffer_map[b])
else:
bufs.append(b)
# Create runner from lib or use existing prg
if item.prg is not None:
runner = item.prg
elif item.ast.op is Ops.SINK:
device = bufs[0].device
if item.lib is not None:
# Create runner from cached lib
prg = get_program(item.ast, Device[device].renderer)
runner = CompiledRunner(prg, item.lib)
else:
# Compile and create runner
runner = get_runner(device, item.ast)
elif item.ast.op is Ops.BUFFER_VIEW:
runner = ViewOp(bufs[0])
elif item.ast.op is Ops.COPY:
if hasattr(Device[bufs[0].device].allocator, '_transfer') and all_same([x.device.split(":")[0] for x in bufs]):
runner = BufferXfer(bufs[0].nbytes, bufs[0].device, bufs[1].device)
else:
runner = BufferCopy(bufs[0].nbytes, bufs[0].device, bufs[1].device)
elif item.ast.op is Ops.ENCDEC:
runner = EncDec(item.ast, bufs[0].nbytes, bufs[1].device)
else:
raise RuntimeError(f"unknown op {item.ast.op}")
self._bound_items.append((runner, bufs, item.metadata, item.fixedvars))
def update(self, buffers: dict[UOp, Buffer]|None = None, var_vals: dict[str, int]|None = None):
"""Update buffer mapping and/or var_vals for next run."""
if buffers is not None:
self.buffer_map.update(buffers)
# Need to rebind if we update buffers
self._bound_items = None
if var_vals is not None:
self.var_vals.update(var_vals)
def __call__(self, var_vals: dict[str, int]|None = None, wait=False, do_update_stats=True, jit=False) -> float|None:
"""Execute all items."""
from tinygrad.engine.realize import CompiledRunner
if var_vals is not None:
self.var_vals.update(var_vals)
# Lazy bind on first call
if self._bound_items is None:
self._bind()
assert self._bound_items is not None
# TODO: create graphs on first run
# if self._first_run:
# self._create_graphs()
# self._first_run = False
# Execute all items
total_et = 0.0
for runner, bufs, metadata, fixedvars in self._bound_items:
merged_var_vals = self.var_vals | fixedvars
# Ensure buffers are allocated (skip if jit - already allocated)
if not jit:
for b in bufs:
b.ensure_allocated()
# Reorder bufs to match program globals if needed
if isinstance(runner, CompiledRunner):
ordered_bufs = [bufs[i] for i in runner.p.globals]
else:
ordered_bufs = bufs
# PROFILE events
if PROFILE:
payload = {"metadata":metadata, "var_vals":merged_var_vals, "bufs":[b.trace_num for b in ordered_bufs], "name":runner.display_name}
payload["outputs"], payload["inputs"] = (runner.p.outs, runner.p.ins) if isinstance(runner, CompiledRunner) else ([0], [1])
cpu_events.append(ProfilePointEvent(runner.device, "exec", len(cpu_events), payload))
et = runner(ordered_bufs, merged_var_vals, wait=wait or DEBUG >= 2)
if et is not None:
total_et += et
# Update stats
if do_update_stats:
GlobalCounters.kernel_count += 1
op_est = sym_infer(runner.estimates.ops, merged_var_vals)
mem_est = sym_infer(runner.estimates.mem, merged_var_vals)
GlobalCounters.global_ops += op_est
GlobalCounters.global_mem += mem_est
if et is not None:
GlobalCounters.time_sum_s += et
if DEBUG >= 2:
lds_est = sym_infer(runner.estimates.lds, merged_var_vals)
mem_est = min(mem_est, lds_est) # there can't be more memory accessed than loads/stores
header_color = 'magenta' if jit else ('green' if runner.first_run else None)
ptm = colored(time_to_str(et, w=9), "yellow" if et > 0.01 else None) if et is not None else ""
flops, membw, ldsbw = op_est/(et or 1e-20), mem_est/(et or 1e-20), lds_est/(et or 1e-20)
flops_str = f"{flops*1e-9:7.0f} GFLOPS" if flops < 1e14 else colored(f"{flops*1e-12:7.0f} TFLOPS", 'green')
mem_str = f"{membw*1e-9:4.0f}|{ldsbw*1e-9:<6.0f} GB/s" if membw < 1e13 and ldsbw < 1e15 else \
colored(f"{membw*1e-12:4.0f}|{ldsbw*1e-12:<6.0f} TB/s", 'green')
print(f"{colored(f'*** {runner.device[:7]:7s} {GlobalCounters.kernel_count:4d}', header_color)}"+
f" {runner.display_name+' '*(46-ansilen(runner.display_name))} arg {len(ordered_bufs):2d} mem {GlobalCounters.mem_used/1e9:6.2f} GB"+
("" if et is None else f" tm {ptm}/{GlobalCounters.time_sum_s*1e3:9.2f}ms ({flops_str} {mem_str})")+
f" {[repr(m) if TRACEMETA >= 2 else str(m) for m in metadata] if metadata else ''}")
runner.first_run = False
return total_et if wait else None
def __add__(self, other: ExecutionUnit) -> ExecutionUnit:
"""Combine two ExecutionUnits, rebuild graph lazily."""
combined = ExecutionUnit(self.items + other.items)
combined.buffer_map = {**self.buffer_map, **other.buffer_map}
combined.var_vals = {**self.var_vals, **other.var_vals}
return combined
def free_intermediates(self):
"""Deallocate internal buffers."""
for buf in self.buffer_map.values():
if buf.is_allocated():
buf.deallocate()
# Reset bound state
self._bound_items = None
self._graphs = None
self._first_run = True
+7 -2
View File
@@ -195,6 +195,8 @@ class CapturedJit(Generic[ReturnType]):
# jit exec
def __call__(self, input_buffers:list[Buffer], var_vals:dict[str, int]) -> ReturnType:
from tinygrad.engine.execution import ExecutionUnit
# assign inputs
for idx, offset, device, size, dtype in self.extra_view_inputs:
input_buffers.append(Buffer(device, size, dtype, base=input_buffers[idx], offset=offset).ensure_allocated())
@@ -213,7 +215,10 @@ class CapturedJit(Generic[ReturnType]):
self._first_run = False
if DEBUG >= 1 and len(self._jit_cache) >= 10: print(f"jit execs {len(self._jit_cache)} kernels")
for ei in self._jit_cache: ei.run(var_vals, jit=True)
# Use ExecutionUnit for execution
unit = ExecutionUnit(self._jit_cache)
unit.update(var_vals=var_vals)
unit(jit=True)
self._clear_inputs()
return self.ret
@@ -251,7 +256,7 @@ class TinyJit(Generic[ReturnType]):
return ret
def add(self, ei:ExecItem):
self._jit_cache.append(ExecItem(ei.ast, [self.add_buffer(buf) for buf in ei.bufs if buf is not None], ei.metadata, ei.fixedvars, ei.prg))
self._jit_cache.append(ExecItem(ei.ast, [self.add_buffer(buf) for buf in ei.bufs if buf is not None], ei.metadata, ei.fixedvars, ei.lib, ei.prg))
def reset(self):
assert self.fxn is not None, "can't reset without function"
+39 -19
View File
@@ -186,12 +186,17 @@ class ExecItem:
bufs: list[Buffer|None] = field(default_factory=list)
metadata: tuple[Metadata, ...] = ()
fixedvars: dict[str, int] = field(default_factory=dict)
lib: bytes|None = None # compiled binary, None for COPY/VIEW/ENCDEC
prg: Runner|None = None
def lower(self):
"""Populate self.prg by lowering the AST."""
"""Populate self.prg and self.lib by lowering the AST."""
if self.prg is not None: return self
try: self.prg = cast(Runner, si_lowerer.rewrite(self.ast, self.bufs))
try:
self.prg = cast(Runner, si_lowerer.rewrite(self.ast, self.bufs))
# Store lib for SINK ops (compiled kernels)
if isinstance(self.prg, CompiledRunner):
self.lib = self.prg.lib
except Exception as e:
if DEBUG >= 2:
print(f"error lowering {self.ast.op}")
@@ -238,23 +243,38 @@ class ExecItem:
capturing: list = [] # put classes with an add method in here
def run_schedule(schedule:list[ExecItem], var_vals:dict[str, int]|None=None, do_update_stats=True):
while len(schedule):
ei = schedule.pop(0).lower()
from tinygrad.engine.execution import ExecutionUnit
# Lower all items first
lowered: list[ExecItem] = []
for ei in schedule:
ei = ei.lower()
if len(capturing) and CAPTURING: capturing[0].add(ei)
if VALIDATE_WITH_CPU and ei.ast.op is Ops.SINK:
# copy in allocated buffers from the GPU
bufs = [b for b in ei.bufs if b is not None]
nb: list[Buffer|None] = [Buffer("CPU", b.size, b.dtype) for b in bufs]
for cpu_b, gpu_b in zip(nb, bufs):
if cpu_b is not None and gpu_b.is_allocated(): cpu_b.ensure_allocated().copyin(gpu_b.as_buffer())
lowered.append(ei)
# run on GPU
ei.run(var_vals, do_update_stats=do_update_stats)
if VALIDATE_WITH_CPU:
# Run item by item with CPU validation
for ei in lowered:
if ei.ast.op is Ops.SINK:
# copy in allocated buffers from the GPU
bufs = [b for b in ei.bufs if b is not None]
nb: list[Buffer|None] = [Buffer("CPU", b.size, b.dtype) for b in bufs]
for cpu_b, gpu_b in zip(nb, bufs):
if cpu_b is not None and gpu_b.is_allocated(): cpu_b.ensure_allocated().copyin(gpu_b.as_buffer())
# validate the output buffers match (NOTE: this is assuming the output is buffer 0)
with Context(BEAM=0): ExecItem(ei.ast, nb, ei.metadata, ei.fixedvars).run(var_vals, do_update_stats=do_update_stats)
import numpy as np
assert nb[0] is not None
np.testing.assert_allclose(bufs[0].numpy(), nb[0].numpy(), rtol=1e-3, atol=1e-3)
else:
ei.run(var_vals, do_update_stats=do_update_stats)
# run on GPU
ei.run(var_vals, do_update_stats=do_update_stats)
# validate the output buffers match (NOTE: this is assuming the output is buffer 0)
with Context(BEAM=0): ExecItem(ei.ast, nb, ei.metadata, ei.fixedvars).run(var_vals, do_update_stats=do_update_stats)
import numpy as np
assert nb[0] is not None
np.testing.assert_allclose(bufs[0].numpy(), nb[0].numpy(), rtol=1e-3, atol=1e-3)
else:
ei.run(var_vals, do_update_stats=do_update_stats)
else:
# Use ExecutionUnit for batched execution
if lowered:
unit = ExecutionUnit(lowered)
unit.update(var_vals=var_vals)
unit(do_update_stats=do_update_stats)