From 67df617fe12b5fbb95abcb0549ce3b85925ddf02 Mon Sep 17 00:00:00 2001 From: Sieds Lykles <93992551+S-Lykles@users.noreply.github.com> Date: Wed, 13 Aug 2025 13:05:39 +0200 Subject: [PATCH] add launch bounds to ptx (#11646) --- tinygrad/renderer/ptx.py | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/tinygrad/renderer/ptx.py b/tinygrad/renderer/ptx.py index 1fd63ee8d9..459b49238a 100644 --- a/tinygrad/renderer/ptx.py +++ b/tinygrad/renderer/ptx.py @@ -6,7 +6,7 @@ from tinygrad.uop.ops import Ops, UOp, PatternMatcher, UPat, GroupOp from tinygrad.dtype import dtypes, DType, PtrDType, AddrSpace from tinygrad.renderer import Renderer from tinygrad.renderer.cstyle import CUDARenderer -from tinygrad.helpers import flatten, get_single_element +from tinygrad.helpers import flatten, get_single_element, prod def render_val(x, dtype): if dtypes.is_float(dtype): @@ -151,11 +151,12 @@ class PTXRenderer(Renderer): mem_types: dict[DType, str] = {**types, dtypes.int8: "s8", dtypes.uint8: "u8", dtypes.bool: "u8", dtypes.float16: "b16"} - def render_kernel(self, kernel, function_name, bufs, regs) -> str: + def render_kernel(self, kernel, function_name, bufs, regs, uops) -> str: def fmt(line): return line if line[0]=="$" else "\t" + line.replace(" ", "\t" if len(line.split(" ")[0]) > 7 else "\t\t", 1) kernel = '\n'.join(map(fmt, [f".reg .{reg.split('_')[-2]} %{reg}<{cnt}>;" for reg,cnt in regs] + kernel + ["ret;"])) + launch_bounds = prod(u.arg[1] for u in uops if u.op is Ops.SPECIAL and u.arg[0][0] == "l") params = ',\n\t'.join([f".param .{'u64' if dtype.__class__ == PtrDType else self.types[dtype]} {name}" for name,dtype in bufs]) - return f"{self.kernel_prefix} {function_name}(\n\t{params}\n)\n{{\n{kernel}\n}}" + return f"{self.kernel_prefix.format(launch_bounds=launch_bounds)} {function_name} (\n\t{params}\n)\n.maxntid {launch_bounds}\n{{\n{kernel}\n}}" def render(self, uops:list[UOp]) -> str: kernel:list[str] = [] @@ -222,4 +223,4 @@ class PTXRenderer(Renderer): kernel.extend([l] if isinstance(l, str) else l) if u.op is Ops.SPECIAL: kernel = [f".reg .u32 %{u.arg[0]};"] + kernel - return self.render_kernel(kernel, name, bufs, c.items()) + return self.render_kernel(kernel, name, bufs, c.items(), uops)