From 58edff61d91b1da18de1216ce086ad495de3c663 Mon Sep 17 00:00:00 2001 From: nimlgen <138685161+nimlgen@users.noreply.github.com> Date: Mon, 17 Aug 2026 16:08:19 +0300 Subject: [PATCH] hcq2: one submitter (#17556) * hcq2: c submitter * x * x * x * simpler * simpler * x * x * Dx * revrt * Dx * x * fst * fix --- test/helpers.py | 4 +- tinygrad/engine/realize.py | 36 +++++++------ tinygrad/runtime/ops_cpu.py | 32 ++++++++---- tinygrad/runtime/support/hcq2.py | 89 +++++++++++++++++++++----------- 4 files changed, 103 insertions(+), 58 deletions(-) diff --git a/test/helpers.py b/test/helpers.py index e90b25c525..e6d9916d4e 100644 --- a/test/helpers.py +++ b/test/helpers.py @@ -86,7 +86,9 @@ def assert_jit_cache_len(fxn, expected_len): if linear is None or not linear.src: if expected_len != 0: raise KernelCountException(expected_len, 0) return - if expected_len and all(call_is_hcq(call) for call in linear.src): expected_len = 4 # HCQ2: fence + reset + merged same-queue calls + finalizer + if expected_len and all(call_is_hcq(call) for call in linear.src): # HCQ2: one batch submitter, or fence + reset + merged calls + finalizer + from tinygrad.runtime.support.hcq2 import HCQ_RUNTIME_DEV + expected_len = 1 if HCQ_RUNTIME_DEV.value == "CPU" else 4 if call_is_graph(linear.src[0]): if len(linear.src) != 1: raise KernelCountException(1, len(linear.src)) inner = linear.src[0].src[0].src[0] # LINEAR UOp inside CUSTOM_FUNCTION diff --git a/tinygrad/engine/realize.py b/tinygrad/engine/realize.py index 11458d7696..cbb2bd459b 100644 --- a/tinygrad/engine/realize.py +++ b/tinygrad/engine/realize.py @@ -3,9 +3,10 @@ from typing import cast, Iterator, Any, Sequence import time, random, itertools, math, contextlib, weakref, array from dataclasses import dataclass, replace, field from tinygrad.helpers import colored, DEBUG, GlobalCounters, ansilen, all_int, prod, flatten, Context, getenv, to_tuple -from tinygrad.helpers import BEAM, size_to_str, time_to_str, VALIDATE_WITH_CPU, PROFILE, ProfilePointEvent, cpu_events, wait_cond +from tinygrad.helpers import BEAM, size_to_str, time_to_str, VALIDATE_WITH_CPU, PROFILE, ProfilePointEvent, cpu_events from tinygrad.uop.ops import Ops, PatternMatcher, UOp, UPat, AxisType, sym_infer, graph_rewrite from tinygrad.device import Device, Buffer, MultiBuffer, ProfileGraphEntry +from tinygrad.dtype import dtypes from tinygrad.renderer import Estimates from tinygrad.codegen import to_program from tinygrad.codegen.opt.postrange import args_from_ast @@ -13,7 +14,9 @@ from tinygrad.codegen.opt.postrange import args_from_ast # **************** Helpers **************** def get_call_arg_uops(call:UOp) -> tuple[UOp, ...]: return tuple(s for s in call.src[1:] if not s.is_bound_var) - +def get_call_var_uops(call:UOp, prg:UOp) -> list[UOp]: + bound = {s.src[0].expr: s.src[1].src[1] for s in call.src[1:] if s.is_bound_var} + return [bound.get(v.expr, v) for v in prg.arg.vars] def get_call_outs_ins(call:UOp) -> tuple[tuple[int, ...], tuple[int, ...]]: ast = call.src[0] if ast.op is Ops.PROGRAM: return tuple(ast.arg.outs), tuple(ast.arg.ins) @@ -166,9 +169,10 @@ def exec_copy(ctx:ExecContext, call:UOp, ast:UOp) -> float|None: def exec_kernel(ctx:ExecContext, call:UOp, ast:UOp) -> float|None: et = None - for device, (bufs, device_vars) in zip(to_tuple(call.src[1].device), unwrap_multi(call, resolve_params(call, ctx.input_uops))): + resolved = resolve_params(call, ctx.input_uops) + for device, (bufs, device_vars) in zip(to_tuple(call.src[1].device), unwrap_multi(call, [resolved[i] for i in ast.arg.globals])): var_vals = {**ctx.var_vals, **device_vars} - prg_bufs = [bufs[i].ensure_allocated() for i in ast.arg.globals] + prg_bufs = [b.ensure_allocated() for b in bufs] rt = get_runtime(device, ast, cache=ctx.cache) global_size, local_size = ast.arg.launch_dims(var_vals) with track_stats(ctx, call, device, prg_bufs, var_vals) as tm: @@ -200,28 +204,26 @@ def exec_graph(ctx:ExecContext, call:UOp, ast:UOp) -> float|None: return t[0] def exec_hcq(ctx:ExecContext, call:UOp, ast:UOp) -> float|None: - if (info:=call.arg.aux).inputs is not None: - bufs = [_resolve(ctx.input_uops[i], ctx.input_uops).buffer for i in call.arg.aux.input_idxs] - table = call.src[1+info.inputs].buffer - for j,dev in enumerate(call.arg.aux.device): - addrs = array.array('Q', [(b.bufs[j] if isinstance(b, MultiBuffer) else b).get_buf(dev).va_addr for b in bufs]) - mv = (table.bufs[j] if isinstance(table, MultiBuffer) else table).ensure_allocated()._buf.cpu_view().view(fmt='Q') - wait_cond(lambda: mv[0], value=0, timeout_ms=ctx.timeout or getenv("HCQDEV_WAIT_TIMEOUT_MS", 30000), msg=f"{dev} hang detected") - mv[:len(addrs)] = addrs + dev = cast(Any, Device[(info:= call.arg.aux).device[0]]) + addrs = [(b.bufs[j] if isinstance(b:=_resolve(ctx.input_uops[k], ctx.input_uops).buffer, MultiBuffer) else b).get_buf(dev_name).va_addr + for devs, idxs in info.input_idxs for j, dev_name in enumerate(devs) for k in idxs] + dev.rt_buffer._buf.cpu_view().view(offset=(base:=dev.rt_allocator.alloc(len(addrs) * 8)), fmt='Q')[:len(addrs)] = array.array('Q', addrs) - exec_kernel(replace(ctx, update_stats=DEBUG>=3), call, ast) + tables = [UOp.from_buffer(dev.rt_buffer.view(len(idxs), dtypes.uint64, base + j*len(idxs)*8), HCQ_RUNTIME_DEV.value) + for devs, idxs in info.input_idxs for j in range(len(devs))] + if info.inputs is not None: call = call.substitute({call.src[1+info.inputs]: UOp.mstack(*tables)}) + exec_kernel(replace(ctx, update_stats=DEBUG>=3, var_vals={**ctx.var_vals, "hcq_inputs_ptr": dev.rt_buffer._buf.va_addr + base}), call, ast) tms = [] - for devices,name,estimates,prof in info.kernels: + for devices, stat_call, prof in info.kernels: for device in devices: tm = None if prof: - (d:=cast(Any, Device[device])).prof_ents[prof[0]] = ProfileGraphEntry(device, name, *prof) + (d:=cast(Any, Device[device])).prof_ents[prof[0]] = ProfileGraphEntry(device, stat_call.arg.name, *prof) if ctx.wait: d.synchronize(timeout=ctx.timeout) st, en = (d.signal(x)._buf.cpu_view().view(fmt='Q')[0] for x in prof) tms.append(tm:=float(en-st)/d.timestamp_divider/1e6) - stat_call = call.replace(arg=replace(call.arg, name=name, aux=replace(info, estimates=estimates, kernels=()))) with track_stats(ctx, stat_call, device, [], ctx.var_vals) as et: et[0] = tm return max(tms) if tms else None @@ -262,7 +264,7 @@ pm_exec = PatternMatcher([ (UPat(Ops.CALL, src=(UPat(Ops.CUSTOM_FUNCTION, arg="validate", name="ast"),), name="call", allow_any_len=True), exec_validate), ]) -if getenv("HCQ2"): from tinygrad.runtime.support.hcq2 import hcq_compile, hcq_link # noqa: E402 # down here, hcq2 imports the helpers above +if getenv("HCQ2"): from tinygrad.runtime.support.hcq2 import hcq_compile, hcq_link, HCQ_RUNTIME_DEV # noqa: E402 # down here, hcq2 imports realize def compile_linear(linear:UOp, beam:int|None=None, validate=False, input_uops:list[UOp]|None=None, profile:bool|None=None) -> UOp: if validate: linear = graph_rewrite(linear, pm_validate, name="validate", walk=True) diff --git a/tinygrad/runtime/ops_cpu.py b/tinygrad/runtime/ops_cpu.py index 1ffd2347f0..0aed3949c3 100644 --- a/tinygrad/runtime/ops_cpu.py +++ b/tinygrad/runtime/ops_cpu.py @@ -13,7 +13,7 @@ from tinygrad.renderer.isa.x86 import X86Renderer from tinygrad.runtime.support.elf import jit_loader from tinygrad.runtime.autogen import libc from tinygrad.codegen import do_to_program -from tinygrad.engine.realize import pm_flatten_linear, get_call_arg_uops, get_runtime +from tinygrad.engine.realize import pm_flatten_linear, get_call_arg_uops, get_call_var_uops, get_runtime from tinygrad import UOp, dtypes from tinygrad.dtype import AddrSpace from tinygrad.uop.ops import KernelInfo, Ops, UPat, PatternMatcher, graph_rewrite @@ -64,7 +64,7 @@ def cpu_cmd(devs:tuple[str, ...], prog, *args:UOp) -> UOp: return UOp(Ops.INS, dtypes.void, words + (UOp.const(0, dtypes.uint64),) * (CMD_SIZE - len(words)), arg="cmd") def cpu_exec(ctx:tuple[str, ...], call:UOp, prg:UOp) -> UOp: - args = [get_call_arg_uops(call)[i].getaddr(ctx) for i in prg.arg.globals] + [v.cast(dtypes.uint64) for v in prg.arg.vars] + args = [get_call_arg_uops(call)[i].getaddr(ctx) for i in prg.arg.globals] + [v.cast(dtypes.uint64) for v in get_call_var_uops(call, prg)] if (core:=prg.arg.runtimevars.get('core_id')) is None: return cpu_cmd(ctx, prg, *args) la = [cpu_cmd(ctx,prg,*args[:(cid:=(len(prg.arg.globals)+core))],UOp.const(t, dtypes.uint64),*args[cid+1:]) for t in range(prg.arg.global_size[0])] @@ -99,10 +99,11 @@ def encode_queue(q:UOp) -> UOp: e = UOp.range(cnt, next(UOp.unique_num), dtype=dtypes.int, src=(cmdbuf, ring)) copy = UOp.group(*[ring.index((base + e*CMD_SIZE + w) % ring_words).store(cmdbuf.index(e*CMD_SIZE + w).load()) for w in range(CMD_SIZE)]) - # wake the worker after each entry, keeping the post with the stores stops it from hoisting out of the loop - wake = copy.end(e) if WIN else make_signal(devs, tag="func:sem_post").after(copy).index(0).load().call(sem.index(0), ret_dtype=dtypes.void).end(e) - bumped = put.after(wake).index(0).store(put.index(0).load() + cnt) - return sysbuf.after(bumped).index(0).store(put.index(0).load() + cnt) if WIN else bumped + bumped = put.after(copy.end(e)).index(0).store(put.index(0).load() + cnt) + if WIN: return sysbuf.after(bumped).index(0).store(put.after(bumped).index(0).load()) + + e = UOp.range(cnt, next(UOp.unique_num), dtype=dtypes.int, src=(bumped,)) + return make_signal(devs, tag="func:sem_post").after(e).index(0).load().call(sem.after(e).index(0), ret_dtype=dtypes.void).end(e) # ***************** @@ -196,11 +197,13 @@ class CPUDevice(HCQ2Compiled): pm_lower = PatternMatcher([(UPat(Ops.CUSTOM_FUNCTION, arg="submit_cmdbuf", src=(UPat(Ops.LINEAR, name="q"),)), encode_queue)]) def __init__(self, device:str=""): + self.workers:list[CPUWorker] = [] super().__init__(device, CPUAllocator(self), [ClangRenderer, CPULLVMRenderer, LVPRenderer, X86Renderer], CPUProgram, arch={'amd64':'x86_64', 'aarch64':'arm64'}.get(m:=platform.machine().lower(), m)+",native") self.pm_bufferize = PatternMatcher( - [(UPat(Ops.PARAM, tag=f"COMPUTE:0_{n}"), lambda ctx, n=n: getattr(ctx[0].worker, n)) for n in ("ring", "put", "sem", "sys", "done")] + + [(UPat(Ops.PARAM, tag=f"{q}_{n}"), lambda ctx, q=q,n=n: getattr(ctx[0].worker(q), n)) + for q in ("COMPUTE:0", "SUBMIT:0") for n in ("ring", "put", "sem", "sys", "done")] + [(UPat(Ops.PARAM, tag=f"func:{f}"), lambda ctx, f=f: ctx[0].func_ptr(f)) for f in FUNCS]) + self.pm_bufferize with Context(EMULATED_DTYPES="", TRACK_MATCH_STATS=0): @@ -210,6 +213,12 @@ class CPUDevice(HCQ2Compiled): def func_ptr(self, name:str) -> Buffer: return self.func_table.view(1, dtypes.uint64, FUNCS.index(name)*8).ensure_allocated() + def synchronize(self, timeout:int|None=None): + for worker in self.workers: + put, done = (getattr(worker, x)._buf.cpu_view().view(fmt='Q') for x in ("put", "done")) + while done[0] < put[0]: self._wait_signal(done, put[0], timeout) + super().synchronize(timeout) + @functools.cached_property def func_table(self) -> Buffer: lib = ctypes.windll.kernel32 if sys.platform == "win32" else libc.dll # type: ignore[attr-defined] @@ -217,8 +226,8 @@ class CPUDevice(HCQ2Compiled): array.array('Q', [unwrap(ctypes.cast(getattr(lib, f), ctypes.c_void_p).value) for f in FUNCS]) return ft - @functools.cached_property - def worker(self) -> CPUWorker: + @functools.cache + def worker(self, queue:str) -> CPUWorker: ring, put, sysbuf, done = (Buffer(self.device, sz, dtypes.uint64, preallocate=True) for sz in (RING_SLOTS*CMD_SIZE, 1, 1, 1)) addr, hsem = 0, None @@ -230,5 +239,6 @@ class CPUDevice(HCQ2Compiled): sem = Buffer(self.device, 1, dtypes.uint64, options=BufferSpec(external_ptr=addr), preallocate=True) worker_args = [ring._buf.va_addr, sysbuf._buf.va_addr if WIN else self.func_ptr('sem_wait')._buf.va_addr, done._buf.va_addr, addr] - (worker:=threading.Thread(target=self.prgs[worker_prog].fxn, daemon=True, args=[ctypes.c_uint64(x) for x in worker_args])).start() - return CPUWorker(ring, put, sem, sysbuf, done, worker) + (thread:=threading.Thread(target=self.prgs[worker_prog].fxn, daemon=True, args=[ctypes.c_uint64(x) for x in worker_args])).start() + self.workers.append(worker:=CPUWorker(ring, put, sem, sysbuf, done, thread)) + return worker diff --git a/tinygrad/runtime/support/hcq2.py b/tinygrad/runtime/support/hcq2.py index 0d42960a39..639b948c2b 100644 --- a/tinygrad/runtime/support/hcq2.py +++ b/tinygrad/runtime/support/hcq2.py @@ -30,9 +30,9 @@ class HCQInfo: device:tuple[str, ...] estimates:Estimates = Estimates() - input_idxs:tuple[int, ...] = () # indexes into input_uops used by this call - inputs:int|None = None - kernels:tuple[tuple[tuple[str, ...], str, Estimates, tuple[int, ...]], ...] = () + input_idxs:tuple[tuple[tuple[str, ...], tuple[int, ...]], ...] = () # per inputs table: (devices, indexes into input_uops) + inputs:int|None = None # index of the inputs table in call.src + kernels:tuple[tuple[tuple[str, ...], UOp, tuple[int, ...]], ...] = () # per kernel: (devices, a call carrying its name and estimates, timestamps) def all_devices_in(d:Any, c:frozenset[str]) -> bool: return {x.split(":")[0] for x in to_tuple(d)} <= c @@ -215,7 +215,7 @@ def _finalize_batch(batch:list[tuple[UOp, tuple[str, ...]]], profile:bool) -> li # and make hcq call name, info = get_call_name(call, get_call_arg_uops(call)), HCQInfo(devices, estimate_uop(call)) ts_ids = [next(UOp.unique_num) for _ in range(2)] if profile else [] - kerns.append((devices, name, info.estimates, tuple(ts_ids))) + kerns.append((devices, make_call(name, call.src[0], info), tuple(ts_ids))) ts_ins = [UOp(Ops.INS, arg="timestamp", src=(make_signal(devices, s),)) for s in ts_ids] q += ts_ins[:1] + [call.replace(arg=replace(call.arg, aux=info))] + ts_ins[1:] @@ -345,14 +345,11 @@ def split_patches(call:UOp) -> UOp|None: scatter = make_scatter_loops(input_patches, tables[0], lt_patches) body = body.substitute({p:p.substitute(scatter | reads) for p in rt_patches}) - if inputs: # fence inputs - fills.append((t:=tables[0][0]).after(make_binary_patch(t, bytes(t.max_numel() * 8)))) # zeroed at link, slot 0 is the host fence - body = body.replace(src=(UOp.sink(*body.src[0].src, t.after(*body.src[0].src).index(0).store(0)),)) # open it once consumed - lt_srcs = collections.defaultdict(list) for p in lt_patches: lt_srcs[p.buf_uop].append(p) return call.replace(src=(body, *call.src[1:], *[b.after(*ps) for b,ps in lt_srcs.items()], *fills), - arg=replace(call.arg, aux=replace(call.arg.aux, input_idxs=tuple(sorted(dedup(b.arg.slot for g in inputs for b in unwrap_mstack(g.buf_uop))))))) + arg=replace(call.arg, aux=replace(call.arg.aux, input_idxs=((call.arg.aux.device, + tuple(sorted(dedup(b.arg.slot for g in inputs for b in unwrap_mstack(g.buf_uop))))),) if inputs else call.arg.aux.input_idxs))) pm_split_patches = PatternMatcher([(UPat(Ops.CALL, src=(UPat(Ops.CUSTOM_FUNCTION, arg="hcq"),), name="call", allow_any_len=True), split_patches)]) # ***************** @@ -371,14 +368,13 @@ def replace_params(call:UOp) -> UOp|None: # keep buffers whose addresses become link-time constants alive and mapped held = args + [r.without_after for r in refhold] - addrs = dedup([g.src[0].without_after for x in call.src for g in x.toposort() if g.op is Ops.GETADDR]) + addrs = dedup([g.src[0].without_after for g in call.toposort() if g.op is Ops.GETADDR]) refhold += [a for a in addrs if a not in held and all(b.op is not Ops.PARAM or b.tag is not None for b in unwrap_mstack(a))] sub = {(b:=u.without_after): UOp.param(i, u.dtype, shape=b.shape, device=HCQ_RUNTIME_DEV.value, volatile=b.op is Ops.PARAM and b.arg.volatile) for i,u in enumerate(c_args)} | {v: v.replace(arg=replace(v.arg, slot=-1)) for v in variables if v.op is Ops.PARAM} | _rank_ranges(tops) - info = replace(call.arg.aux, inputs=next((i for i,u in enumerate(c_args) if u.without_after.tag == "inputs"), None)) - return call.replace(src=(body.substitute(sub).replace(arg="hcq_args"), *c_args, *refhold), - arg=replace(call.arg, aux=info)) # TODO: call.after(*refhold)? + info = replace(call.arg.aux, inputs=next((i for i,u in enumerate(c_args + refhold) if u.without_after.tag == "inputs"), None)) + return call.replace(src=(body.substitute(sub).replace(arg="hcq_args"), *c_args, *refhold), arg=replace(call.arg, aux=info)) pm_replace_params = PatternMatcher([ (UPat(Ops.CALL, src=(UPat(Ops.CUSTOM_FUNCTION, arg="hcq"),), name="call", allow_any_len=True), replace_params)]) @@ -424,8 +420,44 @@ def callify_hcq(call:UOp, cf:UOp) -> UOp: pm_callify_hcq = PatternMatcher([(UPat(Ops.CALL, src=( UPat(Ops.CUSTOM_FUNCTION, arg="hcq_args", src=(UPat(Ops.SINK),), name="cf"),), name="call", allow_any_len=True), callify_hcq)]) +# ***************** +# 9. merge submitters + +def _lane_arg(a:UOp, lane:int, table:UOp) -> UOp: return table if a.tag == "inputs" else a.mselect(lane) if len(to_tuple(a.device)) > 1 else a + +def merge_batch(batch:list[UOp]) -> UOp: + tables = UOp.variable("hcq_inputs_ptr", 0, 2**64-1, dtypes.uint64, param=True) + lanes = [(c, j, sum(len(idxs) * 8 for _, idxs in c.arg.aux.input_idxs)) for c in batch for j in range(len(c.arg.aux.device))] # (call, lane, bytes) + offs = itertools.accumulate((table_bytes for _, _, table_bytes in lanes), initial=0) # every lane owns the next table of the region + cmds = [c.src[0].src[0].call(*[_lane_arg(a.without_after, j, tables + off) for a in c.src[1:]], UOp.variable("_device_num", 0, 1 << 30).bind(j)) + for (c, j, _), off in zip(lanes, offs)] + + info = HCQInfo((HCQ_RUNTIME_DEV.value,), sum((c.arg.aux.estimates for c in batch), start=Estimates()), + input_idxs=tuple(x for c in batch for x in c.arg.aux.input_idxs), kernels=tuple(k for c in batch for k in c.arg.aux.kernels)) + body = UOp.custom_function("hcq", make_submit(*cmds, devs=HCQ_RUNTIME_DEV.value, queue="SUBMIT:0").sink()) + return body.call(*[s for c in batch for s in c.src[1:] if s.without_after.tag != "inputs"], name=f"hcq_submitter ({len(batch)})", aux=info) + +def merge_submitters(linear:UOp) -> UOp: + batches = [(k, list(g)) for k, g in itertools.groupby(linear.src, key=lambda c: isinstance(c.arg.aux, HCQInfo))] + return linear.replace(src=tuple(c for is_hcq, b in batches for c in ([merge_batch(b)] if is_hcq else b))) + +# ***************** +# hcq schedule + hcq_compile_cache:dict[tuple[bytes, bool], UOp] = {} +def hcq_lower(linear:UOp, pm_encode:PatternMatcher) -> UOp: + # lowering to hcq ir + linear = graph_rewrite(linear, pm_encode, walk=True, name="encode and pack", enter_calls=True) + + # patches and runtime uops + linear = graph_rewrite(linear, pm_early_simplify+symbolic, bottom_up=False, name="simplify patches", enter_calls=True) + linear = graph_rewrite(linear, pm_split_patches, walk=True, name="split patches") + + # and compile it + linear = graph_rewrite(linear, pm_replace_params, name="replace params") + return graph_rewrite(linear, pm_callify_hcq, name="callify hcq", enter_calls=True) + @rewrite_group(lambda linear,input_uops,profile,ret: f"HCQ Compile {pluralize('Kernel', len(ret.src))}") def hcq_compile(linear:UOp, input_uops:list[UOp]|None, profile:bool) -> UOp: if input_uops is not None: @@ -440,16 +472,9 @@ def hcq_compile(linear:UOp, input_uops:list[UOp]|None, profile:bool) -> UOp: # schedule linear = graph_rewrite(linear, pm_schedule_and_merge, ctx=({s:p for p,s in back_map.items()}, profile), walk=True, name="schedule and merge hcq") - # lowering to hcq ir - linear = graph_rewrite(linear, pm_encode_cmdbufs+pm_pack_placeholders, walk=True, name="encode and pack", enter_calls=True) - - # patches and runtime uops - linear = graph_rewrite(linear, pm_early_simplify+symbolic, bottom_up=False, name="simplify patches", enter_calls=True) - linear = graph_rewrite(linear, pm_split_patches, walk=True, name="split patches") - - # and compile it - linear = graph_rewrite(linear, pm_replace_params, name="replace params") - final_linear = hcq_compile_cache[cache_key] = graph_rewrite(linear, pm_callify_hcq, name="callify hcq", enter_calls=True) + # lower to hcq programs, then pack the programs of every batch into one C submitter (needs a C runtime device for the program addresses) + linear = hcq_lower(linear, pm_encode_cmdbufs+pm_pack_placeholders) + final_linear = hcq_compile_cache[cache_key] = hcq_lower(merge_submitters(linear), pm_encode_cmdbufs) if HCQ_RUNTIME_DEV.value == "CPU" else linear return final_linear @@ -543,7 +568,6 @@ class HCQ2Compiled(Compiled): super().__init__(device, allocator, compilers, runtime, None, arch=arch) - self.rt_buffer = Buffer(self.device, 64 << 20, dtypes.uint8, options=BufferSpec(uncached=True, cpu_access=True)) self.rt_allocator = BumpAllocator(64 << 20) self.prof_ents:dict[int, ProfileGraphEntry] = {} @@ -567,6 +591,10 @@ class HCQ2Compiled(Compiled): tdiffs.append((st+perf_counter_us())/2 - gpu) Compiled.profile_events.append(ProfileDeviceEvent(self.device, statistics.median(tdiffs), self.device_props())) + @functools.cached_property + def rt_buffer(self) -> Buffer: + return Buffer(self.device, self.rt_allocator.size, dtypes.uint8, options=BufferSpec(uncached=True, cpu_access=True), preallocate=True) + def new_buffer(self, b:UOp, cache:bool) -> Buffer: if cache or b.tag in HCQ_CACHE_TAGS: return Buffer(self.device, b.max_numel(), b.dtype, options=BufferSpec(uncached=True, cpu_access=True, nolru=True)) @@ -578,16 +606,19 @@ class HCQ2Compiled(Compiled): buf.as_memoryview(force_zero_copy=True, no_sync=True).cast('Q')[0] = init_value return buf + def _wait_signal(self, sig:memoryview, value:int, timeout:int|None=None): + timeout = timeout if timeout is not None and self.can_recover else None + st, done = time.perf_counter(), sig[0] + while done < value: + if done != (done:=sig[0]): st = time.perf_counter() + elif time.perf_counter() - st > (timeout or self.wait_timeout_ms) / 1000: self.on_device_hang() + def synchronize(self, timeout:int|None=None): if HCQ_RUNTIME_DEV.value != self.device: Device[HCQ_RUNTIME_DEV.value].synchronize() sig = self.signal("timeline").as_memoryview(force_zero_copy=True, no_sync=True).cast('Q') tl = self.signal("value", 1).as_memoryview(force_zero_copy=True, no_sync=True).cast('Q') - timeout = timeout if timeout is not None and self.can_recover else None - st, done = time.perf_counter(), sig[0] - while done < tl[0] - 1: - if done != (done:=sig[0]): st = time.perf_counter() - elif time.perf_counter() - st > (timeout or self.wait_timeout_ms) / 1000: self.on_device_hang() + self._wait_signal(sig, tl[0] - 1, timeout) if self.prof_ents: self.collect_prof() def on_device_hang(self): raise RuntimeError(f"{self.device} hang detected")