From 7214450c23cca8d7e6e3de57b68af74ec0018c7c Mon Sep 17 00:00:00 2001 From: George Hotz <72895+geohot@users.noreply.github.com> Date: Wed, 2 Oct 2024 17:12:30 +0800 Subject: [PATCH] little symbolic changes [pr] (#6849) * little symbolic changes [pr] * symbolic needs resolve too * no resolve * less change --- test/test_symbolic_shapetracker.py | 14 +++++++++----- test/unit/test_uop_resolve.py | 4 ++++ tinygrad/codegen/kernel.py | 11 ++++++----- tinygrad/ops.py | 3 +++ tinygrad/shape/shapetracker.py | 8 ++++---- tinygrad/shape/view.py | 15 ++++++++------- 6 files changed, 34 insertions(+), 21 deletions(-) diff --git a/test/test_symbolic_shapetracker.py b/test/test_symbolic_shapetracker.py index 3112e6fcb2..a60efe7dda 100644 --- a/test/test_symbolic_shapetracker.py +++ b/test/test_symbolic_shapetracker.py @@ -4,6 +4,9 @@ from tinygrad.shape.symbolic import Variable, NumNode from tinygrad.tensor import Tensor class TestSymbolic(unittest.TestCase): + def assert_tuple_equal(self, x, y): + for a,b in zip(x,y): self.assertFalse(a != b) + def test_symbolic_st(self): x = Variable("x", 1, 100) st = ShapeTracker.from_shape((x, 3)) @@ -31,11 +34,11 @@ class TestSymbolic(unittest.TestCase): k = Variable("k", 1, 5).bind(3) t = Tensor.rand(3, 4).reshape(i, 4).cat(Tensor.rand(3, 4).reshape(j, 4), dim=0).cat(Tensor.rand(3, 4).reshape(k, 4), dim=0) st = t.lazydata.st - assert st.shape == (i+j+k, 4) + self.assert_tuple_equal(st.shape, (i+j+k, 4)) assert st.real_strides() == (4, 1) t = Tensor.rand(3, 3).reshape(i, 3).cat(Tensor.rand(3, 3).reshape(i, 3), dim=0).cat(Tensor.rand(3, 3), dim=0) st = t.lazydata.st - assert st.shape == (2*i+3, 3) + self.assert_tuple_equal(st.shape, (2*i+3, 3)) assert st.real_strides() == (3, 1) def test_cat_dim1_strides(self): @@ -44,10 +47,11 @@ class TestSymbolic(unittest.TestCase): k = Variable("k", 1, 5).bind(4) t = Tensor.rand(3, 4).reshape(3, i).cat(Tensor.rand(3, 4).reshape(3, j), dim=1).cat(Tensor.rand(3, 4).reshape(3, k), dim=1) st = t.lazydata.st - assert st.shape == (3, i+j+k) - assert st.real_strides() == (i+j+k, 1) + self.assert_tuple_equal(st.shape, (3, i+j+k)) + self.assert_tuple_equal(st.real_strides(), (i+j+k, 1)) class TestSymbolicVarVals(unittest.TestCase): + def assert_equal(self, x, y): self.assertFalse(x != y) def test_var_vals_empty(self): assert ShapeTracker.from_shape((3, 4, 5)).var_vals == {} @@ -58,7 +62,7 @@ class TestSymbolicVarVals(unittest.TestCase): def test_var_vals_offset(self): x = Variable("x", 1, 100).bind(3) st = ShapeTracker.from_shape((4, 3)).shrink(((x, x+1), (0, 3))) - assert st.views[-1].offset == x * 3 + self.assert_equal(st.views[-1].offset, x * 3) assert st.var_vals == {Variable("x", 1, 100): 3} def test_var_vals_mask(self): diff --git a/test/unit/test_uop_resolve.py b/test/unit/test_uop_resolve.py index dbb0692c35..3f84d683d8 100644 --- a/test/unit/test_uop_resolve.py +++ b/test/unit/test_uop_resolve.py @@ -39,6 +39,10 @@ class TestUOpResolve(unittest.TestCase): u = UOp.const(dtypes.int, 4) > 7 self.assertFalse(u) + def test_ssimplify(self): + self.assertEqual((8 % UOp.const(dtypes.int, 4)).ssimplify(), 0) + self.assertEqual((8 * UOp.const(dtypes.int, 4)).ssimplify(), 32) + def test_ambiguous_less_than(self): u = UOp.define_var("i", dtypes.pyint, 1, 10) self.assertTrue(resolve(u < 4)) diff --git a/tinygrad/codegen/kernel.py b/tinygrad/codegen/kernel.py index 16ba72358b..9c9381eacb 100644 --- a/tinygrad/codegen/kernel.py +++ b/tinygrad/codegen/kernel.py @@ -6,6 +6,7 @@ from typing import Optional, List, Tuple, cast, Dict, Final, DefaultDict from enum import Enum, auto from tinygrad.ops import BinaryOps, UNSAFE_PAD_OPS, KernelInfo, BUFFER_UOPS, UOp, UOps, print_uops, type_verify, graph_rewrite, PatternMatcher +from tinygrad.ops import resolve from tinygrad.device import Device from tinygrad.renderer import Renderer, TensorCore, Program from tinygrad.dtype import ImageDType, PtrDType @@ -85,7 +86,7 @@ class Kernel: # move all reduce axes to the end reduce = list(enumerate(zip(self.full_shape, self.output_shape))) - permute = tuple([i for i,(s,n) in reduce if s == n] + [i for i,(s,n) in reduce if s != n]) + permute = tuple([i for i,(s,n) in reduce if not resolve(s != n)] + [i for i,(s,n) in reduce if resolve(s != n)]) self.reshape_and_permute(None, permute) # parameters for optimization @@ -137,7 +138,7 @@ class Kernel: @property def first_reduce(self) -> int: - return [x!=y for x,y in zip(self.sts[0].shape[:self.first_upcast]+(0,), self.full_shape[:self.first_upcast]+(1,))].index(True) + return [resolve(x!=y) for x,y in zip(self.sts[0].shape[:self.first_upcast]+(0,), self.full_shape[:self.first_upcast]+(1,))].index(True) @property def first_upcast(self) -> int: return self.shape_len-self.upcasted @@ -297,7 +298,7 @@ class Kernel: if not (axis < len(axis_choices)): return None s0, s1, s2 = axis_choices[-(axis+1)][0][0], axis_choices[-(axis+1)][1][0], axis_choices[-(axis+1)][2] # s0 is n, s1 is m, s2 is k - axis_pads = tuple((x, tc.dims[i]) for i, x in enumerate([s0, s1, s2]) if self.full_shape[x]%tc.dims[i] != 0) + axis_pads = tuple((x, tc.dims[i]) for i, x in enumerate([s0, s1, s2]) if resolve(self.full_shape[x]%tc.dims[i] != 0)) if axis_pads and (opt_level < 2): return None self.bufs_for_tensor_core[reduceop] = (buf0, buf1) if DEBUG >= 3: print("TENSOR CORES", axis_buf0, axis_buf1, tc) @@ -476,7 +477,7 @@ class Kernel: (mulop:=self.reduceop.src[0]).arg is BinaryOps.MUL and mulop.src[0].op is UOps.LOAD and mulop.src[1].op is UOps.LOAD: st0, st1 = self.sts[self.bufs.index(mulop.src[0])], self.sts[self.bufs.index(mulop.src[1])] strides0, strides1 = st0.real_strides(), st1.real_strides() - def has_expanded_axis(shape, strides): return any(s > 1 and st == 0 for s,st in zip(shape,strides)) + def has_expanded_axis(shape, strides): return any(resolve(s > 1) and not resolve(st != 0) for s,st in zip(shape,strides)) if strides0[self.first_reduce] == 1 and not (has_expanded_axis(st0.shape, strides0) and has_expanded_axis(st1.shape, strides1)): for global_idx in range(self.global_dims): if self.full_shape[self.first_reduce]%MV_THREADS_PER_ROW == 0 and self.full_shape[global_idx]%(MV_BLOCKSIZE*MV_ROWS_PER_THREAD) == 0: @@ -625,7 +626,7 @@ class Kernel: reduce_idx = len(self.bufs) + self.reduceops.index(op)*2 alu_op: BinaryOps = op.arg[0] axis = tuple(i for i in range(self.first_reduce+self.group_for_reduces, self.shape_len) - if self.sts[reduce_idx].shape[i] != self.sts[reduce_idx+1].shape[i]) + if resolve(self.sts[reduce_idx].shape[i] != self.sts[reduce_idx+1].shape[i])) if op in self.bufs_for_tensor_core and (tc := self.tensor_core): rsrc = op.src[0] if rsrc.op is UOps.CAST: rsrc = rsrc.src[0] diff --git a/tinygrad/ops.py b/tinygrad/ops.py index 734616d187..22b2e6f6db 100644 --- a/tinygrad/ops.py +++ b/tinygrad/ops.py @@ -61,8 +61,10 @@ class MathTrait: def __truediv__(self, x): return self.alu(BinaryOps.MUL, self.ufix(x).alu(UnaryOps.RECIP)) def __rtruediv__(self, x): return self.ufix(x).alu(BinaryOps.MUL, self.alu(UnaryOps.RECIP)) def __mod__(self, x): return self.alu(BinaryOps.MOD, self.ufix(x)) + def __rmod__(self, x): return self.ufix(x).alu(BinaryOps.MOD, self) def __xor__(self, x): return self.alu(BinaryOps.XOR, self.ufix(x)) def __and__(self, x): return self.alu(BinaryOps.AND, self.ufix(x)) + def __rand__(self, x): return self.ufix(x).alu(BinaryOps.AND, self) def __or__(self, x): return self.alu(BinaryOps.OR, self.ufix(x)) def ne(self, x): return self.alu(BinaryOps.CMPNE, self.ufix(x)) def eq(self, x): return self.ne(x).logical_not() @@ -182,6 +184,7 @@ class UOp(MathTrait): def argstr(self): return f'({", ".join(map(str, self.arg))})' if self.op is UOps.REDUCE_AXIS else self.arg # *** uop evaluation *** def simplify(self): return graph_rewrite(self, simple_pm) + def ssimplify(self) -> Union[UOp, ConstType]: return ret.arg if (ret:=self.simplify()).op is UOps.CONST else ret def _eval(self, dtype, expected_type) -> ConstType: assert self.dtype in dtype, f"eval with wrong dtype {self}" simple_self = self.simplify() diff --git a/tinygrad/shape/shapetracker.py b/tinygrad/shape/shapetracker.py index 577dcbafc4..b524a7498c 100644 --- a/tinygrad/shape/shapetracker.py +++ b/tinygrad/shape/shapetracker.py @@ -7,7 +7,7 @@ from tinygrad.helpers import merge_dicts, getenv from tinygrad.shape.symbolic import Variable, MulNode, SumNode, NumNode, DivNode, ModNode, LtNode, AndNode, sint from tinygrad.shape.view import View, strides_for_shape from tinygrad.dtype import dtypes -from tinygrad.ops import UOp, UOps, BinaryOps, graph_rewrite +from tinygrad.ops import UOp, UOps, BinaryOps, graph_rewrite, resolve from tinygrad.codegen.uopgraph import sym, _get_chain # TODO: this needs to be replaced, there shouldn't be variables in the shapetracker, only ints and UOps @@ -25,10 +25,10 @@ def _uop_view(view:View, idxs:List[UOp], vexpr:UOp) -> Tuple[UOp, UOp]: # TODO: dtypes.realint iexpr = variable_to_uop(view.offset) for idx,sh,st,m in zip(idxs, view.shape, view.strides, view.mask if view.mask is not None else [None]*len(view.shape)): - if sh != 1 and st != 0: iexpr = iexpr + idx*variable_to_uop(st) + if resolve(sh != 1) and resolve(st != 0): iexpr = iexpr + idx*variable_to_uop(st) if m is not None: - if m[0] != 0: vexpr = vexpr * idx.ge(variable_to_uop(m[0])) - if m[1] != sh: vexpr = vexpr * idx.lt(variable_to_uop(m[1])) + if resolve(m[0] != 0): vexpr = vexpr * idx.ge(variable_to_uop(m[0])) + if resolve(m[1] != sh): vexpr = vexpr * idx.lt(variable_to_uop(m[1])) return iexpr, vexpr @dataclass(frozen=True) diff --git a/tinygrad/shape/view.py b/tinygrad/shape/view.py index a72c36bcb4..f5ccd149b4 100644 --- a/tinygrad/shape/view.py +++ b/tinygrad/shape/view.py @@ -2,6 +2,7 @@ from __future__ import annotations import functools, operator, itertools, math from dataclasses import dataclass from typing import Tuple, List, Optional, Dict, Set, cast +from tinygrad.ops import resolve from tinygrad.helpers import prod, all_int, argsort from tinygrad.shape.symbolic import Node, NumNode, Variable, sint, sym_infer @@ -108,8 +109,8 @@ class View: # if any dimension has size >1, but is masked such that only one index in the dimension is unmasked # then its stride can also be set to 0, albeit with a corresponding adjustment required to the offset # TODO: assert comparison with LtNode to avoid mis-using symbolic - if mask and any(elim := [not (b+1 < e) for b,e in mask]): - if any(not (b < e) for b,e in mask): + if mask and any(elim := [not resolve(b+1 < e) for b,e in mask]): + if any(not resolve(b < e) for b,e in mask): strides, offset, mask = (0,) * len(shape), 0, ((0,0),) * len(shape) offset += sum((strides[i] * mask[i][0]) if e else 0 for i, e in enumerate(elim)) strides = tuple(0 if e else st for st,e in zip(strides, elim)) @@ -162,10 +163,10 @@ class View: for term, s, o in zip(reversed(terms), reversed(vm2.shape), reversed(origin)): merged_term += Variable.sum([idxs[d1] * (s1 * merged_size) for d1, s1 in term]) + o * merged_size merged_size *= s - if not (merged_term >= merged_size) and not (merged_term < 0): + if not resolve(merged_term >= merged_size) and not resolve(merged_term < 0): extents.append((merged_size, merged_term)) merged_size, merged_term = 1, NumNode(0) - if merged_term != 0: return None + if resolve(merged_term != 0): return None if (vm2_shape := tuple(s for s,_ in reversed(extents))) != vm2.shape: return (reshaped_vm2 := vm2.reshape(vm2_shape)) and reshaped_vm2 + vm1 @@ -298,10 +299,10 @@ class View: for merged_dim, new_stride, real_dim in reversed(_merge_dims(self.shape, self.strides, self.mask)): acc = 1 # TODO: this <= and != is for symbolic!? - while acc <= merged_dim and acc != merged_dim and (new_dim := next(r_new_shape, 0)) > 0: + while resolve(acc <= merged_dim) and resolve(acc != merged_dim) and (new_dim := next(r_new_shape, 0)) > 0: strides.append(new_stride) - if new_dim != 1: new_stride *= (new_dim if (acc := acc * new_dim) < real_dim else 0) - if acc != merged_dim: break + if resolve(new_dim != 1): new_stride *= (new_dim if resolve((acc := acc * new_dim) < real_dim) else 0) + if resolve(acc != merged_dim): break else: strides += [0,] * (len(new_shape) - len(strides)) new_mask = _reshape_mask(self.mask, self.shape, new_shape)