From 92ff3633b4487ccf9773d48b08e2ce1be6801f4d Mon Sep 17 00:00:00 2001 From: puddingfjz <2811443837@qq.com> Date: Tue, 28 Jul 2026 20:51:02 -0700 Subject: [PATCH] Fix: reserve NEXT_LEVEL workers for ready groups A blocked NEXT_LEVEL group reserves only the head group's targets from later singles until the complete set becomes idle. The reservation flows directly from the group pass to the singles pass, preserving all-or-nothing dispatch while unrelated queues continue. Scheduler regressions cover staged target completion, unrelated workers, and consecutive groups. Peer-waiting collectives submit each complete rank cohort through the group API so they satisfy the same-level placement contract. Fixes #1551 --- docs/directed-next-level-scheduling.md | 17 ++- docs/orchestrator.md | 13 +- docs/scheduler.md | 23 +-- examples/workers/l3/allreduce/README.md | 4 +- examples/workers/l3/allreduce/main.py | 4 +- examples/workers/l3/domain_rank_map/README.md | 4 +- examples/workers/l3/domain_rank_map/main.py | 4 +- .../workers/l3/dual_domain_overlap/README.md | 2 +- .../workers/l3/dual_domain_overlap/main.py | 4 +- .../workers/l3/ep_dispatch_combine/README.md | 1 + .../workers/l3/ep_dispatch_combine/main.py | 4 +- examples/workers/l3/ffn_tp_parallel/README.md | 1 + examples/workers/l3/ffn_tp_parallel/main.py | 4 +- src/common/hierarchical/scheduler.cpp | 20 ++- src/common/hierarchical/scheduler.h | 5 +- tests/st/worker/collectives/_helpers.py | 12 +- tests/ut/cpp/hierarchical/test_scheduler.cpp | 139 +++++++++++++++--- 17 files changed, 196 insertions(+), 65 deletions(-) diff --git a/docs/directed-next-level-scheduling.md b/docs/directed-next-level-scheduling.md index acb368d5b0..cddef648a5 100644 --- a/docs/directed-next-level-scheduling.md +++ b/docs/directed-next-level-scheduling.md @@ -44,14 +44,15 @@ routing operation. Each Scheduler iteration: -- launches the group FIFO head only when every one of its target workers is - idle — all members or none, with no partial reservation; -- dispatches the head of each idle worker's single FIFO; +- checks the group FIFO head and either launches every member together or + reserves all of its target workers until the complete set is idle; +- dispatches the head of each idle, unreserved worker's single FIFO; - dispatches READY SUB work through the existing free-selection path. -Only the group FIFO head is examined, and a launchable group is tried before -conflicting singles. The runtime adds no fairness, aging, priority, or -reservation policy; callers choose worker sets that make acceptable progress. +Only the group FIFO head reserves workers. A reservation blocks new singles on +the group's targets while their running work drains; singles on other workers +continue normally. The reservation is released when the complete group +launches. Later groups remain behind the FIFO head and do not reserve workers. ## Invariants @@ -62,13 +63,15 @@ reservation policy; callers choose worker sets that make acceptable progress. terminal state. - The NEXT_LEVEL worker set is fixed after `init()`; a worker lookup that fails after submit-time validation is an invariant violation, not a fallback. +- A NEXT_LEVEL single does not wait for a not-yet-dispatched peer at the same + scheduler level. Work that requires concurrent placement uses the group API. - No queue or scheduler mutex is held while an endpoint executes. ## Non-goals - No SUB worker selection. - No priorities, work stealing, rebinding, queue scanning, aging, quotas, - preemption, starvation prevention, or partial group reservation. + preemption, or fairness beyond the group-head reservation. - No compatibility flags, environment variables, or macros. ## Related documents diff --git a/docs/orchestrator.md b/docs/orchestrator.md index f1457c7e3f..f027b33d10 100644 --- a/docs/orchestrator.md +++ b/docs/orchestrator.md @@ -239,6 +239,10 @@ A group task is a single DAG node that executes in parallel on N workers. Each worker gets its own `TaskArgs`; the node only reaches COMPLETED when all N finish. +Callers submit tasks that wait for same-level peers as one complete group. +Submitting those members as independent singles can start one member before +its peers are READY, leaving the running member unable to finish. + ```cpp SubmitResult Orchestrator::submit_next_level_group( const CallableIdentity &callable, const std::vector &args_list, @@ -259,10 +263,11 @@ routing. At dispatch time the Scheduler checks the group FIFO head and resolves every entry in `workers` to that exact stable worker ID. It dispatches only if the -entire target set is idle; a blocked group reserves no partial worker set and -does not cause a scan past the FIFO head. Each WorkerThread runs `worker->run` -with its own `task_args_list[i]`. Completion remains aggregated at the group -slot, so downstream consumers are released once after every member is terminal. +entire target set is idle. A blocked group reserves all of its targets against +new singles but does not cause a scan past the FIFO head. Each WorkerThread +runs `worker->run` with its own `task_args_list[i]`. Completion remains +aggregated at the group slot, so downstream consumers are released once after +every member is terminal. --- diff --git a/docs/scheduler.md b/docs/scheduler.md index 3f91b36709..25622cfa5a 100644 --- a/docs/scheduler.md +++ b/docs/scheduler.md @@ -54,8 +54,8 @@ while (true) { on_task_complete(item); } - dispatch_next_level_group(); - dispatch_next_level_singles(); + reserved_workers = dispatch_next_level_group(); + dispatch_next_level_singles(reserved_workers); dispatch_sub_ready(); if (stop_requested && all workers are idle) { @@ -86,8 +86,8 @@ if target worker is idle and its FIFO is non-empty: ``` There is no idle-worker search, rebinding, work stealing, or scan into another -worker's queue. FIFO is independent per worker, so a busy worker A does not -block READY work for worker B. +worker's queue. Outside a group-head reservation, each worker FIFO progresses +independently, so a busy worker A does not block READY work for worker B. ### Group tasks @@ -103,13 +103,15 @@ if every target is idle: dispatch member i to target_worker_ids[i] else: leave the group at the FIFO head + reserve every target against single-task dispatch ``` -The check is all-or-nothing: a blocked group reserves no partial worker set. -The Scheduler does not scan later groups. It continues to single-task queues, -so the runtime adds no fairness, aging, priority, or reservation policy beyond -trying a launchable group before singles in each iteration. Users are -responsible for choosing worker sets that make acceptable progress. +Dispatch remains all-or-nothing: no group member launches until the complete +target set is idle. A blocked FIFO head reserves every target against new +single-task dispatch while existing work drains. Singles on other workers +continue normally, and later groups do not reserve workers because the +Scheduler does not scan past the FIFO head. The reservation is released when +the group launches. ## 4. SUB dispatch @@ -184,6 +186,9 @@ The scheduling invariants are: 5. Only the Scheduler calls `WorkerThread::dispatch`. 6. Only one successful PENDING-to-READY transition enqueues a consumer. 7. A group produces one aggregate DAG completion regardless of member count. +8. A NEXT_LEVEL single does not wait for a not-yet-dispatched peer at the same + scheduler level; work that requires concurrent placement is submitted as a + group. ## 8. Related documents diff --git a/examples/workers/l3/allreduce/README.md b/examples/workers/l3/allreduce/README.md index ccf02afc40..313254e63a 100644 --- a/examples/workers/l3/allreduce/README.md +++ b/examples/workers/l3/allreduce/README.md @@ -9,8 +9,8 @@ Walks through the full flow: 3. **`worker.init()`** — fork chip children; lazy base-communication init 4. **`orch.allocate_domain(...)`** — allocate a communication domain with a `CommBufferSpec` scratch window -5. **`orch.submit_next_level(chip_handle, chip_args, cfg, worker=i)`** — - submit the allreduce task for each rank +5. **`orch.submit_next_level_group(chip_handle, args_list, cfg, workers=...)`** — + submit every mutually waiting allreduce rank as one group 6. **`worker.run(orch_fn, ...)`** — execute the DAG and golden-check against the known expected sum diff --git a/examples/workers/l3/allreduce/main.py b/examples/workers/l3/allreduce/main.py index 8c25f37920..945f7b5f97 100644 --- a/examples/workers/l3/allreduce/main.py +++ b/examples/workers/l3/allreduce/main.py @@ -151,6 +151,7 @@ def orch_fn(orch, _args, cfg): window_size=window_size, buffers=[CommBufferSpec(name="scratch", dtype="float32", count=float_elems, nbytes=scratch_nbytes)], ) as handle: + args_list = [] for i in range(nranks): domain = handle[i] chip_args = TaskArgs() @@ -167,7 +168,8 @@ def orch_fn(orch, _args, cfg): ) chip_args.add_scalar(domain.domain_size) chip_args.add_scalar(domain.device_ctx) - orch.submit_next_level(chip_handle, chip_args, cfg, worker=i) + args_list.append(chip_args) + orch.submit_next_level_group(chip_handle, args_list, cfg, workers=list(range(nranks))) print(f"[allreduce] running {nranks}-chip allreduce DAG...") worker.run(orch_fn, args=None, config=CallConfig()) diff --git a/examples/workers/l3/domain_rank_map/README.md b/examples/workers/l3/domain_rank_map/README.md index 34488435b0..49f11efc5d 100644 --- a/examples/workers/l3/domain_rank_map/README.md +++ b/examples/workers/l3/domain_rank_map/README.md @@ -23,7 +23,9 @@ tail = workers [1, 2] # chip 2 is in both After the inspection pass, each domain runs its own small allreduce — in its **own `worker.run()`**, so chip 2 never juggles two collectives at once. That separation is deliberate; see `dual_domain_overlap` for the case where both -domains are live across the same DAG. +domains are live across the same DAG. The ranks within each allreduce are one +`submit_next_level_group`, so every peer that participates in the device +barrier is dispatched as a complete set. ## Run diff --git a/examples/workers/l3/domain_rank_map/main.py b/examples/workers/l3/domain_rank_map/main.py index ce049e305a..801b54302f 100644 --- a/examples/workers/l3/domain_rank_map/main.py +++ b/examples/workers/l3/domain_rank_map/main.py @@ -210,13 +210,15 @@ def _orch_fn(orch, _args, cfg): window_size=WINDOW_SIZE, buffers=_scratch_buffers(), ) as handle: + args_list = [] for worker_idx in DOMAINS[domain_name]: domain = handle[worker_idx] args = TaskArgs() args.add_tensor(make_tensor_arg(host_inputs[worker_idx]), TensorArgType.INPUT) args.add_tensor(make_tensor_arg(outputs[domain_name][worker_idx]), TensorArgType.OUTPUT_EXISTING) _add_domain_scratch(args, domain) - orch.submit_next_level(allreduce_handle, args, cfg, worker=worker_idx) + args_list.append(args) + orch.submit_next_level_group(allreduce_handle, args_list, cfg, workers=DOMAINS[domain_name]) return _orch_fn diff --git a/examples/workers/l3/dual_domain_overlap/README.md b/examples/workers/l3/dual_domain_overlap/README.md index fa14eb3c7a..5bb40bb025 100644 --- a/examples/workers/l3/dual_domain_overlap/README.md +++ b/examples/workers/l3/dual_domain_overlap/README.md @@ -20,7 +20,7 @@ through both and then computes on the results. | ------- | --- | | **Per-domain identity** | Chip 1 is rank 1 in `left` and rank 0 in `right`. The `workers` list order defines the dense rank, so the same chip legitimately holds two different ranks at once. | | **Domains allocated inside the orch function** | `with orch.allocate_domain(name=..., workers=..., window_size=..., buffers=[CommBufferSpec(...)])` — created and released within one orchestration, not configured on the `Worker`. | -| **`submit_next_level_group`** | The affine stage submits one `TaskArgs` per member in a single call, with `workers=worker_indices`, instead of a loop of `submit_next_level`. | +| **`submit_next_level_group`** | Both the peer-waiting allreduce and the affine stage submit one `TaskArgs` per member in a single call with `workers=worker_indices`. | | **Compute that depends only on its own domain's result** | Each affine task reads `reduce_out[domain][chip]`. The dependency is implicit — same `buffer.addr` as the reduce output — so `left`'s affine work can never consume `right`'s reduction. | ## Run diff --git a/examples/workers/l3/dual_domain_overlap/main.py b/examples/workers/l3/dual_domain_overlap/main.py index 16eb0e1736..9f231c569c 100644 --- a/examples/workers/l3/dual_domain_overlap/main.py +++ b/examples/workers/l3/dual_domain_overlap/main.py @@ -208,6 +208,7 @@ def _orch_fn(orch, _args, cfg): window_size=WINDOW_SIZE, buffers=_scratch_buffers(), ) as handle: + args_list = [] for worker_idx in worker_indices: domain = handle[worker_idx] print( @@ -221,7 +222,8 @@ def _orch_fn(orch, _args, cfg): make_tensor_arg(reduce_out[domain_name][worker_idx]), TensorArgType.OUTPUT_EXISTING ) _add_domain_scratch(args, domain) - orch.submit_next_level(allreduce_handle, args, cfg, worker=worker_idx) + args_list.append(args) + orch.submit_next_level_group(allreduce_handle, args_list, cfg, workers=worker_indices) return _orch_fn diff --git a/examples/workers/l3/ep_dispatch_combine/README.md b/examples/workers/l3/ep_dispatch_combine/README.md index d2874850f0..8f8d09ec50 100644 --- a/examples/workers/l3/ep_dispatch_combine/README.md +++ b/examples/workers/l3/ep_dispatch_combine/README.md @@ -56,6 +56,7 @@ scatters into `routed_y_buf` without clearing it first. | Concept | How | | ------- | --- | | **Three children under one `ChipCallable`** | `children=[(0, dispatch), (1, local_expert), (2, combine)]` — the integers are `func_id`s, matching the `rt_submit_aiv_task(0/1/2, …)` calls in the orchestration. Each child declares only the args it consumes; the orchestration signature is the union. | +| **Collective group dispatch** | The complete per-rank callable is submitted as one NEXT_LEVEL group because its dispatch and combine phases wait on peer ranks. | | **Ordering without dependencies** | The three tasks run back-to-back because `rt_submit_aiv_task` dispatches in submission order, not because any tensor edge forces it. | | **Chaining through host-backed tensors** | `recv_x_out` / `recv_w_out` / `recv_count_out` are `OUTPUT_EXISTING` for dispatch and inputs to `local_expert`; `recv_y` likewise feeds `combine`. | | **A hand-laid-out window** | `SCRATCH_NBYTES` sums every region — counts table, two signal areas, three receive windows, the combine push destination, a third signal — and must match the `kOff*` offsets in the kernels. | diff --git a/examples/workers/l3/ep_dispatch_combine/main.py b/examples/workers/l3/ep_dispatch_combine/main.py index 767fa3cf02..b9e3fe0d33 100644 --- a/examples/workers/l3/ep_dispatch_combine/main.py +++ b/examples/workers/l3/ep_dispatch_combine/main.py @@ -548,6 +548,7 @@ def orch_fn(orch, _args, cfg): ) ], ) as handle: + args_list = [] for i in range(nranks): domain = handle[i] print( @@ -577,7 +578,8 @@ def orch_fn(orch, _args, cfg): ) chip_args.add_scalar(domain.domain_size) chip_args.add_scalar(domain.device_ctx) - orch.submit_next_level(chip_handle, chip_args, cfg, worker=i) + args_list.append(chip_args) + orch.submit_next_level_group(chip_handle, args_list, cfg, workers=list(range(nranks))) print("[ep_dispatch] running 2-chip dispatch DAG...") worker.run(orch_fn, args=None, config=CallConfig()) diff --git a/examples/workers/l3/ffn_tp_parallel/README.md b/examples/workers/l3/ffn_tp_parallel/README.md index f29aae51e9..dcef6d707d 100644 --- a/examples/workers/l3/ffn_tp_parallel/README.md +++ b/examples/workers/l3/ffn_tp_parallel/README.md @@ -17,6 +17,7 @@ Per rank: | Concept | How | | ------- | --- | | **Implicit producer/consumer edge** | `host_partial[i]` is `OUTPUT_EXISTING` on the stage-1 submit and `INPUT` on the stage-2 submit. Both carry the same `buffer.addr`, so TensorMap links the two tasks itself — there is no barrier, no event, and no ordering call in the orch function. | +| **Collective group dispatch** | Stage 2 is one `submit_next_level_group`, so it becomes READY after every rank's stage-1 output exists and dispatches all mutually waiting ranks together. | | **Mixed core types in one DAG** | Stage 1 compiles with `core_type="aic"` and its orchestration calls `rt_submit_aic_task`; stage 2 uses `core_type="aiv"` and `rt_submit_aiv_task`. One `Worker`, one `run()`. | | **`func_id`, not core id** | The integer in `children=[(0, core_callable)]` is the `func_id` the orchestration passes to `rt_submit_*_task(func_id, params)` — here `0` for the matmul and `1` for the reduce. It selects *which child kernel*, not which core type. | | **Cross-rank exchange through a domain buffer** | The stage-2 kernel reduces over a `scratch` buffer in the communication window: a mailbox of `nranks * M * N` floats followed by a signal tail of `nranks` int32 slots. | diff --git a/examples/workers/l3/ffn_tp_parallel/main.py b/examples/workers/l3/ffn_tp_parallel/main.py index 7379dbf28d..c005db185e 100644 --- a/examples/workers/l3/ffn_tp_parallel/main.py +++ b/examples/workers/l3/ffn_tp_parallel/main.py @@ -201,6 +201,7 @@ def orch_fn(orch, _args, cfg): window_size=window_size, buffers=[CommBufferSpec(name="scratch", dtype="float32", count=scratch_count, nbytes=scratch_nbytes)], ) as handle: + allreduce_args = [] for i in range(nranks): domain = handle[i] print( @@ -233,7 +234,8 @@ def orch_fn(orch, _args, cfg): ) a2.add_scalar(domain.domain_size) a2.add_scalar(domain.device_ctx) - orch.submit_next_level(allreduce_handle, a2, cfg, worker=i) + allreduce_args.append(a2) + orch.submit_next_level_group(allreduce_handle, allreduce_args, cfg, workers=list(range(nranks))) print("[ffn_tp_parallel] running 2-chip 2-stage DAG...") worker.run(orch_fn, args=None, config=CallConfig()) diff --git a/src/common/hierarchical/scheduler.cpp b/src/common/hierarchical/scheduler.cpp index acd6e22420..8949b6aac7 100644 --- a/src/common/hierarchical/scheduler.cpp +++ b/src/common/hierarchical/scheduler.cpp @@ -11,7 +11,6 @@ #include "scheduler.h" -#include #include #include @@ -350,8 +349,8 @@ void Scheduler::try_consume(TaskSlot slot) { // sched_thread_ with no surrounding handler, any throw is fatal to the whole // worker tree (std::terminate), not a per-task failure. void Scheduler::dispatch_ready() { - dispatch_next_level_group(); - dispatch_next_level_singles(); + const std::unordered_set reserved_worker_ids = dispatch_next_level_group(); + dispatch_next_level_singles(reserved_worker_ids); dispatch_sub_ready(); } @@ -392,7 +391,7 @@ void Scheduler::dispatch_sub_ready() { } } -void Scheduler::dispatch_next_level_group() { +std::unordered_set Scheduler::dispatch_next_level_group() { TaskSlot slot; while (cfg_.ready_next_level_queues->try_front_group(slot)) { TaskSlotState &s = *cfg_.ring->slot_state(slot); @@ -408,18 +407,22 @@ void Scheduler::dispatch_next_level_group() { const int32_t group_size = s.group_size(); std::vector workers; workers.reserve(static_cast(group_size)); + std::unordered_set target_worker_ids; + target_worker_ids.reserve(static_cast(group_size)); + bool all_workers_idle = true; for (int32_t i = 0; i < group_size; ++i) { const int32_t worker_id = s.target_worker_id(i); WorkerThread *worker = cfg_.manager->get_worker_by_id(WorkerType::NEXT_LEVEL, worker_id); if (worker == nullptr) { throw std::runtime_error("Scheduler::dispatch_next_level_group: invalid target worker"); } - if (std::find(workers.begin(), workers.end(), worker) != workers.end()) { + if (!target_worker_ids.insert(worker_id).second) { throw std::runtime_error("Scheduler::dispatch_next_level_group: duplicate target worker"); } - if (!worker->idle()) return; + if (!worker->idle()) all_workers_idle = false; workers.push_back(worker); } + if (!all_workers_idle) return target_worker_ids; TaskSlot popped; if (!cfg_.ready_next_level_queues->try_pop_group(popped) || popped != slot) { @@ -432,10 +435,13 @@ void Scheduler::dispatch_next_level_group() { workers[static_cast(i)]->dispatch(WorkerDispatch{slot, i}); } } + return {}; } -void Scheduler::dispatch_next_level_singles() { +void Scheduler::dispatch_next_level_singles(const std::unordered_set &reserved_worker_ids) { for (int32_t worker_id : cfg_.ready_next_level_queues->worker_ids()) { + if (reserved_worker_ids.find(worker_id) != reserved_worker_ids.end()) continue; + WorkerThread *worker = cfg_.manager->get_worker_by_id(WorkerType::NEXT_LEVEL, worker_id); if (worker == nullptr) { throw std::runtime_error( diff --git a/src/common/hierarchical/scheduler.h b/src/common/hierarchical/scheduler.h index 00c1abfd27..9eec55be15 100644 --- a/src/common/hierarchical/scheduler.h +++ b/src/common/hierarchical/scheduler.h @@ -40,6 +40,7 @@ #include #include #include +#include #include "types.h" @@ -103,7 +104,7 @@ class Scheduler { void poison_task(TaskSlot slot, const std::string &root_message); void try_consume(TaskSlot slot); void dispatch_ready(); - void dispatch_next_level_group(); - void dispatch_next_level_singles(); + std::unordered_set dispatch_next_level_group(); + void dispatch_next_level_singles(const std::unordered_set &reserved_worker_ids); void dispatch_sub_ready(); }; diff --git a/tests/st/worker/collectives/_helpers.py b/tests/st/worker/collectives/_helpers.py index 4a768340b8..6550cbb9e4 100644 --- a/tests/st/worker/collectives/_helpers.py +++ b/tests/st/worker/collectives/_helpers.py @@ -87,7 +87,7 @@ def _allreduce_scratch_params(mode: str, nranks: int) -> tuple[int, int, int]: def allreduce_orch_fn(orch, callables, task_args, config): - """L3 orch: allocate domain, submit per-rank allreduce tasks. + """L3 orch: allocate a domain and submit all allreduce ranks as one group. Reads nranks and mode_id from task_args scalars. Selects the ChipCallable by mode name (e.g. ``allreduce_onephase``). @@ -120,6 +120,7 @@ def allreduce_orch_fn(orch, callables, task_args, config): window_size=window_size, buffers=[CommBufferSpec(name="scratch", dtype="float32", count=float_elems, nbytes=scratch_nbytes)], ) as handle: + args_list = [] for i in range(nranks): domain = handle[i] chip_args = TaskArgs() @@ -136,7 +137,8 @@ def allreduce_orch_fn(orch, callables, task_args, config): ) chip_args.add_scalar(domain.domain_size) chip_args.add_scalar(domain.device_ctx) - orch.submit_next_level(chip, chip_args, config, worker=i) + args_list.append(chip_args) + orch.submit_next_level_group(chip, args_list, config, workers=list(range(nranks))) # --------------------------------------------------------------------------- @@ -169,7 +171,7 @@ def generic_collective_orch_fn( """Generic L3 orch for single-mode collectives (allgather, reduce_scatter, broadcast, all_to_all). Reads nranks from ``task_args.nranks`` (Scalar). Allocates a comm domain - and submits the ChipCallable named ``chip_name`` for each rank. + and submits the ChipCallable named ``chip_name`` as one all-rank group. Each rank's input/output tensors are named ``in_`` / ``out_``. Optional ``extra_scalars`` are appended after ``domain_size`` and @@ -187,6 +189,7 @@ def generic_collective_orch_fn( window_size=window_size, buffers=[CommBufferSpec(name="scratch", dtype="float32", count=float_elems, nbytes=scratch_nbytes)], ) as handle: + args_list = [] for i in range(nranks): domain = handle[i] chip_args = TaskArgs() @@ -205,7 +208,8 @@ def generic_collective_orch_fn( for s in extras: chip_args.add_scalar(s) chip_args.add_scalar(domain.device_ctx) - orch.submit_next_level(chip, chip_args, config, worker=i) + args_list.append(chip_args) + orch.submit_next_level_group(chip, args_list, config, workers=list(range(nranks))) # --------------------------------------------------------------------------- diff --git a/tests/ut/cpp/hierarchical/test_scheduler.cpp b/tests/ut/cpp/hierarchical/test_scheduler.cpp index 31ca70b436..d56061ce40 100644 --- a/tests/ut/cpp/hierarchical/test_scheduler.cpp +++ b/tests/ut/cpp/hierarchical/test_scheduler.cpp @@ -627,7 +627,7 @@ TEST_F(SchedulerFixture, FailedProducerPoisonsDependentTask) { } // =========================================================================== -// Group task tests -- fixture with 2 MockMailboxWorkers +// Group task tests -- fixture with 3 MockMailboxWorkers // =========================================================================== struct GroupSchedulerFixture : public ::testing::Test { @@ -639,6 +639,7 @@ struct GroupSchedulerFixture : public ::testing::Test { Orchestrator orch; MockMailboxWorker worker_a; MockMailboxWorker worker_b; + MockMailboxWorker worker_c; WorkerManager manager; Scheduler sched; CallConfig cfg; @@ -654,8 +655,10 @@ struct GroupSchedulerFixture : public ::testing::Test { worker_a.start(); worker_b.start(); + worker_c.start(); manager.add_next_level(worker_a.mailbox_ptr()); manager.add_next_level(worker_b.mailbox_ptr()); + manager.add_next_level(worker_c.mailbox_ptr()); manager.start( &allocator, [this](WorkerCompletion completion) { @@ -774,44 +777,134 @@ TEST_F(GroupSchedulerFixture, GroupCompletesOnlyWhenAllDone) { wait_consumed(slot); } -TEST_F(GroupSchedulerFixture, BlockedGroupDoesNotDispatchPartiallyOrReserveIdleWorker) { - auto running = orch.submit_next_level(C(70), single_tensor_args(0xF0, TensorArgType::OUTPUT), cfg, 0); +TEST_F(GroupSchedulerFixture, BlockedGroupReservesTargetsThatBecomeIdleOneAtATime) { + auto running_a = orch.submit_next_level(C(70), single_tensor_args(0xF0, TensorArgType::OUTPUT), cfg, 0); + auto running_b = orch.submit_next_level(C(71), single_tensor_args(0xF1, TensorArgType::OUTPUT), cfg, 1); worker_a.wait_running(); + worker_b.wait_running(); ASSERT_TRUE(worker_a.is_running.load()); + ASSERT_TRUE(worker_b.is_running.load()); - TaskArgs group_a = single_tensor_args(0xF1, TensorArgType::OUTPUT); - TaskArgs group_b = single_tensor_args(0xF2, TensorArgType::OUTPUT); - auto group = orch.submit_next_level_group(C(71), {group_a, group_b}, cfg, {0, 1}); + TaskArgs group_a = single_tensor_args(0xF2, TensorArgType::OUTPUT); + TaskArgs group_b = single_tensor_args(0xF3, TensorArgType::OUTPUT); + auto group = orch.submit_next_level_group(C(72), {group_a, group_b}, cfg, {1, 0}); + auto single_a = orch.submit_next_level(C(73), single_tensor_args(0xF4, TensorArgType::OUTPUT), cfg, 0); + auto single_b = orch.submit_next_level(C(74), single_tensor_args(0xF5, TensorArgType::OUTPUT), cfg, 1); + auto unrelated = orch.submit_next_level(C(75), single_tensor_args(0xF6, TensorArgType::OUTPUT), cfg, 2); - std::this_thread::sleep_for(std::chrono::milliseconds(20)); - EXPECT_EQ(S(group.task_slot).state.load(), TaskState::READY); - EXPECT_FALSE(worker_b.is_running.load()); - EXPECT_EQ(worker_b.dispatched_count(), 0); + auto deadline = std::chrono::steady_clock::now() + std::chrono::milliseconds(500); + while (worker_c.dispatched_count() == 0 && std::chrono::steady_clock::now() < deadline) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + ASSERT_EQ(worker_c.dispatched_count(), 1); + EXPECT_EQ(worker_c.dispatched[0].callable_hash0, 75u); + worker_c.complete(); - auto independent = orch.submit_next_level(C(72), single_tensor_args(0xF3, TensorArgType::OUTPUT), cfg, 1); - worker_b.wait_running(); - ASSERT_TRUE(worker_b.is_running.load()); - EXPECT_EQ(worker_b.dispatched[0].callable_hash0, 72u); + worker_a.complete(); + deadline = std::chrono::steady_clock::now() + std::chrono::milliseconds(50); + while (worker_a.dispatched_count() == 1 && std::chrono::steady_clock::now() < deadline) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + EXPECT_EQ(worker_a.dispatched_count(), 1) << "blocked group must neither dispatch partially nor yield to singles"; EXPECT_EQ(S(group.task_slot).state.load(), TaskState::READY); worker_b.complete(); - wait_consumed(independent.task_slot); + deadline = std::chrono::steady_clock::now() + std::chrono::milliseconds(500); + while ((worker_a.dispatched_count() < 2 || worker_b.dispatched_count() < 2) && + std::chrono::steady_clock::now() < deadline) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + ASSERT_GE(worker_a.dispatched_count(), 2); + ASSERT_GE(worker_b.dispatched_count(), 2); + EXPECT_EQ(worker_a.dispatched[1].callable_hash0, 72u); + EXPECT_EQ(worker_b.dispatched[1].callable_hash0, 72u); + worker_a.complete(); + worker_b.complete(); - auto deadline = std::chrono::steady_clock::now() + std::chrono::milliseconds(500); - while ((!worker_a.is_running.load() || !worker_b.is_running.load()) && + deadline = std::chrono::steady_clock::now() + std::chrono::milliseconds(500); + while ((worker_a.dispatched_count() < 3 || worker_b.dispatched_count() < 3) && std::chrono::steady_clock::now() < deadline) { std::this_thread::sleep_for(std::chrono::milliseconds(1)); } - ASSERT_TRUE(worker_a.is_running.load()); - ASSERT_TRUE(worker_b.is_running.load()); - EXPECT_EQ(worker_a.dispatched[1].callable_hash0, 71u); - EXPECT_EQ(worker_b.dispatched[1].callable_hash0, 71u); - + ASSERT_GE(worker_a.dispatched_count(), 3); + ASSERT_GE(worker_b.dispatched_count(), 3); + EXPECT_EQ(worker_a.dispatched[2].callable_hash0, 73u); + EXPECT_EQ(worker_b.dispatched[2].callable_hash0, 74u); worker_a.complete(); worker_b.complete(); - wait_consumed(running.task_slot); + + wait_consumed(running_a.task_slot); + wait_consumed(running_b.task_slot); wait_consumed(group.task_slot); + wait_consumed(single_a.task_slot); + wait_consumed(single_b.task_slot); + wait_consumed(unrelated.task_slot); +} + +TEST_F(GroupSchedulerFixture, ConsecutiveGroupsReserveOnlyBlockedHeadTargets) { + SubmitResult first_group; + SubmitResult second_group; + SubmitResult single_a; + SubmitResult single_c; + { + std::lock_guard scheduler_pause(sched.loop_mutex()); + first_group = orch.submit_next_level_group( + C(80), {single_tensor_args(0x100, TensorArgType::OUTPUT), single_tensor_args(0x101, TensorArgType::OUTPUT)}, + cfg, {0, 1} + ); + second_group = orch.submit_next_level_group( + C(81), {single_tensor_args(0x102, TensorArgType::OUTPUT), single_tensor_args(0x103, TensorArgType::OUTPUT)}, + cfg, {1, 2} + ); + single_a = orch.submit_next_level(C(82), single_tensor_args(0x104, TensorArgType::OUTPUT), cfg, 0); + single_c = orch.submit_next_level(C(83), single_tensor_args(0x105, TensorArgType::OUTPUT), cfg, 2); + } + + worker_a.wait_running(); + worker_b.wait_running(); + ASSERT_EQ(worker_a.dispatched_count(), 1); + ASSERT_EQ(worker_b.dispatched_count(), 1); + EXPECT_EQ(worker_a.dispatched[0].callable_hash0, 80u); + EXPECT_EQ(worker_b.dispatched[0].callable_hash0, 80u); + EXPECT_EQ(worker_c.dispatched_count(), 0); + + worker_a.complete(); + auto deadline = std::chrono::steady_clock::now() + std::chrono::milliseconds(500); + while (worker_a.dispatched_count() < 2 && std::chrono::steady_clock::now() < deadline) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + ASSERT_EQ(worker_a.dispatched_count(), 2); + EXPECT_EQ(worker_a.dispatched[1].callable_hash0, 82u); + EXPECT_EQ(worker_c.dispatched_count(), 0); + EXPECT_EQ(S(second_group.task_slot).state.load(), TaskState::READY); + + worker_a.complete(); + worker_b.complete(); + deadline = std::chrono::steady_clock::now() + std::chrono::milliseconds(500); + while ((worker_b.dispatched_count() < 2 || worker_c.dispatched_count() < 1) && + std::chrono::steady_clock::now() < deadline) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + ASSERT_EQ(worker_b.dispatched_count(), 2); + ASSERT_EQ(worker_c.dispatched_count(), 1); + EXPECT_EQ(worker_b.dispatched[1].callable_hash0, 81u); + EXPECT_EQ(worker_c.dispatched[0].callable_hash0, 81u); + + worker_b.complete(); + worker_c.complete(); + deadline = std::chrono::steady_clock::now() + std::chrono::milliseconds(500); + while (worker_c.dispatched_count() < 2 && std::chrono::steady_clock::now() < deadline) { + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + ASSERT_EQ(worker_c.dispatched_count(), 2); + EXPECT_EQ(worker_c.dispatched[1].callable_hash0, 83u); + worker_c.complete(); + + wait_consumed(first_group.task_slot); + wait_consumed(second_group.task_slot); + wait_consumed(single_a.task_slot); + wait_consumed(single_c.task_slot); } TEST_F(GroupSchedulerFixture, LaunchableGroupPrecedesConflictingSingles) {