diff --git a/tinygrad/engine/realize.py b/tinygrad/engine/realize.py index ed27debcd9..e918133c38 100644 --- a/tinygrad/engine/realize.py +++ b/tinygrad/engine/realize.py @@ -33,7 +33,7 @@ def get_call_name(call:UOp, bufs:Sequence[Buffer|UOp], var_vals:dict[str, int]|N if ast.op is Ops.COPY: return colored(f"copy {_uop_sz_to_str(arg_uops[0]):>10}, {_dev_str(bufs[0]):>7s} <- {_dev_str(bufs[1]):7s}", "yellow") if ast.op is Ops.CUSTOM_FUNCTION and ast.arg == "encdec": return colored(f"enc/dec {_uop_sz_to_str(arg_uops[0])}", "yellow") if ast.op is Ops.CUSTOM_FUNCTION and ast.arg == "graph": return colored(f"batched {len(ast.src[0].src)}", "cyan") - if ast.op is Ops.CUSTOM_FUNCTION and ast.arg == "hcq": return call.arg.aux.name + if ast.op is Ops.CUSTOM_FUNCTION and ast.arg == "hcq": return cast(str, call.arg.name) raise NotImplementedError("get_call_name is not implemented") # **************** Stat **************** diff --git a/tinygrad/runtime/ops_cpu.py b/tinygrad/runtime/ops_cpu.py index c74363e659..09a77e6a41 100644 --- a/tinygrad/runtime/ops_cpu.py +++ b/tinygrad/runtime/ops_cpu.py @@ -167,10 +167,11 @@ class CPUDevice(HCQCompiled): (UPat(Ops.PARAM, tag="sentinel_signal"), lambda ctx: ctx[0].signal("sentinel", (1 << 64) - 1)), (UPat(Ops.PARAM, tag="timeline_signal"), lambda ctx: ctx[0].signal("timeline")), (UPat(Ops.PARAM, tag="timeline_value"), lambda ctx: ctx[0].signal("value", 1)), + (UPat(Ops.PARAM, tag="signal", name="b"), lambda ctx, b: ctx[0].signal(b.arg.slot)), ]) @functools.cache - def signal(self, name:str, init_value:int=0) -> Buffer: + def signal(self, name:str|int, init_value:int=0) -> Buffer: (buf:=Buffer(self.device, 1, dtypes.uint64, preallocate=True)).as_memoryview(force_zero_copy=True, no_sync=True).cast('Q')[0] = init_value return buf diff --git a/tinygrad/runtime/support/hcq2.py b/tinygrad/runtime/support/hcq2.py index 2b55cc07f6..2933ddfba0 100644 --- a/tinygrad/runtime/support/hcq2.py +++ b/tinygrad/runtime/support/hcq2.py @@ -27,10 +27,8 @@ HCQ_CACHE_TAGS = frozenset(("program", "systems", "template")) @dataclass(frozen=True) class HCQInfo: - name:str - estimates:Estimates device:tuple[str, ...] - queue:str + estimates:Estimates = Estimates() input_idxs:tuple[int, ...] = () # indexes into input_uops used by this call inputs:int|None = None @@ -73,6 +71,8 @@ def make_submit(*cmds, devs:str|tuple[str, ...], queue:str) -> UOp: return UOp.custom_function("submit_cmdbuf", UOp(Ops.LINEAR, src=tuple(cmds), arg=(to_tuple(devs), queue))) def get_submit(ast:UOp) -> UOp: return next(u for u in ast.toposort() if u.op is Ops.CUSTOM_FUNCTION and u.arg == "submit_cmdbuf") +def make_call(name:str, body:UOp, info:HCQInfo) -> UOp: return UOp.custom_function("hcq", body).call(name=name, aux=info) + def encode_kernargs_clike(call:UOp, prg:UOp, devs:str|tuple[str, ...]) -> UOp: data, info = prg.arg buf = UOp.placeholder((data.kernargs_alloc_size // 4,), dtypes.uint32, next(UOp.unique_num), device=devs).rtag("kernargs") @@ -137,12 +137,6 @@ def _build_wait_cmds(slots:dict[str, int], dep_lanes:list[tuple[tuple, int, int] waits.append(UOp(Ops.INS, arg="wait", src=(sig, UOp.const(dtag + 1, dtypes.uint64)))) return waits, {dtag for _, _, dtag in deps} -def make_fence(timeline:UOp, prev:UOp, sigs:list[UOp]) -> UOp: - free = (cur:=timeline.after(loop:=UOp.loop(0)).index(0).load()).end(loop, cur < prev.index(0).load()) - return UOp.sink(*[s.after(free).index(0).store(0) for s in sigs]) - -def _hcq_call(devs, name:str, body:UOp) -> UOp: return UOp.custom_function("hcq", body).call(aux=HCQInfo(name, Estimates(), devs, "COMPUTE:0")) - def _build_finalizers(batch:list[tuple[UOp, tuple[str, ...]]], batch_info:list[tuple[tuple[str, ...], str]], tracker:HCQDepsTracker, slots:dict[str, int]) -> tuple[list[UOp], list[UOp], set[int]]: # collect all buffers which belong to devices @@ -151,46 +145,48 @@ def _build_finalizers(batch:list[tuple[UOp, tuple[str, ...]]], batch_info:list[t for b in itertools.chain.from_iterable(_get_call_bufs_by_lane(call, devices)): for bd in to_tuple(b.device): dev_bufs[bd][id(b)] = b - n, fences, fins, waited = len(batch_info), [], [], set() + n, fences, fins, signal_tags = len(batch_info), [], [], set() for _, devgroup in itertools.groupby(sorted(dev_bufs), key=lambda d: d.split(":")[0]): devs = tuple(devgroup) # to finalize the batch, sync all accesses from other devices to buffers that belong to this device fin_deps = [dl for dl in _get_deps(tracker, [list(dev_bufs[d].values()) for d in devs], None, key=(devs, "COMPUTE:0", n)) if dl[0][2] < n] - waits, cur_waited = _build_wait_cmds(slots, fin_deps, devs, "COMPUTE:0") - waited |= cur_waited + waits, cur_signal_tags = _build_wait_cmds(slots, fin_deps, devs, "COMPUTE:0") + signal_tags |= cur_signal_tags # wait the syncs and signal the device epoch, then bump the timeline on the host - timeline, tl = make_signal(devs, tag="timeline_signal"), make_signal(devs, tag="timeline_value") - submit = make_submit(*waits, UOp(Ops.INS, arg="store", src=(timeline, tl.index(0))), devs=devs, queue="COMPUTE:0") - cur = (bump:=tl.after(submit).index(0)).load() - bumps = [bump.store(cur + 1)] + tl_signal, tl_value = make_signal(devs, tag="timeline_signal"), make_signal(devs, tag="timeline_value") + fin_submit = make_submit(*waits, UOp(Ops.INS, arg="store", src=(tl_signal, tl_value.index(0))), devs=devs, queue="COMPUTE:0") + epoch = (epoch_slot:=tl_value.after(fin_submit).index(0)).load() - # devices running the batch reset their queue signals before each run, fencing on the epoch kept from the previous one - if qs:=dedup([qn for bdevs, qn in batch_info if set(bdevs) & set(devs)]): - prev = make_signal(devs, next(UOp.unique_num)) - fences.append(_hcq_call(devs, "hcq_fence", make_fence(timeline, prev, [make_signal(devs, slots[q]) for q in qs]))) - bumps.append(prev.after(submit).index(0).store(cur)) - fins.append(_hcq_call(devs, "hcq_finalizer", UOp.sink(*bumps))) - return fences, fins, waited + # fence once per device group on this schedule's previous epoch, then reset any queue signals used by the group + qs = dedup([qn for bdevs, qn in batch_info if set(bdevs) & set(devs)]) + sched_epoch = make_signal(devs, next(UOp.unique_num)) + + wait_device_epoch = (done:=tl_signal.after(loop:=UOp.loop(0)).index(0).load()).end(loop, done < sched_epoch.index(0).load()) + resets = [make_signal(devs, slots[q]).after(wait_device_epoch).index(0).store(0) for q in qs] + + fences.append(make_call("hcq_fence", UOp.sink(*(resets or [wait_device_epoch])), HCQInfo(devs))) + fins.append(make_call("hcq_finalizer", UOp.sink(epoch_slot.store(epoch + 1), sched_epoch.after(fin_submit).index(0).store(epoch)), HCQInfo(devs))) + return fences, fins, signal_tags def _finalize_batch(batch:list[tuple[UOp, tuple[str, ...]]]) -> list[UOp]: batch_info = [(devices, "COMPUTE:0" if call.src[0].op is Ops.PROGRAM else "COPY:0") for call, devices in batch] # schedule deps - waited:set[int] = set() + signal_tags:set[int] = set() slots:dict[str, int] = collections.defaultdict(lambda: next(UOp.unique_num)) deps_tracker = HCQDepsTracker() call_waits:list[list[UOp]] = [] for tag, ((call, _), (devices, queue)) in enumerate(zip(batch, batch_info)): deps = _get_deps(deps_tracker, _get_call_bufs_by_lane(call, devices), get_call_outs_ins(call)[0], key=(devices, queue, tag)) - cmds, cur_waited = _build_wait_cmds(slots, deps, devices, queue) + cmds, cur_signal_tags = _build_wait_cmds(slots, deps, devices, queue) call_waits.append(cmds) - waited |= cur_waited + signal_tags |= cur_signal_tags # build fences and finalizers - fences, finalizers, finalizer_waited = _build_finalizers(batch, batch_info, deps_tracker, slots) - waited |= finalizer_waited + fences, finalizers, finalizer_signal_tags = _build_finalizers(batch, batch_info, deps_tracker, slots) + signal_tags |= finalizer_signal_tags src = [] for tag, ((call, _), (devices, queue), q) in enumerate(zip(batch, batch_info, call_waits)): @@ -200,12 +196,12 @@ def _finalize_batch(batch:list[tuple[UOp, tuple[str, ...]]]) -> list[UOp]: q = [UOp(Ops.INS, arg="barrier", src=()), UOp(Ops.INS, arg="wait", src=(make_signal(devices, tag="timeline_signal"), epoch))] + q # and make hcq call - info = HCQInfo(get_call_name(call, get_call_arg_uops(call)), estimate_uop(call), devices, queue) + name, info = get_call_name(call, get_call_arg_uops(call)), HCQInfo(devices, estimate_uop(call)) q += [call.replace(arg=replace(call.arg, aux=info))] # signal the queue if someone waits for us - if tag in waited: q += [UOp(Ops.INS, arg="store", src=(make_signal(devices, slots[queue]), UOp.const(tag + 1, dtypes.uint64)))] - src.append(UOp.custom_function("hcq", make_submit(*q, devs=devices, queue=queue).sink()).call(name="hcq", aux=info)) + if tag in signal_tags: q += [UOp(Ops.INS, arg="store", src=(make_signal(devices, slots[queue]), UOp.const(tag + 1, dtypes.uint64)))] + src.append(make_call(name, make_submit(*q, devs=devices, queue=queue).sink(), info)) return fences + src + finalizers def sched_hcq_batches(l:UOp) -> UOp: @@ -221,10 +217,10 @@ def sched_hcq_batches(l:UOp) -> UOp: def _merged_hcq_call(calls:list[UOp]) -> UOp: # TODO: simplify? if len(calls) == 1: return calls[0] - info = replace(calls[0].arg.aux, name=f"submit {calls[0].arg.aux.queue} ({len(calls)})", - estimates=sum((c.arg.aux.estimates for c in calls), start=Estimates())) - cmds = [cmd for c in calls for cmd in get_submit(c).src[0].src] - return UOp.custom_function("hcq", make_submit(*cmds, devs=info.device, queue=info.queue).sink()).call(name="hcq", aux=info) + devs, queue = get_submit(calls[0]).src[0].arg + body = make_submit(*[cmd for c in calls for cmd in get_submit(c).src[0].src], devs=devs, queue=queue).sink() + return make_call(f"submit {queue} ({len(calls)})", body, + replace(calls[0].arg.aux, estimates=sum((c.arg.aux.estimates for c in calls), start=Estimates()))) def merge_queues(linear:UOp) -> UOp: new_src:list[UOp] = [] @@ -232,24 +228,25 @@ def merge_queues(linear:UOp) -> UOp: limits:dict[tuple[tuple[str, ...], str], int] = collections.defaultdict(lambda: JIT_BATCH_SIZE.value) for call in linear.src: - if not isinstance(info:=call.arg.aux, HCQInfo) or info.name.startswith("hcq_"): # non-hcq call, fence or finalizer: close all open queues + # non-hcq call, fence or finalizer: close all open queues + if not isinstance(call.arg.aux, HCQInfo) or (call.arg.name or "").startswith("hcq_"): new_src += [_merged_hcq_call(opened_qs.pop(k)) for k in list(opened_qs)] + [call] continue - if (old:=opened_qs.pop(key:=(info.device, info.queue), None)) is not None: + devs, queue = get_submit(call).src[0].arg + if (old:=opened_qs.pop(key:=(devs, queue), None)) is not None: if limits[key] and len(old) >= limits[key]: new_src, old, limits[key] = new_src + [_merged_hcq_call(old)], [], limits[key] * 2 new_rec = old + [call] else: # no such queue opened: close every open submit on this queue that shares a device, so submit order is kept - closing = [k for k in opened_qs if k[1] == info.queue and set(k[0]) & set(info.device)] + closing = [k for k in opened_qs if k[1] == queue and set(k[0]) & set(devs)] new_src += [_merged_hcq_call(opened_qs.pop(k)) for k in closing] new_rec = [call] - opened_qs[(info.device, info.queue)] = new_rec + opened_qs[(devs, queue)] = new_rec return linear.replace(src=tuple(new_src + [_merged_hcq_call(c) for c in opened_qs.values()])) -def schedule_and_merge(ctx:dict[UOp, UOp], linear:UOp) -> UOp: - return merge_queues(sched_hcq_batches(linear).substitute(ctx, walk=True, enter_calls=True)) -pm_schedule_and_merge = PatternMatcher([(UPat(Ops.LINEAR, name="linear"), schedule_and_merge)]) +pm_schedule_and_merge = PatternMatcher([(UPat(Ops.LINEAR, name="l"), + lambda ctx, l: merge_queues(sched_hcq_batches(l).substitute(ctx, walk=True, enter_calls=True)))]) # ***************** # 4.2. hcq lowering: ops to ir @@ -312,7 +309,7 @@ def is_input_addr(g:UOp) -> bool: return all(x.op is Ops.PARAM and x.tag is None def split_patches(call:UOp) -> UOp|None: rt_patches:list[UOp] = [] lt_patches:list[UOp] = [] - body = graph_rewrite(call.src[0], pm_trim_link_patches, ctx=(rt_patches, lt_patches), name=f"trim link-time patches ({call.arg.aux.name})") + body = graph_rewrite(call.src[0], pm_trim_link_patches, ctx=(rt_patches, lt_patches), name=f"trim link-time patches ({call.arg.name})") # split patches inputs, internals = partition(dedup(g for p in rt_patches for g in get_getaddrs(p)), is_input_addr)