Brainrot_Muxa/tools/_parallel.py

91 lines
4.5 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Раскладка задач по бэгам на процессы.
Конвейер держит состояние между кадрами — пройденный путь, треки, накопитель
складчатого тела, — поэтому разрезать один бэг нельзя. Зато бэги независимы
друг от друга, и это ровно та зернистость, которая нужна: пять записей на
шесть физических ядер.
Результат совпадает с последовательным прогоном **точно**, а не «примерно»:
генератор случайных чисел создаётся внутри задачи от того же зерна, общей
изменяемой памяти между бэгами нет. Порядок родитель восстанавливает сам —
печатает по готовности, а сохраняет в исходном.
Ускорение упирается не в ядра, а в число бэгов: их пять, и больше пяти
процессов тут просто нечем занять.
"""
from __future__ import annotations
import os
import time
from concurrent.futures import ProcessPoolExecutor, as_completed
def limit_threads() -> None:
"""Один поток BLAS на процесс.
Вызывать ДО создания пула: переменные наследуются дочерним процессом при
запуске, а число потоков BLAS выбирает один раз при импорте numpy и потом
не меняет. Без этого каждый из пяти воркеров разворачивается на все ядра
и они дерутся за те же шесть.
"""
for v in ("OMP_NUM_THREADS", "OPENBLAS_NUM_THREADS", "MKL_NUM_THREADS",
"NUMEXPR_NUM_THREADS", "VECLIB_MAXIMUM_THREADS"):
os.environ.setdefault(v, "1")
def resolve(jobs: int, n_tasks: int) -> int:
"""Сколько процессов поднимать: 0 — по числу задач, но не больше ядер."""
if n_tasks <= 1:
return 1
if jobs > 0:
return max(1, min(jobs, n_tasks))
cpu = os.cpu_count() or 2
return max(1, min(n_tasks, cpu // 2)) # логических ядер вдвое больше физических
class _Timed:
"""Часы держит сам воркер.
В родителе видно только момент готовности, а он у всех задач, стартовавших
разом, почти один и тот же — по такому замеру не понять, какой бэг тяжёлый.
"""
def __init__(self, fn):
self.fn = fn
def __call__(self, task):
t0 = time.time()
return self.fn(task), time.time() - t0
def run(work, tasks, jobs: int):
"""Выполнить `work(task)` по всем задачам, отдавая `(i, task, res, с)`.
Отдаёт по мере готовности, поэтому `i` — исходный номер задачи, и по нему
вызывающий раскладывает результаты обратно в порядок бэгов.
При одном процессе всё считается прямо здесь, без пула: остаётся чем
отлаживать, и трассировка ошибки не проходит через межпроцессную передачу.
"""
tasks = list(tasks)
n = resolve(jobs, len(tasks))
if n <= 1:
for i, t in enumerate(tasks):
t0 = time.time()
yield i, t, work(t), time.time() - t0
return
limit_threads()
with ProcessPoolExecutor(max_workers=n) as ex:
fut = {ex.submit(_Timed(work), t): (i, t) for i, t in enumerate(tasks)}
for f in as_completed(fut):
i, t = fut[f]
res, secs = f.result()
yield i, t, res, secs
def add_argument(ap) -> None:
"""Один и тот же флаг во всех инструментах, чтобы не помнить разные."""
ap.add_argument("--jobs", type=int, default=0,
help="сколько бэгов считать разом; 0 — по числу бэгов, "
"но не больше физических ядер; 1 — в один процесс")