PyTorch Elastic 多进程管理:使用 `torch.distributed.elastic.multiprocessing` 启动与编排多 Worker
2026/9/9 22:23:05 网站建设 项目流程

PyTorch Elastic 多进程管理:使用torch.distributed.elastic.multiprocessing启动与编排多 Worker

【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch

torch.distributed.elastic.multiprocessing(TorchElastic 的 multiprocessing 子模块)是一个用于启动并管理n份 worker 子进程的库:既可以按「函数」(Callable)方式用torch.multiprocessingspawn/fork 出进程,也可以按「二进制可执行程序」(如/usr/bin/echopython train.py)方式用subprocess.Popen拉起子进程。它是torchrun/TorchElastic Agent 在单机上真正拉起训练 worker 的底层实现,也是自定义「一个入口点、多个副本」并行任务时的标准工具。读完本文,你将掌握start_processes的完整参数语义、PContext 进程上下文家族的使用方法、标准输出/错误(stdout/stderr)的 redirect/tee 与日志目录结构、进程返回值与失败错误文件(error.json)的传播机制,并理解函数型与二进制型两种启动方式的差异与适用场景。

本文以 docs/source/elastic/multiprocessing.md 为骨架展开,其全部 API 引用均落在 torch/distributed/elastic/multiprocessing 模块内(入口模块init.py,核心实现 api.py)。

一、模块定位:两种入口点、一套统一上下文

模块 docstring(见init.py)明确了它的职责:

  • 函数入口点(Callable):底层使用torch.multiprocessing(即 Pythonmultiprocessing)通过spawn/fork/forkserver方式启动进程,返回MultiprocessContext
  • 二进制入口点(str):底层使用subprocess.Popen创建子进程,返回SubprocessContext
  • 两者都是父类PContext的具体实现。

这与torch.multiprocessing.start_processes的约定保持一致:start_processes的返回值是一个进程上下文(PContext)。无论用哪种方式启动,调用方都通过同一套 PContext 接口(start()wait()close()pids())管理全部 worker,从而在 TorchElastic 的 agent 代码中做到「调度逻辑与启动机制解耦」。

从源码结构看,整个模块还包含四个支撑文件,分别承担不同职责:

文件职责
api.pyPContext/MultiprocessContext/SubprocessContextRunProcsResultLogsSpecs/DefaultLogsSpecs/LogsDestStd等核心类
errors/基于文件的跨进程错误传播:record装饰器、ProcessFailureerror.json写入
redirects.py在文件描述符层面对 stdout/stderr 的重定向(dup2实现)
tail_log.pyTailLog:模拟 unixtail -f将日志文件内容实时输出到控制台(实现 tee)
subprocess_handler/二进制 worker 的子进程处理器(进程信号发送、关闭等)

二、start_processes:启动多个 Worker 的唯一入口

start_processes是模块对外暴露的唯一「启动入口」,定义在init.py:

def start_processes( name: str, entrypoint: Callable | str, args: dict[int, tuple], envs: dict[int, dict[str, str]], logs_specs: LogsSpecs, log_line_prefixes: dict[int, str] | None = None, start_method: str = "spawn", numa_options: NumaOptions | None = None, duplicate_stdout_filters: list[str] | None = None, duplicate_stderr_filters: list[str] | None = None, ) -> PContext:

2.1 参数语义

核心参数的完整语义来自函数 docstring(init.py):

  • name:人类可读的短名称,描述这些进程的用途(作为 tee 输出到控制台时的行头[{name}{local_rank}]:)。
  • entrypoint:要么是Callable(函数),要么是str(命令/二进制路径)。进程副本数量由args的条目数决定
  • args/envs:以 local rank(副本序号)为 key、分别映射每个副本的入参与环境变量。两者的 key 集合必须一致且完整覆盖{0, 1, ..., nprocs-1},即「所有 local rank 都必须被覆盖」。
  • logs_specsLogsSpecs实例,统一定义log_dirredirectstee(详见第五节)。
  • log_line_prefixes:为每个 local rank 指定日志行的自定义前缀(agent 常用它注入role/rank/local_rank/hostname模板,见 local_elastic_agent.py)。
  • start_method:multiprocessing 启动方式(spawn/fork/forkserver),对二进制入口点无效(忽略)。
  • numa_optionsNumaOptions,用于给 worker 绑定 NUMA 节点(传入后函数入口点会被_maybe_wrap_with_numa_binding包装,见 api.py)。
  • duplicate_stdout_filters/duplicate_stderr_filters:非空时,把被tee选中的 rank 的 stdout/stderr 中命中任意一个过滤子串的行,额外聚合写到filtered_stdout.log/filtered_stderr.log

2.2 用法一:以函数方式启动两个 trainer

docstring 给出了可以直接复制的经典用法(init.py):

from torch.distributed.elastic.multiprocessing import Std, start_processes def trainer(a, b, c): pass # train # 运行两个 trainer: # LOCAL_RANK=0 -> trainer(1,2,3) # LOCAL_RANK=1 -> trainer(4,5,6) ctx = start_processes( name="trainer", entrypoint=trainer, args={0: (1, 2, 3), 1: (4, 5, 6)}, envs={0: {"LOCAL_RANK": 0}, 1: {"LOCAL_RANK": 1}}, logs_specs=DefaultLogsSpecs(log_dir="/tmp/foobar", redirects=Std.ALL), tee={0: Std.ERR}, # 仅把 local rank 0 的 stderr 同步打印到控制台 ) # 等待所有 trainer 结束 ctx.wait()

注意示例中的redirects/tee需通过logs_specsDefaultLogsSpecs)传入;tee语义与 unixtee命令一致——既写入日志文件、又打印到控制台,而redirects则只写文件、不打印。

2.3 用法二:以二进制方式启动 echo

entrypoint为字符串时,等价于直接执行该命令(init.py):

ctx = start_processes( name="echo", entrypoint="echo", # 二进制 / 命令行 args={0: "hello", 1: "world"}, redirects={1: Std.OUT}, logs_specs=DefaultLogsSpecs(log_dir="/tmp/foobar"), ) # 实际执行:echo hello;echo world > /tmp/foobar/1/stdout.log

2.4 三个容易踩的约束

从 docstring 与实现(init.py)可以提炼出以下硬性规则:

  1. 二进制入口点时,args只能传字符串;若传其他类型(如1,2,3[1,2,3]),会被强制str()化后再执行。例如args={0: (1, 2, 3), 1: ([1, 2, 3],)}实际运行的是echo "1" "2" "3"echo "[1, 2, 3]"
  2. argsenvs的 key 集合必须都恰好等于{0,...,nprocs-1}。启动前会调用_validate_full_rank(api.py)校验,key 不齐会抛出RuntimeError: local rank mapping mismatch。例如envs={0:{}}args有两个条目就是非法的。
  3. 函数入口点失败时错误文件默认会被记录,无需额外标注;二进制入口点失败要写出error.json,则入口点主函数必须显式标注@torch.distributed.elastic.multiprocessing.errors.record

三、Std位掩码:redirects 与 tee 的表达方式

Std定义在 api.py,本质是一个IntFlag枚举:

class Std(IntFlag): NONE = 0 OUT = 1 ERR = 2 ALL = OUT | ERR # = 3

它既可以全局作用于所有 rank(传单个Std值),也可以按 rank 选择性作用(传dict[int, Std],未出现的 rank 默认Std.NONE)。redirectsteeLogsSpecs/start_processes的语义如下(init.py):

  • redirects:把指定的 std 流重定向写入log_dir下的日志文件;
  • tee:把指定的 std 流同时写日志文件并打印到控制台(redirect + print);若不想让 worker 输出刷屏,应使用redirects而非tee

3.1Std.from_strto_map

Std.from_str(api.py)支持两种字符串输入(Agent 常把配置从命令行以字符串形式传下来):

Std.from_str("0") # -> Std.NONE Std.from_str("1") # -> Std.OUT Std.from_str("0:3,1:0,2:1") # -> {0: Std.ALL, 1: Std.NONE, 2: Std.OUT}

to_map(api.py)把「单个值或局部映射」统一扩展成覆盖所有 local rank 的映射:

to_map(Std.OUT, local_world_size=2) # {0: Std.OUT, 1: Std.OUT} to_map({1: Std.OUT}, local_world_size=2) # {0: Std.NONE, 1: Std.OUT}

四、Process Context:PContext家族

文档列出的四类进程上下文类集中在 api.py 中。PContext是抽象基类(api.py),名字有意与torch.multiprocessing.ProcessContext区分。

4.1PContext统一生命周期接口

  • start()(api.py):在主线程调用时,会按环境变量TORCHELASTIC_SIGNALS_TO_HANDLE(默认SIGTERM,SIGINT,SIGHUP,SIGQUIT)注册信号处理函数_terminate_process_handler——收到终止信号会抛出SignalException(api.py),该异常不应被吞掉,否则进程永不退出。随后调用_start()真正拉起 worker,并启动TailLog线程。
  • wait(timeout=-1, period=1)(api.py):每隔period秒轮询一次,等待所有进程结束。timeout=0等价于一次 poll(非阻塞查询);timeout<0表示无限等待;超时返回None
  • close(death_sig=None, timeout=30)(api.py):用death_sig(unix 默认SIGTERM、Windows 默认CTRL_C_EVENT)终止所有进程并清理资源;超时后升级为强杀信号(unixSIGKILL、WindowsCTRL_C_EVENT,见_get_kill_signal)。实现中还带SIGKILL后的有界 join,避免 worker 卡在不可中断内核态(如 NCCL/GPU 集合通信挂死)时把 agent 自身拖死。
  • pids():返回{local_rank: pid}映射。

close/wait与信号配合的推荐写法(docstring 内嵌示例,api.py):

pc = start_processes(...) try: pc.wait(1) # ... 做其他工作 except SignalException as e: pc.shutdown(e.sigval, timeout=30) # 优雅退出,超时再强杀

4.2MultiprocessContext:函数型 worker

MultiprocessContext(api.py)用于entrypoint为函数的情形:

  • 构造时按start_method为每个 local rank 建立mp.SimpleQueue_ret_vals)用于回传返回值;
  • _start()内部调用mp.start_processes(fn=_wrap, ..., join=False, daemon=False, start_method=...)(api.py)。_wrap先注入该 rank 的环境变量,再进入 stdout/stderr 重定向上下文,并用record(fn)(*args_)包装执行,最后把返回值put进队列(api.py);
  • _poll()依赖「所有进程全部结束」与「任意进程失败」两种终态,通过ProcessContext.join+ 轮询返回队列避免大返回值导致管道死锁(api.py)。

4.3SubprocessContext:二进制型 worker

SubprocessContext(api.py)用于entrypoint为命令字符串的情形:

  • _start()为每个 local rank 通过get_subprocess_handler创建一个SubprocessHandlerPopen封装);
  • _poll()逐个proc.poll()采集退出码:任意副本失败或全部结束后即收敛结果(all-or-nothing 策略),并对仍在运行的副本执行close()
  • 二进制 worker没有返回值,因此成功时RunProcsResult.return_values会被填充为None占位,以与MultiprocessContext保持一致的 API 形态(api.py)。

4.4RunProcsResult:运行结果容器

RunProcsResult是一个 dataclass(api.py),由PContext的轮询/等待返回,注意以下几点约束:

字段说明
return_values: dict[int, Any]各 rank 的返回值,仅函数型启动时填充;按 local rank 索引
failures: dict[int, ProcessFailure]各 rank 的失败信息(ProcessFailure包含 local_rank、pid、exitcode、error_file、message、timestamp 等)
stdouts/stderrs各 rank 的 stdout.log / stderr.log 路径(未重定向则为空串)
is_failed()len(failures) > 0即为失败

五、LogsSpecsDefaultLogsSpecsLogsDest:日志目录规划

文档列出的日志三件套同样定义在 api.py。

5.1LogsSpecs(抽象基类)

LogsSpecs(api.py)定义日志处理与重定向的抽象协议:

  • 构造参数:log_dir(日志根目录)、redirectstee(均可传单个Std{local_rank: Std}映射);
  • 抽象方法reify(envs) -> LogsDest:根据各 rank 的环境变量,为每个 rank 计算出日志文件目标路径;
  • 抽象属性root_log_dir

5.2DefaultLogsSpecs(默认实现)

DefaultLogsSpecs(api.py)是现成可用、也是 agent 默认采用的实现,行为如下:

  • log_dir不存在则自动创建;为文件时抛NotADirectoryError;未指定时自动tempfile.mkdtemp(prefix="torchelastic_")
  • 日志目录按「运行 + 尝试 + rank」分层组织,reify()依据环境变量TORCHELASTIC_RUN_ID(默认"test_run_id")与TORCHELASTIC_RESTART_COUNT(默认"0")生成目录,并在每次 restart 前用shutil.rmtree清理旧 attempt 目录:
<log_dir>/<rdzv_run_id>/attempt_<attempt>/ ├── <rank>/stdout.log # redirects & OUT 时 ├── <rank>/stderr.log # redirects & ERR 时 ├── <rank>/error.json # 失败时记录错误 ├── filtered_stdout.log # duplicate_stdout_filters 非空时 └── filtered_stderr.log # duplicate_stderr_filters 非空时

实现要点(api.py):

  • tee 的实现方式:先把 tee 目标并入 redirects(先落盘),再靠TailLog把文件内容同步输出到控制台——因此LogsSpecs内部要求 stdouts/stderrs 一定是 tee_stdouts/tee_stderrs 的超集;
  • 每个 rank 的error.json路径会以环境变量TORCHELASTIC_ERROR_FILE注入 worker;
  • 若配置了local_ranks_filter,不在集合内的 rank 即使配置了 tee 也只落盘不 tail、未重定向的流则导向os.devnull,从而实现「只打印选定 rank 的日志到控制台」。

5.3LogsDest

LogsDest(api.py)是reify()的返回值 dataclass,按日志类型保存{local_rank: 文件路径}的映射:

  • stdouts/stderrs:重定向的日志文件路径(未重定向为空串);
  • tee_stdouts/tee_stderrs:被 tail 输出的文件;
  • error_files{local_rank: error.json}
  • filtered_stdout/filtered_stderr:按过滤子串聚合的日志文件路径。

六、Tee 的底层实现:TailLog与文件描述符级重定向

6.1TailLog:无文件等待式 tail

TailLog(tail_log.py)为每个日志文件起一个线程模拟tail -f

  • 日志文件尚未创建时也会优雅等待(轮询interval_sec,默认 0.1 秒),因此可以在 worker 真正写出日志前就启动;
  • 默认输出行头为[{name}{local_rank}]:(可用log_line_prefixes覆盖,tail_log.py);
  • stop()通过 Event 通知各 tail 线程退出并回收线程池;
  • 由于跨文件缓冲,日志行不保证严格按墙钟顺序打印——官方建议业务日志自带时间戳。

6.2redirects.py:重定向的是文件描述符,不只是sys.stdout

redirect_stdout/redirect_stderr(redirects.py)基于 POSIXdup2实现,因此同时覆盖 Python 层与 C 层输出

with redirect_stdout("/tmp/stdout.log"): print("python stdouts are redirected") libc = ctypes.CDLL("libc.so.6") libc.printf(b"c stdouts are also redirected") # 也会被重定向 os.system("echo system stdouts are also redirected") # 也会被重定向 print("stdout restored")

平台限制(源码中可查):macOS 目前不支持该重定向机制(加载 libc 时直接告警并返回None,重定向退化为nullcontext);Windows 则需要通过_dup2/SetStdHandle/CRTFILE*多层重定向(redirects.py)。因此get_std_cm(api.py)在IS_WINDOWS or IS_MACOS时返回空上下文,仅 unix 平台真正执行 fd 级重定向。

七、跨进程错误传播:error.jsonrecordProcessFailure

多进程场景下,异常发生在 worker 进程,agent 无法用 try-catch 直接捕获。TorchElastic 采用基于文件的跨进程错误传播(详见 errors/init.py):

  1. 任何被record装饰器包裹的入口点,在发生未捕获异常时,会连同完整 traceback 写入环境变量TORCHELASTIC_ERROR_FILE指向的 JSON 文件;
  2. 父进程(agent)为每个子进程设置该环境变量,并在子进程失败后聚合各error.json
  3. 父进程选择timestamp 最小(最先发生)的错误作为 root cause 向上传播。

ErrorHandler.record_exception写入的 JSON 结构见 error_handler.py,核心字段为嵌套的message(含py_callstacktimestamp)。ProcessFailure会尝试解析该文件来构造message/timestamp;若error.json不存在(如被信号杀死),则会据退出码生成信息(如Signal -15 (SIGTERM) received by PID ...,见 errors/init.py)。

ChildFailedError(errors/init.py)允许@record包裹的父函数把子进程根异常原样上抛(不包裹父 traceback),适合「父进程只是简单保姆进程、真正的计算都在子进程」的场景。

八、与 TorchElastic Agent 的集成:一切从_start_workers开始

该模块并非孤立 API,它是 TorchElasticLocalElasticAgent拉起 worker 的直接底层。在 local_elastic_agent.py 的_start_workers中,agent 逐 rank 组装好envs(注入LOCAL_RANKTORCHELASTIC_ERROR_FILE等)与args(含宏替换macros.substitute)后,一次性调用:

self._pcontext = start_processes( name=spec.role, entrypoint=spec.entrypoint, args=args, envs=envs, logs_specs=self._logs_specs, log_line_prefixes=log_line_prefixes, start_method=self._start_method, numa_options=spec.numa_options, duplicate_stdout_filters=spec.duplicate_stdout_filters, duplicate_stderr_filters=spec.duplicate_stderr_filters, )

由此可见:torchrun/agent 对日志目录、redirect/tee、错误文件、NUMA 绑定、日志前缀等一系列编排能力,最终都会收敛到start_processesLogsSpecs这套 API 上。需要自定义训练启动器或调试 worker 拉起/退出/日志行为时,这套接口就是直接的切入点。

九、实战小结与最佳实践

  1. 进程数量args决定,且argsenvs的 key 必须恰好是{0,...,nprocs-1}——参数不齐会在进程启动前快速失败;
  2. 想让 worker 输出进文件redirects想既进文件又打印到控制台tee;只想打印个别 rank,可用{rank: Std}映射或local_ranks_filter
  3. 函数型入口点自动获得返回值收集(经mp.SimpleQueue回传)与record错误记录;二进制型入口点没有返回值(RunProcsResult.return_valuesNone占位),错误文件需入口程序自身标注@record
  4. 等待结果优先使用ctx.wait(timeout)(返回值RunProcsResult),注意轮询/等待是 all-or-nothing 语义:全部成功或任一失败即收敛;
  5. 优雅退出务必处理主进程收到的SIGTERM/SIGINT等信号(SignalException),随后调用close(death_sig, timeout),超时后框架会自动升级为SIGKILL
  6. 日志目录建议直接使用DefaultLogsSpecs,它会自动创建目录并按<rdzv_run_id>/attempt_<n>/<rank>/组织,便于失败定位与多轮重试排障。

围绕本文涉及的 API 与语义,还可进一步查阅 docs/source/elastic/multiprocessing.md 及模块内 api.py、tail_log.py、redirects.py 与 errors/ 的源码注释与 docstring,以获得逐类、逐参数的权威细节。

【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询