You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.27 19:42:51