ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

JAX Colocated Python 实战:在多主机与单控制器环境中统一执行设备端 Python 代码

JAX Colocated Python 实战:在多主机与单控制器环境中统一执行设备端 Python 代码 JAX Colocated Python 实战在多主机与单控制器环境中统一执行设备端 Python 代码【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax导读Colocated Python 是 JAX 提供的一套实验性 API它允许开发者把任意可序列化的Python 函数或类实例运送到与目标 JAX 设备位于同一主机的位置上执行——当设备是本地设备时在本地执行当设备是远程设备时则把代码序列化后发送到该设备所属的主机执行。本文基于 docs/notebooks/colocated-python.md 讲解其完整用法如何获取与加速器同主机的 CPU 设备、如何包装函数与类、如何通过specialize提供输入输出规格与设备信息以及它独特的程序顺序 跨线程并发执行模型同时结合 jax/experimental/colocated_python/ 源码与 tests/colocated_python_test.py 测试用例剖析其序列化、IFRT 编译与缓存等底层机制。读完本文你将能够在多主机多控制器 JAX 环境中编写可移植的、执行设备端 Python含文件 I/O的混合系统。注意Colocated Python 目前是实验性 API源码位于jax.experimental.colocated_python其功能与接口可能在不遵循标准 JAX 兼容性策略的情况下发生变化。一、为什么需要 Colocated Python在构建基于 JAX 的多主机机器学习系统时通常会遇到两种部署形态多控制器multi-controller环境每个拥有加速器的主机都运行一份 JAX 代码各自负责自己的设备单控制器single-controller环境只有一个控制主机运行 JAX 代码通过远程调用编排其他带有加速器的主机。这两种形态下在一组设备所在的主机上跑一段 Python 代码这个需求实现方式完全不同。Colocated Python 提供了一个统一抽象colocated_python包装的函数会在输入参数所在设备关联的主机上执行。如果这些设备是本地设备Python 代码就在本地主机运行如果这些设备是远程设备Python 代码会被运送序列化、编译、传输到这些远程设备所属的主机上去运行。这使得同一套代码可以无缝地同时工作于多控制器与单控制器环境无需为两种部署形态分别编写实现。二、第一步获取与加速器同主机的 CPU 设备Colocated CPU devices要用 Colocated Python第一步是拿到与目标加速器设备同居一台主机的 CPU 设备。JAX 提供标准入口import jax import jax.experimental.colocated_python as colocated_python devices jax.devices() cpu_devices colocated_python.colocated_cpu_devices(devices) print(cpu_devices)API 行为细节从源码 jax/experimental/colocated_python/api.py 可以看到colocated_cpu_devices的关键行为输入一组设备Sequence[jax.Device]或一个jax.sharding.Mesh输入为 mesh 时返回一个形状相同、轴名与轴类型相同的 CPU meshapi.py 中通过colocated_cpu_devices(tuple(mesh.devices.flat))得到扁平 CPU 设备后再reshape(mesh.axis_sizes)还原。配对依据一个加速器设备通常伴随一个与其位于同一主机的 CPU 设备当单台主机有多个加速器时可能存在与加速器 1:1 对应的多个 CPU 设备。源码通过设备的colocation_id属性进行匹配device.colocation_id保证输出第 i 个 CPU 设备与输入第 i 个加速器设备相关联的顺序保持语义。已是 CPU 设备若输入本身就是 CPU 设备会原样返回。缓存结果通过util.cache缓存容量 1024重复调用不会重复计算。回退路径当主路径因ValueError/AttributeError失败例如 PjRt-IFRT 后端把 CPU 设备定义在独立 CPU 后端上时会回退到按process_index从jax.devices(backendcpu)中逐个弹出一个 CPU 设备来配对的逻辑api.py源码注释也说明该回退路径将来会被移除。异常找不到任何 CPU 设备时抛出ValueError(No CPU devices found)某个设备没有伴生 CPU 设备时抛出ValueError(fDevice {device} has no colocated devices)一个设备对应多个 CPU 设备歧义时同样报错。像普通 CPU 设备一样使用拿到的 CPU 设备可以完全按常规 JAX API 使用——构造 mesh、创建 sharding、放置数据、执行 jitcpu_mesh jax.sharding.Mesh(cpu_devices, [x]) cpu_sharding jax.sharding.NamedSharding(cpu_mesh, jax.P()) x jax.device_put(1, cpu_sharding) y jax.jit(lambda x: x 1)(x) print(y)注意jax.P()表示PartitionSpec()即整数组件不切分全部设备复制该标量。测试 tests/colocated_python_test.py 验证了传入 mesh 得到的 CPU mesh与传入设备再手动构造 mesh两种方式等价。三、包装函数colocated_python 基础用法CPU 设备还可以用来承载真正的 Python 代码执行。用colocated_python.colocated_python包装一个普通函数即可def f(x): return x 1 f colocated_python.colocated_python(f) y f(x) assert y.sharding x.sharding print(y)注意输出数组y的 sharding 与输入x的 sharding 保持一致。从源码看colocated_python通过 func.py 中的make_callable创建无任何初始特化的 Colocated Python callable并使用api_util.fun_sourceinfo/fun_signature保留原函数的源码信息与签名。在设备端执行 I/O因为 Colocated Python 执行的是普通 Python 代码所以可以在包装函数中做文件读写等 I/O 操作def f(x): with open(/tmp/foo, w) as f: f.write(str(x)) return x f colocated_python.colocated_python(f) jax.block_until_ready(f(x))这里必须使用jax.block_until_ready确保 Python 代码真正执行完毕。原理上Colocated Python 调用与 jit 调用类似可能是异步的调用立即返回 JAX 数组并不会阻塞到输出产生。因此当执行完成这件事本身有意义时比如要确认文件已经写完就必须对某个输出进行阻塞等待。何时会同步执行文档明确列出了 Colocated Python 调用同步执行的两种情况未做特化specialization时的首次调用因为异步执行必须预先知道输出的 shape 与 sharding而首次调用时 Colocated Python 必须先实际运行一次 Python 代码来发现这些信息所以第一次调用会同步执行。源码中这一流程对应 func.py当out_specs_treedef is None且没有out_specs_fn时会先编译并调用一个_make_output_specs_and_push_result_fun名字形如f_output_specs_and_push_result把真实结果暂存到func_backend.SINGLETON_RESULT_STORE随后再通过_make_pop_result_fun名字形如f_pop_result取回结果——这两步都是阻塞式的等价于同步执行。部分 JAX 后端尚未完全支持异步执行这些后端会回退到同步执行。另外包装后的 Python 代码在输入与输出中必须使用完全相同的设备集合——这与表示 SPMD 执行的 jit 函数的约束类似源码_infer_devices_from_args也要求所有参数使用同一个 device list否则报ValueError。四、Specialization向运行时注入额外信息特化specialization是 Colocated Python 的核心机制当某些信息无法提前推断、或者你希望让执行严格按你指定的方式发生时通过specialize方法提供输入、输出与执行相关的额外信息。每个被包装的函数都带有specialize方法它会返回一个携带了新特化信息的新包装函数。def f(x): return x 1 f colocated_python.colocated_python(f) f f.specialize(out_specs_fnlambda x: x) y f(x) assert y.sharding x.shardingspecialize支持三类参数从源码 func.py 的Specialization.updatefunc.py可以确认每个字段只能被设置一次重复设置会抛出ValueError。out_specs_fn提前声明输出规格out_specs_fn是一个接收调用输入构成的jax.ShapeDtypeStructpytree、返回预期输出的jax.ShapeDtypeStructpytree的函数。调用它与 jit 的 tracing 类似但这个函数与原始 Python 代码是分离的——它在调用方caller一侧运行不会被运送到设备上执行也不会在设备端运行。提供out_specs_fn后运行时不再需要先跑一遍真实函数来发现输出规格因此所有调用包括第一次都可以异步执行。实现上out_specs_fn的输出会被tree_flatten后写入Specialization的out_specs_treedef/out_specs_leaves后续每次调用直接复用_make_async_execution_fun编译好的异步可执行体。in_specs固定输入规格in_specs接收一个具体的 pytree顶层是二元组形如(args_specs, kwargs_specs)其中每个叶子是带 sharding 的jax.ShapeDtypeStruct。当某个输入规格必须被固定、或者输出规格函数只能针对某个具体输入规格计算时使用import jax.numpy as jnp def f(x): return x 1 f colocated_python.colocated_python(f) f f.specialize( in_specs( # args ( jax.ShapeDtypeStruct( shape(), dtypejnp.int32, shardingcpu_sharding ), ), # kwargs {}, ), out_specs_fnlambda x: jax.ShapeDtypeStruct( shape(), dtypejnp.int32, shardingcpu_sharding ), ) f(x) # x 必须匹配输入规格。一旦指定了in_specs后续调用的实参必须与规格完全吻合。源码中当输入规格被完整指定时fully_specified_in_spec会跳过按输入多态建立_SpecializedCollection集合的路径直接编译唯一一个特化函数并缓存func.py。devices指定执行设备devices指定该 Colocated Python 函数应在哪些设备上运行。指定devices后没有输入参数的函数也可以执行def f(): with open(/tmp/foo, w) as f: f.write(foo) return f colocated_python.colocated_python(f) f f.specialize(devicescpu_devices) f() # 若 f 未用 devices 特化这里会报错。从源码 func.py 可见无输入调用但未特化devices时会直接抛出ValueError(No devices found. colocated_python function without input arguments must be first specialized with devices.)。测试 tests/colocated_python_test.py 分别覆盖了未特化报错与特化后成功返回jnp.array(0)两个方向。另外devices还用于处理输入设备顺序混合的场景测试test_inputs_with_different_device_orderstests/colocated_python_test.py表明当两次调用的输入以不同设备顺序摆放时应显式specialize(devicescpu_devices)以避免由参数推导设备带来的不确定性。五、包装类colocated_python_classColocated Python 也支持包装 Python 类。包装后真实实例会在与设备关联的主机上创建调用方拿到的是一个包装类wrapper class所有方法调用都会通过 Colocated Python 转发到真实实例class Adder: def __init__(self, increment): print(Adder created) self.increment increment def __del__(self): print(Adder destroyed) def add(self, x): return x self.increment Adder colocated_python.colocated_python_class(Adder) adder Adder(1) x jax.device_put(1, cpu_sharding) y adder.add(x) print(y)当包装实例被销毁时真实实例也会随之销毁且销毁是异步的del adder与普通 Python 的重要语义差异文档明确指出 Colocated Python 类与普通 Python 类的三点关键差异惰性实例化类的真实实例只会在某个非构造方法第一次被调用时在设备关联的主机上创建。上例中Adder(1)只是捕获了构造参数1真正的构造调用要等到第一次adder.add(x)才发生——因为在那之前无法得知Adder实例应该创建在哪些主机上。源码 obj.py 的MethodCallerAtBackend清晰体现了这一点_first_call()中才通过obj_backend.SINGLETON_OBJECT_STORE.get_or_create(uid, initializer)调用cls(*init_args, **init_kwargs)若包装实例从未调用任何方法就被销毁则真实实例根本不会被创建。跨主机的多次创建如果同一包装类的不同方法调用使用了不同设备真实实例可能在不同时间、不同主机上被创建。例如第一次方法调用使用主机 A 的 CPU 设备实例在主机 A 上创建第二次方法调用使用主机 B 的 CPU 设备实例随后在主机 B 上也被创建。方法暂不支持特化目前类方法不支持specialize该能力将在未来加入。从 obj.py 可以看到方法包装器虽然保留了specialize入口转发给底层 callable但注释明确标注了TODO(hyeontaek): Support method specialization similar to function specialization.。实现层面wrap_classobj.py为原类的每个普通方法跳过__init__、__del__、__reduce__、__reduce_ex__生成一个方法包装器并通过_InstanceRegistry单例注册为 JAX 二级缓存为每个实例分配唯一的 63 位随机 uid 并跟踪其在哪些设备上存活。控制器侧还通过_update_instance_devices记录每次方法调用涉及的设备集合供实例生命周期管理使用。六、执行顺序与并发程序顺序program orderColocated Python 提供程序顺序执行保证即使调用可能是异步的返回 JAX 数组而不阻塞调用仍会按照用户程序中发起调用的顺序依次执行。因此默认情况下Colocated Python 调用是串行执行的。测试 tests/colocated_python_test.py 的test_sequential_execution验证了这一保证三个函数依次修改同一个全局状态_testing_global_state100 → 101在完全不做显式阻塞的情况下连续调用断言每步状态正确说明同一线程内的调用严格按序执行。跨线程并发串行语义对一个调用耗时很久、另一个调用却相互独立的场景不利。例如一个 Colocated Python 调用在做耗时的文件读取另一个调用要做与之无关的文件写入——二者本可并发而不互相阻塞。Colocated Python 提供并发执行的能力前提是调用来自不同的线程。下面这个例子会让两个 Colocated Python 调用并发运行import concurrent.futures import time def f(x): time.sleep(1) return x 1 f colocated_python.colocated_python(f) f f.specialize(out_specs_fnlambda x: x) # 使调用变为异步。 with concurrent.futures.ThreadPoolExecutor(2) as executor: fut1 executor.submit(f, x) fut2 executor.submit(f, x) # 大约 1 秒完成而不是 2 秒。 jax.block_until_ready([fut1.result(), fut2.result()])要点这里必须先specialize(out_specs_fn...)使调用成为异步派发否则首次调用会同步执行而丧失并发效果测试 tests/colocated_python_test.py 的test_concurrent_execution用三个线程调用同一个包装函数函数体内使用threading.Barrier(3)等待三方同时到达以此验证并发确实发生无死锁即证明三个调用真正并行执行并发只发生在不同线程之间在同一个线程内部程序顺序保证依然成立。从源码看这一语义的实现位于 func.pyspecialized_func中输出规格的发现/编译过程在threading.Lock保护下进行而异步执行本身在锁外执行注释明确写道 Asynchronous execution runs outside of the mutex to allow concurrent execution从而允许不同线程的调用并发执行。七、底层原理序列化、编译与缓存序列化与 IFRT 程序Colocated Python 之所以能把 Python 代码运送到远程主机核心机制在 func.py 的_compile_to_executable将(函数, in_specs_treedef, in_specs_leaves, out_specs_treedef, out_specs_leaves, devices)六元组通过_serialize序列化调用ifrt_programs.make_colocated_python_program(name, pickled_function, devices, in_specs_leaves, out_specs_leaves)构造 IFRT 程序通过devices[0].client.compile_ifrt_program(program, compile_options)编译为可执行体执行时以execute_sharded(args_leaves, with_tokensFalse)分发到设备端主机。序列化实现在 serialization.py依赖cloudpickle完成函数/闭包/类等 Python 对象的打包测试文件也以HAS_CLOUDPICKLE作为前置条件。序列化器还实现了公共对象去重_CommonObjectState_make_reduce_func_with_common_obj嵌套容器或闭包中反复引用同一对象时只序列化一次引用句柄测试test_serialize_with_shared_objtests/colocated_python_test.py验证了两个共享 mesh 的 sharding 序列化体积小于两个独立 sharding 的 2 倍两个相同 sharding 又小于仅共享 mesh 的两个 sharding。此外serialization.py 与 func.py 还专门处理了 PRNG key 这类特殊 dtypekeythreefry2x32/keyrbg的物理表示转换测试test_prng_key_dtype_serialization系列覆盖了两种 PRNG 实现的序列化往返。特化缓存由于Sharding的等价性比较较慢func.py 的_SpecializedCollection维护了两级缓存快路径WeakSpec缓存以输入叶子的dtype、shape、sharding对象的id()与treedef为键LRU(1)查找极快慢路径StrongSpec缓存以完整输入规格持有 sharding 强引用为键无界命中时同步更新两份缓存。整套特化函数缓存注册为 JAX 二级缓存_JaxSecondLevelCaches名colocated_python_specialized_func_cachefunc.py并可通过jax.clear_caches()类机制清理。测试test_simple_functiontests/colocated_python_test.py通过统计colocated_python_func._get_specialized_func事件验证相同输入规格下反复调用只触发一次特化缓存未命中即只编译一次。同步执行的暂存区同步执行路径首次调用发现输出规格通过 func_backend.py 中的SINGLETON_RESULT_STORE单例实现push(uid, out)暂存结果pop(uid)取回同一 uid 重复 push 或 pop 不存在的结果都会抛ValueError。与之对应的对象实例的跨设备存储使用obj_backend.SINGLETON_OBJECT_STORE。后端回退_compile_to_executable还包含一条后端回退路径func.py当编译报错PjRtCompiler requires an HloProgram时如 McJAX 等尚未实现 Colocated Python 支持的 IFRT 后端会退化为在本地直接反序列化函数并同步调用执行——源码注释说明这一回退将在 McJAX 支持 Colocated Python 后移除。八、能力边界与实战建议综合文档、源码与测试使用 Colocated Python 时有以下几点值得注意实验性 API接口可能不兼容变更请勿在依赖长期稳定接口的生产代码中直接引入。输入输出必须是jax.Array_get_specfunc.py明确要求输入输出为jax.Array否则抛出ValueError源码 TODO 提到未来可能允许 Python 值并自动应用shard_arg。输入输出必须使用同一套设备所有参数必须使用相同的 device list输出也必须落在同一套设备上类似 jit 的 SPMD 约束。无输入函数必须先特化devices否则报错。需要异步/并发的场景优先specialize(out_specs_fn...)让调用异步化跨线程提交以获得并发同线程内保持程序顺序。需要等待副作用完成如文件写入用jax.block_until_ready显式阻塞。设备选择顺序敏感的场景显式specialize(devices...)避免参数驱动的设备推导产生歧义。函数/模块级全局状态测试test_module_variable_accesstests/colocated_python_test.py表明包装函数内访问模块级含测试模块内的全局变量是可以工作的很多缓存机制依赖此行为但这种模式在文档中并不被鼓励用于存放用户自定义状态。字符串与二进制数据处理测试还展示了用 Colocated Python 处理StringDType字符串数组test_string_processing与二进制数据test_binary_data_processing的完整写法包括遍历x.addressable_shards、jax.device_get(shard.data)、处理后用jax.make_array_from_single_device_arrays重建输出数组——这为在设备端做 JAX 不原生支持的数据预处理/后处理提供了现成的可参考模式。九、更多参考官方示例文档本主题原始出处docs/notebooks/colocated-python.md另有配套 colocated-python.ipynb 可交互运行。API 顶层入口与文档字符串jax/experimental/colocated_python/api.py导出三个符号colocated_cpu_devices、colocated_python、colocated_python_class见init.py。函数实现与特化机制jax/experimental/colocated_python/func.py。类包装实现jax/experimental/colocated_python/obj.py 与 obj_backend.py。序列化与编译serialization.py、func_backend.py。单机功能测试tests/colocated_python_test.py多主机场景测试位于 tests/multiprocess/colocated_python_test.py。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表