diff --git a/tinygrad/uop/ops.py b/tinygrad/uop/ops.py index d4edebb174..7f31139937 100644 --- a/tinygrad/uop/ops.py +++ b/tinygrad/uop/ops.py @@ -158,7 +158,6 @@ def dtype_from_uop(op:Ops, src:tuple[UOp,...], arg:Any) -> DType|None: # NOTE: CMPLT, CMPNE, CMPEQ, WHERE, SHL, SHR are handled above if op in GroupOp.Broadcastable: # TODO: support dtype broadcasting (promotion) - if len(src) == 0: return dtypes.void if not all_same([x.dtype for x in src]): raise RuntimeError(f"dtype mismatch in {op}") return src[0].dtype if op in GroupOp.Movement: return src[0].dtype diff --git a/tinygrad/uop/upat.py b/tinygrad/uop/upat.py index 9d87680c75..4a22dcf61c 100644 --- a/tinygrad/uop/upat.py +++ b/tinygrad/uop/upat.py @@ -51,7 +51,7 @@ def _get_clause(self:UPat, base:UOp, depth=0) -> UOp: fork_cond = [UOp(Ops.AND, src=tuple([_get_clause(s, base.index(i), depth) for i,s in enumerate(ss)])) for ss in self.src] and_clause.append(UOp(Ops.OR, src=tuple(fork_cond))) else: raise RuntimeError("broken") - return UOp(Ops.AND, src=tuple(and_clause)) + return UOp(Ops.AND, src=tuple(and_clause)) if and_clause else UOp(Ops.CUSTOMI, arg="True") # *** pattern matcher ***