File size: 885 Bytes
aad9f16 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 | from __future__ import annotations
from dataclasses import dataclass
@dataclass(frozen=True)
class DistributedBatchPlan:
local_item_counts: tuple[int, ...]
sync_batches: int
def distributed_batch_plan(
*,
total_items: int,
num_processes: int,
batch_size: int,
) -> DistributedBatchPlan:
if num_processes < 1:
raise ValueError("num_processes must be >= 1")
if batch_size < 1:
raise ValueError("batch_size must be >= 1")
items_per_proc, extra_items = divmod(total_items, num_processes)
local_item_counts = tuple(
items_per_proc + (1 if process_index < extra_items else 0)
for process_index in range(num_processes)
)
max_local_items = max(local_item_counts, default=0)
sync_batches = (max_local_items + batch_size - 1) // batch_size
return DistributedBatchPlan(local_item_counts, sync_batches)
|