diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 7fe9471032..907d73a2fa 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -522,9 +522,10 @@ jobs: run: HCQ_RUNTIME_DEV=PYTHON HCQ2=1 DEV=MOCKKFD+AMD FORWARD_ONLY=1 PYTHONPATH=. python test/test_tiny.py - name: Run HCQ2 multi-device tests run: | - HCQ_RUNTIME_DEV=PYTHON HCQ2=1 DEV=MOCKKFD+AMD FORWARD_ONLY=1 PYTHONPATH=. python test/unit/test_multitensor.py \ - TestMultiTensor.test_simple_add TestMultiTensor.test_shard_reduce \ - TestMultiTensor.test_backward_sum TestMultiTensor.test_matmul_shard_0_0 + HCQ_RUNTIME_DEV=PYTHON HCQ2=1 DEV=MOCKKFD+AMD FORWARD_ONLY=1 PYTHONPATH=. python -m pytest -n=auto test/backend/test_multitensor.py \ + -k "not (multitensor_jit_input or multitensor_inside_jit or four_add or elementwise_dtype or stack or \ + test_2d_shard_basic or test_2d_shard_elementwise or test_2d_shard_sum_non_sharded_axis or test_2d_shard_matmul or test_numpy or \ + data_parallel_simple_train_step or multi_tensor_jit_graph_assign_updates_each_shard or test_transformer or test_shrink_2d)" - name: Run HCQ2 JIT tests run: HCQ_RUNTIME_DEV=PYTHON HCQ2=1 DEV=MOCKKFD+AMD FORWARD_ONLY=1 PYTHONPATH=. python test/unit/test_jit.py - name: Run HCQ2 unit tests diff --git a/tinygrad/runtime/support/hcq2.py b/tinygrad/runtime/support/hcq2.py index 2464928ca5..dd045c5e19 100644 --- a/tinygrad/runtime/support/hcq2.py +++ b/tinygrad/runtime/support/hcq2.py @@ -141,14 +141,14 @@ def _build_wait_cmds(slots:dict[str, int], dep_lanes:list[tuple[tuple, int, int] # opt2: keep latest dep per (dep device, queue, cur lane) latest = {((dep[0][dlane], dep[1]), lane): (dep, dlane) for dep, dlane, lane in sorted(dep_lanes, key=lambda x: x[0][2])} - deps:dict[tuple, list[int|None]] = collections.defaultdict(lambda: [None]*len(devices)) - for (_, lane), (dep, dlane) in latest.items(): deps[dep][lane] = dlane + deps:dict[tuple, dict[int, list[int]]] = collections.defaultdict(lambda: collections.defaultdict(list)) + for (_, lane), (dep, dlane) in latest.items(): deps[dep][lane].append(dlane) waits = [] - for (ddevs, dqueue, dtag), lanes in deps.items(): - sig = UOp.mstack(*[make_signal(d, tag="sentinel_signal") if dl is None else make_signal(ddevs[dl], slots[dqueue]) - for dl, d in zip(lanes, devices)]) - waits.append(UOp(Ops.INS, arg="wait", src=(sig, UOp.const(dtag + 1, dtypes.uint64)))) + for (ddevs, dqueue, dtag), by_lane in deps.items(): + for ls in itertools.zip_longest(*(by_lane[lane] for lane in range(len(devices)))): + s = UOp.mstack(*[make_signal(d, tag="sentinel_signal") if dl is None else make_signal(ddevs[dl], slots[dqueue]) for dl, d in zip(ls, devices)]) + waits.append(UOp(Ops.INS, arg="wait", src=(s, UOp.const(dtag + 1, dtypes.uint64)))) return waits, {dtag for _, _, dtag in deps} def _build_finalizers(batch:list[tuple[UOp, tuple[str, ...]]], batch_info:list[tuple[tuple[str, ...], str]],