Compare commits

...
21 Commits
Author SHA1 Message Date
geohot 0b00981cd1 fix wmma 2025-10-15 09:38:07 +08:00
geohot 5d485660da reproed failure in emulation 2025-10-15 09:19:56 +08:00
geohot 0e06b5cbb6 support emulate in the NullDevice 2025-10-15 09:11:57 +08:00
George HotzandGitHub 3a4a3e09ea Merge branch 'master' into new_shape 2025-10-14 21:16:44 +08:00
geohot 3b0b3dcff3 oops, i didn't mean to change that 2025-10-14 20:12:47 +08:00
geohot accee5d840 hack for 3 op assign 2025-10-14 20:10:53 +08:00
geohot e6812bbe63 one less st 2025-10-14 19:54:37 +08:00
geohot 18a6492e98 test is broken 2025-10-14 19:50:17 +08:00
geohot 9723b4f1c1 close 2025-10-14 19:46:52 +08:00
geohot 07162df323 size doesn't use st 2025-10-14 19:30:24 +08:00
geohot 7a2e206a0d fix tests 2025-10-14 19:22:16 +08:00
George HotzandGitHub 61855c24a8 Merge branch 'master' into new_shape 2025-10-14 19:15:09 +08:00
geohot 28076d9270 const uses _shape 2025-10-14 19:12:03 +08:00
geohot d51cae1396 shape is good 2025-10-14 18:28:00 +08:00
geohot 0b69698ad4 mostly works 2025-10-14 18:19:02 +08:00
geohot 04ead92ebd _shape is like _device 2025-10-14 17:53:17 +08:00
geohot faddebef07 need to cache it 2025-10-14 17:35:29 +08:00
geohot a659cb18a4 all mops 2025-10-14 17:24:08 +08:00
geohot 8721b6884c more mops 2025-10-14 17:20:04 +08:00
geohot 59512a49fa reshape causing issues 2025-10-14 16:59:25 +08:00
geohot a73b59caa2 work on shape property 2025-10-14 16:50:43 +08:00
5 changed files with 119 additions and 19 deletions
+2
View File
@@ -272,6 +272,8 @@ jobs:
# run: NULL=1 DEBUG=1 python3 examples/sdxl.py --seed 0 --noshow --timing --fakeweights
- name: Run Clip tests for SD MLPerf on NULL backend
run: NULL=1 python -m pytest -n=auto test/external/mlperf_stable_diffusion/external_test_models.py::TestOpenClip --durations=20
- name: Run AMD emulated BERT training on NULL backend
run: EMULATE=AMD_RDNA4 NULL=1 CAPTURE_PROCESS_REPLAY=0 DEFAULT_FLOAT=HALF BENCHMARK=10 BS=66 GPUS=1 BERT_LAYERS=2 MODEL=bert python3 examples/mlperf/model_train.py
# TODO: support fake weights
#- name: Run LLaMA 7B on 4 fake devices
# run: NULL=1 python3 examples/llama.py --gen 1 --size 7B --shard 4 --prompt "Hello." --count 3 --temperature 0 --timing
+1 -1
View File
@@ -658,7 +658,7 @@ class TestMultiTensor(unittest.TestCase):
# it doesn't work like this anymore
# NOTE: this never failed in assign_multi, it failed tensor spec because MULTI was never pushed in the graph
@unittest.expectedFailure
@unittest.skip("this test is broken")
def test_mlb_assign_change_axis(self):
t_none = Tensor.zeros((16, 16)).shard(devices_2).contiguous().realize()
t_zero = Tensor.ones((16, 16)).shard(devices_2, axis=0)
+12 -4
View File
@@ -1,9 +1,11 @@
import functools
from typing import cast
from tinygrad.device import Compiled, Compiler, Allocator
from tinygrad.engine.jit import MultiGraphRunner
from tinygrad.renderer.cstyle import CStyleLanguage
from tinygrad.renderer.cstyle import Renderer, CStyleLanguage
from tinygrad.renderer.llvmir import AMDLLVMRenderer
from tinygrad.uop.ops import Ops
from tinygrad.helpers import cpu_profile
from tinygrad.helpers import cpu_profile, EMULATE
class NullRenderer(CStyleLanguage):
device = "NULL"
@@ -29,5 +31,11 @@ class NullGraph(MultiGraphRunner):
def __call__(self, input_rawbuffers, var_vals, wait=False) -> float|None: return 1e-3
class NullDevice(Compiled):
def __init__(self, device:str): super().__init__(device, NullAllocator(self), [(NullRenderer, Compiler)], functools.partial(NullProgram, device),
NullGraph)
def __init__(self, device:str):
renderer:functools.partial|type[Renderer]
match cast(str, EMULATE.value):
case "AMD": renderer = functools.partial(AMDLLVMRenderer, "gfx1100")
case "AMD_RDNA4": renderer = functools.partial(AMDLLVMRenderer, "gfx1201")
case "": renderer = NullRenderer
case _: raise RuntimeError(f"can't EMULATE device: {EMULATE.value}")
super().__init__(device, NullAllocator(self), [(renderer, Compiler)], functools.partial(NullProgram, device), NullGraph)
+1 -1
View File
@@ -60,7 +60,7 @@ earliest_rewrites = PatternMatcher([
lambda reduce,x: reduce.const_like(identity_element(reduce.arg[0], reduce.dtype)) if x.size == 0 and reduce.size != 0 else None),
# handle size 0
(UPat(GroupOp.All-{Ops.SINK}, name="x"), lambda x: x.const_like(0).rtag(x.tag) if x.st is not None and x.size == 0 else None),
(UPat(GroupOp.All-{Ops.SINK}, name="x"), lambda x: x.const_like(0).rtag(x.tag) if x._shape is not None and x.size == 0 else None),
# remove contiguous on movement ops before a copy on disk
(UPat(GroupOp.Movement-{Ops.SHRINK, Ops.RESHAPE}, name="x").f(Ops.CONTIGUOUS).f(Ops.COPY, allow_any_len=True, name="copy"),
+103 -13
View File
@@ -175,6 +175,7 @@ class UOp(MathTrait, metaclass=UOpMetaClass):
# *** uop shape stuff ***
# TODO: remove this. it's used by the jit and split_reduceop
@recursive_property
def st(self) -> ShapeTracker|None:
if self.op is Ops.INDEX and self.src[0].op in {Ops.DEFINE_GLOBAL, Ops.DEFINE_LOCAL, Ops.DEFINE_REG, Ops.MSTACK,
@@ -223,12 +224,98 @@ class UOp(MathTrait, metaclass=UOpMetaClass):
shape = tuple(1 if i in axis_arg else s for i,s in enumerate(shape))
return ShapeTracker.from_shape(shape)
@recursive_property
def _shape(self) -> tuple[sint, ...]|None:
match self.op:
# late ops don't have shape
case Ops.UNIQUE | Ops.DEVICE | Ops.RANGE | Ops.INDEX | Ops.LOAD | Ops.IF | Ops.BARRIER | \
Ops.VECTORIZE | Ops.VCONST | Ops.SUBSTITUTE | Ops.GEP | Ops.SPECIAL | Ops.UNROLL | Ops.PRECAST:
return None
# some ops init the shape
case Ops.CONST | Ops.DEFINE_VAR | Ops.BIND: return () if self._device is not None else None
case Ops.BUFFER: return (self.arg,)
case Ops.BUFFER_VIEW: return (self.arg[0],)
case Ops.BUFFERIZE: return tuple([int(r.vmax+1) for r in self.src[1:]])
case Ops.DEFINE_GLOBAL | Ops.DEFINE_LOCAL | Ops.DEFINE_REG: return (self.ptrdtype.size,)
# passthrough ops
case Ops.REDUCE | Ops.MSTACK | Ops.MSELECT | Ops.DETACH | Ops.CONTIGUOUS | Ops.CONTIGUOUS_BACKWARD | Ops.FUSE: return self.src[0]._shape
# ops with custom handling
case Ops.KERNEL: return self.arg.ast._shape
case Ops.STORE:
if isinstance(self.dtype, PtrDType): return (self.ptrdtype.size,)
if self.dtype is not dtypes.void: return self.src[0].src[0].shape
return None
# TODO: disallow shape changing bitcast
case Ops.BITCAST:
ps = self.src[0]._shape
if ps is None: return None
if (output_sz:=self.dtype.itemsize) != (input_sz:=self.src[0].dtype.itemsize): return ps[:-1]+(ssimplify((ps[-1]*input_sz) // output_sz),)
return ps
# TODO: disallow reshape from nothing. tested by TestOpenClip.test_multigpu_clip_score
case Ops.RESHAPE:
if self.src[0]._shape is None: return tuple(ssimplify(s) for s in self.arg)
# movement ops change the shape. this is the logic from the old ShapeTracker
# NOTE: ssimplify is required because the shape needs to be canonical for broadcasting and same shape checking
if self.op in GroupOp.Movement.union({Ops.MULTI, Ops.REDUCE_AXIS, Ops.WMMA}):
ps = self.src[0]._shape
# TODO: WMMA is used for both axis WMMA and op WMMA. fix this and remove this hack. tested by BERT on AMD LLVM
if ps is None and self.op is Ops.WMMA: return None
if ps is None: raise RuntimeError(f"movement op {self.op} requires shape")
match self.op:
case Ops.RESHAPE:
if not all(x >= 0 for x in self.arg): raise ValueError(f"shape can't contain negative numbers {self.arg}")
if prod(ps) != prod(self.arg): raise ValueError(f"bad reshape: {ps} -> {self.arg}")
return tuple(ssimplify(s) for s in self.arg)
case Ops.EXPAND:
if len(ps) != len(self.arg) or not all(s==ns or (s==1 and ns>=0) for s,ns in zip(ps, self.arg)):
raise ValueError(f"bad expand: {ps} -> {self.arg}")
return tuple(ssimplify(s) for s in self.arg)
case Ops.PERMUTE:
if sorted(self.arg) != list(range(len(ps))): raise ValueError(f"invalid permutation {self.arg} of len {len(ps)}")
return tuple(ps[i] for i in self.arg)
case Ops.PAD:
# TODO: why do i need resolve here?
if len(ps) != len(self.arg) or not all(resolve(b>=0) and resolve(e>=0) for b,e in self.arg): raise ValueError(f"invalid pad {self.arg}")
return tuple(ssimplify(s+b+e) for s,(b,e) in zip(ps, self.arg))
case Ops.SHRINK:
# TODO: why do i need resolve here?
if len(ps) != len(self.arg) or not all(resolve(0<=b) and resolve(b<=e) and resolve(e<=s) for s,(b,e) in zip(ps, self.arg)):
raise ValueError(f"invalid shrink {self.arg} for {ps}")
return tuple(ssimplify(e-s) for s,e in self.arg)
case Ops.FLIP:
if len(ps) != len(self.arg) or not all(isinstance(x, bool) for x in self.arg): raise ValueError(f"bad flip on {ps}, {self.arg}")
return ps
case Ops.MULTI: return tuple(s*len(self.device) if a == self.axis else s for a,s in enumerate(ps))
case Ops.REDUCE_AXIS | Ops.WMMA:
axis_arg = self.arg[1] if self.op is Ops.REDUCE_AXIS else self.arg[7]
if not isinstance(axis_arg, tuple) or not all(isinstance(x, int) and x>=0 and x<len(ps) for x in axis_arg):
raise ValueError(f"invalid type for axis: {axis_arg}")
return tuple(1 if i in axis_arg else s for i,s in enumerate(ps))
# elementwise ops keep the shape the same. all inputs with shape must match
if self.op in (GroupOp.Elementwise-{Ops.BITCAST}).union({Ops.COPY, Ops.ASSIGN, Ops.NOOP, Ops.SINK, Ops.ALLREDUCE}):
# TODO: remove this hack for 3 op assign
input_shapes = [x._shape for x in (self.src[:2] if self.op is Ops.ASSIGN else self.src) if x._shape is not None]
if len(input_shapes) == 0: return None
if not all_same(input_shapes): raise RuntimeError(f"shape mismatch at {self.op}: {input_shapes}")
return input_shapes[0]
# all Ops must be explicitly handled
raise NotImplementedError(f"no shape handling for {self.op} with {self.dtype}")
@property
def shape(self) -> tuple[sint, ...]:
assert self.st is not None, f"{self.op} doesn't have a shape"
return unwrap(self.st).shape
if (ret:=self._shape) is None: raise RuntimeError(f"shape requested, but {self.op} doesn't have a shape")
return ret
@property
def size(self) -> int: return self.arg[0] if self.op is Ops.BUFFER_VIEW else self.arg if self.op is Ops.BUFFER else unwrap(self.st).size
def size(self) -> int: return prod([int(x.vmax) if isinstance(x, UOp) else x for x in self.shape])
# determine what ranges this is in
@recursive_property
@@ -290,7 +377,7 @@ class UOp(MathTrait, metaclass=UOpMetaClass):
def __getitem__(self, idx): return self.index(idx)
def const_like(self, b:ConstLike):
# constants can optionally have a DEVICE source
return UOp.const(self.dtype, b, device=self._device, shape=self.shape if self.st is not None else None)
return UOp.const(self.dtype, b, device=self._device, shape=self._shape)
def broadcast(self, count:int):
assert self.dtype.count == 1
if count == 1: return self
@@ -428,19 +515,22 @@ class UOp(MathTrait, metaclass=UOpMetaClass):
if self.op is Ops.MULTI: return self.src[0].base # MULTI is really a VIEW
return self
def _mop(self, op:Ops, arg) -> UOp:
def _mop(self, op:Ops, arg, no_reshape_is_no_op:bool=False) -> UOp:
ret = UOp(op, self.dtype, (self,), arg)
if self.st == ret.st: return self # ignore NOOPs, also check ret.st
# for all movement ops, we check shape property
if ret.shape == self.shape and no_reshape_is_no_op: return self
return ret
def forced_reshape(self, arg:tuple[sint, ...], **kwargs): return UOp(Ops.RESHAPE, kwargs.pop("dtype", self.dtype), src=(self,), arg=arg)
def reshape(self, arg:tuple[sint, ...]): return self._mop(Ops.RESHAPE, arg)
def expand(self, arg:tuple[sint, ...]): return self._mop(Ops.EXPAND, arg)
def shrink(self, arg:tuple[tuple[sint, sint], ...]): return self._mop(Ops.SHRINK, arg)
def pad(self, arg:tuple[tuple[sint, sint], ...]): return self._mop(Ops.PAD, arg)
def permute(self, arg:tuple[int, ...]): return self._mop(Ops.PERMUTE, arg)
def flip(self, arg:tuple[bool, ...]): return self._mop(Ops.FLIP, arg)
# in these four, if the shape doesn't change we can return self
def reshape(self, arg:tuple[sint, ...]): return self._mop(Ops.RESHAPE, arg, no_reshape_is_no_op=True)
def expand(self, arg:tuple[sint, ...]): return self._mop(Ops.EXPAND, arg, no_reshape_is_no_op=True)
def shrink(self, arg:tuple[tuple[sint, sint], ...]): return self._mop(Ops.SHRINK, arg, no_reshape_is_no_op=True)
def pad(self, arg:tuple[tuple[sint, sint], ...]): return self._mop(Ops.PAD, arg, no_reshape_is_no_op=True)
# in these two, we have custom logic to check if they are a no-op
def permute(self, arg:tuple[int, ...]): return self._mop(Ops.PERMUTE, arg) if arg != tuple(range(len(self.shape))) else self
def flip(self, arg:tuple[bool, ...]): return self._mop(Ops.FLIP, arg) if any(arg) and len(arg) == len(self.shape) else self
# *** uop UNIQUE ***