forked from tinygrad/tinygrad
func arg order independent from op value
This commit is contained in:
@@ -18,7 +18,8 @@ class IselContext:
|
||||
self.uses = sink.get_consumer_map()
|
||||
self.reg_n = itertools.count()
|
||||
self.stack_size = 0
|
||||
self.func_args = sorted([u for u in self.uses if u.op in (Ops.DEFINE_GLOBAL, Ops.DEFINE_VAR, Ops.SPECIAL)], key=lambda k: (k.op, k.arg))
|
||||
arg_order = {Ops.DEFINE_GLOBAL: 0, Ops.DEFINE_VAR: 1, Ops.SPECIAL: 2}
|
||||
self.func_args = sorted([u for u in self.uses if u.op in arg_order], key=lambda k: (arg_order[k.op], k.arg))
|
||||
|
||||
def inc_stack(self, amt:int):
|
||||
ret = self.stack_size
|
||||
|
||||
@@ -213,16 +213,16 @@ def cmp(x:UOp) -> UOp:
|
||||
return UOp(X86Ops.CMP, src=x.src) if (i:=to_imm(x.src[1])) is None else UOp(X86Ops.CMPi, src=(x.src[0], i))
|
||||
def def_reg(dt:DType, reg:Register|None=None) -> UOp: return UOp(X86Ops.DEFINE_REG, dt, arg=reg)
|
||||
|
||||
# vshufps xmm2, xmm0, xmm1
|
||||
# xmm2 selects its lower 2 32 bits from xmm0 and its upper 2 32 bits from xmm1
|
||||
# vshufps xmm2, xmm0, xmm1, imm
|
||||
# xmm2 selects its lower 2 32 bits from xmm0 and its upper 2 32 bits from xmm1 according to imm
|
||||
def vshufps(x:UOp) -> UOp:
|
||||
def _in(i:int) -> UOp: return s.src[0] if (s:=x.src[i]).op is Ops.GEP else s
|
||||
if len(x.src) != 4 or _in(0) is not _in(1) or _in(2) is not _in(3): return None
|
||||
return UOp(X86Ops.VSHUFPS, x.dtype, (_in(0), _in(2),
|
||||
imm(dtypes.uint8, sum(s.arg[0] << 2*i if s.op is Ops.GEP else 0 for i,s in enumerate(x.src)))))
|
||||
|
||||
# vinsertps xmm2, xmm0, xmm1
|
||||
# inserts any 32 bit element in xmm1 into any position in xmm0, result is written to xmm2
|
||||
# vinsertps xmm2, xmm0, xmm1, imm
|
||||
# inserts any 32 bit element in xmm1 into any position in xmm0 according to immm, result is written to xmm2
|
||||
# this is the fallback slow case for when you can't match more a powerful shuffle
|
||||
def vinsertps(x:UOp) -> UOp:
|
||||
def _insert(ret:UOp, i:int) -> UOp:
|
||||
@@ -232,8 +232,8 @@ def vinsertps(x:UOp) -> UOp:
|
||||
return s if i == v == 0 else UOp(X86Ops.VINSERTPS, x.dtype, (ret, s, imm(dtypes.uint8, v << 6 | i << 4)))
|
||||
return functools.reduce(_insert, range(len(x.src)), def_reg(x.dtype))
|
||||
|
||||
# vpinsq xmm2, xmm0, rax
|
||||
# inserts element in rax into any position in xmm0, result is written to xmm2
|
||||
# vpinsq xmm2, xmm0, rax, imm
|
||||
# inserts element in rax into any position in xmm0, result is written to xmm2 according to imm
|
||||
def vpins(x:UOp) -> UOp:
|
||||
op = {1: X86Ops.VPINSRB, 2: X86Ops.VPINSRW, 4: X86Ops.VPINSRD, 8: X86Ops.VPINSRQ}[x.dtype.scalar().itemsize]
|
||||
return functools.reduce(lambda ret,i: UOp(op, x.dtype, (ret, x.src[i], imm(dtypes.uint8, i))), range(len(x.src)), def_reg(x.dtype))
|
||||
|
||||
Reference in New Issue
Block a user