125 lines
4.1 KiB
Python
125 lines
4.1 KiB
Python
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}")
|