JAX 中文文档(十三)(3)

简介: JAX 中文文档(十三)

JAX 中文文档(十三)(2)https://developer.aliyun.com/article/1559742


GPU 内存分配

原文:jax.readthedocs.io/en/latest/gpu_memory_allocation.html

当第一个 JAX 操作运行时,JAX 将预先分配总 GPU 内存的 75%。 预先分配可以最小化分配开销和内存碎片化,但有时会导致内存不足(OOM)错误。如果您的 JAX 进程因内存不足而失败,可以使用以下环境变量来覆盖默认行为:

XLA_PYTHON_CLIENT_PREALLOCATE=false

这将禁用预分配行为。JAX 将根据需要分配 GPU 内存,可能会减少总体内存使用。但是,这种行为更容易导致 GPU 内存碎片化,这意味着使用大部分可用 GPU 内存的 JAX 程序可能会在禁用预分配时发生 OOM。

XLA_PYTHON_CLIENT_MEM_FRACTION=.XX

如果启用了预分配,这将使 JAX 预分配总 GPU 内存的 XX% ,而不是默认的 75%。减少预分配量可以修复 JAX 程序启动时的内存不足问题。

XLA_PYTHON_CLIENT_ALLOCATOR=platform

这使得 JAX 根据需求精确分配内存,并释放不再需要的内存(请注意,这是唯一会释放 GPU 内存而不是重用它的配置)。这样做非常慢,因此不建议用于一般用途,但可能对于以最小可能的 GPU 内存占用运行或调试 OOM 失败非常有用。

OOM 失败的常见原因

同时运行多个 JAX 进程。

要么使用 XLA_PYTHON_CLIENT_MEM_FRACTION 为每个进程分配适当的内存量,要么设置 XLA_PYTHON_CLIENT_PREALLOCATE=false

同时运行 JAX 和 GPU TensorFlow。

TensorFlow 默认也会预分配,因此这与同时运行多个 JAX 进程类似。

一个解决方案是仅使用 CPU TensorFlow(例如,如果您仅使用 TF 进行数据加载)。您可以使用命令 tf.config.experimental.set_visible_devices([], "GPU") 阻止 TensorFlow 使用 GPU。

或者,使用 XLA_PYTHON_CLIENT_MEM_FRACTIONXLA_PYTHON_CLIENT_PREALLOCATE。还有类似的选项可以配置 TensorFlow 的 GPU 内存分配(gpu_memory_fractionallow_growth 在 TF1 中应该设置在传递给 tf.Sessiontf.ConfigProto 中。参见 使用 GPU:限制 GPU 内存增长 用于 TF2)。

在显示 GPU 上运行 JAX。

使用 XLA_PYTHON_CLIENT_MEM_FRACTIONXLA_PYTHON_CLIENT_PREALLOCATE

提升秩警告

原文:jax.readthedocs.io/en/latest/rank_promotion_warning.html

NumPy 广播规则 允许自动将参数从一个秩(数组轴的数量)提升到另一个秩。当意图明确时,此行为很方便,但也可能导致意外的错误,其中静默的秩提升掩盖了潜在的形状错误。

下面是提升秩的示例:

>>> import numpy as np
>>> x = np.arange(12).reshape(4, 3)
>>> y = np.array([0, 1, 0])
>>> x + y
array([[ 0,  2,  2],
 [ 3,  5,  5],
 [ 6,  8,  8],
 [ 9, 11, 11]]) 

为了避免潜在的意外,jax.numpy 可配置,以便需要提升秩的表达式会导致警告、错误或像常规 NumPy 一样允许。配置选项名为 jax_numpy_rank_promotion,可以取字符串值 allowwarnraise。默认设置为 allow,允许提升秩而不警告或错误。设置为 raise 则在提升秩时引发错误,而 warn 在首次提升秩时引发警告。

可以使用 jax.numpy_rank_promotion() 上下文管理器在本地启用或禁用提升秩:

with jax.numpy_rank_promotion("warn"):
  z = x + y 

这个配置也可以在多种全局方式下设置。其中一种是在代码中使用 jax.config

import jax
jax.config.update("jax_numpy_rank_promotion", "warn") 

也可以使用环境变量 JAX_NUMPY_RANK_PROMOTION 来设置选项,例如 JAX_NUMPY_RANK_PROMOTION='warn'。最后,在使用 absl-py 时,可以使用命令行标志设置选项。

公共 API:jax 包

原文:jax.readthedocs.io/en/latest/jax.html

子包

  • jax.numpy 模块
  • jax.scipy 模块
  • jax.lax 模块
  • jax.random 模块
  • jax.sharding 模块
  • jax.debug 模块
  • jax.dlpack 模块
  • jax.distributed 模块
  • jax.dtypes 模块
  • jax.flatten_util 模块
  • jax.image 模块
  • jax.nn 模块
  • jax.ops 模块
  • jax.profiler 模块
  • jax.stages 模块
  • jax.tree 模块
  • jax.tree_util 模块
  • jax.typing 模块
  • jax.export 模块
  • jax.extend 模块
  • jax.example_libraries 模块
  • jax.experimental 模块

配置

config
check_tracer_leaks jax_check_tracer_leaks 配置选项的上下文管理器。
checking_leaks jax_check_tracer_leaks 配置选项的上下文管理器。
debug_nans jax_debug_nans 配置选项的上下文管理器。
debug_infs jax_debug_infs 配置选项的上下文管理器。
default_device jax_default_device 配置选项的上下文管理器。
default_matmul_precision jax_default_matmul_precision 配置选项的上下文管理器。
default_prng_impl jax_default_prng_impl 配置选项的上下文管理器。
enable_checks jax_enable_checks 配置选项的上下文管理器。
enable_custom_prng jax_enable_custom_prng 配置选项的上下文管理器(临时)。
enable_custom_vjp_by_custom_transpose jax_enable_custom_vjp_by_custom_transpose 配置选项的上下文管理器(临时)。
log_compiles jax_log_compiles 配置选项的上下文管理器。
numpy_rank_promotion jax_numpy_rank_promotion 配置选项的上下文管理器。
transfer_guard(new_val) 控制所有传输的传输保护级别的上下文管理器。

即时编译 (jit)

jit(fun[, in_shardings, out_shardings, …]) 使用 XLA 设置 fun 进行即时编译。
disable_jit([disable]) 禁用其动态上下文下 jit() 行为的上下文管理器。
ensure_compile_time_eval() 确保在追踪/编译时进行评估的上下文管理器(或错误)。
xla_computation(fun[, static_argnums, …]) 创建一个函数,给定示例参数,产生其 XLA 计算。
make_jaxpr([axis_env, return_shape, …]) 创建一个函数,给定示例参数,产生其 jaxpr。
eval_shape(fun, *args, **kwargs) 计算 fun 的形状/数据类型,不进行任何 FLOP 计算。
ShapeDtypeStruct(shape, dtype[, …]) 数组的形状、dtype 和其他静态属性的容器。
device_put(x[, device, src]) x 传输到 device
device_put_replicated(x, devices) 将数组传输到每个指定的设备并形成数组。
device_put_sharded(shards, devices) 将数组片段传输到指定设备并形成数组。
device_get(x) x 传输到主机。
default_backend() 返回默认 XLA 后端的平台名称。
named_call(fun, *[, name]) 在 JAX 计算中给函数添加用户指定的名称。
named_scope(name) 将用户指定的名称添加到 JAX 名称堆栈的上下文管理器。

| block_until_ready(x) | 尝试调用 pytree 叶子上的 block_until_ready 方法。 | ## 自动微分

grad(fun[, argnums, has_aux, holomorphic, …]) 创建一个评估 fun 梯度的函数。
value_and_grad(fun[, argnums, has_aux, …]) 创建一个同时评估 funfun 梯度的函数。
jacfwd(fun[, argnums, has_aux, holomorphic]) 使用正向模式自动微分逐列计算 fun 的雅可比矩阵。
jacrev(fun[, argnums, has_aux, holomorphic, …]) 使用反向模式自动微分逐行计算 fun 的雅可比矩阵。
hessian(fun[, argnums, has_aux, holomorphic]) fun 的 Hessian 矩阵作为稠密数组。
jvp(fun, primals, tangents[, has_aux]) 计算 fun 的(正向模式)雅可比向量乘积。
linearize() 使用 jvp() 和部分求值生成对 fun 的线性近似。
linear_transpose(fun, *primals[, reduce_axes]) 转置一个承诺为线性的函数。
vjp() )) 计算 fun 的(反向模式)向量-Jacobian 乘积。
custom_jvp(fun[, nondiff_argnums]) 为自定义 JVP 规则定义一个可 JAX 化的函数。
custom_vjp(fun[, nondiff_argnums]) 为自定义 VJP 规则定义一个可 JAX 化的函数。
custom_gradient(fun) 方便地定义自定义的 VJP 规则(即自定义梯度)。
closure_convert(fun, *example_args) 闭包转换实用程序,用于与高阶自定义导数一起使用。
checkpoint(fun, *[, prevent_cse, policy, …]) 使 fun 在求导时重新计算内部线性化点。

jax.Array (jax.Array)

Array() JAX 的数组基类
make_array_from_callback(shape, sharding, …) 通过从 data_callback 获取的数据返回一个 jax.Array
make_array_from_single_device_arrays(shape, …) 从每个位于单个设备上的 jax.Array 序列返回一个 jax.Array
make_array_from_process_local_data(sharding, …) 使用进程中可用的数据创建分布式张量。

向量化 (vmap)

vmap(fun[, in_axes, out_axes, axis_name, …]) 向量化映射。
numpy.vectorize(pyfunc, *[, excluded, signature]) 定义一个支持广播的向量化函数。

并行化 (pmap)

pmap(fun[, axis_name, in_axes, out_axes, …]) 支持集体操作的并行映射。
devices([backend]) 返回给定后端的所有设备列表。
local_devices([process_index, backend, host_id]) 类似于 jax.devices(),但仅返回给定进程局部的设备。
process_index([backend]) 返回此进程的整数进程索引。
device_count([backend]) 返回设备的总数。
local_device_count([backend]) 返回此进程可寻址的设备数量。
process_count([backend]) 返回与后端关联的 JAX 进程数。

Callbacks

pure_callback(callback, result_shape_dtypes, …) 调用一个纯 Python 回调函数。
experimental.io_callback(callback, …[, …]) 调用一个非纯 Python 回调函数。
debug.callback(callback, *args[, ordered]) 调用一个可分期的 Python 回调函数。
debug.print(fmt, *args[, ordered]) 打印值,并在分期 JAX 函数中工作。

Miscellaneous

Device 可用设备的描述符。
print_environment_info([return_string]) 返回一个包含本地环境和 JAX 安装信息的字符串。
live_arrays([platform]) 返回后端平台上的所有活动数组。
clear_caches() 清除所有编译和分期缓存。


JAX 中文文档(十三)(4)https://developer.aliyun.com/article/1559745

相关实践学习
在云上部署ChatGLM2-6B大模型(GPU版)
ChatGLM2-6B是由智谱AI及清华KEG实验室于2023年6月发布的中英双语对话开源大模型。通过本实验,可以学习如何配置AIGC开发环境,如何部署ChatGLM2-6B大模型。
相关文章
|
5月前
|
人工智能 运维 安全
别让“龙虾”裸奔!企业规模化“养虾”亟需新一代云网架构护航
本文深入分析了千级 AI Agent 规模化部署时面临的网络架构挑战,包括单 VPC 带宽天花板、安全组规则爆炸、混合云互通复杂等三大硬伤。基于阿里云 ACS+VPC+TR+CEN 的分层隔离架构,提供按业务域划分 VPC、TR 统一路由枢纽、CEN 全球互联的完整解决方案,实现性能隔离、故障隔离、安全合规、弹性扩展和成本优化五大核心价值。适用于 Agent 数量>500、多地域部署、强合规行业及混合云架构的企业场景。
727 6
别让“龙虾”裸奔!企业规模化“养虾”亟需新一代云网架构护航
|
7月前
|
机器学习/深度学习 自然语言处理 并行计算
大模型应用:混合专家模型(MoE):大模型性能提升的关键技术拆解.37
MoE(混合专家模型)是一种高效大模型架构,通过“智能调度+稀疏激活”机制,让多个专业化子网络(专家)按需协作。它兼顾性能与效率:参数规模大但推理仅激活2-4个专家,显著降本提速;既保持通用能力,又在医疗、法律等细分领域更专精,是当前大模型落地的关键技术。
1411 17
|
数据采集 监控 网络协议
【计算机网络】你真的懂学校的校园网吗?
在数字时代,计算机网络已经成为了现代社会不可或缺的一部分。而对于大多数人来说,校园网是我们日常生活中接触最频繁的网络之一,它为学校的师生提供了信息传输、资源共享和互联互通的基础设施。但是,尽管我们每天都在使用校园网,很少有人真正深入了解它的工作原理、安全性和管理细节。
6327 3
|
NoSQL 前端开发 Java
redis的发布/订阅(命令、普通工程、springboot实现)
小美老师给五年级三班上数学课的时候,实现给所在班级进行实时推送数学课程的活动(广播通信)
|
机器学习/深度学习 人工智能 开发者
【AI系统】昇思 MindSpore 关键特性
本文介绍华为自研AI框架昇思MindSpore,一个面向全场景的AI计算框架,旨在提供统一、高效、安全的平台,支持AI算法研究与生产部署。文章详细阐述了MindSpore的定位、架构、特性及在端边云全场景下的应用优势,强调其动静态图统一、联邦学习支持及高性能优化等亮点。
762 7
【AI系统】昇思 MindSpore 关键特性
|
存储 安全 开发工具
App隐私合规评估实务和要点
随着移动互联网的高速发展及监管部门针对移动互联网应用程序(以下简称“App”)隐私合规监管趋严,特别是在个人信息保护法的实施下。本文将深入探讨App隐私合规评估的要点和难点,提供详细的信息,并提供一套轻量级和自动化的App隐私合规治理方案,降低App业务被通报和下架等合规风险,以保障企业App业务正常运营。
1910 0
|
编译器 API C++
JAX 中文文档(三)(3)
JAX 中文文档(三)
416 0
|
Perl
技术笔记:samtools统计重测序数据深度depth、depth
技术笔记:samtools统计重测序数据深度depth、depth
1038 0
|
机器学习/深度学习 生物认证 语音技术
声纹识别入门:原理与基础知识
【10月更文挑战第16天】声纹识别(Voice Biometrics)是生物特征识别技术的一种,它通过分析个人的语音特征来验证身份。与指纹识别或面部识别相比,声纹识别具有非接触性、易于远程操作等特点,因此在电话银行、客户服务、智能家居等领域得到了广泛应用。
3892 0