From d29f0ef7212c3ebac0c61533d37ba092da3a29bd Mon Sep 17 00:00:00 2001 From: qazal <77887910+Qazalin@users.noreply.github.com> Date: Tue, 7 Apr 2026 17:07:09 +0300 Subject: [PATCH 01/21] viz: speed up profiler first render (#15632) * viz: speed up profiler first render * better comment --- tinygrad/viz/js/index.js | 17 ++++++++++------- 1 file changed, 10 insertions(+), 7 deletions(-) diff --git a/tinygrad/viz/js/index.js b/tinygrad/viz/js/index.js index 4ca5f10105..781bac1a6f 100644 --- a/tinygrad/viz/js/index.js +++ b/tinygrad/viz/js/index.js @@ -408,8 +408,8 @@ async function renderProfiler(path, opts) { canvas.addEventListener("wheel", e => (e.stopPropagation(), e.preventDefault()), { passive:false }); const ctx = canvas.getContext("2d"); const canvasTop = rect(canvas).top; - // color by key (name/device) - const colorMap = new Map(); + // map event name to shape and label colors + const colorMap = new Map(), coloredNames = new Map(); // map shapes by event key const shapeMap = new Map(); const heightScale = d3.scaleLinear().domain([0, tracePeak]).range([4,maxheight=100]); @@ -448,11 +448,14 @@ async function renderProfiler(path, opts) { colorMap.set(colorKey, d3.rgb(color)); } const fillColor = colorMap.get(colorKey).brighter(0.3*depth).toString(); - const label = parseColors(e.name).flatMap(({ color, st }) => { - const parts = []; - for (let i=0; i { + const parts = []; + for (let i=0; i Date: Tue, 7 Apr 2026 19:43:51 +0300 Subject: [PATCH 02/21] mlx: graph (#15621) * Dx * Dx * simpler * mypy * x * f * Dx * x * c * x --- tinygrad/helpers.py | 2 + tinygrad/runtime/graph/hcq.py | 152 ++++++++++++++++++------- tinygrad/runtime/ops_rdma.py | 2 +- tinygrad/runtime/support/mlx/mlxdev.py | 4 +- 4 files changed, 117 insertions(+), 43 deletions(-) diff --git a/tinygrad/helpers.py b/tinygrad/helpers.py index 16947a686a..22f2213835 100644 --- a/tinygrad/helpers.py +++ b/tinygrad/helpers.py @@ -61,6 +61,8 @@ def lo32(x:Any) -> Any: return x & 0xFFFFFFFF # Any is sint def hi32(x:Any) -> Any: return x >> 32 # Any is sint def data64(data:Any) -> tuple[Any, Any]: return (data >> 32, data & 0xFFFFFFFF) # Any is sint def data64_le(data:Any) -> tuple[Any, Any]: return (data & 0xFFFFFFFF, data >> 32) # Any is sint +def to_be32(val:Any) -> Any: return ((val & 0xFF) << 24) | (((val >> 8) & 0xFF) << 16) | (((val >> 16) & 0xFF) << 8) | ((val >> 24) & 0xFF) +def to_be64(val:Any) -> Any: return to_be32(val >> 32) | (to_be32(val & 0xFFFFFFFF) << 32) def getbits(value: int, start: int, end: int): return (value >> start) & ((1 << (end - start + 1)) - 1) def i2u(bits: int, value: int): return value if value >= 0 else (1< bool: return str(type(x)) == "" diff --git a/tinygrad/runtime/graph/hcq.py b/tinygrad/runtime/graph/hcq.py index 28473c2e24..ea343dcce9 100644 --- a/tinygrad/runtime/graph/hcq.py +++ b/tinygrad/runtime/graph/hcq.py @@ -7,6 +7,7 @@ from tinygrad.dtype import dtypes from tinygrad.uop.ops import UOp, Ops, Variable from tinygrad.engine.realize import BufferXfer, CompiledRunner, BufferCopy from tinygrad.engine.jit import GraphRunner, MultiGraphRunner +from tinygrad.runtime.ops_rdma import RDMACopyQueue class HCQGraph(MultiGraphRunner): def __init__(self, *args, **kwargs): @@ -50,10 +51,21 @@ class HCQGraph(MultiGraphRunner): self.comp_queues: dict[HCQCompiled, HWQueue] = {dev: unwrap(dev.hw_compute_queue_t)() for dev in self.devices} self.copy_queues: dict[tuple[HCQCompiled, int], HWQueue] = {} # lazy allocation, keyed by (device, queue_idx) + self.rdma_queues: dict[tuple[HCQCompiled, HCQCompiled], RDMACopyQueue] = {} # lazy allocation, keyed by device pair self.num_copy_queues: int = getenv("HCQ_NUM_SDMA", min(len(self.devices), 8) if ALL2ALL >= 1 else 1) + self.num_rdma_ops: dict[tuple[HCQCompiled, HCQCompiled], int] = collections.defaultdict(int) + self.rdma_vars: dict[tuple[HCQCompiled, HCQCompiled], tuple[Variable, Any]] = {} # value is variable and src_qp + self.rdma_deps: dict[int, tuple[HWQueue, list[tuple[HCQSignal, int]], HCQSignal, int]] = {} + self.rdma_last_dest: dict[int, tuple[HWQueue, int]] = {} # per QP id: last (queue, signal_value) for dbell ordering + + # Per-peer-group representative device for signal allocation. For cpu, use devices[0]. + self.pg_dev: dict[Any, HCQCompiled] = {dev.peer_group: self.devices[0] for dev in self.devices if dev._is_cpu()} \ + | {dev.peer_group: dev for dev in self.devices if not dev._is_cpu()} + + self.kick_signals: dict[Any, HCQSignal] = {pg: pg_dev.new_signal(value=0) for pg, pg_dev in self.pg_dev.items()} self.signals: dict[Any, HCQSignal] = {**{dev: dev.new_signal(value=0) for dev in self.devices if not dev._is_cpu()}, - **{"KICK": self.devices[0].new_signal(value=0)}, **{dev: self.devices[0].new_signal(value=0) for dev in self.devices if dev._is_cpu()}} + **{dev: self.pg_dev[dev.peer_group].new_signal(value=0) for dev in self.devices if dev._is_cpu()}} self.kickoff_value: int = 0 self.kickoff_var = UOp.variable("kickoff_var", 0, 0xffffffff, dtype=dtypes.uint32) @@ -63,19 +75,22 @@ class HCQGraph(MultiGraphRunner): self.prof_graph_deps: list[list[int]] = [] self.prof_graph_entries: list[ProfileGraphEntry] = [] - last_j: dict[HWQueue, int|None] = collections.defaultdict(lambda: None) - queue_access: dict[HWQueue, dict[HWQueue, int|None]] = collections.defaultdict(lambda: collections.defaultdict(lambda: None)) - dev_access: dict[HWQueue, set[HCQCompiled]] = collections.defaultdict(set) + self.last_j: dict[HWQueue, int|None] = collections.defaultdict(lambda: None) + self.queue_access: dict[HWQueue, dict[HWQueue, int|None]] = collections.defaultdict(lambda: collections.defaultdict(lambda: None)) + self.dev_access: dict[HWQueue, set[HCQCompiled]] = collections.defaultdict(set) - for dev, queue in self.comp_queues.items(): dev_access[queue].add(dev) + for dev, queue in self.comp_queues.items(): self.dev_access[queue].add(dev) self.input_replace_map: dict[HCQCompiled, set[int]] = collections.defaultdict(set) self.device_vars: dict[HCQCompiled, dict[str, int]] = {} for j,ji in enumerate(self.jit_cache): + ji_devs = [cast(HCQCompiled, Device[cast(Buffer, b).device]) for b in ji.bufs] if isinstance(ji.prg, BufferXfer) else [] + is_rdma = len(ji_devs) > 0 and not any(d._is_cpu() for d in ji_devs) and len(set(d.peer_group for d in ji_devs)) > 1 + if is_exec_prg:=isinstance(ji.prg, CompiledRunner): enqueue_dev: HCQCompiled = ji.prg.dev else: - # For copy ops prioritize enqeueuing on the dest device, so reverse the buffers. + # For copy ops prioritize enqeueuing on the src device, so reverse the buffers. for b in cast(list[Buffer], ji.bufs[::-1]): if (enqueue_dev:=cast(HCQCompiled, Device[b.device])).hw_copy_queue_t is not None: break @@ -85,47 +100,39 @@ class HCQGraph(MultiGraphRunner): if is_exec_prg: enqueue_queue = self.comp_queues[enqueue_dev] + elif is_rdma: + enqueue_queue = self.comp_queues[enqueue_dev] + rdma_key = (cast(HCQCompiled, Device[cast(Buffer, ji.bufs[0]).device]).rdma_dev(), enqueue_dev.rdma_dev()) + self.rdma_queues.setdefault(rdma_key, RDMACopyQueue(enqueue_dev.rdma_dev())) else: assert (enqueue_dev.hw_copy_queue_t is not None), "device must implement a copy queue" queue_idx = self.devices.index(cast(HCQCompiled, Device[cast(Buffer, ji.bufs[0]).device])) % self.num_copy_queues enqueue_queue = self.copy_queues.setdefault((enqueue_dev, queue_idx), - enqueue_dev.hw_copy_queue_t(queue_idx=queue_idx).wait(self.signals['KICK'], self.kickoff_var)) + enqueue_dev.hw_copy_queue_t(queue_idx=queue_idx).wait(self.kick_signals[enqueue_dev.peer_group], self.kickoff_var)) - out_signal = self.signals.setdefault(enqueue_queue, self.devices[0].new_signal(value=0)) + out_signal = self.signals.setdefault(enqueue_queue, self.pg_dev[enqueue_dev.peer_group].new_signal(value=0)) # Get dependencies based on input and output buffers. - rdeps = self._access_resources(ji.bufs, ji.prg.p.outs if is_exec_prg else [0], (enqueue_queue, j + 1)) #type:ignore + if is_rdma: + src_qp, dest_qp = rdma_key[1].iface.connect(rdma_key[0])[:2] + sync_signals, opt_deps, rdeps = self._resolve_deps(ji.bufs[1:], [], enqueue_queue, enqueue_dev, out_signal, j, + is_copy=isinstance(ji.prg, BufferXfer), rdma_qp=src_qp) + peer_queue = self.comp_queues[peer_dev:=cast(HCQCompiled, Device[cast(Buffer, ji.bufs[0]).device])] + peer_out_signal = self.signals.setdefault(peer_queue, self.pg_dev[peer_dev.peer_group].new_signal(value=0)) + peer_sync_signals, peer_opt_deps, peer_rdeps = self._resolve_deps(ji.bufs[:1], [0], peer_queue, peer_dev, peer_out_signal, j, + is_copy=isinstance(ji.prg, BufferXfer), rdma_qp=dest_qp) + self.rdma_deps[j] = (peer_queue, peer_sync_signals + peer_opt_deps, peer_out_signal, j + 1) + self.last_j[peer_queue] = j + else: + sync_signals, opt_deps, rdeps = self._resolve_deps(ji.bufs, cast(CompiledRunner, ji.prg).p.outs if is_exec_prg else [0], enqueue_queue, + enqueue_dev, out_signal, j, is_copy=isinstance(ji.prg, BufferXfer)) - # Update dependencies to include previous kernel in queue. This is required for timeline signals. - opt_deps, deps = [], rdeps + ([(enqueue_queue, prev_ji + 1)] if (prev_ji:=last_j[enqueue_queue]) is not None else []) - - # Optimize dependencies by removing redundant ones. Remove waiting for the value of the queue which is known to be already - # synced with the current queue. - for dep_queue, dep_val in sorted(deps, key=lambda x: x[1], reverse=True): - if (qa:=queue_access[enqueue_queue][dep_queue]) is None or qa < dep_val: - opt_deps.append((self.signals[dep_queue], dep_val)) - queue_access[enqueue_queue][dep_queue] = dep_val - dev_access[enqueue_queue].update(dev_access[dep_queue]) - - # Ensure device is ready for use in current context: the graph has initialized the device and it's safe to operate on it within this graph. - sync_signals = [(self.signals[d], self.kickoff_var) for b in ji.bufs if (d:=Device[cast(Buffer, b).device]) not in dev_access[enqueue_queue]] - dev_access[enqueue_queue].update(cast(HCQCompiled, Device[cast(Buffer, b).device]) for b in ji.bufs) - - # Remove self-dependency for compute and copy queues. - # For compute, in case of NV, optimize when only 1 same-queue dependency exists, since NV chains 2+ executions in this case, - # eliminating dependency need. - dname = enqueue_dev.device.split(":", 1)[0] - can_opt = dname in {"AMD", "QCOM"} or (dname == "NV" and len(sync_signals) == 0 and len(opt_deps) == 1 and id(opt_deps[0][0]) == id(out_signal)) - if can_opt or isinstance(ji.prg, BufferXfer): opt_deps = [x for x in opt_deps if id(x[0]) != id(out_signal)] - - # Enable necessary signals in the schedule by setting the signal value. - for sig, val in opt_deps: self.ji_schedule[val - 1] = self.ji_schedule[val - 1][:5] + (val,) self.ji_schedule[j] = (enqueue_dev, enqueue_queue, sync_signals, opt_deps[::-1], out_signal, None if is_exec_prg else (j + 1)) # Collect profile information if profiling is enabled. if PROFILE: # When execution are chained, we can reuse the end timestamp from the previous command as the start timestamp for the current command. - sig_st = prev_ji * 2 + 1 if len(opt_deps) == 0 and (prev_ji:=last_j[enqueue_queue]) is not None else j * 2 + sig_st = prev_ji * 2 + 1 if len(opt_deps) == 0 and (prev_ji:=self.last_j[enqueue_queue]) is not None else j * 2 # Description based on the command. prof_ji_desc = ji.prg._prg.name if is_exec_prg else TracingKey(f"{ji.bufs[1].device} -> {ji.bufs[0].device}", ret=ji.bufs[0].nbytes) # type: ignore @@ -134,7 +141,7 @@ class HCQGraph(MultiGraphRunner): self.prof_graph_entries.append(ProfileGraphEntry(prof_name, prof_ji_desc, sig_st, j * 2 + 1)) self.prof_graph_deps.append([d - 1 for _, d in rdeps]) - last_j[enqueue_queue] = j + self.last_j[enqueue_queue] = j # Check which signals are used in the profile graph. self.prof_signal_is_used = [any(ent.st_id == j or ent.en_id == j for ent in self.prof_graph_entries) for j in range(len(self.jit_cache) * 2)] @@ -149,7 +156,7 @@ class HCQGraph(MultiGraphRunner): for dev in self.devices: self.comp_queues[dev].memory_barrier().wait(self.virt_timeline_signals[dev], self.virt_timeline_vals[dev]) \ - .wait(self.signals['KICK'], self.kickoff_var).signal(self.signals[dev], self.kickoff_var) + .wait(self.kick_signals[dev.peer_group], self.kickoff_var).signal(self.signals[dev], self.kickoff_var) for j,ji in enumerate(self.jit_cache): enqueue_dev, enqueue_queue, sync_signals, deps, signal, signal_val = self.ji_schedule[j] @@ -165,6 +172,28 @@ class HCQGraph(MultiGraphRunner): # Encode main commands based on ji type. if isinstance(ji.prg, CompiledRunner): enqueue_queue.exec(ji.prg._prg, self.ji_args[j], tuple(ji.prg.p.global_size or (1,1,1)), tuple(ji.prg.p.local_size or (1,1,1))) + elif isinstance(ji.prg, BufferXfer) and len(set(cast(HCQCompiled, Device[cast(Buffer, b).device]).peer_group for b in ji.bufs)) > 1: + dest_queue, dest_deps, dest_out_signal, dest_out_val = self.rdma_deps[j] + for sig, val in dest_deps: dest_queue.wait(sig, val) + + dest, src = [cast(Buffer, x) for x in ji.bufs[0:2]] + dest_dev, src_dev = cast(HCQCompiled, Device[dest.device]), cast(HCQCompiled, Device[src.device]) + dest_rdma, src_rdma = dest_dev.rdma_dev(), src_dev.rdma_dev() + + # get qp info + src_qp, dest_qp, src_cq_buf, dest_cq_buf = src_rdma.iface.connect(dest_rdma) + + # use var for head + head_var = self.rdma_vars.setdefault((dest_rdma, src_rdma), (UOp.variable(f"rdma_var_{j}", 0, 0xffffffff, dtype=dtypes.uint32), src_qp))[0] + next_head = self.num_rdma_ops[(dest_rdma, src_rdma)] + + rdma_queue = self.rdma_queues[(dest_rdma, src_rdma)] + rdma_queue.copy(self.hcq_bufs[j][0], self.hcq_bufs[j][1], dest.nbytes) \ + .encode_ring(enqueue_queue, src_dev, src_rdma.iface, src_qp, src_cq_buf, head_var + next_head, ring_uar=True) \ + .encode_ring(self.comp_queues[dest_dev], dest_dev, dest_rdma.iface, dest_qp, dest_cq_buf, head_var + next_head) + + dest_queue.signal(dest_out_signal, dest_out_val) + self.num_rdma_ops[(dest_rdma, src_rdma)] += 1 elif isinstance(ji.prg, (BufferXfer, BufferCopy)): dest, src = [cast(Buffer, x) for x in ji.bufs[0:2]] for bufid, src in enumerate(cast(list[Buffer], ji.bufs)): @@ -181,7 +210,7 @@ class HCQGraph(MultiGraphRunner): for dev in self.devices: for dep_dev in list(self.copy_to_devs[dev]) + [dev]: for copy_q in self._dev_copy_queues(dep_dev): - if copy_q in self.signals: self.comp_queues[dev].wait(self.signals[copy_q], cast(int, last_j[copy_q]) + 1) + if copy_q in self.signals: self.comp_queues[dev].wait(self.signals[copy_q], cast(int, self.last_j[copy_q]) + 1) self.comp_queues[dev].signal(self.virt_timeline_signals[dev], self.virt_timeline_vals[dev] + 1).bind(dev) for copy_q in self._dev_copy_queues(dev): copy_q.bind(dev) @@ -189,6 +218,44 @@ class HCQGraph(MultiGraphRunner): self.last_timeline: dict[HCQCompiled, tuple[HCQSignal, int]] = {dev: (dev.timeline_signal, 0) for dev in self.devices} self.queue_signals_to_reset = [self.signals[q] for q in list(self.comp_queues.values()) + list(self.copy_queues.values()) if q in self.signals] + def _resolve_deps(self, bufs, outs, enqueue_queue, enqueue_dev, out_signal, j, is_copy, rdma_qp=None): + rdeps = self._access_resources(bufs, outs, (enqueue_queue, j + 1)) #type:ignore + + # Order shared QP doorbell record writes across different compute queues (head+1 must complete before head+2). + if rdma_qp is not None and (prev:=self.rdma_last_dest.get(id(rdma_qp))) is not None and prev[0] is not enqueue_queue: + rdeps = rdeps + [(prev[0], prev[1])] + if rdma_qp is not None: self.rdma_last_dest[id(rdma_qp)] = (enqueue_queue, j + 1) + + # Update dependencies to include previous kernel in queue. This is required for timeline signals. + opt_deps, deps = [], rdeps + ([(enqueue_queue, prev_ji + 1)] if (prev_ji:=self.last_j[enqueue_queue]) is not None else []) + + # Optimize dependencies by removing redundant ones. Remove waiting for the value of the queue which is known to be already + # synced with the current queue. + for dep_queue, dep_val in sorted(deps, key=lambda x: x[1], reverse=True): + if (qa:=self.queue_access[enqueue_queue][dep_queue]) is None or qa < dep_val: + opt_deps.append((self.signals[dep_queue], dep_val)) + self.queue_access[enqueue_queue][dep_queue] = dep_val + self.dev_access[enqueue_queue].update(self.dev_access[dep_queue]) + + # Ensure device is ready for use in current context: the graph has initialized the device and it's safe to operate on it within this graph. + # Only sync with same-peer-group devices; cross-peer-group sync is handled by RDMA. + sync_signals = [(self.signals[d], self.kickoff_var) for b in bufs + if (d:=cast(HCQCompiled, Device[cast(Buffer, b).device])) not in self.dev_access[enqueue_queue] + and (d.peer_group == enqueue_dev.peer_group or rdma_qp is None)] + self.dev_access[enqueue_queue].update(cast(HCQCompiled, Device[cast(Buffer, b).device]) for b in bufs) + + # Remove self-dependency for compute and copy queues. + # For compute, in case of NV, optimize when only 1 same-queue dependency exists, since NV chains 2+ executions in this case, + # eliminating dependency need. For RDMA, keep self-dependency to flush cache. + dname = enqueue_dev.device.split(":", 1)[0] + can_opt = dname in {"AMD", "QCOM"} or (dname == "NV" and len(sync_signals) == 0 and len(opt_deps) == 1 and id(opt_deps[0][0]) == id(out_signal)) + if (can_opt or is_copy) and rdma_qp is None: opt_deps = [x for x in opt_deps if id(x[0]) != id(out_signal)] + + # Enable necessary signals in the schedule by setting the signal value. + for sig, val in opt_deps: self.ji_schedule[val - 1] = self.ji_schedule[val - 1][:5] + (val,) + + return sync_signals, opt_deps, rdeps + def _dev_copy_queues(self, dev): return [q for (d, _), q in self.copy_queues.items() if d == dev] def __call__(self, input_buffers: list[Buffer], var_vals: dict[str, int], wait=False) -> float|None: @@ -209,6 +276,9 @@ class HCQGraph(MultiGraphRunner): for (j,i),input_idx in self.input_replace.items(): hcq_var_vals[self.input_replace_to_var[(j,i)].expr] = input_buffers[input_idx]._buf.va_addr + for (var, qp) in self.rdma_vars.values(): hcq_var_vals[var.expr] = qp.head + for q in self.rdma_queues.values(): q.submit(q.dev, hcq_var_vals) + for dev in self.devices: self.comp_queues[dev].submit(dev, hcq_var_vals_local:=hcq_var_vals|self.device_vars.get(dev, {})) for copy_queue in self._dev_copy_queues(dev): copy_queue.submit(dev, hcq_var_vals_local) @@ -216,7 +286,7 @@ class HCQGraph(MultiGraphRunner): # Launch graph for sig in self.queue_signals_to_reset: sig.value = 0 - self.signals['KICK'].value = self.kickoff_value + for sig in self.kick_signals.values(): sig.value = self.kickoff_value if wait: st = time.perf_counter() @@ -247,8 +317,10 @@ class HCQGraph(MultiGraphRunner): # If all of devices are mapped into CPU address space, can use CPU inside the peer group. cpu_support = all(type(d.timeline_signal.base_buf.view) is MMIOInterface for d in all_devs) - # Check if all devices are within the same peer group. If CPU is supported, don't count it as a separate peer group. - if len(set(d.peer_group for d in all_devs if not (cpu_support and d._is_cpu()))) > 1: return False + # Check if all devices are within the same peer group. Allow cross-peer-group if all peer groups have RDMA devices. + if len(set(d.peer_group for d in all_devs if not (cpu_support and d._is_cpu()))) > 1: + try: [d.rdma_dev() for d in all_devs if not d._is_cpu()] + except RuntimeError: return False if new_call.src[0].op is Ops.COPY: # MOCKGPU is not supported, since it can't execute commands in parallel diff --git a/tinygrad/runtime/ops_rdma.py b/tinygrad/runtime/ops_rdma.py index b4a09e7572..15069a354a 100644 --- a/tinygrad/runtime/ops_rdma.py +++ b/tinygrad/runtime/ops_rdma.py @@ -21,7 +21,7 @@ class RDMACopyQueue(HWQueue): for buf in [iface.dbr_buf, cq_buf] + ([iface.uar_buf] if ring_uar else []): cast(HCQAllocator, dev.allocator).map(buf) hwq.write(iface.dbr_buf.offset(qp.qp_dbr + (4 if ring_uar else 0)), to_be('I', head + 1)) if ring_uar: hwq.write(iface.uar_buf.offset(0x800), to_be('Q', ((head << 8) | 0x0a) << 32 | ((qp.qp_info['qpn'] << 8) | 2)), b64=True) - hwq.poll_bit(cq_buf.offset((head & (qp.cq_size - 1)) * 64 + 60, 4), ((head >> 7) & 1) << 24, mask=0x01000000) + hwq.poll_bit(cq_buf.offset((head & (qp.cq_size - 1)) * 64 + 60, 4), ((head >> (qp.cq_size.bit_length() - 1)) & 1) << 24, mask=0x01000000) hwq.write(iface.dbr_buf.offset(qp.cq_dbr), to_be('I', (head + 1) & 0xFFFFFF)) return self diff --git a/tinygrad/runtime/support/mlx/mlxdev.py b/tinygrad/runtime/support/mlx/mlxdev.py index 739c34f687..ff84e8e463 100644 --- a/tinygrad/runtime/support/mlx/mlxdev.py +++ b/tinygrad/runtime/support/mlx/mlxdev.py @@ -1,6 +1,6 @@ from __future__ import annotations import struct, random, socket, ctypes, functools, itertools -from tinygrad.helpers import getenv, wait_cond, round_up, next_power2, ceildiv, DEBUG, hi32, lo32 +from tinygrad.helpers import getenv, wait_cond, round_up, next_power2, ceildiv, DEBUG, hi32, lo32, to_be32, to_be64 from tinygrad.runtime.support.memory import BumpAllocator from tinygrad.runtime.support.system import PCIDevice from tinygrad.runtime.autogen import mlx5, pci @@ -11,7 +11,7 @@ MLX5_CMD_STRUCTS = {v: (getattr(mlx5, f"struct_mlx5_ifc_{n[12:].lower()}_in_bits getattr(mlx5, f"struct_mlx5_ifc_{n[12:].lower()}_out_bits", None)) for n, v in mlx5.__dict__.items() if n.startswith("MLX5_CMD_OP_")} MLX5_CMD_STRUCTS[mlx5.MLX5_CMD_OP_ACCESS_REG] = (mlx5.struct_mlx5_ifc_access_register_in_bits, mlx5.struct_mlx5_ifc_access_register_out_bits) -def to_be(fmt, val): return struct.unpack('<'+fmt, struct.pack('>'+fmt, val))[0] +def to_be(fmt, val): return to_be32(val) if fmt == 'I' else to_be64(val) def ipv4_to_gid(ip): return bytes(10) + b'\xff\xff' + socket.inet_aton(ip) def udp_sport(lqpn, rqpn): From 890286e8d6e22dc199fcac4b6b5f0aa06c9a94ac Mon Sep 17 00:00:00 2001 From: qazal <77887910+Qazalin@users.noreply.github.com> Date: Tue, 7 Apr 2026 21:18:45 +0300 Subject: [PATCH 03/21] update llama profile.sh (#15633) * update llama profile.sh * BENCHMARK 5 --- .../llama8b/implementations/tinybox_8xMI350X/dev_beam.sh | 2 +- .../llama8b/implementations/tinybox_8xMI350X/profile.sh | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/dev_beam.sh b/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/dev_beam.sh index c345538f6f..6dfbf57089 100755 --- a/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/dev_beam.sh +++ b/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/dev_beam.sh @@ -36,7 +36,7 @@ export DATA_SEED=${DATA_SEED:-5760} export JITBEAM=${JITBEAM:-3} export BEAM_UOPS_MAX=6000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 BEAM_PADTO=1 -export FAKEDATA=1 BENCHMARK=10 +export FAKEDATA=1 BENCHMARK=${BENCHMARK:-10} if [ -z "$FULL_LAYERS" ]; then export LLAMA_LAYERS=2 fi diff --git a/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/profile.sh b/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/profile.sh index cfddfa2601..e8fed36b18 100755 --- a/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/profile.sh +++ b/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/profile.sh @@ -1,5 +1,5 @@ #!/bin/bash export BENCHMARK=5 export EVAL_BS=0 -VIZ=${VIZ:--1} examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/dev_run.sh +VIZ=${VIZ:--1} FULL_LAYERS=1 DEBUG=0 examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/dev_beam.sh extra/viz/cli.py --profile -s "${DEV:-AMD}" From 9c6e925b56e108e7f793a1998ec41e020fdf30f1 Mon Sep 17 00:00:00 2001 From: chenyu Date: Tue, 7 Apr 2026 15:13:00 -0400 Subject: [PATCH 04/21] move lerp to mixin (#15634) last function of math function section --- tinygrad/mixin/elementwise.py | 13 +++++++++++++ tinygrad/tensor.py | 15 --------------- 2 files changed, 13 insertions(+), 15 deletions(-) diff --git a/tinygrad/mixin/elementwise.py b/tinygrad/mixin/elementwise.py index c1cb59cd82..8a4129c495 100644 --- a/tinygrad/mixin/elementwise.py +++ b/tinygrad/mixin/elementwise.py @@ -974,3 +974,16 @@ class ElementwiseMixin(DTypeMixin, CreationMixin): """ if self.dtype != dtypes.bool and not dtypes.is_int(self.dtype): raise RuntimeError(f"{self.dtype} is not supported") return self.logical_not() if self.dtype == dtypes.bool else self ^ -1 + + def lerp(self, end: Self, weight: Self | ConstType) -> Self: + """ + Linearly interpolates between `self` and `end` by `weight`. + + ```python exec="true" source="above" session="tensor" result="python" + print(Tensor([1., 2., 3.]).lerp(Tensor([4., 5., 6.]), 0.5).numpy()) + ``` + """ + if self.dtype == dtypes.uint8 and isinstance(weight, ElementwiseMixin): + w_i = (weight * (1<<(W_PREC:=7)) + 0.5).cast(dtypes.int16) + return (self+(((end - self).cast(dtypes.int8) * w_i + (1<> W_PREC)).cast(dtypes.uint8) + return self + (end - self) * weight diff --git a/tinygrad/tensor.py b/tinygrad/tensor.py index 2f42f83da7..45c45d62ce 100644 --- a/tinygrad/tensor.py +++ b/tinygrad/tensor.py @@ -2385,21 +2385,6 @@ class Tensor(OpMixin): """ return self._apply_uop(UOp.contiguous_backward) - # ***** math functions ***** - - def lerp(self, end:Tensor, weight:Tensor|float) -> Tensor: - """ - Linearly interpolates between `self` and `end` by `weight`. - - ```python exec="true" source="above" session="tensor" result="python" - print(Tensor([1., 2., 3.]).lerp(Tensor([4., 5., 6.]), 0.5).numpy()) - ``` - """ - if self.dtype == dtypes.uint8 and isinstance(weight, Tensor): - w_i = (weight * (1<<(W_PREC:=7)) + 0.5).cast(dtypes.int16) - return (self+(((end - self).cast(dtypes.int8) * w_i + (1<> W_PREC)).cast(dtypes.uint8) - return self + (end - self) * weight - # ***** broadcasted elementwise ops ***** def ufix(self, x) -> Tensor: From a508b8fd2a7b5666278a0a6c2b127e74ed6a4416 Mon Sep 17 00:00:00 2001 From: qazal <77887910+Qazalin@users.noreply.github.com> Date: Wed, 8 Apr 2026 01:18:04 +0300 Subject: [PATCH 05/21] viz: delete redundant things (#15637) * delete that * remove * delete graph config --- extra/viz/cli.py | 1 - tinygrad/viz/js/index.js | 18 +++++++----------- 2 files changed, 7 insertions(+), 12 deletions(-) diff --git a/extra/viz/cli.py b/extra/viz/cli.py index 2fa1459372..25806400d2 100755 --- a/extra/viz/cli.py +++ b/extra/viz/cli.py @@ -59,7 +59,6 @@ def main(args) -> None: events:list = viz.load_pickle(args.profile_path, default=[]) if (profile_bytes:=viz.get_profile(events)) is None: raise RuntimeError(f"empty profile in {args.profile_path}") profile = decode_profile(profile_bytes) - viz.load_amd_counters(viz.ctxs, events) profile["layout"].update([(f'{c["name"]} {s["name"]}', s["data"]) for c in viz.ctxs if c["name"].startswith("SQTT") for s in c["steps"] if "PKTS" in s["name"]]) if args.src is None: diff --git a/tinygrad/viz/js/index.js b/tinygrad/viz/js/index.js index 781bac1a6f..24f8419030 100644 --- a/tinygrad/viz/js/index.js +++ b/tinygrad/viz/js/index.js @@ -16,7 +16,6 @@ const darkenHex = (h, p = 0) => const ANSI_COLORS = ["#b3b3b3", "#ff6666", "#66b366", "#ffff66", "#6666ff", "#ff66ff", "#66ffff", "#ffffff"]; const ANSI_COLORS_LIGHT = ["#d9d9d9","#ff9999","#99cc99","#ffff99","#9999ff","#ff99ff","#ccffff","#ffffff"]; -const colorsCache = new Map(); const parseColors = (name, defaultColor="#ffffff") => Array.from(name.matchAll(/(?:\u001b\[(\d+)m([\s\S]*?)\u001b\[0m)|([^\u001b]+)/g), ([_, code, colored_st, st]) => ({ st: colored_st ?? st, color: code != null ? (code>=90 ? ANSI_COLORS_LIGHT : ANSI_COLORS)[(parseInt(code)-30+60)%60] : defaultColor })); @@ -375,7 +374,6 @@ function setFocus(key) { } const EventTypes = { EXEC:0, BUF:1 }; -const GraphConfig = [{ pcolor:"#c9a8ff", unit:"B", fillColor:"#2B1B72"}, { pcolor:"#4fa3cc", unit:"Hz", fillColor:"#4fa3cc"}]; async function renderProfiler(path, opts) { displaySelection("#profiler"); @@ -487,7 +485,6 @@ async function renderProfiler(path, opts) { div.style("height", levelHeight*levels.length+padding+"px").style("pointerEvents", "none"); } else { const linear = u8(), peak = u64(); - const config = GraphConfig[linear]; const timestamps = [], valueMap = new Map(); // start by unpacking the raw events const memEvents = []; @@ -516,7 +513,7 @@ async function renderProfiler(path, opts) { const yscale = d3.scaleLinear().domain([0, peak]).range([height, 0]); // generic polygon merger const base0 = yscale(0); - const sum = {x:[], y0:[], y1:[], fillColor:config.fillColor}; + const sum = {x:[], y0:[], y1:[], fillColor:linear ? null : "#2b1b72"}; for (let i=0; i 0) data.first = data.first == null ? timestamps[0] : Math.min(data.first, timestamps[0]); - data.tracks.set(k, { shapes:[sum], eventType, linear, visible, offsetY, pcolor:config.pcolor, height, peak, scaleFactor:maxheight*4/height, - get views() { return [[sum], linear ? null : buildBufShapes()]; }, valueMap, rowBorderColor }); + data.tracks.set(k, { shapes:[sum], eventType, linear, visible, offsetY, pcolor:linear ? "#4fa3cc" : "#c9a8ff", height, peak, scaleFactor:maxheight*4/height, + get views() { return [[sum], linear ? null : buildBufShapes()]; }, valueMap, rowBorderColor, unit:linear ? "Hz" : "B" }); div.style("height", height+padding+"px").style("cursor", "pointer").on("click", (e) => { if (linear) return; const newFocus = e.currentTarget.id === focusedDevice ? null : e.currentTarget.id; @@ -566,7 +563,7 @@ async function renderProfiler(path, opts) { if (tid === newFocus) { track.shapes = track.views[1]; offset += rescaleTrack(track, tid, track.scaleFactor); } else if (tid === focusedDevice) { track.shapes = track.views[0]; offset += rescaleTrack(track, tid, 1/track.scaleFactor); } } - data.axes.y = newFocus != null ? { domain:[0, (t=data.tracks.get(newFocus)).peak], range:[t.offsetY+t.height, t.offsetY], fmt:config.unit } : null; + data.axes.y = newFocus != null ? { domain:[0, (t=data.tracks.get(newFocus)).peak], range:[t.offsetY+t.height, t.offsetY], fmt:t.unit } : null; toggleCls(document.getElementById(focusedDevice), document.getElementById(newFocus), "expanded"); focusedDevice = newFocus; return resize(); @@ -611,12 +608,11 @@ async function renderProfiler(path, opts) { const visibleYStart = profilerEl.scrollTop-canvasTop + rect(profilerEl).top, visibleYEnd = visibleYStart+profilerEl.clientHeight; ctx.textBaseline = "middle"; // draw shapes - for (const [k, { shapes, eventType, linear, visible, offsetY, valueMap, pcolor, scolor, rowBorderColor }] of data.tracks) { + for (const [k, { shapes, eventType, linear, visible, offsetY, valueMap, pcolor, scolor, unit, rowBorderColor }] of data.tracks) { visible.length = 0; const trackHeight = rect(document.getElementById(k)).height; if (offsetY+trackHeight < visibleYStart || offsetY > visibleYEnd) continue; const addBorder = scolor != null ? (w) => { if (w > 10) { ctx.strokeStyle = scolor; ctx.stroke(); } } : null; - const config = GraphConfig[linear]; for (const e of shapes) { if (eventType === EventTypes.BUF) { // generic polygon if (e.x[0]>et || e.x.at(-1)=0; i--) ctx.lineTo(x[i], offsetY+e.y0[i]); ctx.closePath(); ctx.fillStyle = e.fillColor; ctx.fill(); } } else { // contiguous rect From bf3763526ab5f08f5ef1fe5598c084189ef56ac8 Mon Sep 17 00:00:00 2001 From: b1tg <33436708+b1tg@users.noreply.github.com> Date: Wed, 8 Apr 2026 10:26:23 +0800 Subject: [PATCH 06/21] llm: buffer SSE chunks to fix parse errors from split reads (#15641) --- tinygrad/apps/llm.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/tinygrad/apps/llm.py b/tinygrad/apps/llm.py index 6fe4a19201..9ffa500d4f 100644 --- a/tinygrad/apps/llm.py +++ b/tinygrad/apps/llm.py @@ -313,10 +313,14 @@ CHAT_HTML = b'''tinygrad chat