拓冰建站拓冰建站
首页 / 资讯中心 / 正文

JAX 内存空间与主机卸载(Memory Spaces Host Offloading):用 pinned_host 腾出设备显存/内存的完整实战指南

JAX 内存空间与主机卸载Memory Spaces Host Offloading用 pinned_host 腾出设备显存/内存的完整实战指南【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax导读JAX 的每个分片sharding都携带一个memory kind它决定了数组究竟活在又快又小的设备内存device还是又慢又大的主机内存pinned_host。基于这一抽象JAX 提供了Host offloading主机卸载模式把模型参数、激活值与优化器状态暂存在主机内存里仅在计算真正需要它们时按需搬回设备从而用一次显式的 host↔device 传输换取设备内存容量的大幅释放。读完本文你将掌握 memory kind 的基本概念与with_memory_kind派生方法、jax.device_put与out_shardings的跨空间搬运技巧以及分别针对激活值、参数、优化器状态的三种主机卸载实战方案并学会用jax.stages.Compiled.memory_analysis预先量化收益。说明本文示例输出取自加速器平台原文数值以 TPU 为基准。memory kind 的支持度因平台而异示例旨在说明原理而非逐字可复现且卸载会引入真实的 host↔device 传输开销——投入生产前务必先测量。背景为什么需要第二个内存空间加速器的设备内存带宽高、容量有限主机内存带宽低但容量大。训练大模型时参数、反向传播所需的激活残差residuals与优化器动量都争抢同一块设备内存常常成为瓶颈。JAX 的解法不是压缩数据而是为数组提供两个乃至更多合法的栖身之所每个Sharding都带一个memory_kind属性取值为device默认或pinned_host数组存放的空间完全由 sharding 决定——把一个数组从设备挪到主机等价于把它device_put到另一个 memory kind 的 sharding 上主机端使用pinned页锁定内存以便与设备做 DMA 传输。从源码结构可以印证这一设计JAX 在底层把 memory kind 翻译成内存空间枚举见 jax/_src/core.py 中的mem_kind_to_spacepinned_host→MemorySpace.Host其余一律 →MemorySpace.Device与反向的mem_space_to_kindMemorySpace.Host→pinned_host。也就是说device/pinned_host只是对 XLA 内存空间概念的字符串别名数组的 sharding 一经设定其 aval 的memory_space也随之确定。积木一memory kind 与 with_memory_kindmemory_kind是 sharding 的一个只读属性with_memory_kind(kind)则返回一个换到另一内存空间的新 sharding原 sharding 不变。基类 jax/_src/sharding.py 声明了这对接口各实现类如 NamedSharding、SingleDeviceSharding/GSPMDSharding见 jax/_src/sharding_impls.py均在其上实现了具体逻辑import jax import jax.numpy as jnp from jax.sharding import Mesh, NamedSharding, PartitionSpec as P mesh Mesh(jax.devices()[:1], x) s_dev NamedSharding(mesh, P(x), memory_kinddevice) s_host s_dev.with_memory_kind(pinned_host) print(s_dev) # NamedSharding(..., memory_kinddevice) print(s_host) # NamedSharding(..., memory_kindpinned_host)要点NamedSharding(mesh, P(x), memory_kinddevice)明确写出了设备分片s_dev.with_memory_kind(pinned_host)生成了同一网格、同一分片规则、但数组改放主机的分片。对单设备场景jax.sharding.SingleDeviceSharding(device, memory_kinddevice)同样支持memory_kind与with_memory_kind见 jax/_src/sharding_impls.py。积木二jax.device_put 跨空间搬运jax.device_put把数组放到或移动到sharding 所指向的空间是卸载/回载的基本搬运工具其完整签名与异步语义见 jax/_src/api.pyarr jnp.arange(8.0).reshape(2, 4) arr_host jax.device_put(arr, s_host) arr_dev jax.device_put(arr, s_dev) print(arr_host.sharding.memory_kind) # pinned_host print(arr_dev.sharding.memory_kind) # device关键观察device_put的第二个参数可以是Device、Sharding也可以是它们组成的树前缀若目标 sharding 与数组当前所在空间一致它近似恒等不同则触发真实传输该操作是异步的会立即返回而不会阻塞 Python 线程直到传输完成device_put同样可在 jitted 函数内部调用用于计算中途搬运数值——这正是下文各卸载模式的核心抓手。积木三用 out_shardings 让编译函数直接产出主机数组jit的out_shardings接受带 memory kind 的 sharding因此一个编译函数可以消费设备输入、直接向主机内存产出输出device-to-host消费主机上的输入、向设备产出结果host-to-device。f jax.jit(lambda x: x, out_shardingss_host) out_host f(arr_dev) # inputs on device, outputs land in host memory g jax.jit(lambda x: x 1, out_shardingss_dev) out_dev g(arr_host) # host-resident input, device-resident result注意out_shardings需要与函数输出结构对应——单个 sharding 对单输出元组对多输出见下文优化器状态示例。模式一卸载激活值activations反向传播保存的残差往往是训练中最大的临时内存开销。与其把它们留在设备内存里、或干脆在反向时重算autodiff 可以卸载它们前向把被命名的残差搬到主机反向再取回设备。这属于 rematerialization重物化机制的范畴入口是jax.checkpoint配合策略jax.checkpoint_policies.save_and_offload_only_these_names。策略参数语义该策略共有四个参数源码见 jax/_src/ad_checkpoint.py参数含义反向时行为names_which_can_be_saved允许保存在设备上的命名值残差直接留在设备内存Saveablenames_which_can_be_offloaded允许卸载到目标空间的命名值前向被搬到目标空间Offloadable(src, dst)反向取回offload_src卸载源空间通常为deviceoffload_dst卸载目标空间通常为pinned_host其中names_which_can_be_saved与names_which_can_be_offloaded不允许有交集源码会直接抛ValueError其余未命名或不在名单中的值一律Recompute反向重算。真正执行搬运时实现会按api.device_put(x, core.mem_kind_to_space(policy.offload_dst))把值放到目标空间、再在反向前按offload_src取回见 jax/_src/ad_checkpoint.py。一个直观效果对本文反复使用的10 层 scanned MLP示例仅把各层激活卸载到主机临时内存便从17.25 MB 降到 6.50 MB。完整策略写法为policy jax.checkpoint_policies.save_and_offload_only_these_names( names_which_can_be_saved[], names_which_can_be_offloaded[x], offload_srcdevice, offload_dstpinned_host, )配套测试可参考 tests/memories_test.py 的ActivationOffloadingTest它验证了前向 jaxpr 中卸载残差以MemorySpace.Host输出、反向 jaxpr 以MemorySpace.Device输入并检查编译产物中copy-start/copy-done的复制指令确实存在。有关该机制与 rematerialization 的完整讨论含与仅保存/仅重算策略的对比参见重物化专题文档 docs/301/remat.mdOffloading instead of recomputing一节。模式二卸载参数just-in-time 取权重模型参数可以长期驻留主机内存每个 layer 在真正使用权重前按需取用。该模式的通用骨架是先在主机内存初始化/加载参数再在 layer 函数内部、权重使用前调用jax.device_put(w, s_dev)搬上设备from jax.ad_checkpoint import checkpoint_name from jax import checkpoint_policies as cp policy cp.save_and_offload_only_these_names( names_which_can_be_saved[], names_which_can_be_offloaded[x], offload_srcdevice, offload_dstpinned_host, ) def hybrid_layer(x, w): # Move this layers parameters to device memory just in time. w1, w2 jax.tree.map(lambda w: jax.device_put(w, s_dev), w) x checkpoint_name(x, x) # offload this activation (see the remat docs) y x w1 return y w2, None def hybrid_scanned(w, x): remat_layer jax.remat(hybrid_layer, policypolicy, prevent_cseFalse) result jax.lax.scan(remat_layer, x, w)[0] return jnp.sum(result) input jnp.ones((256, 256), dtypejnp.float32) * 0.001 w1 jnp.ones((10, 256, 1024), dtypejnp.float32) * 0.001 w2 jnp.ones((10, 1024, 256), dtypejnp.float32) * 0.001 # Parameters live in host memory... wh1 jax.device_put(w1, s_host) wh2 jax.device_put(w2, s_host) # ...and the input stays on the device. f jax.jit(jax.grad(hybrid_scanned)) result f((wh1, wh2), input)量化收益memory_analysis 读数对上述示例jax.stages.Compiled.memory_analysisTPU 上报告Temp size: 4.75 MB Argument size: 0.25 MB Total size: 25.00 MB而不卸载的基线为 17.25 MB 临时内存 20.25 MB 参数内存。收益来自三个效应的叠加参数卸载把权重从设备参数内存中移除20.25 MB → 0.25 MB设备上只剩输入激活卸载把临时内存从 17.25 MB 压到 6.50 MB两者交互还能再省一点6.50 MB → 4.75 MB——remat 策略避免了 JAX 在反向过程中始终持有权重的设备副本。两个必须知道的限制jax.lax.scan是这个模式的关键承载若改用显式 Python 循环参数会持续占用设备内存无法获得节省参数卸载目前只在按 axis 0 扫描时有效对其他轴扫描时把参数送回设备需要插入代价高昂的transpose且并非所有平台都支持。模式三卸载优化器状态optimizer state优化器状态如 Adam 的一阶/二阶动量在每一步中只被短暂使用却长期霸占设备内存。同样的模式依然适用步骤间状态驻留主机步骤内搬到设备更新后的状态通过out_shardings发回主机import optax s_dev jax.sharding.SingleDeviceSharding(jax.devices()[0], memory_kinddevice) s_host jax.sharding.SingleDeviceSharding(jax.devices()[0], memory_kindpinned_host) optimizer optax.chain(optax.clip_by_global_norm(1.0), optax.adam(learning_rate0.1)) # (network and loss definitions elided) def step(params, opt_state, inputs): grads jax.grad(lambda p: compute_loss(p, inputs))(params) opt_state jax.device_put(opt_state, s_dev) # fetch state to the device updates, new_opt_state optimizer.update(grads, opt_state, params) new_params optax.apply_updates(params, updates) return new_params, new_opt_state params init_params() # on device opt_state optimizer.init(params) opt_state jax.device_put(opt_state, s_host) # state lives on the host step jax.jit( step, donate_argnums(0,), out_shardings(s_dev, s_host), # params to device, state back to host ) new_params, new_opt_state step(params, opt_state, input)技术细节单设备场景下使用SingleDeviceSharding(jax.devices()[0], memory_kind...)显式声明空间out_shardings(s_dev, s_host)是一个与返回元组结构对齐的元组params落到设备、更新后的opt_state直接发回主机省去一次显式搬运donate_argnums(0,)捐献输入参数缓冲区配合传输尽可能复用内存。量化收益与开销结构对一个四层 7168×7168 的 MLP Adam 训练步memory analysis 报告不卸载总计4.59 GB卸载后2.87 GB——省下约1.72 GB几乎全部来自优化器状态退出设备参数内存。需要正视这笔 trade-off 的结构卸载可能增加临时内存更新后的状态在复制到主机前仍需设备缓冲区暂存且 XLA 的 latency-hiding 调度会为了传输与计算重叠而延长缓冲区存活区间但只要参数内存的节省大于新增临时内存净收益通常依然可观——所以先测量再决定是铁律。如何量化测量memory_analysis 与运行时剖析编译期预估jax.stages.Compiled.memory_analysis本文所有数字均来自jax.stages.Compiled.memory_analysis实现见 jax/_src/stages.py。它在你运行之前就报告一个编译函数的内存分解compiled f.lower((wh1, wh2), input).compile() print(compiled.memory_analysis())总量估算方式临时内存Temp 参数内存Argument 输出内存Output再减去别名alias部分。该 API 面向可视化和调试返回结构在不同 JAX/jaxlib 版本间可能不一致若后端/编译器不支持会抛出NotImplementedError源码中的降级路径会尝试cost_analysis。运行时测量profiling 与 tracingmemory_analysis只回答编译后预计占多少无法回答传输是否真的与计算重叠。要观测真实的运行时设备内存占用与传输调度请使用 JAX 的设备内存剖析与 tracing 工具详见 docs/201/profiling.md如 Memory Profile / Memory Viewer 等内存视图工具并配合 docs/201/gpu-memory.md 了解 GPU 内存分配机制。平台支持与注意事项平台差异memory kind 支持度因平台而异。仓库测试 tests/memories_test.py 在 CPU 后端直接跳过Memories do not work on CPU backend.仅在 TPU / GPU 上运行——pinned_host卸载本质上是加速器场景的能力。先测量再承诺卸载会引入真实 host↔device 传输成本且 latency-hiding 可能延长部分缓冲区存活时间。是否值得以memory_analysis与运行剖析的数字为准。收益结构参数/优化器状态类卸载主要削减Argument size激活卸载主要削减Temp size两者配合remat 策略还能进一步消除反向阶段的设备权重副本。结构性前提批量化的jax.lax.scanaxis 0是参数卸载能否生效的前提Python 显式循环无法享受该收益。从示例到生产模式组合建议把三种模式拼在一起即可得到一条完整的低设备内存训练流水线初始化阶段参数与优化器状态创建后立即device_put到pinned_host空间s_host训练输入保持在设备s_dev前向/反向阶段用jax.checkpointsave_and_offload_only_these_names策略标记可卸载的激活名让残差落主机每步内部在 layer / step 函数中按需device_put取用参数与状态每步输出通过out_shardings(s_dev, s_host)让更新后的参数回设备、新状态直接落主机避免额外拷贝验证阶段先compile()后用memory_analysis预估再对拍 profiling 数据确认传输确实与计算重叠。这套基于 memory kind 的抽象把数组放在哪块内存从隐式约定提升为一等公民它不改变任何数学语义只改变数据的存放与移动时机——当你被设备内存逼到墙角时这往往比重计算remat或梯度压缩更直接、更可预测地释放容量。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门