From 0e0cba2cfcd4f23cdbe0a7adfb71f622491e0e6a Mon Sep 17 00:00:00 2001 From: qazal <77887910+Qazalin@users.noreply.github.com> Date: Sun, 5 Jan 2025 13:02:41 +0200 Subject: [PATCH] move llvm_bf16_cast to the renderer [pr] (#8502) * move llvm_bf16_cast to the renderer [pr] * cast to half is fine too * delete the old one * wish i could just cast the ptr --- tinygrad/renderer/llvmir.py | 6 ++++++ tinygrad/tensor.py | 2 +- 2 files changed, 7 insertions(+), 1 deletion(-) diff --git a/tinygrad/renderer/llvmir.py b/tinygrad/renderer/llvmir.py index b6706ec802..66967aaee6 100644 --- a/tinygrad/renderer/llvmir.py +++ b/tinygrad/renderer/llvmir.py @@ -73,6 +73,10 @@ llvm_rewrite = PatternMatcher([ (UPat(Ops.ENDIF, name="x"), lambda ctx,x: f" br label %ifskip_{ctx[x.src[0]][1:]}\nifskip_{ctx[x.src[0]][1:]}:"), ]) +def llvm_bf16_cast(buf:UOp, idx:UOp, root:UOp): + u16_buf = buf.replace(dtype=dtypes.ushort.ptr(size=cast(PtrDType,buf.dtype).size)) + return UOp.load(UOp.index(u16_buf, idx), dtype=dtypes.ushort).cast(dtypes.uint).mul(1<<16).bitcast(dtypes.float32).cast(root.dtype) + class LLVMRenderer(Renderer): device = "LLVM" supports_float4 = False @@ -87,6 +91,8 @@ class LLVMRenderer(Renderer): (UPat(Ops.CAST, dtype=dtypes.bool, name="x"), lambda x: x.src[0] != x.src[0].const_like(0)), # rewrite MAX to CMPLT + WHERE (UPat(Ops.MAX, name="m"), lambda m: (m.src[0] < m.src[1]).where(m.src[1], m.src[0])), + # rewrite bf16 CAST(LOAD) to CAST(BITCAST) + (UPat(Ops.CAST, name="root", src=(UPat.load(UPat.index(UPat.var("buf"), UPat.var("idx")), dtype=dtypes.bfloat16),)), llvm_bf16_cast), ]) def render(self, name: str, uops: list[UOp]) -> str: diff --git a/tinygrad/tensor.py b/tinygrad/tensor.py index 60eb56d41e..f859297c42 100644 --- a/tinygrad/tensor.py +++ b/tinygrad/tensor.py @@ -3709,7 +3709,7 @@ class Tensor(SimpleMathTrait): def llvm_bf16_cast(self, dtype:DTypeLike): # hack for devices that don't support bfloat16 assert self.dtype == dtypes.bfloat16 - return self.to("LLVM").bitcast(dtypes.uint16).cast(dtypes.uint32).mul(1<<16).bitcast(dtypes.float32).cast(dtype) + return self.to("LLVM").cast(dtype) def cast(self, dtype:DTypeLike) -> Tensor: """