diff --git a/tinygrad/renderer/cstyle.py b/tinygrad/renderer/cstyle.py index 15592b8ece..b626e09b91 100644 --- a/tinygrad/renderer/cstyle.py +++ b/tinygrad/renderer/cstyle.py @@ -114,7 +114,7 @@ class CStyleLanguage(Renderer): [") {\n" + tmp] + ['\n'.join(kernel), "\n}"]) return prg if prefix is None else "\n".join(prefix)+f"\n{prg}" - def render_cast(self, dt:DType, val: str) -> str: return f"({self.render_dtype(dt)})({val})" + def render_cast(self, dt:DType, val: int|str) -> str: return f"({self.render_dtype(dt)})({val})" def render_dtype(self, dt:DType, mutable=True) -> str: if isinstance(dt, ImageDType): return f"{'write_only' if mutable else 'read_only'} image2d_t" if isinstance(dt, PtrDType): diff --git a/tinygrad/renderer/wgsl.py b/tinygrad/renderer/wgsl.py index 3083e50d85..885dcccbe4 100644 --- a/tinygrad/renderer/wgsl.py +++ b/tinygrad/renderer/wgsl.py @@ -73,7 +73,7 @@ class WGSLRenderer(CStyleLanguage): (UPat.var("a") != UPat.var("a"), lambda ctx,a: f"(min({ctx[a]}, 1.0) == 1.0 && max({ctx[a]}, -1.0) == -1.0)"), ]) + base_rewrite - def render_cast(self, dt:DType, val: str) -> str: return f"{self.type_map[dt]}({val})" + def render_cast(self, dt:DType, val: int|str) -> str: return f"{self.type_map[dt]}({val})" def render_dtype(self, dt:DType, mutable=True) -> str: return "var" def render_load(self, x:str, dt:DType) -> str: return f"atomicLoad(&{x})" if is_packed(dt) else x def buf_map(self, dt:DType) -> str: return "atomic" if is_packed(dt) else self.type_map[dt.base]