如何高效在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
相关产品推荐
相关产品推荐

