diff --git a/examples/llm.c/export.py b/examples/llm.c/export.py index 40fe46548e..776a28b8a7 100755 --- a/examples/llm.c/export.py +++ b/examples/llm.c/export.py @@ -5,8 +5,8 @@ from tinygrad import Device, nn, Tensor, dtypes, Variable Device.DEFAULT = "CLANG" from train_gpt2 import GPT, GPTConfig from tinygrad.helpers import dedup, to_function_name, flatten, getenv, GRAPH, GlobalCounters, ansilen, to_function_name -from tinygrad.engine.schedule import create_schedule, memory_planner -from tinygrad.engine.realize import get_kernel, run_schedule +from tinygrad.engine.schedule import create_schedule +from tinygrad.engine.realize import get_kernel, memory_planner, run_schedule from tinygrad.ops import BufferOps, MetaOps TIMING = getenv("TIMING") diff --git a/examples/openpilot/compile2.py b/examples/openpilot/compile2.py index 967cdb1544..f3f484a0d5 100644 --- a/examples/openpilot/compile2.py +++ b/examples/openpilot/compile2.py @@ -16,8 +16,8 @@ from tinygrad import Tensor, Device, GlobalCounters, dtypes from tinygrad.dtype import ImageDType from tinygrad.device import Buffer from tinygrad.helpers import partition, Context, fetch, getenv, DEBUG, tqdm -from tinygrad.engine.realize import run_schedule, lower_schedule, ExecItem, CompiledRunner -from tinygrad.engine.schedule import ScheduleItem, create_schedule, memory_planner +from tinygrad.engine.realize import run_schedule, lower_schedule, ExecItem, CompiledRunner, memory_planner +from tinygrad.engine.schedule import ScheduleItem, create_schedule from tinygrad.ops import MetaOps from tinygrad.tensor import _to_np_dtype Device.DEFAULT = "GPU" diff --git a/tinygrad/engine/jit.py b/tinygrad/engine/jit.py index f572b6575d..eead7678c3 100644 --- a/tinygrad/engine/jit.py +++ b/tinygrad/engine/jit.py @@ -8,8 +8,7 @@ from tinygrad.device import Buffer, Compiled, Device from tinygrad.dtype import DType from tinygrad.shape.shapetracker import ShapeTracker from tinygrad.shape.symbolic import Variable, sint, sym_infer -from tinygrad.engine.realize import ExecItem, capturing, EmptyOp, ViewOp, BufferXfer, CompiledRunner, Runner -from tinygrad.engine.schedule import _internal_memory_planner +from tinygrad.engine.realize import ExecItem, capturing, EmptyOp, ViewOp, BufferXfer, CompiledRunner, Runner, _internal_memory_planner from tinygrad.nn.state import get_parameters from dataclasses import dataclass from weakref import WeakKeyDictionary diff --git a/tinygrad/engine/realize.py b/tinygrad/engine/realize.py index f58d7767a3..18b2f905a9 100644 --- a/tinygrad/engine/realize.py +++ b/tinygrad/engine/realize.py @@ -1,7 +1,8 @@ -from typing import List, Dict, Optional, cast, Generator, Tuple +from typing import List, Dict, Optional, cast, Generator, Tuple, Union import time, pprint +from collections import defaultdict from dataclasses import dataclass, replace -from tinygrad.helpers import colored, getenv, DEBUG, GlobalCounters, ansilen, BEAM, NOOPT, all_int, CAPTURING, Metadata, Context, TRACEMETA +from tinygrad.helpers import colored, getenv, DEBUG, GlobalCounters, ansilen, BEAM, NOOPT, all_int, CAPTURING, Metadata, Context, TRACEMETA, dedup from tinygrad.ops import MetaOps, LazyOp from tinygrad.dtype import dtypes from tinygrad.device import Device, Buffer @@ -219,3 +220,48 @@ def run_schedule(schedule:List[ScheduleItem], var_vals:Optional[Dict[Variable, i for ei in lower_schedule(schedule): if len(capturing) and CAPTURING: capturing[0].add(ei) ei.run(var_vals, do_update_stats=do_update_stats) + +# **************** memory planning **************** + +def _internal_memory_planner(buffers:List[Union[List[Buffer], Tuple[Buffer, ...]]], noopt_buffers=None, debug_prefix="") -> Dict[Buffer, Buffer]: + if getenv("NO_MEMORY_PLANNER"): return {} + first_appearance, last_appearance = {}, {} + for i,u in enumerate(buffers): + for buf in u: + if buf.is_allocated() or buf.lb_refcount > 0 or (noopt_buffers is not None and buf.base in noopt_buffers): continue + if buf.base not in first_appearance: first_appearance[buf.base] = i + last_appearance[buf.base] = i + + # Sort buffers by size in descending order, prioritizing largest buffers for allocation first. + # Track free segments, each containing (start, stop, and buffer that could be reused on this segment). + free_segs: Dict[Tuple, List[Tuple[int, int, Buffer]]] = defaultdict(list) # Dict[buffer key, Tuple[start, end, buffer to reuse on the seg]] + def find_replace_buffer(buf, st, en): + key = (buf.device, buf.dtype, buf.options) + ((buf.nbytes,) if not hasattr(Device[buf.device].allocator, "offset") else tuple()) + + default_buf = (0, len(buffers) - 1, buf) # will return the buffer itself if the replace one is not found. + seg_st, seg_en, seg_buf = next((free_segs[key].pop(i) for i,(sst,sen,_) in enumerate(free_segs[key]) if sst <= st and en <= sen), default_buf) + + free_segs[key] += [(seg_st, st - 1, seg_buf)] if st - 1 >= seg_st else [] + free_segs[key] += [(en + 1, seg_en, seg_buf)] if seg_en >= en + 1 else [] + + return seg_buf if seg_buf.nbytes == buf.nbytes else Buffer(buf.device, buf.size, buf.dtype, base=seg_buf) + + buffer_requests = sorted([(first_appearance[buf], last_appearance[buf], buf) for buf in first_appearance.keys()], key=lambda x: -x[2].nbytes) + assigned = {buf:find_replace_buffer(buf, st, en) for st, en, buf in buffer_requests} + + for i,u in enumerate(buffers): + for buf in u: + if buf.is_allocated() or buf.lb_refcount > 0 or (noopt_buffers is not None and buf.base in noopt_buffers): continue + if buf._base is not None: assigned[buf] = Buffer(buf.device, buf.size, buf.dtype, base=assigned.get(buf.base, buf.base).base, offset=buf.offset) + else: assigned[buf] = assigned.get(buf, buf) + + if DEBUG >= 1 and len(ak:=dedup(x for x in assigned.keys() if x._base is None)) != len(av:=dedup(x for x in assigned.values() if x._base is None)): + print(debug_prefix+f"memory reduced from {sum([x.nbytes for x in ak])/1e6:.2f} MB -> {sum([x.nbytes for x in av])/1e6:.2f} MB,", + f"{len(ak)} -> {len(av)} bufs") + return assigned + +def memory_planner(schedule:List[ScheduleItem]) -> List[ScheduleItem]: + # Exclude buffers involved in load ops (e.g transfers) to preserve parallelism in graphs. + assigned = _internal_memory_planner([si.bufs for si in schedule], + noopt_buffers={b for si in schedule if si.ast.op is not MetaOps.KERNEL for b in si.bufs}) + return [ScheduleItem(si.ast, tuple(assigned.get(x, x) for x in si.bufs), si.metadata) for si in schedule] diff --git a/tinygrad/engine/schedule.py b/tinygrad/engine/schedule.py index cfc3b3d03e..997ebba193 100644 --- a/tinygrad/engine/schedule.py +++ b/tinygrad/engine/schedule.py @@ -1,7 +1,7 @@ import sys, pickle, atexit from collections import defaultdict, deque from dataclasses import dataclass -from typing import Tuple, List, Dict, Optional, Set, DefaultDict, Union, cast, get_args +from typing import Tuple, List, Dict, Optional, Set, DefaultDict, cast, get_args from tinygrad.ops import MetaOps, BufferOps, LazyOp, Op, ReduceOps, ConstBuffer, MemBuffer, UNSAFE_PAD_OPS, UnaryOps, reduce_st from tinygrad.engine.graph import log_lazybuffer, realized_lazybuffer from tinygrad.helpers import GRAPH, DEBUG, MULTIOUTPUT, SAVE_SCHEDULE, FUSE_CONV_BW, FUSE_ARANGE, Context, GlobalCounters, colored, prod, dedup,\ @@ -10,7 +10,7 @@ from tinygrad.shape.symbolic import Variable, sint from tinygrad.dtype import ConstType, ImageDType, dtypes from tinygrad.lazy import LazyBuffer from tinygrad.shape.shapetracker import ShapeTracker -from tinygrad.device import Buffer, Device +from tinygrad.device import Buffer from tinygrad.shape.view import View, strides_for_shape # creation can recurse a lot @@ -392,48 +392,3 @@ def create_schedule(outs:List[LazyBuffer], seen:Optional[Set[LazyBuffer]]=None) schedule, var_vals = create_schedule_with_vars(outs, seen) assert len(var_vals) == 0 return schedule - -# *** memory planning *** - -def _internal_memory_planner(buffers:List[Union[List[Buffer], Tuple[Buffer, ...]]], noopt_buffers=None, debug_prefix="") -> Dict[Buffer, Buffer]: - if getenv("NO_MEMORY_PLANNER"): return {} - first_appearance, last_appearance = {}, {} - for i,u in enumerate(buffers): - for buf in u: - if buf.is_allocated() or buf.lb_refcount > 0 or (noopt_buffers is not None and buf.base in noopt_buffers): continue - if buf.base not in first_appearance: first_appearance[buf.base] = i - last_appearance[buf.base] = i - - # Sort buffers by size in descending order, prioritizing largest buffers for allocation first. - # Track free segments, each containing (start, stop, and buffer that could be reused on this segment). - free_segs: Dict[Tuple, List[Tuple[int, int, Buffer]]] = defaultdict(list) # Dict[buffer key, Tuple[start, end, buffer to reuse on the seg]] - def find_replace_buffer(buf, st, en): - key = (buf.device, buf.dtype, buf.options) + ((buf.nbytes,) if not hasattr(Device[buf.device].allocator, "offset") else tuple()) - - default_buf = (0, len(buffers) - 1, buf) # will return the buffer itself if the replace one is not found. - seg_st, seg_en, seg_buf = next((free_segs[key].pop(i) for i,(sst,sen,_) in enumerate(free_segs[key]) if sst <= st and en <= sen), default_buf) - - free_segs[key] += [(seg_st, st - 1, seg_buf)] if st - 1 >= seg_st else [] - free_segs[key] += [(en + 1, seg_en, seg_buf)] if seg_en >= en + 1 else [] - - return seg_buf if seg_buf.nbytes == buf.nbytes else Buffer(buf.device, buf.size, buf.dtype, base=seg_buf) - - buffer_requests = sorted([(first_appearance[buf], last_appearance[buf], buf) for buf in first_appearance.keys()], key=lambda x: -x[2].nbytes) - assigned = {buf:find_replace_buffer(buf, st, en) for st, en, buf in buffer_requests} - - for i,u in enumerate(buffers): - for buf in u: - if buf.is_allocated() or buf.lb_refcount > 0 or (noopt_buffers is not None and buf.base in noopt_buffers): continue - if buf._base is not None: assigned[buf] = Buffer(buf.device, buf.size, buf.dtype, base=assigned.get(buf.base, buf.base).base, offset=buf.offset) - else: assigned[buf] = assigned.get(buf, buf) - - if DEBUG >= 1 and len(ak:=dedup(x for x in assigned.keys() if x._base is None)) != len(av:=dedup(x for x in assigned.values() if x._base is None)): - print(debug_prefix+f"memory reduced from {sum([x.nbytes for x in ak])/1e6:.2f} MB -> {sum([x.nbytes for x in av])/1e6:.2f} MB,", - f"{len(ak)} -> {len(av)} bufs") - return assigned - -def memory_planner(schedule:List[ScheduleItem]) -> List[ScheduleItem]: - # Exclude buffers involved in load ops (e.g transfers) to preserve parallelism in graphs. - assigned = _internal_memory_planner([si.bufs for si in schedule], - noopt_buffers={b for si in schedule if si.ast.op is not MetaOps.KERNEL for b in si.bufs}) - return [ScheduleItem(si.ast, tuple(assigned.get(x, x) for x in si.bufs), si.metadata) for si in schedule] diff --git a/tinygrad/tensor.py b/tinygrad/tensor.py index 308c510afd..7e15ca64d2 100644 --- a/tinygrad/tensor.py +++ b/tinygrad/tensor.py @@ -15,8 +15,8 @@ from tinygrad.multi import MultiLazyBuffer from tinygrad.ops import MetaOps, truncate from tinygrad.device import Device, Buffer, BufferOptions from tinygrad.shape.symbolic import sint, Variable, MulNode, SumNode, NumNode, Node -from tinygrad.engine.realize import run_schedule -from tinygrad.engine.schedule import ScheduleItem, create_schedule_with_vars, memory_planner +from tinygrad.engine.realize import run_schedule, memory_planner +from tinygrad.engine.schedule import ScheduleItem, create_schedule_with_vars # **** start with two base classes, Tensor and Function ****