forked from tinygrad/tinygrad
remove UOp cast and bitcast override [PR] (#17011)
This commit is contained in:
@@ -803,7 +803,7 @@ def _compile_sopp(inst: ir3.SOPP | ir4.SOPP, ctx: _Ctx) -> UOp:
|
||||
pcode = get_pcode(inst.op)
|
||||
pc_bytes = ctx.rpc() # PC is already 64-bit byte address
|
||||
vcc, exec_val = ctx.rmask(_c(VCC_LO.offset)), ctx.rexec()
|
||||
srcs = {'PC': pc_bytes.cast(dtypes.int64), 'SIMM16': simm16, 'SCC': ctx.rsgpr_dyn(_c(SCC.offset)), 'VCC': vcc,
|
||||
srcs: dict[str, UOp|int] = {'PC': pc_bytes.cast(dtypes.int64), 'SIMM16': simm16, 'SCC': ctx.rsgpr_dyn(_c(SCC.offset)), 'VCC': vcc,
|
||||
'VCCZ': vcc.eq(UOp.const(vcc.dtype, 0)).cast(dtypes.uint32),
|
||||
'EXECZ': exec_val.eq(UOp.const(exec_val.dtype, 0)).cast(dtypes.uint32)}
|
||||
for dest, val in parse_pcode(pcode, srcs)[1]:
|
||||
@@ -858,7 +858,7 @@ def _compile_sop(inst: ir3.SOP1|ir3.SOP2|ir3.SOPC|ir3.SOPK|ir4.SOP1|ir4.SOP2|ir4
|
||||
if isinstance(inst, ir4.SOPK): s0 = simm16
|
||||
elif isinstance(inst, irc.SOPK) and 'CMPK' not in op_name and 'SETREG' not in op_name: s0 = simm16_sext
|
||||
else: s0 = ctx.rsgpr_dyn(sdst_off)
|
||||
srcs = {'S0': s0, 'S1': simm16_sext, 'SIMM16': simm16_sext, 'D0': ctx.rsgpr_dyn(sdst_off)}
|
||||
srcs: dict[str, UOp|int] = {'S0': s0, 'S1': simm16_sext, 'SIMM16': simm16_sext, 'D0': ctx.rsgpr_dyn(sdst_off)}
|
||||
dst_off, dst_size = sdst_off, 1
|
||||
# S_GETREG_B32: extract bits from HW register. Handle as special case since HW_REGISTERS is not a normal variable.
|
||||
# HW register values are stored at SGPR[SGPR_COUNT-16 + hwRegId] by _init_wave.
|
||||
|
||||
@@ -688,10 +688,10 @@ class Parser:
|
||||
return _extract_bits(base, hi, lo)
|
||||
# Dynamic bit slice: (base >> lo) & ((1 << (hi - lo + 1)) - 1)
|
||||
dt = dtypes.uint64 if base.dtype in (dtypes.uint64, dtypes.int64) else dtypes.uint32
|
||||
hi, lo = first.cast(dt), second.cast(dt)
|
||||
width = hi - lo + _const(dt, 1)
|
||||
hi_u, lo_u = first.cast(dt), second.cast(dt)
|
||||
width = hi_u - lo_u + _const(dt, 1)
|
||||
mask = (_const(dt, 1) << width) - _const(dt, 1)
|
||||
return (base.cast(dt) >> lo) & mask
|
||||
return (base.cast(dt) >> lo_u) & mask
|
||||
self.eat('RBRACKET')
|
||||
dt_suffix = None
|
||||
if self.try_eat('DOT'):
|
||||
@@ -1123,11 +1123,11 @@ def parse_block(lines: list[str], start: int, env: dict[str, VarVal], funcs: dic
|
||||
val = parse_tokens(toks[j:], env, funcs)
|
||||
lo_dt, hi_dt = DTYPES.get(lo_type, dtypes.uint64), DTYPES.get(hi_type, dtypes.uint32)
|
||||
lo_bits = 64 if lo_dt in (dtypes.uint64, dtypes.int64) else 32
|
||||
lo_val = val.cast(lo_dt) if val.dtype.itemsize * 8 <= lo_bits else (val & _const(val.dtype, (1 << lo_bits) - 1)).cast(lo_dt)
|
||||
hi_val = (val >> _const(val.dtype, lo_bits)).cast(hi_dt)
|
||||
block_assigns[lo_var] = env[lo_var] = lo_val
|
||||
block_assigns[hi_var] = env[hi_var] = hi_val
|
||||
if assigns is not None: assigns.extend([(f'{lo_var}.{lo_type}', lo_val), (f'{hi_var}.{hi_type}', hi_val)])
|
||||
lo_u = val.cast(lo_dt) if val.dtype.itemsize * 8 <= lo_bits else (val & _const(val.dtype, (1 << lo_bits) - 1)).cast(lo_dt)
|
||||
hi_u = (val >> _const(val.dtype, lo_bits)).cast(hi_dt)
|
||||
block_assigns[lo_var] = env[lo_var] = lo_u
|
||||
block_assigns[hi_var] = env[hi_var] = hi_u
|
||||
if assigns is not None: assigns.extend([(f'{lo_var}.{lo_type}', lo_u), (f'{hi_var}.{hi_type}', hi_u)])
|
||||
i += 1
|
||||
continue
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from typing import TYPE_CHECKING, Self
|
||||
from tinygrad.dtype import DType, DTypeLike, dtypes, to_dtype
|
||||
from tinygrad.uop import Ops
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from tinygrad.uop.ops import UOp
|
||||
@@ -29,7 +30,7 @@ class DTypeMixin:
|
||||
print(t.dtype, t.numpy())
|
||||
```
|
||||
"""
|
||||
return self if self.dtype == (dt:=to_dtype(dtype)) else self._wrap_uop(self._uop.cast(dt))
|
||||
return self if self.dtype == (dt:=to_dtype(dtype)) else self._wrap_uop(self._uop.alu(Ops.CAST, arg=dt))
|
||||
|
||||
def bitcast(self, dtype:DTypeLike) -> Self:
|
||||
"""
|
||||
@@ -44,7 +45,7 @@ class DTypeMixin:
|
||||
print(t.dtype, t.numpy())
|
||||
```
|
||||
"""
|
||||
return self if self.dtype == (dt:=to_dtype(dtype)) else self._wrap_uop(self._uop.bitcast(dt))
|
||||
return self if self.dtype == (dt:=to_dtype(dtype)) else self._wrap_uop(self._uop.alu(Ops.BITCAST, arg=dt))
|
||||
|
||||
def element_size(self) -> int:
|
||||
"""
|
||||
|
||||
+1
-8
@@ -4,7 +4,7 @@ import sys, time, functools, itertools, math, operator, hashlib, os, types, pick
|
||||
from dataclasses import dataclass, replace
|
||||
from enum import Enum, auto
|
||||
from tinygrad.uop import Ops, GroupOp
|
||||
from tinygrad.dtype import ConstType, dtypes, DType, DTypeLike, to_dtype, truncate, least_upper_dtype, Invalid, AddrSpace
|
||||
from tinygrad.dtype import ConstType, dtypes, DType, DTypeLike, truncate, least_upper_dtype, Invalid, AddrSpace
|
||||
from tinygrad.dtype import ConstFloat, PyConst, InvalidType, storage_fmt_for_dtype, to_storage_scalar, from_storage_scalar
|
||||
from tinygrad.device import Buffer, MultiBuffer, canonicalize_device
|
||||
from tinygrad.helpers import ContextVar, all_int, prod, getenv, all_same, Context, partition, temp, unwrap, T, argfix, Metadata, flatten, TRACEMETA
|
||||
@@ -560,13 +560,6 @@ class UOp(RandMixin, metaclass=UOpMetaClass):
|
||||
def broadcast(self, count:int):
|
||||
if count == 1: return self
|
||||
return UOp(Ops.STACK, src=(self,)*count)
|
||||
def cast(self, dtype:DTypeLike):
|
||||
dtype = to_dtype(dtype)
|
||||
if self.dtype == dtype: return self
|
||||
return UOp(Ops.CAST, arg=dtype, src=(self,))
|
||||
def bitcast(self, dtype:DTypeLike):
|
||||
dtype = to_dtype(dtype)
|
||||
return self if self.dtype == dtype else UOp(Ops.BITCAST, arg=dtype, src=(self,))
|
||||
def load(self, *src:UOp, **kwargs): return UOp(Ops.LOAD, src=(self,)+src, **kwargs)
|
||||
def store(self, src:UOp|ConstType, gate:UOp|None=None, **kwargs):
|
||||
srcs = (self, self.const_like(src) if not isinstance(src, UOp) else src) + ((gate,) if gate is not None else ())
|
||||
|
||||
Reference in New Issue
Block a user