from __future__ import annotations
from typing import Any, Callable, cast, TYPE_CHECKING, Type, Sequence, Iterable, Final, Iterator
import sys, time, functools, itertools, math, operator, hashlib, os, types, pickle, pathlib, inspect, weakref, collections, struct
from dataclasses import dataclass, replace
from enum import Enum, auto
from tinygrad.uop import Ops, GroupOp
from tinygrad.dtype import ConstType, dtypes, DType, DTypeLike, truncate, least_upper_dtype, least_upper_float, Invalid, AddrSpace, strong_dtype
from tinygrad.dtype import PyConst, InvalidType, bitcast
from tinygrad.device import Buffer, BufferSpec, MultiBuffer, canonicalize_device, is_disk_device, TinyELF
from tinygrad.helpers import ContextVar, all_int, prod, getenv, all_same, Context, partition, temp, unwrap, T, argfix, Metadata, flatten, TRACEMETA
from tinygrad.helpers import PROFILE, dedup, cdiv, cmod, floordiv, floormod, diskcache_put, to_function_name, cpu_profile, TracingKey
from tinygrad.helpers import VIZ, SPEC, CAPTURE_PROCESS_REPLAY, DISALLOW_BROADCAST, get_shape, fully_flatten, to_tuple
from tinygrad.helpers import colored, ansilen, printable, Target, is_image_shape, strides_for_shape
if TYPE_CHECKING:
  from tinygrad.renderer import Estimates

class AxisType(Enum):
  def __repr__(self): return str(self)
  def __lt__(self, other:AxisType): return self.value < other.value
  # Nesting order: RANGE args sort by axis type, then axis id.
  DEVICE = auto(); GLOBAL = auto(); LOCAL = auto(); WARP = auto(); WEAK = auto(); LOOP = auto() # noqa: E702
  UPCAST = auto(); PLACEHOLDER = auto() # noqa: E702

@dataclass(frozen=True, order=True)
class ParamArg:
  slot: int
  dtype: DType
  vmin_vmax: tuple[PyConst, PyConst]|None = None
  multiple_of: int|None = None
  name: str|None = None
  addrspace: AddrSpace|None = AddrSpace.GLOBAL
  device: str|tuple[str, ...]|None = None
  volatile: bool = False
  # (h, w) if this is an image2d buffer, then the size CONST is h*w*4
  image: tuple[int, int]|None = None
  # the device Buffer for a realized BUFFER. the UOp is the owner of the Buffer: they live and die together (1:1)
  buffer: Buffer|MultiBuffer|None = None
  bind_on_realize: bool = False
  # the bound value of a Variable (an ALU PARAM with a value range); None means unbound
  val: PyConst|None = None
  # the memory the linker gives an ALLOC
  spec: BufferSpec|None = None
  def __repr__(self):
    fields = (("vmin_vmax", None), ("multiple_of", None), ("name", None), ("addrspace", AddrSpace.GLOBAL), ("device", None),
              ("volatile", False), ("image", None), ("bind_on_realize", False), ("val", None), ("spec", None))
    args = [repr(self.slot), repr(self.dtype)] + [f"{k}={v!r}" for k,default in fields if (v:=getattr(self, k)) != default]
    if self.buffer is not None:
      args.append(f"buffer=UOp.new_buffer({self.device!r}, {self.buffer.size}, {self.dtype!r}, {self.slot}).buffer")
    return f"ParamArg({', '.join(args)})"
axis_letters = {AxisType.DEVICE: "d", AxisType.GLOBAL: "g", AxisType.LOCAL: "l", AxisType.WARP: "w", AxisType.WEAK: "L",
                AxisType.LOOP: "L", AxisType.UPCAST: "u"}
axis_colors = {AxisType.DEVICE: "green", AxisType.GLOBAL: "blue", AxisType.LOCAL: "cyan", AxisType.WARP: "CYAN",
               AxisType.WEAK: "red", AxisType.LOOP: "red", AxisType.UPCAST: "yellow"}

range_start = {Ops.STAGE: 1, Ops.REDUCE: 1, Ops.END: 1, Ops.CALL: 1, Ops.LINEAR: 0}

# https://en.wikipedia.org/wiki/Identity_element
def identity_element(op:Ops, dt:DType) -> PyConst: return dt.const({Ops.ADD:0, Ops.MUL:1, Ops.MAX:dt.min}[op])

# With True as the default, this matches the old symbolic behavior
def resolve(x:UOp|bool, default:bool=True):
  if isinstance(x, bool): return x
  assert x.dtype == dtypes.bool, "UOp in resolve must be bool"
  # NOTE: generating the text for the exception is expensive, so we do this
  return bool(sx.vmin) if (sx:=x.simplify()).vmin == sx.vmax else default

# smax/smin are replacements for max/min that preserve symbolic
def _suop(lst, uop_fxn, python_fxn):
  uops, nums = partition(lst, lambda x: isinstance(x, UOp))
  return ssimplify(functools.reduce(uop_fxn, uops + ([python_fxn(nums)] if nums else [])))
def smax(*lst) -> sint: return _suop(argfix(*lst), UOp.maximum, max)
def smin(*lst) -> sint: return _suop(argfix(*lst), UOp.minimum, min)
def srender(x:sint) -> str: return x.render() if isinstance(x, UOp) else str(x)
def _align_left(*shapes:tuple[sint, ...]) -> tuple[tuple[sint, ...], ...]:
  max_dim = max(len(s) for s in shapes)
  return tuple((1,)*(max_dim-len(s))+s for s in shapes)
def _broadcast_shape(*shapes:tuple[sint, ...]) -> tuple[sint, ...]:
  if all_same(shapes): return shapes[0]
  # per right-aligned dim: sizes of 1 broadcast to the others, which must all agree
  ret = []
  for sizes in zip(*_align_left(*shapes)):
    if len(rest:=dedup([s for s in sizes if isinstance(s, UOp) or s != 1])) > 1:
      raise IndexError(f"shape mismatch: objects cannot be broadcast to a single shape {shapes}")
    ret.append(rest[0] if rest else 1)
  return tuple(ret)
def broadcast_axes(src_shape:tuple[sint, ...], out_shape:tuple[sint, ...]) -> tuple[int, ...]:
  # out axes that are added or expanded
  if (nleft:=len(out_shape)-len(src_shape)) < 0: raise RuntimeError(f"cannot broadcast {src_shape} into {out_shape}")
  return tuple(range(nleft)) + tuple(nleft+i for i,s in enumerate(src_shape) if resolve(s == 1, default=False) and resolve(out_shape[nleft+i] != 1))

def ssimplify(uop:sint): return uop.ssimplify() if isinstance(uop, UOp) else uop
def sym_infer(uop: UOp|int, var_vals: dict[str, int]) -> int: return uop.sym_infer(var_vals) if isinstance(uop, UOp) else uop

def range_str(u:UOp, color=False) -> str:
  ret = '_'.join([str(x) if x >= 0 else "m"+str(-x) for x in u.axis_id])
  return colored(ret, axis_colors[u.axis_type]) if color else ret

def multirange_str(rngs:Iterable[UOp], color=False, pad=None) -> str:
  ret = ','.join([range_str(x, color=color) for x in sorted(rngs, key=lambda x: x.arg)])
  if pad is not None: ret += " " * (pad-ansilen(ret))
  return ret

def shape_to_shape_arg(arg:tuple[sint, ...]) -> UOp:
  src = tuple(x if isinstance(x, UOp) else UOp.const(x) for x in arg)
  for x in src:
    if not dtypes.is_int(x.dtype): raise RuntimeError(f"shape must be int, got {x.dtype} in {arg}")
  return src[0] if len(src) == 1 else UOp(Ops.STACK, src=src)

def promo_dtype(src:tuple[UOp,...]) -> DType:
  dts = [x.dtype for x in src]
  return dts[0] if all_same(dts) else least_upper_dtype(*dts)

def dtype_from_uop(op:Ops, src:tuple[UOp,...], arg:Any) -> DType:
  # here are the dtype production rules, total over all Ops
  match op:
    case Ops.STORE | Ops.LINEAR | Ops.SINK | Ops.PROGRAM | Ops.SOURCE | \
         Ops.BACKEDGE | Ops.BARRIER | Ops.IF | Ops.ENDIF | Ops.NOOP | \
         Ops.REWRITE_ERROR | Ops.PYLITERAL:
      # always void
      return dtypes.void
    case Ops.CALL:
      # a call has the dtype of its body, void for opaque bodies
      return src[0].dtype
    case Ops.CUSTOM_FUNCTION:
      # an external function states its return dtype in the arg
      return arg.dtype
    case Ops.CUSTOM | Ops.CUSTOMI:
      assert isinstance(arg, tuple) and len(arg) == 2 and isinstance(arg[1], DType), f"CUSTOM/CUSTOMI arg must be (str, DType), got {arg}"
      return arg[1]
    case Ops.INS:
      # arg is (instruction, dtype), a queue command or an asm line is void
      assert isinstance(arg, tuple) and len(arg) == 2 and isinstance(arg[1], DType), f"INS arg must be (instruction, DType), got {arg}"
      return arg[1]
    case Ops.INDEX:
      # an image access is always float, no matter the storage dtype
      # TODO: should there be a CAST so src[0].dtype just work?
      if (b:=src[0]).op is Ops.PARAM and is_image_shape(b.shape): return dtypes.float
      return b.dtype
    case Ops.LOAD | Ops.UNSHARD | Ops.REDUCE | Ops.AFTER | Ops.RANGE | \
         Ops.CONTIGUOUS_BACKWARD | Ops.COPY | Ops.STAGE | Ops.DETACH | \
         Ops.MSTACK | Ops.MSELECT | Ops.ALLREDUCE | Ops.SPECIAL | Ops.END:
      # pass through first
      return src[0].dtype
    case Ops.CMPLT | Ops.CMPNE | Ops.CMPEQ:
      return dtypes.bool
    case Ops.SIN | Ops.LOG2 | Ops.EXP2 | Ops.SQRT | Ops.RECIPROCAL:
      return dtypes.bool if src[0].base.is_invalid else least_upper_float(src[0].dtype)
    case Ops.WHERE:
      if src[0].dtype != dtypes.bool: raise RuntimeError(f"where cond must be bool, got {src[0].dtype}")
      return promo_dtype(src[1:])
    case Ops.STACK:
      if len(src) == 0: return dtypes.void
      return promo_dtype(src)
    case Ops.WMMA:
      # WMMA output dtype is the accumulator dtype (src[2])
      return src[2].dtype
    case Ops.GETADDR:
      return dtypes.uint64
    case Ops.THREEFRY:
      return dtypes.uint64
    case Ops.FDIV:
      return least_upper_float(promo_dtype(src))
    case Ops.SHL | Ops.SHR:
      if not all(dtypes.is_int(x.dtype) or x.base.is_invalid for x in src):
        raise RuntimeError(f"shift operands must be int, got {[x.dtype for x in src]}")
      return src[0].dtype
    case Ops.BUFFER | Ops.ALLOC | Ops.PARAM:
      assert isinstance(arg, ParamArg), f"{op} must have ParamArg"
      return arg.dtype
    case Ops.BINARY:
      return dtypes.uint8
    case Ops.CAST | Ops.BITCAST:
      assert isinstance(arg, DType), f"CAST/BITCAST arg must be DType, got {arg}"
      return arg
    case Ops.CONST:
      # derived from the value. order matters: bool is an int subclass, ConstFloat is a float subclass
      if isinstance(arg, InvalidType): return dtypes.bool  # Invalid is always bool, the promo lattice bottom
      if isinstance(arg, bool): return dtypes.bool
      if isinstance(arg, int): return dtypes.weakint
      if isinstance(arg, float): return dtypes.weakfloat
      raise TypeError(f"no dtype for CONST with arg {arg}")
  if op in GroupOp.Unary: return src[0].dtype
  # NOTE: CMPLT, CMPNE, CMPEQ, WHERE, SHL, SHR are handled above
  if op in GroupOp.Broadcastable: return promo_dtype(src)
  if op in GroupOp.Movement: return src[0].dtype
  raise RuntimeError(f"no dtype for {op} with arg {arg}")

class UOpMetaClass(type):
  ucache:dict[tuple, weakref.ReferenceType[UOp]] = {}
  def __call__(cls, op:Ops, src:tuple[UOp,...]=tuple(), arg:Any=None, tag:Any=None, metadata:tuple[Metadata,...]|None=None):
    # NOTE: the key must separate nodes of different dtype: a CONST's dtype is the type of its arg, and True == 1 as dict keys
    if (wret:=UOpMetaClass.ucache.get(key:=(op, src, arg, tag, type(arg)), None)) is not None and (ret:=wret()) is not None: return ret
    UOpMetaClass.ucache[key] = weakref.ref(created:=super().__call__(op, src, arg, tag))
    if metadata is not None: all_metadata[created] = metadata
    if SPEC > 1:
      from tinygrad.uop.spec import spec_full
      if SPEC > 2:
        # SPEC=3 checks the shape
        _ = created._shape
      with Context(CHECK_OOB=0): fret = cast(bool|None, spec_full.rewrite(created))
      if fret is not True: raise RuntimeError(f"SPEC ISSUE {fret}: {created}")
    return created

# some uops map to other stuff
all_metadata:weakref.WeakKeyDictionary[UOp, tuple[Metadata, ...]] = weakref.WeakKeyDictionary() # TODO: should this be here?

# recursive_property replaces functools.cached_property in recursive UOp functions to prevent RecursionError
class recursive_property:
  def __init__(self, fxn):
    self.fxn = fxn
    self.nm = fxn.__name__
    self.__doc__ = fxn.__doc__
  def __get__(self, x:UOp|None, owner=None):
    if x is None: return self
    for node in x.toposort(gate=lambda node: self.nm not in node.__dict__): node.__dict__[self.nm] = self.fxn(node)
    return x.__dict__[self.nm]

class _UOpTuple(tuple):
  # UOps are hash-consed, so structurally identical UOps are the same object and equality can be identity.
  # This keeps tuple < from rescanning deep equal prefixes with O(n) value equality at every level of the walk.
  __hash__ = tuple.__hash__
  def __eq__(self, other): return self is other
  def __ne__(self, other): return self is not other

# we import this late so we can use resolve/smax in mixins
from tinygrad.mixin.rand import RandMixin

# NOTE: this should be frozen, but frozen is slower
@dataclass(eq=False, slots=True)
class UOp(RandMixin, metaclass=UOpMetaClass):
  op:Ops
  src:tuple[UOp, ...] = tuple()
  arg:Any = None
  tag:Any = None
  @recursive_property
  def dtype(self) -> DType: return dtype_from_uop(self.op, self.src, self.arg)
  def __del__(self):
    # NOTE: getattr because this object may be partially constructed (e.g. if __init__ raised, like the BEAM timeout SIGALRM)
    try: del UOpMetaClass.ucache[(self.op, self.src, self.arg, self.tag, type(self.arg))]
    except (AttributeError, KeyError): pass
  def __reduce__(self): return UOp, (self.op, self.src, self.arg, self.tag, self.metadata)
  def replace(self, **kwargs) -> UOp:
    new_args = (kwargs.pop("op", self.op), kwargs.pop("src", self.src), kwargs.pop("arg", self.arg), kwargs.pop("tag", self.tag))
    assert len(kwargs) == 0, f"unused kwargs in replace {list(kwargs)}"
    if (self.op, self.src, self.arg, self.tag) == new_args: return self
    return UOp(*new_args)
  def rtag(self, tag=True): return self.replace(tag=tag)
  @property
  def val(self):
    if self.op is Ops.CONST: return self.arg
    if self.op is Ops.GETADDR: return cast("Buffer", self.src[0].buffer).get_buf(to_tuple(self.arg)[0])
    # a casted const CAST(dt, CONST(v)) is one const: .val reads the value through the CAST
    assert self.op is Ops.CAST and self.src[0].op is Ops.CONST, f"val is only for consts, got {self.op}"
    return self.src[0].val
  @property
  def is_invalid(self) -> bool: return self.op is Ops.CONST and self.val is Invalid
  @recursive_property
  def key(self) -> bytes:
    return hashlib.sha256(str((self.op, self.dtype, self.arg)).encode() + b"".join([s.key for s in self.src])).digest()
  def __repr__(self):
    from tinygrad.uop.render import pretty_print
    return pretty_print(self)
  def argstr(self):
    if self.op is Ops.REDUCE: return f'({", ".join(map(str, self.arg))})'
    return repr(self.arg)
  def tagstr(self): return f", tag={self.tag}" if self.tag is not None else ""

  @functools.cached_property
  def backward_slice(self:UOp) -> dict[UOp, None]:
    res: dict[UOp, None] = self.toposort(enter_calls=False)
    res.pop(self)
    return res

  @property
  def backward_slice_with_self(self:UOp) -> dict[UOp, None]: return {self:None, **self.backward_slice}
  def op_in_backward_slice_with_self(self, *ops:Ops) -> bool:
    # Check self first, then iterate backward_slice (avoids creating intermediate dict)
    return self.op in ops or any(x.op in ops for x in self.backward_slice)

  @recursive_property
  def _bool_slice(self) -> frozenset[UOp]: return frozenset().union(*[s.bool_slice for s in self.src])
  # NOTE: self is added outside the cache, a cached self-reference is a cycle the refcounter can't free
  @property
  def bool_slice(self) -> frozenset[UOp]: return self._bool_slice | {self} if self.dtype is dtypes.bool else self._bool_slice

  def toposort(self, gate:Callable|None=None, enter_calls=True) -> dict[UOp, None]:
    cache: dict[UOp, None] = {}
    stack: list[tuple[UOp, bool]] = [(self, False)] # each stack entry is (node, visited_flag)
    while stack:
      node, visited = stack.pop()
      if node in cache: continue
      if not visited:
        if gate is None or gate(node):
          stack.append((node, True))  # push node back on stack to process after its srcs
          for s in reversed(node.src if enter_calls else node.src_without_body):
            stack.append((s, False)) # push srcs on the stack
      else: cache[node] = None # second time i'm seeing this node, add it to returned toposort
    return cache

  def topovisit(self, visitor:Callable[[UOp], T], cache:dict[UOp, T]) -> T:
    # NOTE: this shares a lot of code with toposort
    stack: list[tuple[UOp, bool]] = [(self, False)]
    while stack:
      node, visited = stack.pop()
      if node in cache: continue
      if not visited:
        stack.append((node, True))
        for s in reversed(node.src): stack.append((s, False))
      else: cache[node] = visitor(node)
    return cache[self]

  @functools.cached_property
  def tuplize(self) -> _UOpTuple:
    # arg goes through repr: args of different types (None, str, tuple) must stay mutually comparable for the sort
    return _UOpTuple((self.op.value, repr(self.arg), self.dtype,)+tuple([x.tuplize for x in self.src]))

  # *** uop shape stuff ***

  @recursive_property
  def _shape(self) -> tuple[sint, ...]|None:
    match self.op:
      # late ops don't have shape
      case Ops.IF | Ops.BARRIER | Ops.SINK | Ops.REWRITE_ERROR | Ops.ENDIF | Ops.BACKEDGE | \
           Ops.LINEAR | Ops.PROGRAM | Ops.SOURCE:
        return None

      # INS shape is always scalar, vector width is in the instruction encoding
      case Ops.INS:
        if self.dtype is dtypes.void: return None
        return ()

      # special (terrible) case for RESHAPE on NOOP
      case Ops.RESHAPE:
        if self.src[0].op is Ops.NOOP: return self.marg

      # hacks for NOOP
      case Ops.NOOP:
        return self.src[0]._shape if len(self.src) >= 1 else None

      case Ops.INDEX:
        shp:list[sint] = []
        for s in self.src[1:]: shp.extend(list(s.shape))
        return tuple(shp) + self.src[0].shape[len(self.src[1:]):]

      case Ops.STACK:
        if len(self.src) == 0: return ()
        return (len(self.src),) + self.src[0].shape
      case Ops.CONST:
        return ()

      # some ops init the shape
      case Ops.GETADDR: return ()
      case Ops.RANGE | Ops.SPECIAL: return ()
      case Ops.BINARY: return (len(self.arg),)
      case Ops.BUFFER | Ops.ALLOC | Ops.PARAM:
        assert len(self.src[0].as_shape) <= 1
        if (img:=self.arg.image) is not None: return (img[0], img[1], 4)
        return self.src[0].as_shape
      case Ops.CUSTOM | Ops.CUSTOMI:
        if self.dtype is dtypes.void: return None
        input_shapes = [x._shape for x in self.src if x._shape is not None]
        return _broadcast_shape(*input_shapes) if input_shapes else None
      case Ops.CUSTOM_FUNCTION: return None if self.dtype is dtypes.void else ()
      case Ops.PYLITERAL: return None
      case Ops.STAGE:
        # STAGE adds the existing shape to the front, opposite of INDEX
        return tuple([int(r.vmax+1) for r in self.src[1:]])+self.src[0].shape

      # wmma output shape = accumulator shape (src[2])
      case Ops.WMMA:
        wmma_b = _broadcast_shape(self.src[0].shape[:-1], self.src[1].shape[:-1], self.src[2].shape[:-1])
        return wmma_b + (self.src[2].shape[-1],)

      # passthrough ops
      case Ops.MSTACK | Ops.MSELECT | Ops.DETACH | Ops.CONTIGUOUS_BACKWARD | Ops.AFTER | Ops.LOAD | \
           Ops.COPY | Ops.ALLREDUCE | Ops.STORE | Ops.END | Ops.CALL:
        return self.src[0]._shape

      case Ops.BITCAST:
        ps = self.src[0]._shape
        if ps is None: return None
        if (output_sz:=self.dtype.itemsize) != (input_sz:=self.src[0].dtype.itemsize) and len(ps) > 0:
          if isinstance(ps[-1], int) and (ps[-1]*input_sz) % output_sz: raise RuntimeError("unsupported size in bitcast")
          return ps[:-1]+(ssimplify((ps[-1]*input_sz) // output_sz),)
        return ps

    # movement ops change the shape
    if self.op in GroupOp.Movement.union({Ops.UNSHARD, Ops.REDUCE}):
      ps = self.src[0]._shape
      if ps is None: raise RuntimeError(f"movement op {self.op} requires shape, {self.src[0].op} doesn't have one")
      match self.op:
        case Ops.RESHAPE:
          if not all(x >= 0 for x in self.marg): raise ValueError(f"shape can't contain negative numbers {self.marg}")
          # with symbolic views prod equality can be true at runtime but unprovable, only reject provably unequal products
          if resolve(prod(ps) != prod(self.marg), False): raise ValueError(f"bad reshape: {ps} -> {self.marg}")
          return self.marg
        case Ops.EXPAND:
          return tuple(self.marg) + ps
        case Ops.PERMUTE:
          if sorted(self.marg) != list(range(len(ps))): raise ValueError(f"invalid permutation {self.marg} of len {len(ps)}")
          return tuple(ps[i] for i in self.marg)
        case Ops.PAD:
          # TODO: why do i need resolve here?
          if len(ps) != len(self.marg) or not all(resolve(sz>=0) and resolve(0<=o) and resolve(o+s<=sz) for s,(o,sz) in zip(ps, self.marg)):
            raise ValueError(f"invalid pad {self.marg} for {ps}")
          return tuple(sz for _,sz in self.marg)
        case Ops.SHRINK:
          # TODO: why do i need resolve here?
          if len(ps) != len(self.marg) or not all(resolve(0<=o) and resolve(sz>=0) and resolve(o+sz<=s) for s,(o,sz) in zip(ps, self.marg)):
            raise ValueError(f"invalid shrink {self.marg} for {ps}")
          return tuple(sz for _,sz in self.marg)
        case Ops.FLIP:
          if len(ps) != len(self.marg) or not all(isinstance(x, bool) for x in self.marg): raise ValueError(f"bad flip on {ps}, {self.marg}")
          return ps
        case Ops.UNSHARD: return tuple(s*(int(self.src[1:][self.arg.index(a)].vmax)+1) if a in self.arg else s for a,s in enumerate(ps))
        case Ops.REDUCE:
          num_axes = self.arg[1]
          if not isinstance(num_axes, int) or num_axes < 0 or num_axes > len(ps):
            raise ValueError(f"invalid type for axis: {num_axes}")
          return ps[num_axes:]

    if self.op in GroupOp.Unary.union({Ops.CAST}):
      assert len(self.src) == 1, "unary ops must have 1 src"
      return self.src[0]._shape

    # elementwise ops keep the shape the same. all inputs with shape must match
    if self.op in GroupOp.Broadcastable:
      input_shapes = [x._shape for x in self.src]
      assert len(self.src) > 0 and all(x is not None for x in input_shapes), f"None input shape not supported for {self.op}"
      if DISALLOW_BROADCAST and not all_same(input_shapes):
        raise RuntimeError(f"shape mismatch at {self.op}: {input_shapes} {[x.op for x in self.src]}")
      # broadcasting lives in _shape property now
      return _broadcast_shape(*input_shapes)

    # all Ops must be explicitly handled
    raise NotImplementedError(f"no shape handling for {self.op} with {self.dtype}")

  @property
  def shape(self) -> tuple[sint, ...]:
    if (ret:=self._shape) is None: raise RuntimeError(f"shape requested, but {self.op} doesn't have a shape")
    return ret

  @property
  def shard_shape(self) -> tuple[sint, ...]:
    if not isinstance(self.device, tuple) or self.axis is None: return self.shape
    dcount = int(self.src[1].vmax)+1 if self.op is Ops.UNSHARD else len(self.device)
    return tuple(x//dcount if i == self.axis else x for i,x in enumerate(self.shape))

  @property
  def max_shard_shape(self) -> tuple[int, ...]: return to_max_shape(self.shard_shape)

  @functools.cached_property
  def ended_ranges(self) -> tuple[UOp, ...]:
    if self.op is Ops.CALL and self.body.op is Ops.CUSTOM_FUNCTION: return ()
    if self.op is Ops.END: return tuple(r for r in self.src[1:] if r.op is Ops.RANGE)
    if self.op is Ops.BACKEDGE: return self.src[1:2]  # the condition's other ranges remain live
    if self.op in range_start: return self.src[range_start[self.op]:]
    if self.op is Ops.AFTER: return tuple(flatten([x.ended_ranges for x in self.src[1:]]))
    if self.op is Ops.BARRIER: return tuple(flatten([x.ended_ranges for x in self.src]))
    # UNSHARD ends the DEVICE range: its src is per-device index math, the device axis is carried by the axis metadata
    if self.op is Ops.UNSHARD: return self.src[1:]
    return ()

  # determine what ranges this is in
  @recursive_property
  def _ranges(self) -> dict[UOp, None]:
    ret: dict[UOp, None] = {}
    for s in self.src: ret.update(s.ranges)
    for er in self.ended_ranges:
      if er.op is Ops.RANGE:
        # if it's a single RANGE, we don't flow through it.
        ret.pop(er, None)
      else:
        # if it's not a RANGE, we include all ranges in srcs.
        # technically we shouldn't flow through these ranges either, but this is pre pm_add_control_flow so it's the same.
        for s in er.ranges: ret.pop(s, None)
    return ret

  @property
  def ranges(self) -> dict[UOp, None]:
    if self.op is Ops.RANGE: return {self:None} | self._ranges
    return self._ranges

  @property
  def axis_id(self) -> tuple[int, ...]:
    assert self.op is Ops.RANGE, f"axis_id is only for RANGE, not {self.op}"
    return self.arg[1:]

  @property
  def axis_type(self) -> AxisType:
    assert self.op is Ops.RANGE, f"axis_type is only for RANGE, not {self.op}"
    return self.arg[0]

  # *** uop evaluation ***

  def simplify(self, tracked=False):
    if self.op is Ops.CONST: return self
    if self.op is Ops.SINK and all(s.op is Ops.CONST or (s.op is Ops.STACK and len(s.src) == 0) for s in self.src): return self
    # late import!
    from tinygrad.uop.symbolic import symbolic
    with Context(TRACK_MATCH_STATS=0 if not tracked else TRACK_MATCH_STATS.value):
      return graph_rewrite(self, symbolic, name="simplify")
  def ssimplify(self) -> UOp|ConstType:
    if (ret := self.simplify()).op is Ops.CAST and ret.src[0].op is Ops.CONST: return ret.dtype.const(ret.src[0].val)
    return ret.val if ret.op is Ops.CONST else ret
  def _eval(self, dtype, expected_type:Type[T]) -> T:
    assert self.dtype in dtype, f"eval with wrong dtype {self}"
    vmin, vmax = (simple_self:=self.simplify())._min_max
    if vmin != vmax: raise ValueError(f"eval failed to be a single number, range is {vmin} to {vmax} in {simple_self.render()}")
    assert isinstance(vmin, expected_type), f"vmin is wrong dtype {type(vmin)} != {expected_type}"
    return vmin
  def __bool__(self): return self._eval((dtypes.bool,), bool)
  def __int__(self): return self._eval(dtypes.ints+(dtypes.weakint,), int)
  def __float__(self): return float(self._eval(dtypes.floats+(dtypes.weakfloat,), float))
  def substitute(self, dvars:dict[UOp, UOp], name:str|None=None, extra_pm:PatternMatcher|None=None, walk:bool=False, enter_calls:bool=False):
    dvars = {k:v for k,v in dvars.items() if k is not v}
    if len(dvars) == 0: return self
    with Context(TRACK_MATCH_STATS=(0 if name is None else TRACK_MATCH_STATS.value)):
      return graph_rewrite(self, (extra_pm+_substitute) if extra_pm is not None else _substitute, dvars,
                           bottom_up=True, walk=walk, enter_calls=enter_calls, name=name)
  # NOTE: this is not called by Tensor slice (Tensor handles UOps directly), but satisfies SupportsIndex for type checking
  def __index__(self): return self.__int__()

  # *** uop tracing stuff ***

  @recursive_property
  def trace_num(self):
    num = next(ucount)
    # tags can contain UOps (callify tags nodes with their originals): store them as trace_nums, same as srcs
    tag = tuple(t.trace_num if isinstance(t, UOp) else t for t in self.tag) if isinstance(self.tag, tuple) else self.tag
    # the trace must not retain the device Buffer: store a placeholder instead (the real one would pin memory and fail pickling)
    arg = replace(self.arg, buffer=cast("Buffer", object())) if isinstance(self.arg, ParamArg) and self.arg.buffer is not None else self.arg
    uop_fields[num] = (self.op, tuple(s.trace_num for s in self.src), arg, tag)+((self.metadata,) if TRACEMETA>=2 else ())
    return num

  # *** uop syntactic sugar ***

  def sink(*srcs:UOp|None, **kwargs):  # pylint: disable=no-self-argument
    return UOp(Ops.SINK, src=tuple([x for x in srcs if x is not None]), **kwargs)
  def group(*srcs:UOp|None, **kwargs):  # pylint: disable=no-self-argument
    if len(srcs) == 1 and isinstance(srcs[0], UOp): return srcs[0]
    return UOp(Ops.STACK).after(*[x for x in srcs if x is not None], **kwargs)
  @property
  def body(self) -> UOp:
    """the body of a CALL: the program, copy or function reference being called (its first src)"""
    if self.op is not Ops.CALL: raise RuntimeError(f"body requested, but {self.op} is not a CALL")
    return self.src[0]
  @property
  def is_inline_call(self) -> bool:
    return self.op is Ops.CALL and self.body.op is Ops.SINK and self.body.arg is None and not self.arg.precompile
  @property
  def has_unbound_outputs(self) -> bool:
    """does this call still have unresolved outputs: ALLOCs among its inputs (minted by call_with_outputs,
    resolved when the call is inlined or the outputs are materialized). a lifecycle query, not a call type"""
    return self.op is Ops.CALL and any((b:=x.unsharded_base).op is Ops.ALLOC and not b.arg.bind_on_realize for x in self.src[1:])
  @property
  def unbound_outputs(self) -> tuple[UOp, ...]:
    """the unresolved outputs of this call: an AFTER on each ALLOC input, usable like a normal buffer"""
    return tuple(x.after(self) for x in self.src[1:] if (b:=x.unsharded_base).op is Ops.ALLOC and not b.arg.bind_on_realize)
  def index(self, *srcs:UOp|int|None, **kwargs):
    new_srcs: list[UOp] = [UOp.const(x) if isinstance(x, int) else x for x in srcs if x is not None]
    if len(new_srcs) == 1 and new_srcs[0].op is Ops.CONST and self.op is Ops.STACK: return self.src[new_srcs[0].val]
    return UOp(Ops.INDEX, src=(self,)+tuple(new_srcs), **kwargs)
  def __getitem__(self, idx):
    # buffers index into INDEX UOps (scalar lookup); everything else uses the shared mixin view path
    if self.addrspace in (None, AddrSpace.ALU) or self.device is not None: return super(UOp, self).__getitem__(idx)
    idx = self._normalize_indices(list(argfix(idx)))
    if len(slice_idx:=[i for i,x in enumerate(idx) if isinstance(x, slice)]):
      # apply SHRINK for slices that aren't the full range
      bounds = tuple((s.start or 0, s.stop if s.stop is not None else self.shape[i]) if isinstance(s, slice) else (0, self.shape[i])
                     for i, s in enumerate(idx))
      src = self.shrink(bounds)
      non_slice_args = [x for x in idx if not isinstance(x, slice)]
      if not non_slice_args: return src  # all dims are slices, no indexing needed
      perm = src.permute(tuple([i for i in range(src.ndim) if i not in slice_idx] + slice_idx))
      return perm.index(*non_slice_args)
    return self.index(*idx)
  @property
  def _uop(self) -> UOp: return self
  @classmethod
  def _wrap_uop(cls, u:UOp) -> UOp: return u
  def const_like(self, b:ConstLike, dtype:DType|None=None):
    ret = UOp.const(b, dtype or self.dtype)
    return ret._mop(Ops.EXPAND, arg=self._shape) if self._shape and ret._shape != self._shape else ret
  def vconst_like(self, b:ConstLike):
    # for use after movement ops have been removed
    return UOp.const(b, self.dtype).broadcast(self.max_numel())
  def ufix(self, x):
    if isinstance(x, UOp): return x
    return UOp.const(x)
  def broadcast(self, count:int):
    if count == 1: return self
    return UOp(Ops.STACK, src=(self,)*count)
  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 ())
    return UOp(Ops.STORE, src=srcs, **kwargs)
  def end(self, *src:UOp): return UOp(Ops.END, src=(self,)+src) if len(src) else self
  def backedge(self, loop:UOp, cond:UOp): return UOp(Ops.BACKEDGE, src=(self, loop, cond))
  def after(self, *src:UOp, **kwargs): return UOp(Ops.AFTER, src=(self,)+src, **kwargs) if len(src) else self
  @property
  def without_after(self) -> UOp: return self.src[0].without_after if self.op is Ops.AFTER else self
  def barrier(self, *src:UOp): return UOp(Ops.BARRIER, src=(self,)+src)
  def ins(self, arg, **kwargs): return UOp(Ops.INS, kwargs.pop("src", self.src), (arg, kwargs.pop("dtype", self.dtype)), kwargs.pop("tag", self.tag))
  def contract(self, *rngs:UOp):
    assert all(x.axis_type == AxisType.UPCAST for x in rngs), "all contract ranges must be upcast"
    return UOp.stack(*[self.substitute(dict(zip(rngs, [r.const_like(i) for r,i in zip(rngs, idx)])))
                           for idx in itertools.product(*[range(int(r.vmax)+1) for r in rngs])])
  def alu(self, op, *src:UOp, **kwargs): return UOp(op, src=(self, *src), **kwargs)
  @staticmethod
  def const(b:ConstLike, dtype:DType|None=None):
    if dtype is None or b is Invalid: dtype = dtypes.from_py(b)
    if isinstance(b, UOp): return b.cast(dtype)
    # NOTE: it always has to be STACK now, even if they are all the same
    if isinstance(b, tuple): return UOp.stack(*[UOp.const(c, dtype) for c in b])
    # .cast folds away at exactly the dtypes a CONST derives (bool/weakint/weakfloat): bare there, the pair everywhere else
    return UOp(Ops.CONST, arg=dtype.const(b), src=()).cast(dtype)
  # cast, except for CONST, in which case rebuild a new CONST at the dtype
  def ccast(self, dtype:DType): return UOp.const(self.val, dtype) if self.op is Ops.CONST else self.cast(dtype)
  # a forced CAST for bool: .cast(bool) folds, so UOp.const cannot state the width
  @staticmethod
  def cconst(b:ConstLike, dtype:DType): return UOp(Ops.CAST, src=(UOp.const(b),), arg=dtype)
  @staticmethod
  def range(end:sint, axis_id, axis_type=AxisType.WEAK, *, dtype=dtypes.weakint, **kwargs):
    return UOp(Ops.RANGE, src=(sint_to_uop(end, dtype),), arg=(axis_type, axis_id), **kwargs)
  @staticmethod
  def loop(axis_id:int): return UOp(Ops.RANGE, src=(UOp(Ops.NOOP),), arg=(AxisType.WEAK, axis_id))
  @staticmethod
  def special(end:sint, name:str): return UOp(Ops.SPECIAL, src=(sint_to_uop(end),), arg=name)
  @staticmethod
  def wmma(a:UOp, b:UOp, acc:UOp, dims:tuple[int, int, int], threads:int, tc_upcast_axes=None):
    # dtype_in is stored in the arg (not derived from src[0].dtype) because bitcast rewrites change src dtypes
    return UOp(Ops.WMMA, src=(a, b, acc), arg=(dims, a.dtype, threads, tc_upcast_axes))
  def _rop(self, op:Ops, axis:tuple[int, ...]):
    # NOTE: we don't allow reduce on 1s axis
    axis = tuple(sorted(axis))
    reduce_axis = tuple(x for x in axis if resolve(self.shape[x] != 1))
    if not len(reduce_axis):
      return self.reshape(tuple(s for i,s in enumerate(self.shape) if i not in axis))
    # permute so reduced axes are at the front
    perm = reduce_axis + tuple(i for i in range(len(self.shape)) if i not in reduce_axis)
    ret = UOp(Ops.REDUCE, src=(self.permute(perm),), arg=(op, len(reduce_axis)))
    return ret.reshape(tuple(s for i,s in enumerate(self.shape) if i not in axis)) if axis != reduce_axis else ret
  @staticmethod
  def invalid() -> UOp: return UOp.const(Invalid)
  def valid(self, cond) -> UOp: return cond.where(self, self.const_like(Invalid))
  def get_idx(self) -> UOp:
    if self.op is Ops.STACK: return UOp.stack(*(x.get_idx() for x in self.src))
    return self.src[1] if self.op is Ops.WHERE and self.src[2].is_invalid else self
  def get_valid(self) -> UOp:
    if self.op is Ops.STACK: return UOp.stack(*(x.get_valid() for x in self.src))
    return self.src[0] if self.op is Ops.WHERE and self.src[2].is_invalid else UOp.const(not self.is_invalid)
  def reduce(self, *src:UOp, **kwargs):
    arg = kwargs.pop('arg', None)
    if isinstance(arg, Ops): arg = (arg, 0)
    return UOp(Ops.REDUCE, src=(self,)+src, arg=arg, **kwargs)

  def bufferize(self, *args, **kwargs): return UOp(Ops.STAGE, src=(self,)+args, **kwargs)
  def allreduce(self, op, device:str|tuple[str, ...]):
    assert isinstance(self.device, tuple), f"allreduce must be on tuple {self.device} isn't"
    return UOp(Ops.ALLREDUCE, src=(self,), arg=(op, device))
  def overflows(self, dtype:DType) -> bool: return self.vmin < dtype.min or dtype.max < self.vmax

  def split_uop(self:UOp, sep:Ops) -> Iterator[UOp]:
    if self.op is sep:
      for s in self.src: yield from s.split_uop(sep)
    else: yield self

  # *** multi-device helpers ***

  def unshard(self, axis:int|tuple[int, ...]|None, device_range:UOp|tuple[UOp, ...]|None=None):
    assert axis is not None, "multi None is no longer supported"
    # an UNSHARD carries the value and one sharding range per sharded axis (arg is the tuple of sharded axes,
    # sorted). the single-axis axis form defaults the range to a DEVICE range over the devices; a range need not
    # be DEVICE, e.g. a LOCAL range shards a kernel tile into per-thread fragments
    if isinstance(axis, int): axis = (axis,)
    if device_range is None:
      assert isinstance(self.device, tuple), f"multi device must be tuple, {self.device} isn't"
      device_range = (UOp.range(len(self.device), 0, AxisType.DEVICE),)
    if isinstance(device_range, UOp): device_range = (device_range,)
    assert isinstance(device_range, tuple) and len(axis) == len(device_range) and len(set(axis)) == len(axis)
    axis, device_range = map(tuple, zip(*sorted(zip(axis, device_range))))
    return UOp(Ops.UNSHARD, src=(self, *device_range), arg=axis)

  @property
  def sharding(self) -> tuple[tuple[int, UOp], ...]:
    """(axis, RANGE) pairs this value is sharded over (the source of truth for shard bounds/counts)."""
    return tuple(zip(self.arg, self.src[1:])) if self.op is Ops.UNSHARD else ()

  @property
  def bounds(self):
    if self.axis is None: raise RuntimeError("bounds is not defined when axis is None")
    dcount = int(self.src[1].vmax)+1 if self.op is Ops.UNSHARD else len(self.device)
    return tuple(itertools.pairwise(itertools.accumulate([self.src[0].shape[self.axis] for _ in range(dcount)], initial=0)))

  @functools.cached_property
  def axis(self) -> int|None:
    # COPY removes axis. TODO: add more tests for this, and consider MSELECT/MSTACK
    if self.op is Ops.COPY: return None
    if self.op is Ops.UNSHARD:
      if len(self.arg) != 1: raise RuntimeError(f"UOp is sharded on multiple axes {self.arg}, use .sharding")
      return self.arg[0]
    # NOTE: they all have to share an axis, we always choose [-1]. src axes are right-aligned into the output shape
    if self.op in GroupOp.ALU.union({Ops.STACK}):
      return axes[-1] if (axes := dedup([x.axis+len(self.shape)-len(x.shape) for x in self.src if x.axis is not None])) else None
    if len(self.src) == 0: return None
    src_axis = self.src[0].axis
    if self.op is Ops.SHRINK and src_axis is not None and self.marg[src_axis] != (0, self.src[0].shape[src_axis]):
      return None # SHRINK will remove the sharding if it's on axis
    if self.op is Ops.REDUCE:
      if src_axis is None: return None
      if src_axis < self.arg[1]: return None
      return src_axis - self.arg[1]
    if self.op is Ops.RESHAPE:
      if src_axis is None: return None
      arg_acc:list[sint] = [ssimplify(x) for x in itertools.accumulate(self.marg, operator.mul, initial=1)]
      # new_axis is the last one that preserves prod(prior to new_axis) and must not move items between shards
      target = ssimplify(prod(self.src[0].shape[:src_axis]))
      if target not in arg_acc: raise RuntimeError(f"reshape {self.src[0].shape} -> {self.shape} moved items between shards")
      new_axis = len(arg_acc) - arg_acc[::-1].index(target) - 1
      dcount = len(self.device) if isinstance(self.device, tuple) else \
        int(next(u.src[1] for u in self.src[0].toposort() if u.op is Ops.UNSHARD).vmax)+1
      if self.shape[new_axis] % dcount != 0: raise RuntimeError(f"reshape {self.src[0].shape} -> {self.shape} moved items between shards")
      return new_axis
    if self.op is Ops.PERMUTE: return self.marg.index(src_axis) if src_axis is not None else None
    if self.op is Ops.EXPAND: return src_axis + len(self.marg) if src_axis is not None else None
    return src_axis

  def _shard(self, axis:int, rng:UOp) -> UOp:
    if len(self.shape) == 0: return self  # scalars broadcast, no sharding needed
    dcount = int(rng.vmax)+1
    if self.shape[axis] % dcount != 0: raise RuntimeError(f"multi axis uneven: {self.shape[axis]=} {axis=} {dcount=}")
    sz = self.shape[axis] // dcount
    return self.shrink(tuple((0,s) if i != axis else (rng*sz,rng*sz+sz) for i,s in enumerate(self.shape)))
  def shard(self, devices:tuple[str, ...], axis:int|None=None) -> UOp:
    copied = self.copy_to_device(devices)
    return copied if axis is None else copied._shard(axis, UOp.range(len(devices), 0, AxisType.DEVICE)).unshard(axis)

  def copy_to_device(self, device:str|tuple[str, ...], arg=None):
    if is_disk_device(device):
      raise RuntimeError("COPY to DISK is not allowed; use STORE to an explicit disk buffer")
    assert arg is None or isinstance(self.device, tuple)
    inp = self if arg is None else UOp(Ops.MSELECT, src=(self,), arg=arg)
    if inp.dtype in dtypes.weaks: raise RuntimeError(f"cannot create storage for weak dtype {inp.dtype}")
    # multi-device COPYs carry the DEVICE range as src[1] (like UNSHARD's sharding ranges)
    return UOp(Ops.COPY, src=(inp, *UOp.device_range_src(device)), arg=device)
  def store_call(self, src:UOp) -> UOp:
    """Executable bulk transfer into this buffer."""
    return self.param_like(0).store(src.param_like(1)).call(self, src)
  def mselect(self, arg:int) -> UOp: return UOp(Ops.MSELECT, src=(self,), arg=arg)
  def mstack(self, *srcs: UOp) -> UOp: return UOp(Ops.MSTACK, src=(self,)+srcs) if len(srcs) else self
  @property
  def metadata(self) -> tuple[Metadata, ...]|None: return all_metadata.get(self, None)

  # little helpers
  def on_disk(self:UOp): return isinstance(self.device, str) and self.device.startswith("DISK")
  def on_creation_device(self:UOp): return isinstance(self.device, str) and self.device.startswith(("DISK", "NPY", "PYTHON"))
  def needs_storage(self:UOp) -> bool: return not self.is_virtual and (self.storage_base.op is Ops.ALLOC or not self.has_buffer_identity())

  # *** uop movement ops ***

  @property
  def base(self) -> UOp:
    if self.op in GroupOp.Movement: return self.src[0].base
    if self.op is Ops.DETACH: return self.src[0].base  # DETACH can't change base
    return self

  # base with UNSHARD
  @property
  def unsharded_base(self) -> UOp:
    if self.op in GroupOp.Movement: return self.src[0].base
    if self.op is Ops.DETACH: return self.src[0].base  # DETACH can't change base
    # TODO: why can't this be in normal base?
    if self.op is Ops.UNSHARD: return self.src[0].base
    return self

  # the storage this uop ultimately targets: base with UNSHARD, BITCAST and AFTER stripped
  @property
  def storage_base(self) -> UOp:
    b = self.unsharded_base
    while b.op in {Ops.BITCAST, Ops.AFTER, Ops.UNSHARD}: b = b.src[0].unsharded_base
    return b

  # cached property here makes external_uop_gc fail, why?
  @property
  def as_shape(self) -> tuple[sint, ...]:
    if self.op is Ops.CONST: return (self.val,)
    if self.op is not Ops.STACK: return (ssimplify(self),)
    return tuple(s.val if s.op is Ops.CONST else ssimplify(s) for s in self.src)

  @functools.cached_property
  def marg(self):
    match self.op:
      case Ops.RESHAPE | Ops.EXPAND: return self.src[1].as_shape
      case Ops.PAD | Ops.SHRINK: return tuple(zip(self.src[1].as_shape, self.src[2].as_shape))
      case Ops.PERMUTE | Ops.FLIP: return self.arg
      case _: raise RuntimeError(f"{self.op} is not a MovementOp")

  def _mop(self, op:Ops, arg) -> UOp:
    # early NOOP
    if op is Ops.EXPAND and len(arg) == 0: return self
    if op in {Ops.SHRINK, Ops.PAD} and len(arg) == 0:
      assert len(self.shape) == 0, "0 len arg only valid on zero length shape"
      return self
    match op:
      case Ops.RESHAPE | Ops.EXPAND: src_args = [arg]
      case Ops.PAD | Ops.SHRINK: src_args = list(zip(*arg))
      case Ops.PERMUTE | Ops.FLIP: src_args = []
      case Ops.STACK:
        srcs = (self,)+tuple(arg)
        dtype = dtype_from_uop(Ops.STACK, srcs, None)
        return UOp(Ops.STACK, src=tuple(u if u.base.is_invalid else u.ccast(dtype) for u in srcs))
      case _: raise RuntimeError(f"{op} is not a MovementOp")
    usrcs = [shape_to_shape_arg(arg) for arg in src_args]
    if len(usrcs) == 0: return UOp(op, src=(self,), arg=arg)
    return UOp(op, src=(self,)+UOp.sink(*usrcs).simplify().src)

  # *** uop Buffer stuff ***

  unique_num = itertools.count(0)

  def getaddr(self, device=None) -> UOp:
    if self.without_after.op not in {Ops.BUFFER, Ops.ALLOC, Ops.SHRINK, Ops.BITCAST, Ops.BINARY,
                                    Ops.MSTACK, Ops.MSELECT, Ops.PARAM, Ops.LINEAR}: return self
    return UOp(Ops.GETADDR, src=(self,), arg=device or to_tuple(self.device)[0])
  @staticmethod
  def device_range_src(device:str|tuple[str, ...]|None) -> tuple[UOp, ...]:
    # BUFFER/ALLOC/COPY carry a DEVICE range when targeting multiple devices
    return (UOp.range(len(device), 0, AxisType.DEVICE),) if isinstance(device, tuple) else ()
  @staticmethod
  def new_buffer(device:str|tuple[str, ...], size:int, dtype:DType, num=None):
    if dtype in dtypes.weaks: raise RuntimeError(f"cannot create storage for weak dtype {dtype}")
    assert isinstance(size, int), f"new_buffer size must be a concrete int, got {size}"
    slot = next(UOp.unique_num) if num is None else num
    buf = MultiBuffer(device, size, dtype) if isinstance(device, tuple) else Buffer(device, size, dtype)
    return UOp(Ops.BUFFER, src=(UOp.const(size),)+UOp.device_range_src(device), arg=ParamArg(slot, dtype, device=device, buffer=buf))
  @staticmethod
  def from_buffer(opaque:Buffer|MultiBuffer, device:str|tuple[str, ...]|None=None):
    # the opaque Buffer goes straight in the arg: the ucache dedups because the arg (and thus the Buffer) is part of the key
    return UOp(Ops.BUFFER, src=(UOp.const(opaque.size),)+UOp.device_range_src(device or opaque.device),
               arg=ParamArg(-id(opaque), opaque.dtype, device=device or opaque.device, buffer=opaque))
  def empty_like(self, dtype:DTypeLike|None=None, device:str|tuple[str, ...]|None=None) -> UOp:
    device = canonicalize_device(self.device if device is None else device)
    axis = self.axis if isinstance(device, tuple) else None
    ret = UOp.empty(self.shard_shape if axis is not None else self.shape, dtype=self.commit_dtype() if dtype is None else dtype, device=device)
    return ret.unshard(axis) if axis is not None else ret
  @staticmethod
  def _frompy(x:list|tuple|bytes, dtype:DType) -> UOp:
    if isinstance(x, bytes): ret, data = UOp.new_buffer("PYTHON", len(x)//dtype.itemsize, dtype), x
    else:
      # bfloat16 and fp8 have no struct format, so pack a float32 buffer and cast
      bdtype = dtypes.float32 if dtype in [dtypes.bfloat16, *dtypes.fp8s] else dtype
      assert bdtype.fmt is not None, f"{bdtype=} has None fmt"
      ret = UOp.new_buffer("PYTHON", prod(shape:=get_shape(x)), bdtype).reshape(shape)
      data = struct.pack(f"{prod(shape)}{bdtype.fmt}", *[truncate[bdtype](bdtype.const(xi)) for xi in fully_flatten(x)])
    if not data: ret.buffer.allocate(memoryview(bytearray()))
    else: (buf:=ret.buffer.ensure_allocated()).allocator._copyin(buf._buf, memoryview(data))
    if ret.dtype != dtype: ret = ret.cast(dtype)
    return ret
  def clone(self, device=None) -> UOp:
    device = canonicalize_device(device or self.device)
    if is_disk_device(device):
      raise RuntimeError("cannot clone DISK storage; use STORE to an explicit disk buffer")
    ret = self.empty_like(device=device)
    src = self if self.device is None or self.device == device else self.copy_to_device(device)
    return ret.after(ret.store(src.cast(ret.dtype)))
  @recursive_property
  def device(self) -> str|tuple[str, ...]|None:
    if self.op is Ops.PARAM: return self.arg.device
    if self.op is Ops.STAGE: return self.src[0].device if self.arg is None else self.arg.device
    if self.op is Ops.AFTER: return self.src[0].device
    if self.op is Ops.MSELECT:
      assert isinstance(self.src[0].device, tuple), f"mselect must be on tuple device, getting {self.src[0].device}"
      return self.src[0].device[self.arg]
    if self.op is Ops.MSTACK: return tuple(cast(str, x.device) for x in self.src)
    if self.op in {Ops.BUFFER, Ops.ALLOC}: return self.arg.device
    if self.op is Ops.COPY: return self.arg
    if self.op is Ops.ALLREDUCE: return self.arg[1]
    for x in self.src:
      if x.device is not None: return x.device
    return None
  @property
  def is_virtual(self) -> bool:
    # NOTE: no device means no place to store, weak means no width to store. neither can back a buffer as-is
    # TODO: unify with has_buffer_identity
    return self.device is None or self.dtype in dtypes.weaks
  @recursive_property
  def addrspace(self) -> AddrSpace|None:
    if self.op is Ops.PARAM: return self.arg.addrspace
    if self.op in {Ops.BUFFER, Ops.ALLOC}: return self.arg.addrspace
    if self.op in {Ops.SPECIAL, Ops.RANGE, Ops.CONST}: return AddrSpace.ALU
    if self.op is Ops.BINARY: return AddrSpace.GLOBAL
    if self.op is Ops.LOAD: return AddrSpace.ALU # LOAD brings things into the ALU
    if self.op in {Ops.INDEX, Ops.CAST, Ops.AFTER, Ops.REDUCE, Ops.STORE, Ops.MSTACK, Ops.MSELECT, Ops.END, Ops.UNSHARD}:
      return self.src[0].addrspace
    if self.op in GroupOp.Movement: return self.src[0].addrspace
    if self.op in {Ops.STACK, Ops.WMMA} or self.op in GroupOp.Elementwise:
      ad = [x.addrspace for x in self.src if x.addrspace is not None]
      if not len(ad) or not all_same(ad): return None
      return ad[0]
    return None
  @property
  def buf_uop(self) -> UOp:
    if self.op in GroupOp.Defines: return self
    if self.op is Ops.MSELECT: return self.src[0].buf_uop.mselect(self.arg)
    if self.op is Ops.MSTACK: return UOp(Ops.MSTACK, src=tuple(x.buf_uop for x in self.src))
    if self.base.op is Ops.AFTER: return self.base.src[0].buf_uop.base
    s = self
    while len(s.src) and s.op not in {Ops.BUFFER, Ops.ALLOC, Ops.PARAM, Ops.STAGE, Ops.MSTACK}: s = s.src[0]
    return s

  def contiguous_view(self) -> tuple[UOp, int]|None:
    from tinygrad.schedule.prepare import pm_mops
    from tinygrad.uop.symbolic import symbolic

    # WEBGPU and CL do not support views.
    # WEBGPU requires that minUniformBufferOffsetAlignment be at least 32 bytes: https://gpuweb.github.io/gpuweb/#adapter-capability-guarantees
    # CL 1.1 provides the clCreateSubBuffer API, but at the time of writing, relevant CL runtimes (rusticl, adreno, nvidia, amd) do not provide
    # reasonable values for CL_DEVICE_MEM_BASE_ADDR_ALIGN. cl_ext_buffer_device_address could potentially help, but this extension is not provided
    # by relevant CL runtimes at time of writing.
    if (dev:=self.device) is not None and any(d.startswith(("WEBGPU", "CL")) for d in ((dev,) if isinstance(dev, str) else dev)): return None

    rng = UOp.range(self.numel(), 0)
    out = graph_rewrite(self.flatten().index(rng), pm_mops+symbolic, name="contiguous_view_offset")
    if out.op is not Ops.INDEX or len(out.src)-1 != len((b:=out.src[0]).shape): return None
    offset = (sum(i*s for i,s in zip(out.src[1:], strides_for_shape(b.shape))) - rng).ssimplify()
    if not isinstance(offset, int): return None
    if b.op is not Ops.BITCAST: return b, offset
    osz, isz = b.element_size(), b.src[0].element_size()
    if (offset*osz) % isz or (self.numel()*osz) % isz: return b, offset
    return b.src[0].flatten()[offset*osz//isz:(offset+self.numel())*osz//isz].contiguous_view()

  def contiguous_view_offset(self) -> int|None: return None if (view := self.contiguous_view()) is None else view[1]

  def has_buffer_identity(self, after_ok=False):
    """Check if this UOp has a storage identity in the graph, whether or not its buffer is bound."""
    # TODO: this is confusing because UOp.variable('v', 0, 1, dtypes.weakfloat) is True for jit to work, but it doesn't have a buffer
    if self.op in {Ops.RESHAPE, Ops.UNSHARD, Ops.MSELECT}: return self.src[0].has_buffer_identity(after_ok)
    if after_ok and self.op == Ops.AFTER: return self.src[0].has_buffer_identity(after_ok)
    return self.op in GroupOp.Defines

  def _base_buffer_is_realized(self) -> bool:
    """Walk through AFTER chain to find if the underlying buffer is realized (has allocated memory)."""
    u = self.base
    while u.op is Ops.AFTER: u = u.src[0]
    return u.is_realized

  @functools.cached_property
  def _buffer_view(self) -> tuple[UOp, int]:
    # Cache only the UOp and byte offset, never an allocated Buffer view.
    if (cv := self.contiguous_view()) is None: raise RuntimeError(f"non-contiguous view is not supported for {self.device} buffer")
    return cv[0], cv[1]*cv[0].dtype.itemsize

  @property
  def buffer(self) -> Buffer|MultiBuffer:
    # a bare STAGE (same-device materialization) keeps the source's buffer
    if self.op is Ops.STAGE and self.arg is None: return self.src[0].buffer
    if self.op in {Ops.CONTIGUOUS_BACKWARD, Ops.RESHAPE, Ops.UNSHARD, Ops.DETACH, Ops.AFTER}: return self.src[0].buffer
    # this buffer can process disk tensors and simple movement ops.
    # NOTE: the view Buffer returned here is transient (short-lived), it only wraps an offset into the base BUFFER's storage
    if self is not self.base or self.op is Ops.BITCAST:
      src, offset = self._buffer_view
      if isinstance(buf:=src.buffer, MultiBuffer):
        mbuf = MultiBuffer.__new__(MultiBuffer)
        mbuf.bufs = [x.view(prod(self.max_shape), self.dtype, offset) for x in buf.bufs]
        return mbuf
      return buf.view(prod(self.max_shape), self.dtype, offset)
    if self.op is Ops.MSELECT:
      ret = self.src[0].buffer
      assert isinstance(ret, MultiBuffer)
      return ret.bufs[self.arg]
    if self.op is Ops.MSTACK:
      ret = MultiBuffer.__new__(MultiBuffer)
      ret.bufs = [cast(Buffer, x.buffer) for x in self.src]
      assert all_same([(x.size, x.dtype) for x in ret.bufs]), "multibuffers mismatch buffers"
      return ret
    assert self.op is Ops.BUFFER and self.arg.buffer is not None, f"must be a realized BUFFER {self}"
    return self.arg.buffer
  @property
  def realized(self) -> Buffer|MultiBuffer|None:
    if self.op is Ops.UNSHARD: return self.src[0].realized
    # only these can be realized
    if self.op not in (Ops.BUFFER, Ops.MSTACK): return None
    # LOCAL/REG scratch buffers are never realized
    if self.op is Ops.BUFFER and self.addrspace in (AddrSpace.LOCAL, AddrSpace.REG): return None
    # an ALLOC (directly or as an MSTACK source) is not realized
    if self.op_in_backward_slice_with_self(Ops.ALLOC): return None
    # NOTE: this is used by the JIT to determine which inputs we capture
    return self.buffer if self.buffer.is_allocated() else None
  @property
  def is_realized(self) -> bool: return self.base.realized is not None

  # *** uop Variable stuff ***

  @staticmethod
  def variable(name:str, min_val:PyConst, max_val:PyConst, dtype:DType=dtypes.weakint, multiple_of:int=1) -> UOp:
    # a Variable is a scalar ALU PARAM with a name and a value range; binding it sets the val payload on the arg
    return UOp(Ops.PARAM, src=(UOp(Ops.STACK),), arg=ParamArg(-1, dtype, name=name, vmin_vmax=(min_val, max_val),
                                                          multiple_of=multiple_of, addrspace=AddrSpace.ALU))
  @property
  def is_variable(self) -> bool:
    # a Variable is a scalar ALU PARAM that carries a value range
    return self.op is Ops.PARAM and self.arg.vmin_vmax is not None and self.arg.addrspace is AddrSpace.ALU and self._shape == ()
  @property
  def is_bound_var(self) -> bool:
    # a bound Variable is a Variable with the val payload set
    return self.is_variable and self.arg.val is not None
  @property
  def expr(self) -> str:
    assert self.op in {Ops.PARAM, Ops.BUFFER}
    return unwrap(self.arg.name)
  def bind(self, val:int|UOp):
    assert self.is_variable and not self.is_bound_var, f"op is {self.op}, need an unbound Variable"
    uval = UOp.const(val) if isinstance(val, int) else val
    assert uval.op is Ops.CONST, f"bind value must be a CONST, not {uval.op}"
    assert self.vmin <= uval.vmin and uval.vmax <= self.vmax, f"bind {val} not in range [{self.vmin}, {self.vmax}]"
    assert uval.divides(self.arg.multiple_of) is not None, f"bind {val} not divisible by {self.arg.multiple_of}"
    return self.replace(arg=replace(self.arg, val=uval.val))
  def unbound(self) -> Variable:
    assert self.is_variable, f"op is {self.op}, need Variable"
    # strip the tag too: tags are kernel-graph processing state, the unbound Variable is the canonical node
    return self.replace(arg=replace(self.arg, val=None), tag=None)
  def unbind(self) -> tuple[Variable, int]:
    assert self.is_bound_var, f"can't unbind {self}"
    return self.unbound(), self.arg.val
  def unbind_all(self) -> tuple[UOp, dict[Variable, int]]:
    bound = {x: x.unbound() for x in self.backward_slice_with_self if x.is_bound_var}
    return self.substitute(bound, walk=True), {v: cast(int, x.arg.val) for x, v in bound.items()}
  def variables(self) -> list[Variable]:
    ret = set()
    for x in self.backward_slice_with_self:
      if x.op is Ops.PARAM and x.addrspace is AddrSpace.ALU:
        ret.add(x.unbound() if x.is_variable else x)
      elif x.op is Ops.RANGE and x.axis_type is AxisType.DEVICE:
        ret.add(UOp.variable("_device_num", 0, x.vmax, dtype=x.dtype))
    return sorted(ret, key=lambda v: (v.arg.name or "", v.arg.slot))

  # *** uop symbolic stuff ***

  def const_factor(self) -> int:
    """largest known int that divides self"""
    # TODO: for negatives it's not the largest
    if self.op is Ops.CONST: return self.val
    if self.op is Ops.STACK: return math.gcd(*[x.const_factor() for x in self.src])
    if self.op is Ops.ADD: return math.gcd(self.src[0].const_factor(), self.src[1].const_factor())
    if self.op is Ops.MUL: return self.src[0].val if self.src[0].op is Ops.CONST else self.src[1].val if self.src[1].op is Ops.CONST else 1
    if self.op in GroupOp.Defines and self.arg.multiple_of is not None: return self.arg.multiple_of
    return 1
  def divides(self, v:int) -> UOp|None:
    if v==1: return self
    if self.op is Ops.CONST: return self.const_like(self.val//v) if self.val%v == 0 else None
    if self.op is Ops.STACK:
      srcs = tuple(s.divides(v) for s in self.src)
      return None if any(s is None for s in srcs) else UOp(Ops.STACK, src=cast(tuple[UOp, ...], srcs))
    if self.op is Ops.ADD: return d0+d1 if (d0:=self.src[0].divides(v)) is not None and (d1:=self.src[1].divides(v)) is not None else None
    if self.op is Ops.MUL:
      if (d0:=self.src[0].divides(v)) is not None: return d0 * self.src[1]
      if (d1:=self.src[1].divides(v)) is not None: return self.src[0] * d1
    if self.op in GroupOp.Defines and self.arg.multiple_of is not None:
      return self // v if self.arg.multiple_of%v == 0 else None
    return None # generic None if we aren't sure
  def pop_const(self, op=Ops.ADD) -> tuple[UOp, PyConst]:  # NOTE: assume Invalid ALU is resolved
    return (self.src[0], self.src[1].val) if self.op is op and self.src[1].op is Ops.CONST else (self, identity_element(op, self.dtype))
  @staticmethod
  def gcd(*uops: UOp) -> UOp:
    # Strip explicit coefficients, not multiple_of, so symbolic factors stay recognizable by divide_exact.
    terms, factors = zip(*[u.pop_const(Ops.MUL) for u in uops])
    count = functools.reduce(operator.and_, [collections.Counter(term.split_uop(Ops.MUL)) for term in terms])
    if not count: factors = tuple(u.const_factor() for u in uops)
    return math.prod(count.elements(), start=uops[0].const_like(math.gcd(*factors)))
  def divide_exact(self, v:UOp) -> UOp|None:
    if self is v: return self.const_like(1)
    if v.op is Ops.CONST: return self.divides(v.val)
    if self.op is Ops.ADD: return None if (s0:=self.src[0].divide_exact(v)) is None or (s1:=self.src[1].divide_exact(v)) is None else s0+s1
    if self.op is Ops.MUL:
      (fac, const), (div_fac, div_const) = self.pop_const(Ops.MUL), v.pop_const(Ops.MUL)
      new_count = collections.Counter(fac.split_uop(Ops.MUL))
      new_count.subtract(div_fac.split_uop(Ops.MUL))
      if const%div_const==0 and all(v>=0 for v in new_count.values()): return math.prod(new_count.elements(), start=self.const_like(const//div_const))
    return None # generic None if we aren't sure
  @property
  def vmin(self) -> PyConst: return self._min_max[0]
  @property
  def vmax(self) -> PyConst: return self._min_max[1]
  @functools.cached_property
  def _min_max(self) -> tuple[PyConst, PyConst]:
    if self.op in GroupOp.Binary and not dtypes.is_float(self.dtype):
      (s0_vmin, s0_vmax), (s1_vmin, s1_vmax) = self.src[0]._min_max, self.src[1]._min_max
      if self.op is Ops.ADD: return s0_vmin+s1_vmin, s0_vmax+s1_vmax
      if self.op is Ops.SUB: return s0_vmin-s1_vmax, s0_vmax-s1_vmin
      if self.op is Ops.AND and dtypes.is_int(self.dtype) and s1_vmin == s1_vmax >= 0:
        return 0, s1_vmax if s0_vmin < 0 else min(s0_vmax, int(s1_vmax) & ((1 << int(s0_vmax).bit_length()) - 1))
      if self.op is Ops.MUL: return min(vals:=(s0_vmin*s1_vmin, s0_vmin*s1_vmax, s0_vmax*s1_vmin, s0_vmax*s1_vmax)), max(vals)
      # SHL/SHR on consts only
      if self.op is Ops.SHL and s1_vmin == s1_vmax and all_int(t:=(s0_vmin, s0_vmax, s1_vmin)): return t[0] << t[2], t[1] << t[2]
      if self.op is Ops.SHR and s1_vmin == s1_vmax and all_int(t:=(s0_vmin, s0_vmax, s1_vmin)): return t[0] >> t[2], t[1] >> t[2]
      if self.op is Ops.CMOD:
        if (c:=s1_vmin) == s1_vmax > 0:
          return (0 if s0_vmin > 0 else s0_vmin if 0 >= s0_vmin > -c else -(s1_vmax-1), 0 if s0_vmax < 0 else s0_vmax if 0 <= s0_vmax < c else c-1)
        if s1_vmin > 0: return (0, s1_vmax-1) if s0_vmin >= 0 else (-(s1_vmax-1), 0) if s0_vmax <= 0 else (-(s1_vmax-1), s1_vmax-1)
        if s1_vmax < 0: return (0, -s1_vmin-1) if s0_vmin >= 0 else (-(-s1_vmin-1), 0) if s0_vmax <= 0 else (-(-s1_vmin-1), -s1_vmin-1)
      if self.op is Ops.CDIV:
        assert isinstance(s0_vmin, int) and isinstance(s0_vmax, int) and isinstance(s1_vmin, int) and isinstance(s1_vmax, int)
        if s1_vmin*s1_vmax>0:
          return min(vals:=(cdiv(s0_vmin, s1_vmin), cdiv(s0_vmin, s1_vmax), cdiv(s0_vmax, s1_vmin), cdiv(s0_vmax, s1_vmax))), max(vals)
      if self.op is Ops.FLOORDIV:
        assert isinstance(s0_vmin, int) and isinstance(s0_vmax, int) and isinstance(s1_vmin, int) and isinstance(s1_vmax, int)
        if s0_vmin > s0_vmax: return 0, 0  # numerator range is empty (e.g. RANGE with end=0)
        if s1_vmin*s1_vmax>0: return min(vals:=(s0_vmin//s1_vmin, s0_vmin//s1_vmax, s0_vmax//s1_vmin, s0_vmax//s1_vmax)), max(vals)
      if self.op is Ops.FLOORMOD:
        assert isinstance(s0_vmin, int) and isinstance(s0_vmax, int) and isinstance(s1_vmin, int) and isinstance(s1_vmax, int)
        if s0_vmin > s0_vmax: return 0, 0  # numerator range is empty (e.g. RANGE with end=0)
        if (c:=s1_vmin) == s1_vmax > 0: return (s0_vmin%c, s0_vmax%c) if s0_vmin//c == s0_vmax//c else (0, c-1)
        if (c:=s1_vmin) == s1_vmax < 0: return (s0_vmin%c, s0_vmax%c) if s0_vmin//c == s0_vmax//c else (c+1, 0)
        if s1_vmin > 0: return (0, s1_vmax-1)
        if s1_vmax < 0: return (s1_vmin+1, 0)
      if self.op is Ops.XOR and s1_vmin == s1_vmax == -1 and isinstance(s0_vmin, int) and isinstance(s0_vmax, int):
        return ~int(s0_vmax), ~int(s0_vmin)
      if self.op is Ops.MAX: return max(s0_vmin, s1_vmin), max(s0_vmax, s1_vmax)
      if self.op is Ops.CMPLT: return (s0_vmax<s1_vmin, s0_vmin<s1_vmax)
      if self.op is Ops.CMPNE: return ((s0_vmax < s1_vmin) or (s1_vmax < s0_vmin), not (s0_vmin == s0_vmax == s1_vmin == s1_vmax))
      if self.op is Ops.OR and self.dtype == dtypes.bool: return s0_vmin or s1_vmin, s0_vmax or s1_vmax
      if self.op is Ops.AND and self.dtype == dtypes.bool: return s0_vmin and s1_vmin, s0_vmax and s1_vmax
    if self.op is Ops.WHERE: return min(self.src[1].vmin, self.src[2].vmin), max(self.src[1].vmax, self.src[2].vmax)
    # NOTE: returned UOp is assumed to be CONST
    if self.op in GroupOp.Defines and self.arg.vmin_vmax is not None: return self.arg.vmin_vmax
    if self.op in (Ops.RANGE, Ops.SPECIAL) and self.dtype is not dtypes.void: return 0, (self.src[0]-1).vmax
    if self.op is Ops.STACK: return min(x.vmin for x in self.src), max(x.vmax for x in self.src)
    # a load from a constant table is one of its values
    if self.op is Ops.LOAD and (b:=self.src[0].buf_uop).op is Ops.BINARY:
      return min(e:=memoryview(b.arg).cast(unwrap(self.dtype.fmt)).tolist()), max(e)
    # a NAN is outside every interval
    if self.op is Ops.CONST and self.val is not Invalid and not (isinstance(self.val, float) and math.isnan(self.val)): return self.val, self.val
    if self.op is Ops.PAD: return min(self.src[0].vmin, 0), max(self.src[0].vmax, 0)  # PAD adds zeros
    if self.op in GroupOp.Movement|{Ops.INDEX, Ops.STAGE, Ops.AFTER, Ops.DETACH, Ops.COPY, Ops.CONTIGUOUS_BACKWARD}: return self.src[0]._min_max
    if self.op is Ops.CAST:
      # rounding is monotone (truncation toward zero into an int, to-nearest onto the value grid into a float)
      smin, smax = self.src[0]._min_max
      trunc = truncate.get(self.dtype) if dtypes.is_float(self.dtype) else math.trunc if dtypes.is_int(self.dtype) else None
      if trunc is not None: smin, smax = (trunc(v) if math.isfinite(v) else v for v in (smin, smax))
      if dtypes.is_unsigned(self.dtype) and 0 <= smin and smax <= self.dtype.max: return smin, smax
      # a signed or float destination holds the part of the source that overlaps it: overflow is undefined, a nan bound overlaps nothing
      if self.dtype in dtypes.floats+dtypes.sints+dtypes.weaks and smin <= self.dtype.max and self.dtype.min <= smax:
        return max(self.dtype.min, smin), min(smax, self.dtype.max)
    return self.dtype.min, self.dtype.max

  @functools.cached_property
  def _sym_fxn(self):
    from tinygrad.uop.render import _render_with_splits, renderer_infer
    sself = self.simplify()
    varnames = tuple(dedup(x.expr for x in sself.toposort() if x.op is Ops.PARAM and x.arg.addrspace == AddrSpace.ALU))
    # TODO: sanitize varnames, or don't use naked eval while staying fast
    ret = _render_with_splits(list(sself.toposort()), renderer_infer, {sself})
    lines = [f"  {k}={v}" for k,v in ret.items() if k != "ast"] + [f"  return {ret['ast']}"]
    ns: dict[str, Any] = {"max": max, "cdiv": cdiv, "cmod": cmod, "floordiv": floordiv, "floormod": floormod, "bitcast": bitcast, "dtypes": dtypes}
    exec(f"def _f({','.join(varnames)}):\n"+'\n'.join(lines), ns)  # pylint: disable=exec-used
    return ns["_f"], varnames

  def sym_infer(self, var_vals:dict[str, int]):
    fxn, varnames = self._sym_fxn
    return fxn(**{k:v for k,v in var_vals.items() if k in varnames})

  def render(self, simplify=True, pm:PatternMatcher|None=None) -> str:
    ctx: dict[UOp, str] = {}
    from tinygrad.uop.render import renderer
    pm = renderer if pm is None else pm
    for u in (s:=self.simplify() if simplify else self).toposort():
      ctx[u] = cast(str, pm.rewrite(u, ctx=ctx))
    return ctx[s]

  def render_uir(self) -> str:
    from tinygrad.uop.render import render_uir
    return render_uir(self)

  # *** uop high level syntactic sugar ***

  @staticmethod
  def alloc(shape:tuple[sint, ...], dtype:DType, slot:int|None=None, addrspace=AddrSpace.GLOBAL, device=None, axis:int|None=None,
            spec:BufferSpec|None=None):
    ret = UOp(Ops.ALLOC, src=(UOp.const(prod(to_max_shape(shape))),)+UOp.device_range_src(device),
              arg=ParamArg(next(UOp.unique_num) if slot is None else slot, strong_dtype(dtype),
                           addrspace=addrspace, device=device, spec=spec))
    return ret.reshape(()) if not shape else ret.view_as(shape, axis)
  def alloc_like(self, slot:int|None=None, addrspace=AddrSpace.GLOBAL): return UOp.alloc(self.max_shard_shape, self.dtype, slot, addrspace)

  @staticmethod
  def placeholder(shape:tuple[int, ...], dtype:DType, slot:int|None=None, addrspace=AddrSpace.GLOBAL, device=None):
    dtype = strong_dtype(dtype)  # storage is never weak: a placeholder commits the width of what's put in it
    if slot is None: slot = next(UOp.unique_num)
    if addrspace is AddrSpace.GLOBAL:
      ret = UOp(Ops.PARAM, src=(UOp.const(prod(shape)),), arg=ParamArg(slot, dtype, addrspace=addrspace, device=device))
    else:
      assert addrspace in (AddrSpace.LOCAL, AddrSpace.REG)
      assert device is None, "LOCAL and REG placeholders cannot have a device"
      ret = UOp(Ops.BUFFER, src=(UOp.const(prod(shape)),), arg=ParamArg(slot, dtype, addrspace=addrspace))
    if len(shape) > 1: ret = ret.reshape(shape)
    return ret
  def placeholder_like(self, slot:int, addrspace=AddrSpace.GLOBAL):
    assert all_int(self.shape), "no placeholder-like on symbolic shape"
    return UOp.placeholder(self.max_shard_shape, self.dtype, slot, addrspace)

  # set is store+end+after
  def set(self:UOp, val:UOp|ConstType, end:UOp|tuple[UOp, ...]|list[UOp]=()) -> UOp:
    return self.src[0].after(self.store(val).end(*argfix(end)))

  # TODO: this should replace placeholder
  @staticmethod
  def param(slot:int, dtype:DType, shape:tuple[sint, ...]|sint|None=None, device=None, vmin_vmax:tuple[PyConst, PyConst]|None=None,
            multiple_of:int|None=None, name=None, addrspace=AddrSpace.GLOBAL, volatile:bool=False):
    """create a PARAM: a single sint or 1-d shape gives a flat param of that size, a None shape gives a scalar param.
    src[0] stores the concrete max size (never symbolic): a multi-dim shape is a RESHAPE on top of the flat param,
    a symbolic shape is a max-size param shrunk to the real shape"""
    if dtype in dtypes.weaks: raise RuntimeError(f"cannot create param for weak dtype {dtype}")
    shape = (shape,) if isinstance(shape, (int, UOp)) else shape or ()
    ret = UOp(Ops.PARAM, src=(shape_to_shape_arg((prod(to_max_shape(shape)),) if shape else ()),),
              arg=ParamArg(slot, dtype, vmin_vmax, multiple_of, name, addrspace, device, volatile))
    return ret.view_as(shape)
  def param_like(self, slot:int, name:str|None=None):
    # Scalar arguments bind by slot; names and values stay at the call site, not in schedule cache keys.
    if self.op is Ops.PARAM and self.addrspace is AddrSpace.ALU:
      return self.replace(arg=replace(self.arg, slot=slot, name=name, val=None))
    # multi-device values become a per-shard sized param wrapped in UNSHARD: the sharding lives in the graph, not the arg
    if self.axis is not None and isinstance(self.device, tuple):
      return UOp(Ops.PARAM, src=(UOp.const(prod(to_max_shape(self.shard_shape))),),
                 arg=ParamArg(slot, self.dtype, name=name, device=self.device)).view_as(self.shard_shape, self.axis)
    return UOp.param(slot, self.dtype, self._shape, self.device, name=name)
  def view_as(self:UOp, shape:tuple[sint, ...], axis:int|None=None) -> UOp:
    """view flat storage as the given (possibly symbolic) shape, optionally sharded on axis, the UNSHARD gives back the multiplied shape"""
    max_shape = to_max_shape(shape)
    ret = self.reshape(max_shape) if len(shape) > 1 else self
    if tuple(max_shape) != tuple(shape): ret = ret.shrink_to(shape)
    return ret if axis is None else ret.unshard(axis)

  @staticmethod
  def custom_function(name:str, *src:UOp, dtype:DType=dtypes.void) -> UOp:
    return UOp(Ops.CUSTOM_FUNCTION, src=src, arg=CustomFunction(name, dtype))

  def call(self, *srcs:UOp, grad_fxn:Callable|None=None,
           name:str|None=None, precompile:bool=False, precompile_backward:bool=False, aux:Any=None) -> UOp:
    """call a body with the given args: a plain CallInfo CALL. all inputs must be ready (buffers/params), this never
    creates ALLOCs: use call_with_outputs for calls that produce values"""
    assert self.op in OPAQUE_CALL_BODIES, f"cannot call a {self.op} body, use call_with_outputs for value-producing bodies"
    # calls are launched per device, so an open DEVICE range is allowed to cross the call boundary
    assert all(r.axis_type is AxisType.DEVICE for r in self.ranges), \
      f"ranges {self.ranges} are leaking out of the call in {self.render_uir()}"
    # an external C call is a CALL on a CUSTOM_FUNCTION body stating the (possibly void) return dtype, the callee
    # (a function pointer) in source, rendered as an indirect call
    return UOp(Ops.CALL, src=(self,)+srcs, arg=CallInfo(grad_fxn, name, precompile, precompile_backward, aux))

  @staticmethod
  def call_with_outputs(values:tuple[UOp, ...], *srcs:UOp, grad_fxn:Callable|None=None,
                        name:str|None=None, precompile:bool=False, precompile_backward:bool=False, aux:Any=None,
                        output_pos:tuple[int, ...]|None=None) -> tuple[UOp, ...]:
    """call a body producing the given values, returning the outputs. the body stores into output PARAMs, and the
    outputs are ALLOCs passed as extra inputs to the call (you AFTER on them like normal buffers).
    the buffers are bound to the output PARAMs positionally wherever the call is resolved, just like the args.
    output_pos gives the position of each output in the arg list (default: a block after the inputs), the inputs take
    the remaining positions in order; when it's given, input params must already be slotted at their final positions.
    output_pos must be strictly ascending: the body's stores and the call args pair positionally by values order"""
    # the device defaults to the first device in the values or args, like srcs-based device resolution
    default_dev = next((x.device for x in itertools.chain(values, srcs) if x.device is not None), None)
    pos = tuple(range(len(srcs), len(srcs)+len(values))) if output_pos is None else output_pos
    assert len(pos) == len(values) and len(set(pos)) == len(pos), "output_pos must be one distinct position per output"
    assert all(a < b for a, b in zip(pos, pos[1:])), f"output_pos {output_pos} must be strictly ascending"
    assert all(0 <= p < len(srcs)+len(values) for p in pos), f"output_pos {output_pos} must be within the arg list"
    # the inputs take the slots not in pos, in order: symbolic output shapes resolve against the final argument slots
    param_map: list[UOp|None] = [None] * (len(srcs) + len(values))
    it = iter(srcs)
    for i in range(len(param_map)):
      if i not in pos: param_map[i] = next(it)
    def mint(o:UOp, p:int) -> tuple[UOp, UOp]:
      # Declare output storage in the callee's shape, then bind its shape in the caller's scope.
      dev = o.device if o.device is not None else default_dev
      axis = o.axis if isinstance(o.device, tuple) else None
      buf = UOp.alloc(o.shard_shape, o.dtype, device=dev, axis=axis)
      shp = tuple(graph_rewrite(s, _pm_resolve_params, param_map, walk=True) if isinstance(s, UOp) else s for s in o.shard_shape)
      return UOp.alloc(shp, o.dtype, slot=buf.buf_uop.arg.slot, device=dev, axis=axis), buf.param_like(p)
    outputs = tuple(mint(o, p) for o, p in zip(values, pos))
    rets = tuple(r for r, _ in outputs)
    body = UOp.sink(*[p.store(v) for v, (_, p) in zip(values, outputs)])
    args: list[UOp|None] = [None] * (len(srcs) + len(values))
    for p, r in zip(pos, rets): args[p] = r
    it = iter(x.contiguous() if precompile else x for x in srcs)
    call = body.call(*[r if r is not None else next(it) for r in args], grad_fxn=grad_fxn, name=name, precompile=precompile,
                     precompile_backward=precompile_backward, aux=aux)
    return tuple(r.after(call) for r in rets)

  # one-line convenience for the single-output case: self is the value
  def call_with_output(self, *srcs:UOp, **kwargs) -> UOp: return UOp.call_with_outputs((self,), *srcs, **kwargs)[0]
  def custom_kernel(*srcs:UOp, fxn:Callable, grad_fxn:Callable|None=None) -> list[UOp]:
    placeholders = [UOp.placeholder_like(s, slot=i) for i,s in enumerate(srcs)]
    kernel = fxn(*placeholders).call(*srcs, grad_fxn=grad_fxn)
    return [s.after(kernel) for s in srcs]

  def to_elf(self) -> TinyELF:
    assert self.op is Ops.PROGRAM and isinstance(self.arg, ProgramInfo), "to_elf should only be called on a PROGRAM ast"
    params = tuple(u for u in self.src[1].src if u.op is Ops.PARAM and u.addrspace != AddrSpace.ALU)
    # sig slots are compact: buffers in globals order (runtimes launch buffers in that order), then vars. raw call-arg
    # positions skip buffers for kernels using a sparse subset of the call's buffers (CL binds bufs[slot])
    gmap = {s:j for j, s in enumerate(self.arg.globals)}
    sig = tuple((u.arg.name, gmap[u.arg.slot], u.dtype, u._shape) for u in params) + \
          tuple((v.arg.name, len(self.arg.globals)+j, v.dtype, v._shape) for j, v in enumerate(self.arg.vars))
    return TinyELF(self.src[3].arg, self.src[0].arg.function_name, self.arg.target, sig, self.key)

  @property
  def src_without_body(self) -> tuple[UOp, ...]: return self.src[1:] if self.op is Ops.CALL else self.src

def uopfunc(fn:Callable[..., UOp]) -> Callable[..., UOp]: # sugar for body.call(*args): uop args become params
  def param(i:int, n:str, a:UOp) -> UOp:
    if a.addrspace in (None, AddrSpace.ALU): return UOp.param(i, a.commit_dtype(dtypes.int), name=n, addrspace=AddrSpace.ALU)
    return UOp.param(i, a.dtype, 1 if a.op is Ops.INDEX else a.max_numel(), a.device, name=n, addrspace=a.addrspace)
  def outlined(*args, **kwargs) -> UOp:
    bound = inspect.signature(fn).bind(*args, **kwargs).arguments
    ins = {n: a for n, a in bound.items() if isinstance(a, UOp)}
    f = graph_rewrite(fn(**(bound|{n: param(i, n, a) for i, (n, a) in enumerate(ins.items())})), pm_renumber_slots, ctx=itertools.count(), walk=True)
    with Context(TRACK_MATCH_STATS=0): return f.call(*ins.values(), name=fn.__name__)
  return functools.wraps(fn)(outlined)

@dataclass(frozen=True)
class KernelInfo:
  name: str = "test"            # name of the kernel
  applied_opts: tuple = tuple()
  opts_to_apply: tuple|None = None
  estimates: Estimates|None = None
  beam: int = 0
  @property
  def function_name(self): return to_function_name(self.name)

@dataclass(frozen=True)
class ProgramInfo:
  global_size: tuple[int|float, ...] = (1, 1, 1)
  local_size: tuple[int, ...] = (1, 1, 1)
  vars: tuple[UOp, ...] = ()
  globals: tuple[int, ...] = ()
  outs: tuple[int, ...] = ()
  ins: tuple[int, ...] = ()
  target: Target = Target()

  def launch_dims(self, var_vals:dict[str, int]) -> tuple[tuple[int, ...], tuple[int, ...]]:
    global_size = tuple([sym_infer(sz, var_vals) for sz in self.global_size])  # type: ignore[arg-type]
    local_size = tuple([sym_infer(sz, var_vals) for sz in self.local_size])
    return global_size, local_size

  def vals(self, var_vals:dict[str, int]) -> tuple[int, ...]:
    try: return tuple(var_vals[k.expr] for k in self.vars)
    except KeyError as e: raise RuntimeError(f"unbound Variable {e}") from None

  @staticmethod
  def from_sink(sink:UOp, target:Target=Target()) -> ProgramInfo:
    _vars: list[UOp] = []
    _globals: list[int] = []
    outs: list[int] = []
    ins: list[int] = []
    global_size: list[int] = [1, 1, 1]
    local_size: list[int] = [1, 1, 1]
    for u in sink.toposort(enter_calls=False):
      if u.op is Ops.PARAM and u.addrspace == AddrSpace.ALU: _vars.append(u)
      if u.op is Ops.PARAM and u.addrspace != AddrSpace.ALU: _globals.append(u.arg.slot)
      if u.op in (Ops.STORE, Ops.LOAD):
        if (idx:=u.src[0]).op in (Ops.INDEX, Ops.SHRINK) or (u.src[0].op is Ops.CAST and (idx:=u.src[0].src[0]).op is Ops.INDEX):
          if (buf:=idx.src[0].buf_uop).op is Ops.PARAM: (outs if u.op is Ops.STORE else ins).append(buf.arg.slot)
      if u.op is Ops.SPECIAL: (local_size if u.arg[0] == 'l' else global_size)[int(u.arg[-1])] = cast(int, u.src[0].ssimplify())
    if not outs and not ins: outs = ins = _globals # if neither is inferred, default to all buffers
    return ProgramInfo(tuple(global_size), tuple(local_size),
                       tuple(sorted(_vars, key=lambda v: v.arg.slot)), tuple(sorted(dedup(_globals))), tuple(sorted(dedup(outs))),
                       tuple(sorted(dedup(ins))), target)

# the body of a CALL is always one of these: programs (SINK/PROGRAM/LINEAR), bulk stores, and function references
OPAQUE_CALL_BODIES = {Ops.SINK, Ops.PROGRAM, Ops.LINEAR, Ops.STORE, Ops.CUSTOM_FUNCTION}

# the arg of CUSTOM_FUNCTION: an external function symbol and the dtype of the value it returns
@dataclass(frozen=True)
class CustomFunction:
  name: str
  dtype: DType = dtypes.void

@dataclass(frozen=True)
class CallInfo:
  grad_fxn: Callable|None = None
  name: str|None = None
  precompile: bool = False
  precompile_backward: bool = False
  aux: Any = None
  # grad_fxn can't be pickled
  def __reduce__(self): return (CallInfo, (None, self.name, self.precompile, self.precompile_backward, self.aux))
  def __repr__(self):
    gf = id(self.grad_fxn) if self.grad_fxn else None
    return f"CallInfo({gf}, {repr(self.name)}, {self.precompile}, {self.precompile_backward})"

# ******** ops in python ********

def safe_exp2(x):
  try: return 2 ** x
  except OverflowError: return math.inf

def safe_pow(x, y):
  if isinstance(x, int) and isinstance(y, int) and y < 0: return x**(y%2) if abs(x) == 1 else 0
  try: return math.nan if isinstance(p:=pow(x, y), complex) else p
  except ZeroDivisionError: return math.inf
  except ValueError: return math.inf if x > 0 else -math.inf

python_alu: dict[Ops, Callable]  = {
  Ops.LOG2: lambda x: math.log2(x) if x > 0 else -math.inf if x == 0 else math.nan, Ops.EXP2: safe_exp2,
  Ops.SQRT: lambda x: math.sqrt(x) if x >= 0 else math.nan, Ops.RECIPROCAL: lambda x: 1/x if x != 0 else math.copysign(math.inf, x),
  Ops.SIN: lambda x: math.sin(x) if not math.isinf(x) else math.nan, Ops.POW: safe_pow,
  Ops.TRUNC: lambda x: math.trunc(x) if math.isfinite(x) else x,
  Ops.NEG: operator.neg, Ops.ADD: operator.add, Ops.SUB: operator.sub, Ops.MUL: operator.mul, Ops.CMPNE: operator.ne, Ops.CMPLT: operator.lt,
  Ops.XOR: operator.xor, Ops.OR: operator.or_, Ops.AND: operator.and_, Ops.SHR: operator.rshift, Ops.SHL: operator.lshift, Ops.MAX: max,
  Ops.CMOD: cmod, Ops.CDIV: cdiv, Ops.FLOORDIV: floordiv, Ops.FLOORMOD: floormod,
  Ops.MULACC: lambda x,y,z: (x*y)+z, Ops.WHERE: lambda x,y,z: y if x else z, Ops.CMPEQ: operator.eq}

def exec_alu(op:Ops, dtype:DType, operands, truncate_output=True):
  if any(isinstance(x, tuple) for x in operands):
    count = max(len(x) for x in operands if isinstance(x, tuple))
    return tuple([exec_alu(op, dtype, [x[i] if isinstance(x, tuple) else x for x in operands]) for i in range(count)])
  if op in GroupOp.Binary and Invalid in operands: return Invalid
  alu = python_alu[op](*operands)
  if truncate_output and (truncate_fxn:=truncate.get(dtype)) is not None: return truncate_fxn(alu)
  return alu

# ***** pattern matcher *****

def get_location() -> tuple[str, int]:
  frm = sys._getframe(1)
  # skip over ops.py and anything in mixin
  while frm.f_back is not None and not frm.f_back.f_code.co_filename.startswith("<frozen"):
    fn = frm.f_code.co_filename.replace("\\", "/")
    if not (fn.endswith("/ops.py") or "/mixin/" in fn): break
    frm = frm.f_back
  return frm.f_code.co_filename, frm.f_lineno

class UPat(RandMixin):
  __slots__ = ("op", "match_dtype", "match_tag", "arg", "name", "src", "is_any")
  def __init__(self, op:Ops|tuple[Ops, ...]|set[Ops]|None=None, dtype:DType|tuple[DType, ...]|set[DType]|None=None,
               src:tuple[UPat, ...]|list[UPat]|UPat|None=None, arg:Any=None,
               name:str|None=None, allow_any_len:bool=False, custom_early_reject:set[Ops]|None=None, location=None, is_any:bool=False,
               tag:Any=None):
    assert op is None or isinstance(op, (Ops, tuple, set)), f"op must be Ops or tuple of Ops, not {op!r}"
    self.op: tuple[Ops, ...]|None = (op,) if isinstance(op, Ops) else (tuple(op) if isinstance(op, set) else op)
    self.match_dtype: tuple[DType, ...]|None = (dtype,) if isinstance(dtype, DType) else (tuple(dtype) if isinstance(dtype, set) else dtype)
    self.match_tag: tuple[Any, ...]|None = (tag,) if isinstance(tag, str) else (tuple(tag) if isinstance(tag, set) else tag)
    self.arg, self.name, self._in_src, self.custom_early_reject = arg, name, src, custom_early_reject
    self.src: Any = None
    self.is_any = is_any
    assert self.name != "ctx", "UPat can't be named ctx"
    assert dtype is None or isinstance(dtype, DType) or all(isinstance(x, DType) for x in dtype), f"invalid dtype {dtype}"

    # try all permutations if it's a list
    if isinstance(src, list): self.src = list(itertools.permutations(src)) if not all_same(src) else [tuple(src)]
    # only one if it's a tuple
    elif isinstance(src, tuple): self.src = [src]
    # repeat if it's a UPat
    elif isinstance(src, UPat): self.src = [itertools.repeat(src)]

    self.strict_length = not (allow_any_len or isinstance(src, UPat) or src is None)
    self.required_len: int = 0 if isinstance(src, UPat) or src is None else len(src)
    self.location = location or get_location()

    if custom_early_reject is not None: self.early_reject = custom_early_reject
    else:
      upat_match = [src] if isinstance(src, UPat) else ([] if src is None else self.src[0])
      self.early_reject = {pp.op[0] for pp in upat_match if pp.op is not None and len(pp.op) == 1}

  @property
  def dtype(self) -> DType: return self.match_dtype[0] if self.match_dtype is not None else dtypes.void

  def __reduce__(self):
    return UPat, (self.op, self.match_dtype, self._in_src, self.arg, self.name, not self.strict_length, self.custom_early_reject, self.location,
                  self.is_any, self.match_tag)
  def named(self, name:str):
    return UPat(self.op, self.match_dtype, self._in_src, self.arg, name, not self.strict_length, self.custom_early_reject, tag=self.match_tag)

  @staticmethod
  def any(*src): return UPat(src=src, is_any=True)
  def or_casted(self, name:str|None=None): return UPat.any(self if name is None else self.named(name), UPat(Ops.CAST, name=name, src=(self,)))
  def or_bitcasted(self, name:str|None=None): return UPat.any(self if name is None else self.named(name), UPat(Ops.BITCAST, name=name, src=(self,)))
  def or_after(self, name:str|None=None):
    return UPat.any(self if name is None else self.named(name), UPat(Ops.AFTER, name=name, src=(self,), allow_any_len=True))
  @staticmethod
  @functools.cache
  def var(name:str|None=None, dtype:DType|tuple[DType, ...]|None=None): return UPat(dtype=dtype, name=name)
  @staticmethod
  @functools.cache
  def cvar(name:str|None=None, dtype:DType|tuple[DType, ...]|None=None, arg=None): return UPat(Ops.CONST, dtype, name=name, arg=arg)
  @staticmethod
  def const(b:ConstType, dtype:DType|tuple[DType, ...]|None=None): return UPat(Ops.CONST, dtype=dtype, arg=b)
  @staticmethod
  def custom_function(fn:str, **kwargs): return UPat(Ops.CUSTOM_FUNCTION, arg=CustomFunction(fn), **kwargs)

  # lil helper
  def f(self, op, **kwargs): return UPat(op, src=(self,), **kwargs)

  # copied from UOp
  def sink(*srcs:UPat|None, **kwargs):  # pylint: disable=no-self-argument
    return UPat(Ops.SINK, src=tuple([x for x in srcs if x is not None]), **kwargs)
  def index(self, *srcs:UPat|None, **kwargs):
    return UPat(Ops.INDEX, src=(self,)+tuple(x for x in srcs if x is not None), **kwargs)
  def cast(self, dtype=None, **kwargs):
    if dtype is not None and self.match_dtype == (dtype,): return self
    return UPat(Ops.CAST, dtype, (self,), **kwargs)
  def bitcast(self, dtype=None): return UPat(Ops.BITCAST, dtype, (self,))
  def load(self, *src:UPat, **kwargs): return UPat(Ops.LOAD, src=(self,)+src, **kwargs)
  def store(self, *src:UPat, **kwargs): return UPat(Ops.STORE, src=(self,)+src, **kwargs)
  def reduce(self, *src:UPat, **kwargs):
    arg = kwargs.pop('arg', None)
    if isinstance(arg, Ops): arg = (arg, 0)
    return UPat(Ops.REDUCE, self.match_dtype, src=(self,)+src, arg=arg, **kwargs)
  def broadcast(self, **kwargs): return UPat(Ops.STACK, self.match_dtype, src=self, **kwargs)
  def after(self, *src:UPat, **kwargs): return UPat(Ops.AFTER, self.match_dtype, (self,)+src, **kwargs)
  def end(self, *src:UPat, **kwargs): return UPat(Ops.END, src=(self,)+src, **kwargs)
  def backedge(self, loop:UPat, cond:UPat, **kwargs): return UPat(Ops.BACKEDGE, src=(self, loop, cond), **kwargs)

  def _broadcasted(self, y, reverse=False) -> tuple[UPat, UPat]:
    y = self.ufix(y)
    return (y, self) if reverse else (self, y)
  def ufix(self, x): return UPat.cvar(arg=x) if not isinstance(x, UPat) else x
  def __floordiv__(self, x): return self._binop(Ops.FLOORDIV, x, False)
  def __rfloordiv__(self, x): return self._binop(Ops.FLOORDIV, x, True)
  def mod(self, x, reverse=False): return self._binop(Ops.FLOORMOD, x, reverse)
  def alu(self, op:Ops, *src:UPat):
    asrc = (self,)+src
    return UPat(op, dtypes.bool if op in {Ops.CMPLT, Ops.CMPNE} else asrc[-1].match_dtype, list(asrc) if op in GroupOp.Commutative else asrc)

  def match(self:UPat, uop:UOp, store:dict[str, UOp]) -> list[dict[str, UOp]]:
    if self.is_any: return flatten([x.match(uop, store.copy()) for x in self.src[0]])
    if (self.op is not None and uop.op not in self.op) or \
       (self.name is not None and store.setdefault(self.name, uop) is not uop) or \
       (self.match_dtype is not None and uop.dtype not in self.match_dtype) or \
       (self.arg is not None and self.arg != uop.arg) or \
       (self.match_tag is not None and uop.tag not in self.match_tag) or \
       (len(uop.src) < self.required_len) or \
       (self.strict_length and len(uop.src) != self.required_len): return []
    if self.src is None: return [store]
    res: list[dict[str, UOp]] = []
    for vp in self.src:
      stores, new_stores = [store.copy()], []
      for uu, vv in zip(uop.src, vp):
        for s in stores: new_stores.extend(vv.match(uu, s))
        stores, new_stores = new_stores, []
      res.extend(stores)
    return res

def deconstruct_function(fxn:Callable) -> tuple:
  # globals can be referenced from arbitrarily nested code objects (comprehensions/lambdas, pre PEP 709)
  def names(co:types.CodeType) -> set: return set(co.co_names).union(*(names(c) for c in co.co_consts if isinstance(c, types.CodeType)))
  new_globals = {k:v for k,v in fxn.__globals__.items() if k in names(fxn.__code__)}
  # NOTE: optional round trip through pickle!
  assert fxn.__closure__ is None, "closures are not supported in pattern matchers"
  ret = fxn.__code__, new_globals, fxn.__name__, fxn.__defaults__
  return pickle.loads(pickle.dumps(ret)) if getenv("TEST_PICKLE") else ret

@functools.cache
def upat_interpret(p:UPat, fxn:Callable) -> Callable:
  real_fxn = types.FunctionType(*deconstruct_function(fxn))
  if 'ctx' in inspect.signature(real_fxn).parameters:
    def universal_match(uop, ctx):
      for match in p.match(uop, {}):
        if (ret:=real_fxn(ctx=ctx, **match)) is not None: return ret  # pylint: disable=not-callable
      return None
  else:
    def universal_match(uop, _):
      for match in p.match(uop, {}):
        if (ret:=real_fxn(**match)) is not None: return ret  # pylint: disable=not-callable
      return None
  return universal_match

def upat_deferred_compile(p:UPat, fxn:Callable, entry:list) -> Callable:
  def lazy_compile(uop, ctx):
    from tinygrad.uop.upat import upat_compile
    entry[1] = upat_compile(p, fxn)
    return entry[1](uop, ctx)
  return lazy_compile

class PatternMatcher:
  def __init__(self, patterns:Sequence[tuple[UPat, Callable|tuple]], compiled=bool(getenv("UPAT_COMPILE", 1))):
    # if this comes from a pickle, we reconstruct the lambda functions here
    self.patterns:list[tuple[UPat, Callable]] = [(p,types.FunctionType(*fxn) if isinstance(fxn, tuple) else fxn) for p,fxn in patterns]
    # NOTE: use of DefaultDict here is very dangerous! all keys will live for the lifetime of the PatternMatcher!
    self.pdict: dict[Ops, list[list]] = {}
    # uop is required, arg is optional
    for p,fxn in self.patterns:
      assert p.op is not None
      entry: list = [p, None, p.early_reject]
      entry[1] = upat_deferred_compile(p, fxn, entry) if compiled else upat_interpret(p, fxn)
      for uop in p.op: self.pdict.setdefault(uop, []).append(entry)

  def __reduce__(self): return PatternMatcher, ([(x,deconstruct_function(fxn) if fxn.__name__ == "<lambda>" else fxn) for x,fxn in self.patterns],)

  @functools.cache  # pylint: disable=method-cache-max-size-none
  def __add__(self, more:PatternMatcher) -> PatternMatcher: return PatternMatcher(self.patterns+more.patterns)

  def rewrite(self, uop:UOp, ctx=None):
    if len(pats:=self.pdict.get(uop.op, [])):
      if (ler:=uop.__dict__.get('_src_ops')) is None: uop.__dict__['_src_ops'] = ler = {u.op for u in uop.src}
      for _,match,early_reject in pats:
        if not early_reject.issubset(ler): continue
        if (ret:=match(uop, ctx)) is not None and ret is not uop: return ret
    return None

# *** tracking pattern matcher ***

TRACK_MATCH_STATS = ContextVar("TRACK_MATCH_STATS", 2 if VIZ else 0)
REWRITE_STACK_LIMIT = ContextVar("REWRITE_STACK_LIMIT", 250000)
match_stats:dict[UPat, list[int|float]] = dict()

# TRACK_MATCH_STATS>=2 or VIZ=1 saves all matches
ucount = itertools.count()
uop_fields:dict[int, tuple] = {}

@dataclass(frozen=True)
class TrackedGraphRewrite:
  loc:tuple[str, int]                           # location that called graph_rewrite
  sink:int                                      # the sink input to graph_rewrite
  matches:list[tuple[int, int, tuple, float]]   # before/after UOp, UPat location and time
  name:str                                      # name of the rewrite
  depth:int                                     # depth if it's a subrewrite
  bottom_up:bool
  walk:bool
  enter_calls:bool

tracked_keys:list[TracingKey] = []
tracked_ctxs:list[list[TrackedGraphRewrite]] = []
_name_cnt:dict[str, itertools.count] = {}

if CAPTURE_PROCESS_REPLAY:
  replay_capture: list[bytes] = []
  import atexit, uuid
  @atexit.register
  def save_to_diskcache():
    uid = uuid.uuid4() # one id per process
    for i,v in enumerate(replay_capture): diskcache_put("process_replay", f"{uid}_{i}", v, prepickled=True)

def add_trace_group(kt:TracingKey) -> None:
  tracked_keys.append(kt)
  tracked_ctxs.append([])

active_group:list[int] = []
active_rewrites:list[TrackedGraphRewrite] = []
def rewrite_group(name:Callable[..., str|TracingKey]|bool=True, replay:bool=False, new_ctx:bool=True):
  if not new_ctx: assert not callable(name) and not replay, "name fxn and replay are only supported for new_ctx groups"
  def _decorator(func):
    def __wrapper(*args, **kwargs):
      # without tracking, we just call the function (unless top-level, which always profiles)
      if TRACK_MATCH_STATS < 2 and not new_ctx: return func(*args, **kwargs)
      fn = key = func.__name__
      idx = -1
      if TRACK_MATCH_STATS >= 2:
        if new_ctx:
          add_trace_group(key:=TracingKey(n:=f"{fn} n{next(_name_cnt.setdefault(fn, itertools.count(1)))}", (n,)))
          active_group.append(idx:=len(tracked_keys)-1)
        else:
          rewrite_name = str(kwargs.get("name", None) or fn)
          assert args and isinstance(args[0], UOp), f"invalid match tracing inputs for {rewrite_name} with {args}"
          loc = ((frm:=sys._getframe(1)).f_code.co_filename, frm.f_lineno)
          depth = len(active_rewrites)
          if not tracked_ctxs: add_trace_group(TracingKey(f"default {fn}"))
          dest_group = active_group[-1] if active_group else len(tracked_ctxs)-1
          tracked_ctxs[dest_group].append(ctx:=TrackedGraphRewrite(loc, args[0].trace_num, [], rewrite_name, depth, kwargs.get("bottom_up", False),
                                                                   kwargs.get("walk", False), kwargs.get("enter_calls", False)))
          active_rewrites.append(ctx)
          key = rewrite_name  # profile spans are named after the rewrite step
      with cpu_profile(key, "TINY") as e:
        ret = func(*args, **kwargs)
      if TRACK_MATCH_STATS >= 2:
        if new_ctx: active_group.pop()
        else: active_rewrites.pop()
        if callable(name):
          name_ret = name(*args, **kwargs, ret=ret)
          assert isinstance(name_ret, (TracingKey, str)), f"name function returned {type(name_ret)}"
          tracked_keys[idx] = k = TracingKey(n:=tracked_keys[idx].display_name.replace(fn, name_ret), (n,)) if isinstance(name_ret, str) else name_ret
          e.name = TracingKey(k.display_name if isinstance(name_ret, str) else f"{fn} for {k.display_name}", k.keys)
      if CAPTURE_PROCESS_REPLAY and replay:
        # find the unittest frame we're capturing in
        frm = sys._getframe(1)
        while (f_back:=frm.f_back) is not None and f_back.f_globals.get("__name__", "").split(".")[0] not in ("unittest", "_pytest"):
          frm = f_back
        replay_loc = f"{frm.f_code.co_filename.split('/')[-1]}:{frm.f_lineno} {frm.f_code.co_name}"
        # capture global context vars and all the args passed in
        inputs = (fn, args, kwargs, ContextVar._cache)
        replay_capture.append(pickle.dumps(inputs+(replay_loc, ret)))
      return ret
    return __wrapper
  return _decorator

class TrackedPatternMatcher(PatternMatcher):
  def rewrite(self, uop:UOp, ctx=None):
    if len(pats:=self.pdict.get(uop.op, [])):
      ret = None
      ler = {u.op for u in uop.src}
      for p,match,early_reject in pats:
        if p not in match_stats: match_stats[p] = [0,0,0.0,0.0]
        st = time.perf_counter()
        if not early_reject.issubset(ler):
          match_stats[p][2] += time.perf_counter()-st
          continue
        match_stats[p][1] += 1
        try: ret = match(uop, ctx)
        except Exception as e:
          if TRACK_MATCH_STATS >= 2 and active_rewrites:
            err_str = f"{type(e).__name__}\n{sys.exc_info()[1]}"
            active_rewrites[-1].matches.append((uop.trace_num, UOp(Ops.REWRITE_ERROR, src=uop.src, arg=err_str).trace_num, p.location, 0))
          raise
        if ret is not None and ret is not uop:
          match_stats[p][0] += 1
          match_stats[p][3] += (et:=time.perf_counter()-st)
          if TRACK_MATCH_STATS >= 3: print(f"{et*1e6:7.2f} us -- ", printable(p.location))
          if TRACK_MATCH_STATS >= 2 and isinstance(ret, UOp) and active_rewrites:
            active_rewrites[-1].matches.append((uop.trace_num, ret.trace_num, p.location, et))
          return ret
        match_stats[p][2] += time.perf_counter()-st
    return None

@dataclass(frozen=True)
class RewriteTrace: keys:list[TracingKey]; rewrites:list[list[TrackedGraphRewrite]]; uop_fields:dict[int, tuple] # noqa: E702

if TRACK_MATCH_STATS or PROFILE:
  PatternMatcher = TrackedPatternMatcher  # type: ignore
  import atexit
  @atexit.register
  def print_match_stats():
    if TRACK_MATCH_STATS >= 2:
      with open(fn:=temp("rewrites.pkl", append_user=True), "wb") as f:
        print(f"rewrote {len(tracked_ctxs)} graphs and matched {sum(len(r.matches) for x in tracked_ctxs for r in x)} times, saved to {fn}")
        pickle.dump(RewriteTrace(tracked_keys, tracked_ctxs, uop_fields), f)
    if getenv("PRINT_MATCH_STATS", int(TRACK_MATCH_STATS.value and not VIZ)):
      ret = [0,0,0.0,0.0]
      for k,v in sorted(list(match_stats.items()), key=lambda x: x[1][2]+x[1][3]):
        loc_str = f"{k.location[0].split('/')[-1]}:{k.location[1]}"
        if v[1] != 0: print(f"{v[0]:6d} / {v[1]:7d} -- {v[3]*1000.:9.2f} / {(v[2]+v[3])*1000.:9.2f} ms -- {loc_str:20s}", printable(k.location))
        ret = [x+y for x,y in zip(ret, v)]
      print(f"{ret[0]:6d} / {ret[1]:7d} -- {ret[3]*1000.:9.2f} / {(ret[2]+ret[3])*1000.:9.2f} ms -- TOTAL")
      print(f"{len(match_stats)} rules, {sum(v[0] > 0 for v in match_stats.values())} matched once")
    TRACK_MATCH_STATS.value = 0
    launch_viz("REWRITE_DATA", temp("rewrites.pkl", append_user=True))

  def launch_viz(env_str:str, data:str):
    os.environ[f"{env_str}_DATA"] = data
    if not TRACK_MATCH_STATS and not PROFILE:
      os.environ["VIZ"], os.environ["PROFILE"], os.environ["TRACK_MATCH_STATS"] = "0", "0", "0"
      args = ['--rewrites-path', os.getenv("REWRITE_DATA", "")] if os.getenv("REWRITE_DATA", "") else []
      args += ['--profile-path', os.getenv("PROFILE_DATA", "")] if os.getenv("PROFILE_DATA", "") else []
      viz_path = pathlib.Path(__file__).resolve().parent.parent / "viz" / "serve.py"
      if VIZ > 0 and sys.stdout.isatty(): os.execv(sys.executable, [sys.executable, viz_path.as_posix()] + args)
      if VIZ: print("saved viz files, view using: python -m tinygrad.viz.cli")
      VIZ.value = 0

# *** simple graph rewrite engine ***

# A pure Python sentinel, but *typed* as UOp so it fits all the dict annotations
SENTINEL: Final[UOp] = cast(UOp, object())
class BottomUpGate(Exception): pass
class RewriteContext:
  def __init__(self, pm, bpm, ctx=None, enter_calls=False):
    self.pm: PatternMatcher|None = pm
    self.bpm: PatternMatcher|None = bpm
    self.bpm_cache: dict[UOp, UOp|None] = {}
    self.ctx = ctx
    self.replace: dict[UOp, UOp] = {}
    self.enter_calls = enter_calls

  # no cache needed: pm_rewrite is called at most once per UOp due to the replace dict check in unified_rewrite
  def pm_rewrite(self, x:UOp) -> UOp|None: return unwrap(self.pm).rewrite(x, self.ctx)

  def cached_bpm_rewrite(self, x:UOp) -> UOp|None:
    if (ret:=self.bpm_cache.get(x,SENTINEL)) is not SENTINEL: return ret
    ret = self.bpm_cache[x] = unwrap(self.bpm).rewrite(x, self.ctx)
    return ret

  def walk_rewrite(self, root:UOp) -> UOp:
    """MLIR-style Walk Pattern Rewrite Driver: single-pass, no re-traversal into rewritten subtrees."""
    stack: list[tuple[UOp, bool]] = [(root, False)]
    while stack:
      n, processed = stack.pop()
      if n in self.replace: continue
      if not processed:
        # bottom-up: try bpm on original node first, if it rewrites, use result as-is (no traversal into replacement)
        if self.bpm is not None and (rewritten:=self.cached_bpm_rewrite(n)) is not None:
          self.replace[n] = rewritten
          continue
        # no rewrite, process children then come back to rebuild
        stack.append((n, True))
        # CALL bodies are never rewritten separately, rewrites that need them pass enter_calls=True
        for x in reversed(n.src[1:] if n.op is Ops.CALL and not self.enter_calls else n.src):
          if x not in self.replace: stack.append((x, False))
      else:
        # rebuild node with rewritten srcs
        skip = int(n.op is Ops.CALL and not self.enter_calls)
        new_src = n.src[:skip] + tuple(self.replace.get(x, x) for x in n.src[skip:])
        new_n = UOp(n.op, new_src, n.arg, n.tag) if new_src != n.src else n
        # top-down: try pm on rebuilt node, use result as-is (no re-traversal)
        if self.pm is not None and (rewritten:=self.pm_rewrite(new_n)) is not None: new_n = rewritten
        self.replace[n] = new_n
    return self.replace.get(root, root)

  def unified_rewrite(self, root:UOp) -> UOp:
    stack: collections.deque[tuple[UOp, int, UOp]] = collections.deque([(root, 0, root)])
    on_stack = {root}  # all UOps either on the stack or in self.replace, i.e. dont have to be placed again
    waitlist: dict[UOp, list[tuple[UOp, int, UOp]]] = {}  # UOps waiting on a dependency to be in self.replace
    while stack:
      if len(stack) > REWRITE_STACK_LIMIT: raise RuntimeError("infinite loop in graph_rewrite (stack too big)")
      n, stage, new_n = stack.pop()
      if n in self.replace: continue  # skip any nodes we have seen
      if stage == 0:
        # if bottom up, we rewrite this node early. in both cases, we add its srcs to the stack
        if self.bpm is not None:
          # apply rewrite rules until a fixed point is reached. may return `uop` itself if PatternMatcher doesn't match
          test_n: UOp|None = n
          seen = set()
          try:
            while test_n is not None:
              if test_n in seen: raise RuntimeError("infinite loop in fixed_point_rewrite")
              seen.add(test_n)
              new_n, test_n = test_n, self.cached_bpm_rewrite(test_n)
          except BottomUpGate:
            # if the bpm matching raised a gate, we are done with this node and dont continue down the srcs
            self.replace[n] = unwrap(test_n)
            if n in waitlist: stack.extend(waitlist.pop(n))
            continue
        stack.append((n, 1, new_n))
        # NOTE: CALLs are handled as a special case: their bodies are not included in the graph_rewrite,
        # rewrites that need them pass enter_calls=True
        for x in reversed(new_n.src[1:] if new_n.op is Ops.CALL and not self.enter_calls else new_n.src):
          if x in on_stack: continue
          stack.append((x, 0, x))
          on_stack.add(x)
      elif stage == 1:
        skip = int(new_n.op is Ops.CALL and not self.enter_calls)
        tmp = list(new_n.src[:skip])
        for x in new_n.src[skip:]:
          if (rx:=self.replace.get(x, SENTINEL)) is SENTINEL:
            # source not ready: register in waitlist instead of spinning
            waitlist.setdefault(x, []).append((n, 1, new_n))
            break
          tmp.append(rx)
        else:
          # in stage 1, once all srcs are rewritten, rebuild (if changed) or run top-down rewrite
          if (new_src:=tuple(tmp)) == new_n.src:
            # if top down, do the rewrite. if no rewrite or bottom up, we are done rewriting this node so we add it to the dict
            if self.pm is None or (new_src_n:=self.pm_rewrite(new_n)) is None:
              self.replace[n] = new_n
              if n in waitlist: stack.extend(waitlist.pop(n))
              continue
          else:
            # if srcs changed from rewrites, construct a new UOp with the new srcs
            new_src_n = UOp(new_n.op, new_src, new_n.arg, new_n.tag)
          # trigger a rewrite of new_src_n, then after that rewrite is done, link it back to n
          stack.append((n, 2, new_src_n))
          stack.append((new_src_n, 0, new_src_n))
      else:
        # in stage 2, we link the result of new_n to the result of n
        if (replaced_new_n:=self.replace.get(new_n, SENTINEL)) is SENTINEL:
          # not ready: register in waitlist instead of spinning
          waitlist.setdefault(new_n, []).append((n, 2, new_n))
        else:
          # otherwise we are done
          self.replace[n] = replaced_new_n
          if n in waitlist: stack.extend(waitlist.pop(n))
    if root not in self.replace:
      def label(u:UOp) -> str: return f"{u.op.name}@{id(u):x}"
      details = [f"  {label(n)} -> {label(new_n)} waits for {label(dep)}"
                 for dep, waiters in itertools.islice(waitlist.items(), 5) for n, _, new_n in waiters[:1]]
      raise RuntimeError("graph_rewrite stalled: unresolved rewrite dependencies (possible cycle). "
                         "A replacement may depend on the node being rewritten.\n" + "\n".join(details))
    return self.replace[root]

@rewrite_group(new_ctx=False)
def graph_rewrite(sink:UOp, pm:PatternMatcher, ctx=None, bottom_up=False, name=None, bpm=None, walk=False, enter_calls=False) -> UOp:
  rewrite_ctx = RewriteContext(pm if not bottom_up else None, pm if bottom_up else bpm, ctx, enter_calls)
  return rewrite_ctx.walk_rewrite(sink) if walk else rewrite_ctx.unified_rewrite(sink)


def sint_to_uop(x:sint, dtype=dtypes.weakint) -> UOp: return UOp.const(x, dtype)
def to_max_shape(shape:tuple[sint, ...]) -> tuple[int, ...]: return tuple(int(x.vmax) if isinstance(x, UOp) else x for x in shape)

_substitute = PatternMatcher([(UPat(tuple(Ops), name="x"), lambda ctx,x: ctx.get(x,None))])
_pm_resolve_params = PatternMatcher([(UPat(Ops.PARAM, name="p"), lambda ctx,p: ctx[p.arg.slot] if p.arg.slot >= 0 else None)])

def resolve_returned_after(r:UOp, t:UOp) -> UOp|None:
  """Extract a call output's matching store, preserving writes to the enclosing scope's output PARAMs."""
  stores = [st for st in t.src if st.op is Ops.STORE and st.src[0].unsharded_base is r.unsharded_base]
  if len(stores) != 1: return None
  return r.after(stores[0]) if r.unsharded_base.op is Ops.PARAM else stores[0].src[1]
remove_all_tags = PatternMatcher([(UPat(GroupOp.All, name="x"), lambda x: x.replace(tag=None) if x.tag is not None else None)])

# a store's storage keeps the views and drops AFTERs (they only sequence stores)
pm_drop_after = PatternMatcher([(UPat(Ops.AFTER, name="a"), lambda a: a.src[0])])

pm_renumber_slots = PatternMatcher([
  (UPat(Ops.RANGE, name="u"), lambda ctx, u: u.replace(arg=(u.axis_type, next(ctx))+u.axis_id[1:])),
  (UPat(Ops.BUFFER, name="u"), lambda ctx, u: u.replace(arg=replace(u.arg, slot=next(ctx))) if u.addrspace is AddrSpace.REG else None),
])

def gate_kernel_sink(x:UOp) -> bool:
  if x.op is Ops.LINEAR: return False
  if x.op is Ops.SINK and isinstance(x.arg, KernelInfo): return False
  return True


# *** what was symbolic.py ***

sint = int|UOp
Variable = UOp

ConstLike = ConstType|Variable|tuple[ConstType, ...]
