From 76dce1eb8d20698cfa8a43c1598efdc44f8aff72 Mon Sep 17 00:00:00 2001 From: nimlgen <138685161+nimlgen@users.noreply.github.com> Date: Fri, 28 Aug 2026 11:29:00 +0300 Subject: [PATCH] tiny hcq2 changes (#17797) * tiny hcq2 changes * x * x --- test/backend/test_jit.py | 3 ++- test/device/test_hcq2.py | 4 ++-- tinygrad/engine/realize.py | 10 +++++----- tinygrad/helpers.py | 2 +- tinygrad/runtime/ops_amd.py | 4 ++-- tinygrad/runtime/support/hcq.py | 14 +++----------- tinygrad/runtime/support/hcq2.py | 12 ++++++------ tinygrad/runtime/support/memory.py | 13 +++++++++++-- 8 files changed, 32 insertions(+), 30 deletions(-) diff --git a/test/backend/test_jit.py b/test/backend/test_jit.py index 23db717e26..a9c34ad678 100644 --- a/test/backend/test_jit.py +++ b/test/backend/test_jit.py @@ -6,7 +6,7 @@ from test.helpers import assert_jit_cache_len, call_is_graph, not_support_multi_ from test.unit.test_jit import _simple_test from tinygrad import Tensor, Variable, TinyJit, Device, dtypes from tinygrad.engine.jit import graph_class -from tinygrad.helpers import JIT, DEV, GlobalCounters +from tinygrad.helpers import JIT, DEV, GlobalCounters, HCQ2 from tinygrad.uop.ops import Ops from tinygrad.renderer.isa.x86 import X86Renderer @@ -235,6 +235,7 @@ class TestJitPrune(unittest.TestCase): assert_jit_cache_len(w2_prune, 1) class TestJitFree(unittest.TestCase): + @unittest.skipIf(HCQ2, "hcq2 keeps refs to intermediate buffers") def test_free_intermediates(self): ext_tensor = Tensor([1,24,23,45,1]) @TinyJit diff --git a/test/device/test_hcq2.py b/test/device/test_hcq2.py index e312e980d3..43c90f93c5 100644 --- a/test/device/test_hcq2.py +++ b/test/device/test_hcq2.py @@ -3,10 +3,10 @@ from unittest.mock import patch from tinygrad import Device, Tensor from tinygrad.device import Buffer from tinygrad.dtype import dtypes -from tinygrad.helpers import getenv +from tinygrad.helpers import HCQ2 from tinygrad.runtime.support.hcq2 import HCQ_DEVS, all_devices_in -@unittest.skipUnless(getenv("HCQ2") and all_devices_in(Device.DEFAULT, HCQ_DEVS), "hcq2 device required") +@unittest.skipUnless(HCQ2 and all_devices_in(Device.DEFAULT, HCQ_DEVS), "hcq2 device required") class TestHCQ2(unittest.TestCase): def test_copy_without_copy_queue(self): with patch.object(Device[Device.DEFAULT], "has_copy_queue", False): diff --git a/tinygrad/engine/realize.py b/tinygrad/engine/realize.py index 31b3a14b87..baac566d76 100644 --- a/tinygrad/engine/realize.py +++ b/tinygrad/engine/realize.py @@ -2,8 +2,8 @@ from __future__ import annotations from typing import cast, Iterator, Any, Sequence import random, itertools, math, weakref, array, decimal from dataclasses import dataclass, replace, field -from tinygrad.helpers import colored, DEBUG, GlobalCounters, ansipad, all_int, prod, flatten, Context, getenv, to_tuple, tqdm, dedup -from tinygrad.helpers import BEAM, size_to_str, time_to_str, VALIDATE_WITH_CPU, PROFILE, ProfilePointEvent, cpu_events, perf_counter_us +from tinygrad.helpers import colored, DEBUG, GlobalCounters, ansipad, all_int, prod, flatten, Context, to_tuple, tqdm, dedup +from tinygrad.helpers import BEAM, size_to_str, time_to_str, VALIDATE_WITH_CPU, HCQ2, PROFILE, ProfilePointEvent, cpu_events, perf_counter_us from tinygrad.uop.ops import Ops, PatternMatcher, UOp, UPat, AxisType, sym_infer, graph_rewrite, ProgramInfo from tinygrad.device import Device, Buffer, MultiBuffer, ProfileGraphEntry from tinygrad.dtype import dtypes @@ -305,17 +305,17 @@ 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, HCQ_RUNTIME_DEV # noqa: E402 # down here, hcq2 imports realize +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) if (beam_val:=BEAM.value if beam is None else beam) >= 1: linear = graph_rewrite(linear, pm_beam, ctx=beam_val, walk=True) linear = lower_and_compile(linear) linear = graph_rewrite(linear, pm_optimize_local_size, name="optimize local size", walk=True) - if getenv("HCQ2"): linear = hcq_compile(linear, input_uops, bool(PROFILE or DEBUG >= 2) if profile is None else profile) + if HCQ2: linear = hcq_compile(linear, input_uops, bool(PROFILE or DEBUG >= 2) if profile is None else profile) return linear -def link_linear(linear:UOp, cache=True) -> UOp: return hcq_link(linear, cache=cache) if getenv("HCQ2") else linear +def link_linear(linear:UOp, cache=True) -> UOp: return hcq_link(linear, cache=cache) if HCQ2 else linear def run_linear(linear:UOp, var_vals:dict[str, int]|None=None, input_uops:Sequence[UOp]=(), update_stats=True, jit=False, wait=False): inputs = list(input_uops) diff --git a/tinygrad/helpers.py b/tinygrad/helpers.py index 6b9276a2fa..b3e7fdded8 100644 --- a/tinygrad/helpers.py +++ b/tinygrad/helpers.py @@ -240,7 +240,7 @@ TRANSCENDENTAL, NOLOCALS = ContextVar("TRANSCENDENTAL", 1), ContextVar("NOLOCALS SPLIT_REDUCEOP, NO_MEMORY_PLANNER, LRU = ContextVar("SPLIT_REDUCEOP", 1), ContextVar("NO_MEMORY_PLANNER", 0), ContextVar("LRU", 1) RING, ALL2ALL, ALLREDUCE_CAST = ContextVar("RING", 1), ContextVar("ALL2ALL", 0), ContextVar("ALLREDUCE_CAST", 1) CACHELEVEL, IGNORE_BEAM_CACHE = ContextVar("CACHELEVEL", 2), ContextVar("IGNORE_BEAM_CACHE", 0) -VALIDATE_WITH_CPU = ContextVar("VALIDATE_WITH_CPU", 0) +VALIDATE_WITH_CPU, HCQ2 = ContextVar("VALIDATE_WITH_CPU", 0), ContextVar("HCQ2", 0) # TODO: this is broken for some indexing DISABLE_FAST_IDIV = ContextVar("DISABLE_FAST_IDIV", 1) FUSE_OPTIM = ContextVar("FUSE_OPTIM", 0) diff --git a/tinygrad/runtime/ops_amd.py b/tinygrad/runtime/ops_amd.py index 3a30cb6791..b3094e11ac 100644 --- a/tinygrad/runtime/ops_amd.py +++ b/tinygrad/runtime/ops_amd.py @@ -8,7 +8,7 @@ from tinygrad.runtime.support.hcq import MMIOInterface, BumpAllocator, hcq_filte from tinygrad.uop.ops import sint from tinygrad.device import Compiled, BufferSpec, TinyELF from tinygrad.helpers import getenv, round_up, data64_le, DEBUG, PROFILE, ProfileEvent, lo32, hi32, colored, prod, ContextVar, TracingKey -from tinygrad.helpers import VIZ, ceildiv, unwrap, pluralize +from tinygrad.helpers import VIZ, HCQ2, ceildiv, unwrap, pluralize from tinygrad.renderer.cstyle import HIPRenderer, HIPCCRenderer from tinygrad.renderer.llvmir import AMDLLVMRenderer from tinygrad.runtime.autogen import kfd, hsa, sqtt, amdgpu_kd, amdgpu_drm @@ -1153,4 +1153,4 @@ class AMDDevice(HCQCompiled): def hw_copy_queues(self): return [(f"SDMA:{i}", functools.partial(unwrap(self.hw_copy_queue_t), queue_idx=i)) for i in self.sdma_queues] -if getenv("HCQ2"): from extra.hcq2.ops_amd2 import * # noqa: F401, F403 # pylint: disable=unused-import +if HCQ2: from extra.hcq2.ops_amd2 import * # noqa: F401, F403 # pylint: disable=unused-import diff --git a/tinygrad/runtime/support/hcq.py b/tinygrad/runtime/support/hcq.py index 1481fe7e0c..14db571e3a 100644 --- a/tinygrad/runtime/support/hcq.py +++ b/tinygrad/runtime/support/hcq.py @@ -1,24 +1,16 @@ from __future__ import annotations from typing import cast, Callable, Type, TypeVar, Generic, Any -import contextlib, decimal, statistics, time, ctypes, array, os, struct, collections, itertools +import contextlib, decimal, statistics, time, ctypes, array, os, collections, itertools try: import fcntl # windows misses that except ImportError: fcntl = None #type:ignore[assignment] -from tinygrad.helpers import DEV, PROFILE, getenv, to_mv, from_mv, cpu_profile, ProfileRangeEvent, unwrap +from tinygrad.helpers import DEV, PROFILE, getenv, from_mv, cpu_profile, ProfileRangeEvent, unwrap from tinygrad.helpers import suppress_finalizing, pluralize, TracingKey from tinygrad.device import Device, BufferSpec, Compiled, LRUAllocator, ProfileDeviceEvent, ProfileProgramEvent, Program, TinyELF from tinygrad.uop.ops import sym_infer, sint, UOp from tinygrad.runtime.autogen import libc -from tinygrad.runtime.support.memory import BumpAllocator +from tinygrad.runtime.support.memory import BumpAllocator, MMIOInterface from tinygrad.renderer import Renderer -class MMIOInterface: - def __init__(self, addr:int, nbytes:int, fmt='B'): self.mv, self.addr, self.nbytes, self.fmt = to_mv(addr, nbytes).cast(fmt), addr, nbytes, fmt - def __len__(self): return self.nbytes // struct.calcsize(self.fmt) - def __getitem__(self, k): return (self.mv[k] if self.fmt == 'B' else self.mv[k].tolist()) if isinstance(k, slice) else self.mv[k] - def __setitem__(self, k, v): self.mv[k] = v - def view(self, offset:int=0, size:int|None=None, fmt=None) -> MMIOInterface: - return MMIOInterface(self.addr+offset, (self.nbytes - offset) if size is None else size, fmt=fmt or self.fmt) - class FileIOInterface: """ Hardware Abstraction Layer for HCQ devices. The class provides a unified interface for interacting with hardware devices. diff --git a/tinygrad/runtime/support/hcq2.py b/tinygrad/runtime/support/hcq2.py index 3bc3806b29..f91a4b9cec 100644 --- a/tinygrad/runtime/support/hcq2.py +++ b/tinygrad/runtime/support/hcq2.py @@ -1,5 +1,5 @@ from __future__ import annotations -from typing import cast, TypeVar, Generic, Any, Sequence, Iterable +from typing import cast, TypeVar, Generic, Any, Sequence, Iterable, TYPE_CHECKING import struct, functools, time, collections, itertools, decimal, statistics from dataclasses import replace, dataclass, field from tinygrad.helpers import suppress_finalizing, dedup, pluralize, JIT_BATCH_SIZE, unwrap, PROFILE @@ -9,11 +9,11 @@ from tinygrad.device import ProfileDeviceEvent, ProfileGraphEntry, ProfileGraphE from tinygrad.uop.ops import Ops, sint, UOp, UPat, PatternMatcher, KernelInfo, graph_rewrite, rewrite_group, GroupOp from tinygrad.uop.symbolic import symbolic from tinygrad.dtype import dtypes, truncate, DType -from tinygrad.runtime.support.hcq import MMIOInterface, HCQBuffer -from tinygrad.runtime.support.memory import BumpAllocator +from tinygrad.runtime.support.memory import BumpAllocator, MMIOInterface from tinygrad.renderer import Renderer, Estimates -from tinygrad.engine.realize import to_program, get_call_arg_uops, get_call_name, get_call_outs_ins, estimate_uop -from tinygrad.engine.realize import pm_flatten_linear, lower_and_compile +from tinygrad.engine.realize import to_program, get_call_arg_uops, get_call_name, get_call_outs_ins, estimate_uop, pm_flatten_linear,lower_and_compile + +if TYPE_CHECKING: from tinygrad.runtime.support.hcq import HCQBuffer # TODO: remove that # ***************** # 0. helpers @@ -462,7 +462,7 @@ def hcq_lower(linear:UOp, pm_encode:PatternMatcher) -> UOp: linear = graph_rewrite(linear, pm_split_patches, walk=True, name="split patches") # and compile it - return lower_and_compile(graph_rewrite(linear, pm_replace_params, walk=True, name="replace params")) + with Context(EMULATED_DTYPES=""): return lower_and_compile(graph_rewrite(linear, pm_replace_params, walk=True, name="replace params")) @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: diff --git a/tinygrad/runtime/support/memory.py b/tinygrad/runtime/support/memory.py index 30d0fbc5eb..317c2e2121 100644 --- a/tinygrad/runtime/support/memory.py +++ b/tinygrad/runtime/support/memory.py @@ -1,6 +1,15 @@ -import collections, functools, dataclasses, enum +from __future__ import annotations +import collections, functools, dataclasses, enum, struct from typing import Any, ClassVar -from tinygrad.helpers import round_up, getenv +from tinygrad.helpers import round_up, getenv, to_mv + +class MMIOInterface: + def __init__(self, addr:int, nbytes:int, fmt='B'): self.mv, self.addr, self.nbytes, self.fmt = to_mv(addr, nbytes).cast(fmt), addr, nbytes, fmt + def __len__(self): return self.nbytes // struct.calcsize(self.fmt) + def __getitem__(self, k): return (self.mv[k] if self.fmt == 'B' else self.mv[k].tolist()) if isinstance(k, slice) else self.mv[k] + def __setitem__(self, k, v): self.mv[k] = v + def view(self, offset:int=0, size:int|None=None, fmt=None) -> MMIOInterface: + return MMIOInterface(self.addr+offset, (self.nbytes - offset) if size is None else size, fmt=fmt or self.fmt) class BumpAllocator: def __init__(self, size:int, base:int=0, wrap:bool=True): self.size, self.ptr, self.base, self.wrap = size, 0, base, wrap