如何将Jax实现的鸢尾花数据归一化函数改为纯函数风格?
符合Jax风格的Iris数据集归一化实现
问题背景
我正在学习Jax并编写一个简单分类器,需要对sklearn的Iris数据集做归一化。当前函数可运行,但根据Jax文档提示,应避免lambda函数与向量迭代。我不熟悉函数式编程,想了解更符合Jax风格的写法。
当前代码如下:
import jax.numpy as jnp from jax import jit, vmap # lots of imports ... iris = load_it('data', 'iris.pkl') def normalize(data): return jnp.apply_along_axis(lambda x: x/jnp.linalg.norm(x), 1, data) # TODO: use a functional style, maybe use partial # and get rid of the lambda ... tic = time.perf_counter() iris_data_normal = normalize(iris.data) toc = time.perf_counter() print(f"It took jax {toc - tic:0.4f} seconds.")
运行耗时:It took jax 0.0677 seconds.
优化方案
Jax的核心优势在于向量化操作与XLA编译加速,jnp.apply_along_axis本质是底层循环,lambda会限制JIT的优化空间。以下是两种更贴合Jax设计理念的实现方式:
方法1:使用vmap实现批处理映射
vmap是Jax处理批量数据的标准工具,能自动将单样本函数映射到整个数据集,无需手动迭代:
import jax.numpy as jnp from jax import jit, vmap import time # 定义单样本归一化逻辑 def normalize_single(x): return x / jnp.linalg.norm(x) # 将单样本函数映射到批次维度(默认按第一个维度映射,适配(n_samples, n_features)格式) normalize = vmap(normalize_single) # 可选:用jit进一步加速编译 normalize_jit = jit(normalize) iris = load_it('data', 'iris.pkl') tic = time.perf_counter() iris_data_normal = normalize_jit(iris.data) toc = time.perf_counter() print(f"It took jax {toc - tic:0.4f} seconds.")
方法2:利用广播机制(性能最优)
直接通过Jax的广播特性实现全量向量化运算,完全避免循环与映射,是效率最高的实现方式:
import jax.numpy as jnp from jax import jit import time def normalize(data): # 计算每个样本的范数,保留维度以支持广播运算 sample_norms = jnp.linalg.norm(data, axis=1, keepdims=True) return data / sample_norms # JIT编译优化 normalize_jit = jit(normalize) iris = load_it('data', 'iris.pkl') tic = time.perf_counter() iris_data_normal = normalize_jit(iris.data) toc = time.perf_counter() print(f"It took jax {toc - tic:0.4f} seconds.")
方案优势说明
- 消除lambda函数,让JIT编译器能更充分地分析和优化代码逻辑
- 基于向量化操作而非隐式循环,最大化利用Jax的XLA编译加速能力,性能会显著优于原
apply_along_axis版本 - 代码符合函数式编程风格,逻辑清晰直观,便于后续扩展与维护
内容的提问来源于stack exchange,提问作者Leigh Gable
相关产品推荐
相关产品推荐

