# JAX 入门示例代码 # 演示 JAX 的基本用法,包括 jnp 数组、jit 加速 import numpy as np import jax import jax.numpy as jnp from jax import jit, grad, vmap import time # --- 1. 像 NumPy 一样简单 --- print("--- 1. NumPy-like API ---") key = jax.random.PRNGKey(0) # JAX 处理随机数需要一个密钥 x_jnp = jnp.arange(10) y_jnp = jax.random.normal(key, (10,)) print(f"JAX 数组 x: {x_jnp}") print(f"JAX 数组 y (随机): {y_jnp}") print(f"x 和 y 的点积: {jnp.dot(x_jnp, y_jnp)}") # --- 2. jax.jit: 极致加速 --- print("\n--- 2. @jit 加速 ---") # 定义一个包含大量运算的函数 def slow_function(x): # 这是一个比较耗时的操作 (矩阵乘法循环) for _ in range(100): x = jnp.dot(x, x.T) return x # 创建 JIT 编译版本的函数 fast_function = jit(slow_function) # 创建数据 big_matrix = jax.random.normal(key, (200, 200)) # --- 计时对比 --- # a) 运行普通 NumPy/Python 版本 start_time = time.time() slow_function(big_matrix).block_until_ready() # .block_until_ready() 确保 JAX 计算完成 numpy_time = time.time() - start_time print(f"普通 NumPy/Python 版本耗时: {numpy_time:.6f} 秒") # b) 运行 JIT 编译版本 # 第一次运行会包含编译时间 start_time = time.time() fast_function(big_matrix).block_until_ready() compile_time = time.time() - start_time print(f"JIT 版本 (首次,含编译) 耗时: {compile_time:.6f} 秒") # 第二次运行,将只体现执行速度 start_time = time.time() fast_function(big_matrix).block_until_ready() jit_time = time.time() - start_time print(f"JIT 版本 (第二次) 耗时: {jit_time:.6f} 秒") print(f"加速比 (第二次运行 vs 普通版本): {numpy_time / jit_time:.2f} 倍") # --- 3. jax.grad: 自动求导 --- print("\n--- 3. grad 自动求导 ---") # 定义一个简单的函数 f(x) = x^3 + 2x^2 + 5 def my_func(x): return x**3 + 2*x**2 + 5 # 使用 grad 创建一个计算 my_func 导数的函数 # f'(x) = 3x^2 + 4x grad_my_func = grad(my_func) x_val = 2.0 derivative = grad_my_func(x_val) expected_derivative = 3 * x_val**2 + 4 * x_val print(f"函数 f(x) = x^3 + 2x^2 + 5") print(f"在 x = {x_val} 处的导数是: {derivative}") print(f"理论上的导数值是: {expected_derivative}") # --- 4. jax.vmap: 自动向量化 --- print("\n--- 4. vmap 自动向量化 ---") # 定义一个只能处理单个向量的函数 (向量点积) def single_dot_product(a, b): return jnp.dot(a, b) # 创建一批 (batch) 数据 # 假设我们有 5 对向量,每对向量长度为 3 batch_a = jnp.arange(15).reshape(5, 3) batch_b = jnp.arange(15, 30).reshape(5, 3) # 使用 vmap 将函数向量化 # in_axes=(0, 0) 表示对 a 和 b 的第 0 轴 (批次轴) 进行映射 # out_axes=0 表示输出结果也沿着第 0 轴堆叠 batch_dot_product = vmap(single_dot_product, in_axes=(0, 0), out_axes=0) results = batch_dot_product(batch_a, batch_b) print("批次数据 a:\n", batch_a) print("批次数据 b:\n", batch_b) print("vmap 向量化计算的点积结果:\n", results) # 对比手动 for 循环 manual_results = jnp.array([single_dot_product(a, b) for a, b in zip(batch_a, batch_b)]) print("手动 for 循环的结果:\n", manual_results) print(f"vmap 结果与手动循环结果是否一致: {jnp.allclose(results, manual_results)}") # --- 5. JAX 的随机数处理 --- print("\n--- 5. JAX 的随机数 ---") # 1. 创建一个初始密钥 key = jax.random.PRNGKey(42) print(f"初始密钥: {key}") # 2. 使用密钥生成随机数 random_data_1 = jax.random.normal(key, (3,)) print(f"第一次生成的随机数据: {random_data_1}") # 3. 再次使用同一个密钥,会得到完全相同的结果! random_data_2 = jax.random.normal(key, (3,)) print(f"第二次使用相同密钥生成的数据: {random_data_2}") # 4. 正确的做法:分割密钥 key, subkey = jax.random.split(key) # key 更新为新的主密钥, subkey 用于本次操作 random_data_3 = jax.random.normal(subkey, (3,)) print(f"\n分割密钥后,第一次生成的随机数据: {random_data_3}") key, subkey = jax.random.split(key) # 再次分割 random_data_4 = jax.random.normal(subkey, (3,)) print(f"分割密钥后,第二次生成的随机数据: {random_data_4}")