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

如何高效在Nx张量切片上映射要求特定形状输入的函数?

高效批量处理Nx张量切片的方案

方案1:用Nx.map在defn内批量处理

Nx.map是Nx专门用来对张量指定维度做逐元素映射的工具,在defn里使用时,底层会被优化成批量操作,完全避免while循环的串行开销,还能减少数据进出defn的次数,完美适配你的场景。

比如针对你提到的Nx.Random.randint_split/4,假设输入密钥是形状{6,2}的张量,要给每个{2}切片生成{131,131}的矩阵,最后拼接成向量,代码可以这么写:

import Nx.Defn

defn batch_randint_split(keys, target_shape, low, high) do
  # 对keys的第0维度(每个{2}的密钥切片)做映射
  batch_results = Nx.map(keys, fn single_key ->
    Nx.Random.randint_split(single_key, target_shape, low, high)
  end)
  # 把批量结果拼接成向量,最终形状是{6*131*131}
  Nx.flatten(batch_results)
end

这个写法在默认后端就能比while循环快很多,因为Nx.map会把循环展开成批量计算逻辑,而非逐次迭代。等你迁移到EXLA后端时,EXLA还会自动把这个映射操作编译成并行执行的代码,性能提升会更明显。

方案2:纯向量化改造(性能最优)

其实Nx.Random.randint_split本身支持批量密钥输入,只要调整输出形状匹配批量维度,就能实现完全无循环的向量化操作,性能和Nx.Random.uniform_split看齐。

比如密钥是{6,2},要给每个密钥生成{131,131}的矩阵,代码可以这么写:

import Nx.Defn

defn batch_randint_split_vectorized(keys, single_shape, low, high) do
  batch_size = Nx.axis_size(keys, 0)
  # 构造批量输出形状:[批量数] + 单个结果的形状
  output_shape = [batch_size] ++ single_shape
  # 直接传入批量密钥,randint_split会自动处理每个切片
  full_result = Nx.Random.randint_split(keys, output_shape, low, high)
  # 转成目标向量
  Nx.reshape(full_result, [batch_size * Nx.prod(single_shape)])
end

这种纯向量化的方式没有任何循环开销,是性能最优的方案,不管是默认后端还是EXLA都能最大化利用计算资源。

为什么你的while方案性能差?

Nx.Defn.Kernel.while/4是串行迭代逻辑,每次循环只能处理一个切片,而且defn内部的while循环无法被编译器并行化,当批量数大的时候(比如生成131x131矩阵的场景),逐次处理的延迟会不断累积,导致整体耗时飙升。

适配EXLA的小提示

  • 所有张量操作都放在defn函数内部完成,别在defn外面做切片、循环这类操作,减少Elixir和Nx后端之间的数据拷贝。
  • 用Nx.map或向量化写法时,EXLA会自动把代码编译成适合CPU的并行指令,不需要额外修改。
  • 可以提前用EXLA.compile/1预编译defn函数,进一步提升重复调用的性能。

内容的提问来源于stack exchange,提问作者David Maze

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 16:47:28