diff --git a/test/backend/test_isel.py b/test/backend/test_isel.py index e454431bd4..cc2a022174 100644 --- a/test/backend/test_isel.py +++ b/test/backend/test_isel.py @@ -26,7 +26,7 @@ class TestIselX86(unittest.TestCase): with self.subTest(dtype=dt, count=count): v = [UOp.variable(str(i), 0, 0, dt) if count == 1 else vector(str(i), dt, count) for i in range(nargs)] n = self.isel_rewrite(expr(*v)) - self.assertIs(n.arg.op, op) + self.assertEqual(n.arg, op) self.assertIs(n.dtype, dt) self.assertEqual(n.shape, () if count == 1 else (count,)) @@ -37,9 +37,9 @@ class TestIselX86(unittest.TestCase): d = (a != b).where(a, b) f = c + d n = self.isel_rewrite(f) - self.assertTrue(n.src[0].arg.op is X86Ops.CMOVL and n.src[1].arg.op is X86Ops.CMOVNE) + self.assertTrue(n.src[0].arg == X86Ops.CMOVL and n.src[1].arg == X86Ops.CMOVNE) # both comparisons become the same instruction - self.assertTrue(n.src[0].src[2] == n.src[1].src[2] and n.src[0].src[2].arg.op is X86Ops.CMP) + self.assertTrue(n.src[0].src[2] == n.src[1].src[2] and n.src[0].src[2].arg == X86Ops.CMP) def test_vmax(self): dt_op = [(dtypes.float32, 1, X86Ops.VMAXSS), (dtypes.float64, 1, X86Ops.VMAXSD), @@ -66,22 +66,22 @@ class TestIselX86(unittest.TestCase): a = UOp.variable("a", 0, 0, dtypes.int32) n = self.isel_rewrite(a.broadcast(4)) # need to move src from gpr to xmm before broadcasting - self.assertTrue(n.arg.op is X86Ops.VPBROADCASTD and n.src[0].arg.op is X86Ops.VMOVD) + self.assertTrue(n.arg == X86Ops.VPBROADCASTD and n.src[0].arg == X86Ops.VMOVD) # if we can fuse a load we can skip the move and access memory directly load = UOp.param(0, dtypes.int32, (16,)).index(UOp.const(dtypes.int32, 0)).load() n = self.isel_rewrite(load.broadcast(4)) - self.assertTrue(n.arg.op is X86Ops.VPBROADCASTD and len(n.src) == 4) + self.assertTrue(n.arg == X86Ops.VPBROADCASTD and len(n.src) == 4) def test_narrow_load_fold(self): load = UOp.param(0, dtypes.uint8, (1,)).index(UOp.const(dtypes.index, 0)).load().cast(dtypes.uint16) n = self.isel_rewrite(load) - self.assertIs(n.arg.op, X86Ops.MOVZX) + self.assertEqual(n.arg, X86Ops.MOVZX) self.assertEqual(len(n.src), 4) def test_vbroadcastss(self): a = UOp.variable("a", 0, 0, dtypes.float32) valid = [UOp.vectorize(a, a, a, a), UOp.vectorize(a, a, a, a, a, a, a, a)] - for shuf in valid: self.assertIs(self.isel_rewrite(shuf).arg.op, X86Ops.VBROADCASTSS) + for shuf in valid: self.assertEqual(self.isel_rewrite(shuf).arg, X86Ops.VBROADCASTSS) def test_vshufps(self): a, b = vector("a", dtypes.float32, 8), vector("b", dtypes.float32, 8) @@ -94,14 +94,14 @@ class TestIselX86(unittest.TestCase): UOp.vectorize(lane(a, 1), lane(a, 2), lane(a, 3), lane(a, 0)), UOp.vectorize(lane(a, 3), lane(a, 2), lane(a, 1), lane(a, 0), lane(a, 7), lane(a, 6), lane(a, 5), lane(a, 4)), UOp.vectorize(lane(a, 0), lane(a, 0), lane(b, 1), lane(b, 1), lane(a, 4), lane(a, 4), lane(b, 5), lane(b, 5))] - for shuf in valid: self.assertIs(self.isel_rewrite(shuf).arg.op, X86Ops.VSHUFPS) + for shuf in valid: self.assertEqual(self.isel_rewrite(shuf).arg, X86Ops.VSHUFPS) invalid = [UOp.vectorize(lane(a, 0), lane(b, 1), lane(a, 2), lane(b, 3)), UOp.vectorize(lane(a, 0), lane(a, 1), lane(b, 4), lane(b, 5)), UOp.vectorize(lane(a, 0), lane(a, 5), lane(b, 2), lane(b, 3)), UOp.vectorize(lane(a, 0), lane(a, 0), lane(a, 0), lane(a, 0), lane(a, 4), lane(a, 4), lane(a, 4), lane(a, 5)), UOp.vectorize(lane(a, 0), lane(a, 0), lane(b, 0), lane(b, 0), lane(a, 4), lane(a, 4), lane(b, 4), lane(a, 4))] - for shuf in invalid: self.assertIsNot(getattr(self.isel_rewrite(shuf).arg, "op", None), X86Ops.VSHUFPS) + for shuf in invalid: self.assertNotEqual(self.isel_rewrite(shuf).arg, X86Ops.VSHUFPS) def test_vshufpd(self): a, b = vector("a", dtypes.float64, 4), vector("b", dtypes.float64, 4) @@ -113,25 +113,25 @@ class TestIselX86(unittest.TestCase): UOp.vectorize(lane(a, 1), lane(b, 1)), UOp.vectorize(lane(a, 0), lane(b, 1), lane(a, 2), lane(b, 3)), UOp.vectorize(lane(a, 1), lane(a, 1), lane(a, 3), lane(a, 3))] - for shuf in valid: self.assertIs(self.isel_rewrite(shuf).arg.op, X86Ops.VSHUFPD) + for shuf in valid: self.assertEqual(self.isel_rewrite(shuf).arg, X86Ops.VSHUFPD) invalid = [UOp.vectorize(c, c, c, c), UOp.vectorize(lane(a, 0), lane(a, 1), lane(b, 2), lane(b, 3)), UOp.vectorize(lane(a, 2), lane(b, 3), lane(a, 2), lane(b, 3)), UOp.vectorize(lane(a, 0), lane(b, 1), lane(a, 0), lane(b, 1))] - for shuf in invalid: self.assertIsNot(getattr(self.isel_rewrite(shuf).arg, "op", None), X86Ops.VSHUFPD) + for shuf in invalid: self.assertNotEqual(self.isel_rewrite(shuf).arg, X86Ops.VSHUFPD) def test_vinsertps(self): a, b, c = vector("a", dtypes.float32, 4), vector("b", dtypes.float32, 4), vector("c", dtypes.float32, 4) d = UOp.variable("e", 0, 0, dtypes.float32) # moving 0th element to position 0 does nothing so only 1 vinsertps is generated n = self.isel_rewrite(UOp.vectorize(lane(a, 0), d)) - self.assertIs(n.arg.op, X86Ops.VINSERTPS) - self.assertIsNot(n.src[0].arg.op if n.src[0].op is Ops.INS else None, X86Ops.VINSERTPS) + self.assertEqual(n.arg, X86Ops.VINSERTPS) + self.assertNotEqual(n.src[0].arg if n.src[0].op is Ops.INS else None, X86Ops.VINSERTPS) valid = [UOp.vectorize(lane(a, 0), lane(b, 1), lane(a, 2), lane(b, 3)), UOp.vectorize(lane(a, 3), lane(b, 2), lane(c, 1), d)] - for shuf in valid: self.assertIs(self.isel_rewrite(shuf).arg.op, X86Ops.VINSERTPS) + for shuf in valid: self.assertEqual(self.isel_rewrite(shuf).arg, X86Ops.VINSERTPS) # complex address is [base + index*scale + displacement] def test_complex_address(self): diff --git a/tinygrad/codegen/late/regalloc.py b/tinygrad/codegen/late/regalloc.py index 3af670ed9b..584b23e134 100644 --- a/tinygrad/codegen/late/regalloc.py +++ b/tinygrad/codegen/late/regalloc.py @@ -132,6 +132,7 @@ def regalloc_rewrite(ctx:LinearScanRegallocContext, x:UOp): return nx, before + [nx] + after +# match every op so ctx.idx stays aligned with the linearized uop list pm_regalloc_rewrite = PatternMatcher([ (UPat(set(Ops), name="x"), regalloc_rewrite), ]) diff --git a/tinygrad/renderer/isa/x86.py b/tinygrad/renderer/isa/x86.py index 40b4b244de..5a4d4de397 100644 --- a/tinygrad/renderer/isa/x86.py +++ b/tinygrad/renderer/isa/x86.py @@ -136,8 +136,8 @@ def is_address(x:UOp) -> bool: if x.op is Ops.PARAM: return x.arg.addrspace is AddrSpace.GLOBAL if x.op is Ops.BUFFER: return True if x.op is Ops.INS: - if x.arg.op is X86Ops.LEA or x.arg.op is X86Ops.DEFINE and x.tag == (RSP,): return True - return x.dtype is dtypes.uint64 and x.arg.op in {X86Ops.MOV, X86Ops.CMOVB, X86Ops.CMOVL, X86Ops.CMOVE, X86Ops.CMOVNE} and \ + if x.arg == X86Ops.LEA or (x.arg == X86Ops.DEFINE and x.tag == (RSP,)): return True + return x.dtype is dtypes.uint64 and x.arg in {X86Ops.MOV, X86Ops.CMOVB, X86Ops.CMOVL, X86Ops.CMOVE, X86Ops.CMOVNE} and \ (x.shape == () or any(is_address(s) for s in x.src[:2])) if x.op in {Ops.INDEX, Ops.SHRINK, Ops.AFTER, Ops.NOOP} and x.src: return is_address(x.src[0]) return x.op is Ops.WHERE and is_address(x.src[1]) @@ -190,7 +190,7 @@ def gated_load(ctx, addr:UOp, alt:UOp, gate:UOp, x:UOp): # the selected address is a 64bit value, the AFTER orders the load after the scratch store and carries the element dtype for the encoder sel = gate.where(addr.replace(dtype=dtypes.uint64), local_idx) ptr = UOp(Ops.AFTER, addr.dtype, (sel, (local_idx if count == 1 else local).store(alt))) - return ptr.load(dtype=x.dtype, arg=count) + return ptr.load(dtype=x.dtype) def gated_store(addr:UOp, gate:UOp, val:UOp): local = UOp.placeholder((val.max_numel(),), addr.src[0].dtype, -1, AddrSpace.LOCAL) @@ -199,9 +199,6 @@ def gated_store(addr:UOp, gate:UOp, val:UOp): # legalize the new style graph for isel. NOTE: this runs after the spec is verified, some of these rewrites violate it pre_isel_matcher = PatternMatcher([ - # preserve coalesced load width when its address is selected before the load - (UPat(Ops.SHRINK, src=(UPat(), UPat(), UPat.cvar("c"))).load(allow_any_len=True, name="x"), - lambda x,c: x.replace(arg=c.arg) if x.arg is None else None), # zero extending scalar 32bit int is a noop (UPat.var("y", dtypes.uint32).cast(dtypes.int64s, name="x"), lambda y,x: x.replace(op=Ops.NOOP, arg=None) if y.max_numel() == 1 else None), # cast between signed and unsigned int is a noop @@ -255,11 +252,12 @@ def is_foldable(ctx:IselContext, x:UOp, s:UOp) -> bool: return len(ctx.uses.get( def base(x:UOp, i:int) -> UOp: return s.src[0] if (s:=x.src[i]).op is Ops.INDEX else s def const_arg(x:UOp) -> int|None: if x.op is Ops.CONST: return x.arg - return x.src[0].arg if x.op is Ops.INS and x.arg.op is X86Ops.MOVi and x.src[0].op is Ops.CONST else None + return x.src[0].arg if x.op is Ops.INS and x.arg == X86Ops.MOVi and x.src[0].op is Ops.CONST else None def lane(x:UOp, i:int) -> int: if (s:=x.src[i]).op is not Ops.INDEX: return 0 return unwrap(const_arg(s.src[1])) def to_int(dt:DType): return {dtypes.float16: dtypes.int16, dtypes.float32: dtypes.int32, dtypes.float64: dtypes.int64}[dt] +def nbytes(x:UOp) -> int: return x.dtype.itemsize * x.max_numel() def def_reg(dt:DType, reg:Register|None=None, shape:tuple=()) -> UOp: return UOp(Ops.INS, dt, arg=Insn(X86Ops.DEFINE, shape), tag=None if reg is None else (reg,)) def imm(dt:DType, v:int) -> UOp: return UOp.const(dt, truncate[dt](v)).rtag() @@ -269,14 +267,22 @@ def to_imm(c:UOp) -> UOp|None: if c.dtype is dtypes.uint64: return imm(dtypes.uint32, c.arg) if not c.overflows(dtypes.uint32) else None if c.dtype in dtypes.ints+(dtypes.bool,): return imm(c.dtype, c.arg) return None +# scalar/packed float opcode pairs: (ss, sd, ps, pd) +def fop(x:UOp, ss, sd, ps, pd, **kwargs) -> UOp: + scalar, dt = x.max_numel() == 1, x.dtype if x.dtype in (dtypes.float32, dtypes.float64) else x.src[0].dtype + return x.ins((ss if scalar else ps) if dt is dtypes.float32 else (sd if scalar else pd), **kwargs) def cmp(x:UOp) -> UOp: if x.src[0].dtype is dtypes.float32: return x.ins(X86Ops.VUCOMISS, dtype=dtypes.void) if x.src[0].dtype is dtypes.float64: return x.ins(X86Ops.VUCOMISD, dtype=dtypes.void) return x.ins(X86Ops.CMP, dtype=dtypes.void) if (i:=to_imm(x.src[1])) is None else x.ins(X86Ops.CMPi, dtype=dtypes.void, src=(x.src[0], i)) def vcmp(x:UOp) -> UOp: v = imm(dtypes.uint8, {Ops.CMPLT: 1, Ops.CMPNE: 4, Ops.CMPEQ: 0}[x.op]) - if x.src[0].dtype is dtypes.float32: return x.ins(X86Ops.VCMPSS if x.max_numel() == 1 else X86Ops.VCMPPS, src=x.src + (v,)) - return x.ins(X86Ops.VCMPSD if x.max_numel() == 1 else X86Ops.VCMPPD, src=x.src + (v,)) + return fop(x, X86Ops.VCMPSS, X86Ops.VCMPSD, X86Ops.VCMPPS, X86Ops.VCMPPD, src=x.src + (v,)) + +# size -> simd move opcodes +SIMD_LOAD = {2: X86Ops.VPINSRW, 4: X86Ops.VMOVSS, 8: X86Ops.VMOVSD, 16: X86Ops.VMOVUPS, 32: X86Ops.VMOVUPS} +SIMD_STORE = {2: X86Ops.VPEXTRW, 4: X86Ops.VMOVSSm, 8: X86Ops.VMOVSDm, 16: X86Ops.VMOVUPSm, 32: X86Ops.VMOVUPSm} +SIMD_COPY = {2: X86Ops.VMOVSS, 4: X86Ops.VMOVSS, 8: X86Ops.VMOVSD, 16: X86Ops.VMOVUPS, 32: X86Ops.VMOVUPS} # vshufps xmm2, xmm0, xmm1, imm # for 128 bit xmm2 selects its lower 2 32 bits from xmm0 and its upper 2 32 bits from xmm1 according to imm @@ -307,7 +313,7 @@ def vinsertps(x:UOp) -> UOp: s, v = base(x, i), lane(x, i) # moving the 0th element into the 0th position does nothing return s if i == v == 0 else x.ins(X86Ops.VINSERTPS, src=(ret, s, imm(dtypes.uint8, v << 6 | i << 4))) - return functools.reduce(_insert, range(len(x.src)), def_reg(x.dtype)) + return functools.reduce(_insert, range(len(x.src)), def_reg(x.dtype, shape=x.max_shape)) # vpinsq xmm2, xmm0, rax, imm # inserts element in rax into any position in xmm0, result is written to xmm2 according to imm @@ -359,38 +365,34 @@ def fold_address(x:UOp) -> tuple[UOp, UOp, UOp, UOp]: if idx.op is Ops.CONST: return (base, UOp(Ops.NOOP), _disp(idx.arg * scale), sz) return (base, _cast(idx), _disp(0), sz) +def simd_count(x:UOp) -> int: + # only treat as packed when the value fits a single xmm/ymm move; structural shapes are scalar + n = x.max_numel() + return n if n > 1 and x.dtype.itemsize * n in SIMD_LOAD else 1 + def lower_copy(x:UOp) -> UOp: - if is_address(x.src[0]) or x.max_numel() == 1 and x.dtype in dtypes.ints+(dtypes.bool,): return x.ins(X86Ops.MOV, shape=()) - size = x.dtype.itemsize * x.max_numel() - if size in (16, 32): return x.ins(X86Ops.VMOVUPS) - if size == 8: return x.ins(X86Ops.VMOVSD) - if size in (2, 4): return x.ins(X86Ops.VMOVSS) - raise RuntimeError(f"unsupported x86 copy size {size}") + if is_address(x.src[0]) or simd_count(x) == 1 and x.dtype in dtypes.ints+(dtypes.bool,): return x.ins(X86Ops.MOV, shape=()) + if (size:=x.dtype.itemsize * simd_count(x)) not in SIMD_COPY: raise RuntimeError(f"unsupported x86 copy size {size}") + return x.ins(SIMD_COPY[size]) def lower_load(ctx:IselContext|None, x:UOp, address:UOp) -> UOp|None: if ctx is not None and any(u.op is Ops.STACK for u in ctx.uses.get(x, ())): return None - count, src = x.arg if isinstance(x.arg, int) else 1, fold_address(address) + count, src = simd_count(x), fold_address(address) shape = () if count == 1 else (count,) if count == 1 and x.dtype in dtypes.ints+(dtypes.bool,): return x.ins(X86Ops.MOV, shape=shape, src=src) - size = x.dtype.itemsize * count + if (size:=x.dtype.itemsize * count) not in SIMD_LOAD: raise RuntimeError(f"unsupported x86 load size {size}") if size == 2: - return x.ins(X86Ops.VPINSRW, shape=shape, + return x.ins(SIMD_LOAD[size], shape=shape, src=(def_reg(x.dtype, x.tag if isinstance(x.tag, Register) else None, shape),) + src + (imm(dtypes.uint8, 0),)) - if size in (16, 32): return x.ins(X86Ops.VMOVUPS, shape=shape, src=src) - if size == 8: return x.ins(X86Ops.VMOVSD, shape=shape, src=src) - if size == 4: return x.ins(X86Ops.VMOVSS, shape=shape, src=src) - raise RuntimeError(f"unsupported x86 load size {size}") + return x.ins(SIMD_LOAD[size], shape=shape, src=src) def lower_store(x:UOp, address:UOp, value:UOp) -> UOp: - src = fold_address(address) - if value.max_numel() == 1 and value.dtype in dtypes.ints+(dtypes.bool,): + src, count = fold_address(address), simd_count(value) + if count == 1 and value.dtype in dtypes.ints+(dtypes.bool,): return x.ins(X86Ops.MOVm, src=src+(value,)) if (immv:=to_imm(value)) is None else x.ins(X86Ops.MOVi, src=src+(immv,)) - size = value.dtype.itemsize * value.max_numel() - if size == 2: return x.ins(X86Ops.VPEXTRW, src=src+(value, imm(dtypes.uint8, 0))) - if size in (16, 32): return x.ins(X86Ops.VMOVUPSm, src=src+(value,)) - if size == 8: return x.ins(X86Ops.VMOVSDm, src=src+(value,)) - if size == 4: return x.ins(X86Ops.VMOVSSm, src=src+(value,)) - raise RuntimeError(f"unsupported x86 store size {size}") + if (size:=value.dtype.itemsize * count) not in SIMD_STORE: raise RuntimeError(f"unsupported x86 store size {size}") + if size == 2: return x.ins(SIMD_STORE[size], src=src+(value, imm(dtypes.uint8, 0))) + return x.ins(SIMD_STORE[size], src=src+(value,)) def select_index(ctx, x:UOp) -> UOp|None: if not is_address(x.src[0]): return None @@ -417,9 +419,9 @@ def abi(ctx:IselContext, x:UOp) -> UOp|None: def alloc_vregs(ctx:IselContext, x:UOp) -> UOp|None: # register placeholders with real registers - if x.op is Ops.INS and x.arg.op is X86Ops.DEFINE and x.tag is not None: return None + if x.op is Ops.INS and x.arg == X86Ops.DEFINE and x.tag is not None: return None # this is an immediate - if x.op is Ops.INS and x.arg.op is X86Ops.FRAME_INDEX: return None + if x.op is Ops.INS and x.arg == X86Ops.FRAME_INDEX: return None # no register definition if x.dtype is dtypes.void: return None # already allocated vregs @@ -429,7 +431,7 @@ def alloc_vregs(ctx:IselContext, x:UOp) -> UOp|None: if isinstance(x.tag, tuple): defs = [ctx.vreg(x.tag)] elif is_address(x): defs = [ctx.vreg(WGPR)] elif x.max_numel() > 1 or x.dtype in dtypes.floats: - if x.dtype.itemsize * x.max_numel() > 32: raise RuntimeError(f"x86 only supports SIMD values up to 32 bytes, got {x.dtype}{x.shape}") + if nbytes(x) > 32: raise RuntimeError(f"x86 only supports SIMD values up to 32 bytes, got {x.dtype}{x.shape}") defs = [ctx.vreg(XMM)] elif x.dtype in dtypes.ints+(dtypes.bool,): defs = [ctx.vreg(WGPR)] # TODO: add this once the scheduler can track register pressure @@ -455,7 +457,7 @@ isel_matcher = PatternMatcher([ (UPat(Ops.SINK, name="x"), lambda x: x.replace(src=(x.ins(X86Ops.RET, src=x.src + tuple(def_reg(dtypes.uint64, r) if r in GPR else def_reg(dtypes.float64, r, (2,)) for r in CALLEE_SAVED)),)) \ - if not x.src or x.src[0].op is not Ops.INS or x.src[0].arg.op is not X86Ops.RET else None), + if not x.src or x.src[0].op is not Ops.INS or x.src[0].arg != X86Ops.RET else None), # function abi constraints (UPat((Ops.PARAM, Ops.SPECIAL), name="x"), abi), # constants that can't be immediates, move them to registers @@ -464,14 +466,10 @@ isel_matcher = PatternMatcher([ (UPat.cvar("x", dtypes.floats), lambda x: UOp.const(dt:=to_int(x.dtype), struct.unpack(dt.fmt, struct.pack(x.dtype.fmt, x.arg))[0]).bitcast(x.dtype) if not x.tag else None), # TODO: these should use a.maximum(b) / a.minimum(b) - ((UPat.var("a") < UPat.var("b")).where(UPat.var("b", dtypes.float32), UPat.var("a")), lambda a,b: - a.ins(X86Ops.VMAXSS if a.max_numel() == 1 else X86Ops.VMAXPS, src=(a, b))), - ((UPat.var("a") < UPat.var("b")).where(UPat.var("b", dtypes.float64), UPat.var("a")), lambda a,b: - a.ins(X86Ops.VMAXSD if a.max_numel() == 1 else X86Ops.VMAXPD, src=(a, b))), - ((UPat.var("a") < UPat.var("b")).where(UPat.var("a", dtypes.float32), UPat.var("b")), lambda a,b: - a.ins(X86Ops.VMINSS if a.max_numel() == 1 else X86Ops.VMINPS, src=(a, b))), - ((UPat.var("a") < UPat.var("b")).where(UPat.var("a", dtypes.float64), UPat.var("b")), lambda a,b: - a.ins(X86Ops.VMINSD if a.max_numel() == 1 else X86Ops.VMINPD, src=(a, b))), + ((UPat.var("a") < UPat.var("b")).where(UPat.var("b", (dtypes.float32, dtypes.float64)), UPat.var("a")), lambda a,b: + fop(a, X86Ops.VMAXSS, X86Ops.VMAXSD, X86Ops.VMAXPS, X86Ops.VMAXPD, src=(a, b))), + ((UPat.var("a") < UPat.var("b")).where(UPat.var("a", (dtypes.float32, dtypes.float64)), UPat.var("b")), lambda a,b: + fop(a, X86Ops.VMINSS, X86Ops.VMINSD, X86Ops.VMINPS, X86Ops.VMINPD, src=(a, b))), # conditional moves that use masks NOTE: these currently assume a mask producing cmp exists (UPat.var("m").where(UPat.var("a", dtypes.ints), UPat.var("b")), lambda m,a,b: a.ins(X86Ops.VPBLENDVB, src=(b, a, m.replace(dtype=m.src[0].dtype))) if a.max_numel() > 1 and not is_address(a) else None), @@ -510,12 +508,11 @@ isel_matcher = PatternMatcher([ (UPat(Ops.CMPLT, src=(UPat.var("a", dtypes.int32s), UPat.var("b")), name="x"), lambda a,b,x: x.ins(X86Ops.VPCMPGTD, src=(b, a))), (UPat(Ops.CMPLT, src=(UPat.var("a", dtypes.int64s), UPat.var("b")), name="x"), lambda a,b,x: x.ins(X86Ops.VPCMPGTQ, src=(b, a))), # float unary - (UPat.var("y", dtypes.float32).sqrt().named("x"), lambda y,x: x.ins(X86Ops.VSQRTSS, src=(y, y)) if x.max_numel() == 1 else x.ins(X86Ops.VSQRTPS)), - (UPat.var("y", dtypes.float64).sqrt().named("x"), lambda y,x: x.ins(X86Ops.VSQRTSD, src=(y, y)) if x.max_numel() == 1 else x.ins(X86Ops.VSQRTPD)), - (UPat.var("y", dtypes.float32).trunc().named("x"), lambda y,x: - x.ins(X86Ops.VROUNDSS, src=(y, y, imm(dtypes.uint8, 3))) if x.max_numel() == 1 else x.ins(X86Ops.VROUNDPS, src=(y, imm(dtypes.uint8, 3)))), - (UPat.var("y", dtypes.float64).trunc().named("x"), lambda y,x: - x.ins(X86Ops.VROUNDSD, src=(y, y, imm(dtypes.uint8, 3))) if x.max_numel() == 1 else x.ins(X86Ops.VROUNDPD, src=(y, imm(dtypes.uint8, 3)))), + (UPat.var("y", (dtypes.float32, dtypes.float64)).sqrt().named("x"), lambda y,x: + fop(x, X86Ops.VSQRTSS, X86Ops.VSQRTSD, X86Ops.VSQRTPS, X86Ops.VSQRTPD, src=(y, y) if x.max_numel() == 1 else (y,))), + (UPat.var("y", (dtypes.float32, dtypes.float64)).trunc().named("x"), lambda y,x: + fop(x, X86Ops.VROUNDSS, X86Ops.VROUNDSD, X86Ops.VROUNDPS, X86Ops.VROUNDPD, + src=((y, y, imm(dtypes.uint8, 3)) if x.max_numel() == 1 else (y, imm(dtypes.uint8, 3))))), # shufles (UPat.var("y", dtypes.float32).broadcast(name="x"), lambda y,x: x.ins(X86Ops.VBROADCASTSS, src=(y,))), # for float16 we route the srcs through gprs unless we can fold them, this is suboptimal for values in xmms, in that case we want vpunpcklwd @@ -532,10 +529,8 @@ isel_matcher = PatternMatcher([ X86Ops.VPSRLDQ if y.dtype in dtypes.floats else {1:X86Ops.VPEXTRB, 2:X86Ops.VPEXTRW, 4:X86Ops.VPEXTRD, 8:X86Ops.VPEXTRQ}[y.dtype.itemsize], shape=(), src=(y, imm(dtypes.uint8, ci * x.dtype.itemsize if y.dtype in dtypes.floats else ci)))), # fused multiply add - ((UPat(Ops.MUL, dtypes.float32, name="a") + UPat.var("b")).named("c"), lambda ctx,a,b,c: - a.ins(X86Ops.VFMADD213SS if a.max_numel() == 1 else X86Ops.VFMADD213PS, src=(*a.src, b)) if is_foldable(ctx, c, a) else None), - ((UPat(Ops.MUL, dtypes.float64, name="a") + UPat.var("b")).named("c"), lambda ctx,a,b,c: - a.ins(X86Ops.VFMADD213SD if a.max_numel() == 1 else X86Ops.VFMADD213PD, src=(*a.src, b)) if is_foldable(ctx, c, a) else None), + ((UPat(Ops.MUL, (dtypes.float32, dtypes.float64), name="a") + UPat.var("b")).named("c"), lambda ctx,a,b,c: + fop(a, X86Ops.VFMADD213SS, X86Ops.VFMADD213SD, X86Ops.VFMADD213PS, X86Ops.VFMADD213PD, src=(*a.src, b)) if is_foldable(ctx, c, a) else None), # packed bitwise ((UPat() & UPat()).named("x"), lambda x: x.ins(X86Ops.VPAND) if x.max_numel() > 1 else None), ((UPat() | UPat()).named("x"), lambda x: x.ins(X86Ops.VPOR) if x.max_numel() > 1 else None), @@ -579,14 +574,14 @@ isel_matcher = PatternMatcher([ (UPat.var("a", dtypes.ints+(dtypes.bool,)) ^ UPat.var("b"), lambda a,b: a.ins(X86Ops.XOR, src=(a, b))), (UPat(Ops.SUB, dtypes.ints, (UPat.var("a"), UPat.var("b"))), lambda a,b: a.ins(X86Ops.SUB, src=(a, b))), # float binary - ((UPat(dtype=dtypes.float32) + UPat()).named("x"), lambda x: x.ins(X86Ops.VADDSS if x.max_numel() == 1 else X86Ops.VADDPS)), - ((UPat(dtype=dtypes.float64) + UPat()).named("x"), lambda x: x.ins(X86Ops.VADDSD if x.max_numel() == 1 else X86Ops.VADDPD)), - ((UPat(dtype=dtypes.float32) * UPat()).named("x"), lambda x: x.ins(X86Ops.VMULSS if x.max_numel() == 1 else X86Ops.VMULPS)), - ((UPat(dtype=dtypes.float64) * UPat()).named("x"), lambda x: x.ins(X86Ops.VMULSD if x.max_numel() == 1 else X86Ops.VMULPD)), - (UPat(Ops.SUB, dtypes.float32, name="x"), lambda x: x.ins(X86Ops.VSUBSS if x.max_numel() == 1 else X86Ops.VSUBPS)), - (UPat(Ops.SUB, dtypes.float64, name="x"), lambda x: x.ins(X86Ops.VSUBSD if x.max_numel() == 1 else X86Ops.VSUBPD)), - (UPat(Ops.FDIV, dtypes.float32, name="x"), lambda x: x.ins(X86Ops.VDIVSS if x.max_numel() == 1 else X86Ops.VDIVPS)), - (UPat(Ops.FDIV, dtypes.float64, name="x"), lambda x: x.ins(X86Ops.VDIVSD if x.max_numel() == 1 else X86Ops.VDIVPD)), + ((UPat(dtype=(dtypes.float32, dtypes.float64)) + UPat()).named("x"), + lambda x: fop(x, X86Ops.VADDSS, X86Ops.VADDSD, X86Ops.VADDPS, X86Ops.VADDPD)), + ((UPat(dtype=(dtypes.float32, dtypes.float64)) * UPat()).named("x"), + lambda x: fop(x, X86Ops.VMULSS, X86Ops.VMULSD, X86Ops.VMULPS, X86Ops.VMULPD)), + (UPat(Ops.SUB, (dtypes.float32, dtypes.float64), name="x"), + lambda x: fop(x, X86Ops.VSUBSS, X86Ops.VSUBSD, X86Ops.VSUBPS, X86Ops.VSUBPD)), + (UPat(Ops.FDIV, (dtypes.float32, dtypes.float64), name="x"), + lambda x: fop(x, X86Ops.VDIVSS, X86Ops.VDIVSD, X86Ops.VDIVPS, X86Ops.VDIVPD)), # casts (UPat(dtype=dtypes.int32).cast(dtypes.float32, name="x"), lambda x: x.ins(X86Ops.VCVTDQ2PS) if x.max_numel() > 1 else None), (UPat(dtype=dtypes.int32).cast(dtypes.float64, name="x"), lambda x: x.ins(X86Ops.VCVTDQ2PD) if x.max_numel() > 1 else None), @@ -635,11 +630,11 @@ isel_matcher = PatternMatcher([ # **** X86Op -> X86Op **** # fold loads into X86Ops that allow it, if beneficial (UPat(Ops.INS, src=(UPat(Ops.LOAD, src=(UPat(name="a"),), name="y"),), allow_any_len=True, name="x"), lambda ctx,y,a,x: - x.replace(src=fold_address(a) + x.src[1:]) if x.arg.op in X86GroupOp.ReadMem1st and is_foldable(ctx, x, y) else None), + x.replace(src=fold_address(a) + x.src[1:]) if x.arg in X86GroupOp.ReadMem1st and is_foldable(ctx, x, y) else None), (UPat(Ops.INS, src=(UPat(), UPat(Ops.LOAD, src=(UPat(name="a"),), name="y")), allow_any_len=True, name="x"), lambda ctx,y,a,x: - x.replace(src=x.src[:1] + fold_address(a) + x.src[2:]) if x.arg.op in X86GroupOp.ReadMem2nd and is_foldable(ctx, x, y) else None), + x.replace(src=x.src[:1] + fold_address(a) + x.src[2:]) if x.arg in X86GroupOp.ReadMem2nd and is_foldable(ctx, x, y) else None), (UPat(Ops.INS, src=(UPat(), UPat(), UPat(Ops.LOAD, src=(UPat(name="a"),), name="y")), allow_any_len=True, name="x"), lambda ctx,y,a,x: - x.replace(src=x.src[:2] + fold_address(a) + x.src[3:]) if x.arg.op in X86GroupOp.ReadMem3rd and is_foldable(ctx, x, y) else None), + x.replace(src=x.src[:2] + fold_address(a) + x.src[3:]) if x.arg in X86GroupOp.ReadMem3rd and is_foldable(ctx, x, y) else None), # allocate virtual registers (UPat((Ops.INS, Ops.BUFFER), name="x"), alloc_vregs), ]) @@ -649,8 +644,8 @@ isel_matcher = PatternMatcher([ # so we rematerialize. This is different from rematerialization you might want to do in regalloc because it is not optional, # regalloc shouldn't rematerialize if a src of the instruction is dead, but here you need to as there's no fallback load from stack def flag_rematerialize(ctx:PreRegAllocContext, x:UOp): - flag_def = x if x.op in (Ops.RANGE, Ops.END) or x.op is Ops.INS and x.arg.op in X86GroupOp.WriteFlags \ - else x.src[-1] if x.op is Ops.INS and x.arg.op in X86GroupOp.ReadFlags else None + flag_def = x if x.op in (Ops.RANGE, Ops.END) or x.op is Ops.INS and x.arg in X86GroupOp.WriteFlags \ + else x.src[-1] if x.op is Ops.INS and x.arg in X86GroupOp.ReadFlags else None if flag_def is None: return None if ctx.lock is not None and ctx.lock is not flag_def: ctx.clobbered.add(ctx.lock) ctx.lock = flag_def @@ -676,7 +671,7 @@ def lower_range(ctx, x:UOp) -> tuple[UOp, list[UOp]]: # final rewrite to match the isa spec post_regalloc_matcher = PatternMatcher([ # rewrite FRAME_INDEX to IMM now that the stack size is known - (UPat(Ops.INS, name="x"), lambda ctx,x: (nx:=x.const_like(ctx.stack_size + x.tag), [nx]) if x.arg.op is X86Ops.FRAME_INDEX else None), + (UPat(Ops.INS, name="x"), lambda ctx,x: (nx:=x.const_like(ctx.stack_size + x.tag), [nx]) if x.arg == X86Ops.FRAME_INDEX else None), # rewrite RANGE to ACC = 0 -> LABEL -> JUMP if ACC >= loop bound (UPat(Ops.RANGE, name="x"), lambda ctx,x: lower_range(ctx, x)), # rewrite END to ACC + 1 -> JUMP -> LABEL, also add the out of loop JUMP to the src so this becomes the jump target @@ -685,7 +680,7 @@ post_regalloc_matcher = PatternMatcher([ UOp(Ops.INS, arg=Insn(X86Ops.LABEL), tag=f".LOOP_OUT_{ctx.loop_label[x.src[1]]}")])), # rewrite two address instructions to two address form, if reused src wasn't coalesced insert a move (UPat(Ops.INS, name="x"), lambda ctx,x: (nx:=x.replace(src=x.src[1:]), - [ctx.ren.copy(x.src[0], greg(x)), nx] if greg(x) != greg(x.src[0]) else [nx]) if x.arg.op in X86GroupOp.TwoAddress else None), + [ctx.ren.copy(x.src[0], greg(x)), nx] if greg(x) != greg(x.src[0]) else [nx]) if x.arg in X86GroupOp.TwoAddress else None), ]) # ***** X86 instruction encoding ***** @@ -699,8 +694,8 @@ def encode(x:UOp, opc:int, reg:int|None=None, pp:int=0, sel:int=0, we:int=0) -> rm = cast(Register, greg(rm_uop)).index idx = cast(Register, greg(idx_uop)).index if idx_uop is not None and greg(idx_uop) is not None else 4 # for a memory operand the rm size is the element size from the address, otherwise it's the size of the value in the register - rm_sz = sz_uop.arg if sz_uop is not None else 8 if is_address(rm_uop) else rm_uop.dtype.itemsize * rm_uop.max_numel() - reg_sz = (8 if is_address(reg_uop) else reg_uop.dtype.itemsize * reg_uop.max_numel()) if reg_uop is not None else 0 + rm_sz = sz_uop.arg if sz_uop is not None else 8 if is_address(rm_uop) else nbytes(rm_uop) + reg_sz = (8 if is_address(reg_uop) else nbytes(reg_uop)) if reg_uop is not None else 0 sz = reg_sz or rm_sz # encode instruction @@ -722,7 +717,7 @@ def encode(x:UOp, opc:int, reg:int|None=None, pp:int=0, sel:int=0, we:int=0) -> # REX byte is required when 64 bit or an extended reg is used (index 8 - 15) or lower 8 bits of (rsp, rbp, rsi, rdi) are accessed if w | r | _x | b | (reg_sz == 1 & reg >> 2) | (rm_sz == 1 & rm >> 2): inst += bytes([0b0100 << 4 | w << 3 | r << 2 | _x << 1 | b]) # legacy 8bit opcode is 1 less than 16-64bit variants - if (rm_sz == 1 or reg_sz == 1) and x.arg.op not in X86GroupOp.ReadFlags | {X86Ops.LEA}: opc -= 1 + if (rm_sz == 1 or reg_sz == 1) and x.arg not in X86GroupOp.ReadFlags | {X86Ops.LEA}: opc -= 1 # OPCODE byte inst += opc.to_bytes((opc.bit_length() + 7) // 8, 'big') # MODRM byte @@ -760,18 +755,18 @@ def encode(x:UOp, opc:int, reg:int|None=None, pp:int=0, sel:int=0, we:int=0) -> # get the encoding structure of the uop # when a uop writes to memory it takes the form of a store, dtype is void, no definition address:tuple[UOp|None, ...] - if x.arg.op in X86GroupOp.WriteMem: + if x.arg in X86GroupOp.WriteMem: if len(x.src) > 4: address, rest = x.src[:4], x.src[4:] else: address, rest = (x, None, None, None), x.src return _encode(rest[0], *address, *(None, *rest[1:])) if reg is None else _encode(None, *address, *(None, *rest[:1])) - if x.arg.op in X86GroupOp.Rm1st: + if x.arg in X86GroupOp.Rm1st: if len(x.src) > 3: address, rest = x.src[:4], x.src[4:] else: address, rest = (x.src[0], None, None, None), x.src[1:] imm_uop = rest[:1] if rest and rest[0].op is Ops.CONST else (None,) return _encode(x, *address, *(None, *imm_uop)) if reg is None else _encode(None, *address, *(x if sel else None, *imm_uop)) - if x.arg.op in X86GroupOp.Rm2nd: + if x.arg in X86GroupOp.Rm2nd: if len(x.src) > 4: address, rest = x.src[1:5], x.src[:1] + x.src[5:] else: address, rest = (x.src[1], None, None, None), x.src[:1] + x.src[2:] # cmp/vucomiss reg, rm don't define a new register @@ -908,7 +903,7 @@ class X86Renderer(ISARenderer): super().__init__(target) from tinygrad.runtime.support.compiler_cpu import X86Compiler self.compiler = X86Compiler() - def is_two_address(self, x:UOp) -> bool: return x.op is Ops.INS and x.arg.op in X86GroupOp.TwoAddress + def is_two_address(self, x:UOp) -> bool: return x.op is Ops.INS and x.arg in X86GroupOp.TwoAddress def stack_pointer(self) -> UOp: return def_reg(dtypes.uint64, RSP) # the value of a BUFFER is its address, it moves through registers and the stack as a 64bit int def copy(self, x:UOp, reg:Register): @@ -929,32 +924,34 @@ class X86Renderer(ISARenderer): def fill(self, disp:UOp, x:UOp, reg:Register) -> UOp: if is_address(x): return UOp(Ops.INS, dtypes.uint64, arg=Insn(X86Ops.MOV), src=fold_address(self.stack_pointer().index(disp)), tag=reg) - address, sz = fold_address(self.stack_pointer().index(disp)), x.dtype.itemsize*x.max_numel() - if x.max_numel() == 1 and x.dtype in dtypes.ints+(dtypes.bool,): return UOp(Ops.INS, x.dtype, address, Insn(X86Ops.MOV, x.shape), reg) - if sz == 2: return UOp(Ops.INS, x.dtype, (def_reg(x.dtype, reg, x.max_shape),) + address + (imm(dtypes.uint8, 0),), - Insn(X86Ops.VPINSRW, x.shape), reg) - return UOp(Ops.INS, x.dtype, address, Insn({4:X86Ops.VMOVSS, 8:X86Ops.VMOVSD, 16:X86Ops.VMOVUPS, 32:X86Ops.VMOVUPS}[sz], x.shape), reg) + src, shape = fold_address(self.stack_pointer().index(disp)), () if x.max_numel() == 1 else x.max_shape + if x.max_numel() == 1 and x.dtype in dtypes.ints+(dtypes.bool,): + return UOp(Ops.INS, x.dtype, src, Insn(X86Ops.MOV, shape), reg) + if (size:=nbytes(x)) not in SIMD_LOAD: raise RuntimeError(f"unsupported x86 fill size {size}") + if size == 2: + return UOp(Ops.INS, x.dtype, (def_reg(x.dtype, reg, shape),) + src + (imm(dtypes.uint8, 0),), Insn(SIMD_LOAD[size], shape), reg) + return UOp(Ops.INS, x.dtype, src, Insn(SIMD_LOAD[size], shape), reg) def asm_str(self, uops:list[UOp], function_name:str) -> str: - def _format_op(x:UOp) -> str: return f" {(o[7:-1] if (o:=str(x.arg.op))[-1] in ('i', 'm') else o[7:]).lower():7s}" + def _format_op(x:UOp) -> str: return f" {(o[7:-1] if (o:=str(x.arg))[-1] in ('i', 'm') else o[7:]).lower():7s}" def _format_operands(x:UOp) -> str: def _format(src:tuple[UOp, ...]) -> list[str]: - return [str(s.arg) if s.op is Ops.CONST else reg_strs[o].get(8 if is_address(s) else s.dtype.itemsize*s.max_numel(), o) if \ + return [str(s.arg) if s.op is Ops.CONST else reg_strs[o].get(8 if is_address(s) else nbytes(s), o) if \ (o:=str(greg(s))) in reg_strs else o for s in src if greg(s) is not None] def _mem_adress(base:UOp, idx:UOp, disp:UOp, sz:UOp) -> list[str]: return [f"[{greg(base)}" + (f" + {greg(idx)}*{sz.arg}" if greg(idx) else "") + (f" + {disp.arg}" if disp.arg else "") + "]"] - if len(x.src) > 4 and x.arg.op in X86GroupOp.WriteMem: ret = _mem_adress(*x.src[:4]) + _format(x.src[4:]) - elif len(x.src) > 3 and x.arg.op in X86GroupOp.Rm1st: ret = _format((x,)) + _mem_adress(*x.src[:4]) + _format(x.src[4:]) - elif len(x.src) > 4 and x.arg.op in X86GroupOp.Rm2nd: ret = _format((x, x.src[0])) + _mem_adress(*x.src[1:5]) + _format(x.src[5:]) + if len(x.src) > 4 and x.arg in X86GroupOp.WriteMem: ret = _mem_adress(*x.src[:4]) + _format(x.src[4:]) + elif len(x.src) > 3 and x.arg in X86GroupOp.Rm1st: ret = _format((x,)) + _mem_adress(*x.src[:4]) + _format(x.src[4:]) + elif len(x.src) > 4 and x.arg in X86GroupOp.Rm2nd: ret = _format((x, x.src[0])) + _mem_adress(*x.src[1:5]) + _format(x.src[5:]) else: ret = _format((x,) + x.src) return ", ".join(ret) asm = [f".{function_name}:"] for u in uops: - if u.op is not Ops.INS or u.arg.op is X86Ops.DEFINE: continue - if u.arg.op is X86Ops.LABEL: asm.append(f"{str(u.tag)}:") - elif u.arg.op is X86Ops.RET: asm.append(_format_op(u)) + if u.op is not Ops.INS or u.arg == X86Ops.DEFINE: continue + if u.arg == X86Ops.LABEL: asm.append(f"{str(u.tag)}:") + elif u.arg == X86Ops.RET: asm.append(_format_op(u)) else: asm.append(_format_op(u) + " " + _format_operands(u)) return "\n".join(asm) @@ -963,14 +960,14 @@ class X86Renderer(ISARenderer): jumps: dict[UOp, int] = {} binary = bytearray() for u in uops: - if u.op is not Ops.INS or u.arg.op is X86Ops.DEFINE: continue - if u.arg.op is X86Ops.LABEL: + if u.op is not Ops.INS or u.arg == X86Ops.DEFINE: continue + if u.arg == X86Ops.LABEL: targets[u.tag] = len(binary) continue - if (op:=u.arg.op) not in encodings or (l:=encodings[op](u)) is None: - raise RuntimeError(f"failed to encode {op} with {u.dtype} srcs {[x.dtype for x in u.src]}") + if u.arg not in encodings or (l:=encodings[u.arg](u)) is None: + raise RuntimeError(f"failed to encode {u.arg} with {u.dtype} srcs {[x.dtype for x in u.src]}") binary.extend(l) - if op in (X86Ops.JL, X86Ops.JB, X86Ops.JE, X86Ops.JNE, X86Ops.JGE, X86Ops.JMP): jumps[u] = len(binary) + if u.arg in (X86Ops.JL, X86Ops.JB, X86Ops.JE, X86Ops.JNE, X86Ops.JGE, X86Ops.JMP): jumps[u] = len(binary) # fixup jump targets now that encoding size is known for u in uops: if (t:=jumps.get(u)) is not None: binary[t-4:t] = (targets[u.tag] - t).to_bytes(4, 'little', signed=True) diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index 59a76f8951..4bd88ad88a 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -37,6 +37,9 @@ class ParamArg: class Insn: op: Any shape: tuple[sint, ...] = () + def __eq__(self, other): return (self.op, self.shape) == (other.op, other.shape) if isinstance(other, Insn) else self.op == other + def __hash__(self): return hash(self.op) + def __str__(self): return str(self.op) axis_letters = {AxisType.GLOBAL: "g", AxisType.THREAD: "t", AxisType.LOCAL: "l", AxisType.WARP: "w", AxisType.LOOP: "L", AxisType.UPCAST: "u", AxisType.GROUP_REDUCE: "G", AxisType.REDUCE: "R", AxisType.UNROLL: "r"} axis_colors = {AxisType.GLOBAL: "blue", AxisType.THREAD: "BLUE", AxisType.LOCAL: "cyan", AxisType.WARP: "CYAN", AxisType.LOOP: "WHITE",