From 681a5e0cfdcd205e9df250c55cde00fea9907960 Mon Sep 17 00:00:00 2001 From: chenyu Date: Mon, 13 Jul 2026 14:23:11 -0400 Subject: [PATCH] remove UOp cast and bitcast override [PR] (#17011) --- test/mockgpu/amd/emu.py | 4 ++-- test/mockgpu/amd/pcode.py | 16 ++++++++-------- tinygrad/mixin/dtype.py | 5 +++-- tinygrad/uop/ops.py | 9 +-------- 4 files changed, 14 insertions(+), 20 deletions(-) diff --git a/test/mockgpu/amd/emu.py b/test/mockgpu/amd/emu.py index 815ed7f85a..0dea6e175a 100644 --- a/test/mockgpu/amd/emu.py +++ b/test/mockgpu/amd/emu.py @@ -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. diff --git a/test/mockgpu/amd/pcode.py b/test/mockgpu/amd/pcode.py index 975d3ed5fd..8958979bd2 100644 --- a/test/mockgpu/amd/pcode.py +++ b/test/mockgpu/amd/pcode.py @@ -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 diff --git a/tinygrad/mixin/dtype.py b/tinygrad/mixin/dtype.py index 09c14cd67c..3a4499f90c 100644 --- a/tinygrad/mixin/dtype.py +++ b/tinygrad/mixin/dtype.py @@ -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: """ diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index 91459f9259..03ca73deec 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -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 ())