Elixir Nx 与数值计算:张量、Axon 与模型推理

深入 Elixir 数值计算栈 Nx:张量形状与类型提升、命名维度与广播归约语义、EXLA/XLA 惰性编译与 JIT 缓存、Axon 模型定义与自动微分训练循环、Nx.Serving 自动批处理推理,以及与 ONNX/PyTorch 权重互操作和 NIF 阻塞调度器、进程外显存等生产陷阱。

BEAM 虚拟机以并发、容错和软实时见长,但它的数值计算能力长期被视为短板:没有原生多维数组,浮点运算要经过 tagged term 拆箱,循环密集型代码很快撞上调度器瓶颈。Elixir 生态用 Nx(Numerical Elixir) 补上了这块拼图——它不是要取代 Python 的 NumPy/PyTorch,而是让「数值计算」与「OTP 并发」在同一个运行时里协作:模型推理可以作为监督树里的一个普通进程,失败隔离、热更新、背压治理全部沿用既有基础设施。

本文要回答三个问题:Nx 的张量与后端抽象究竟如何工作;Axon 怎样把一个模型定义编译成可执行计算图;把推理放进生产服务时,NIF、调度器与内存搬运会带来哪些坑。

Nx 的定位与张量基础

Nx 的核心数据结构是 张量(Tensor),它是一个不可变(immutable)的、带形状与元素类型描述的多维数组。与 Erlang 的列表和元组不同,张量在内存中是连续布局的,这正是能高效交给底层数值库的前提。

# mix.exs
defp deps do
  [
    {:nx, "~> 0.7"},
    {:exla, "~> 0.7"},
    {:axon, "~> 0.7"}
  ]
end
iex> t = Nx.tensor([[1, 2, 3], [4, 5, 6]])
#Nx.Tensor<
  s64[2][3]
  [
    [1, 2, 3],
    [4, 5, 6]
  ]
>

iex> Nx.shape(t)
{2, 3}

iex> Nx.type(t)
{:s64}

iex> Nx.size(t)
6

三个属性决定了一个张量的一切行为:shape(维度)、type(元素类型)、backend(后端)。Nx 支持的元素类型覆盖 {:s, 8|16|32|64}、{:u, ...}、{:f, 32|64} 与 {:bf, 16}。类型默认跟随输入推断,整数列表得到 s64,浮点得到 f32;在混合运算里 Nx 会向上提升类型,Nx.add/2 遇到 s64 与 f32 会得到 f32。

类型提升规则容易在训练中埋雷:如果权重初始化为 f32 而输入数据是整数,第一次矩阵乘法就会把权重提升为 f64,显存与算力翻倍。显式固定类型是纪律:

x = Nx.tensor([1, 2, 3], type: :f32)
w = Nx.tensor([0.5, 0.5, 0.5], type: :f32)
Nx.multiply(x, w) |> Nx.type()
# {:f32}

形状(shape)支持命名维度,这在处理 batch/sequence/feature 三类轴时极大降低出错概率:

x = Nx.tensor([[[1.0]], [[2.0]]], names: [:batch, :seq, :feat])
Nx.shape(x)
# {:batch, :seq, :feat}

命名维度会参与广播与归约的语义校验——Nx.sum(x, axes: [:seq]) 比 Nx.sum(x, axes: [1]) 更难写错。代价是每次运算多一层元数据开销,纯性能敏感的循环里可以改用整数轴。

惰性后端与立即求值

Nx 把「构造表达式」与「执行计算」分离。默认后端 Nx.BinaryBackend 是立即求值的纯 Elixir 实现,好处是零依赖、可调试、结果确定;坏处是它逐元素走 Elixir 函数,矩阵乘法比原生库慢几个数量级。把后端换成 EXLA 后,同一段代码会被构造成计算图并交给 XLA 编译执行。

# config/config.exs
import Config

config :nx, default_backend: EXLA.Backend

也可以按表达式临时指定:

Nx.default_backend(EXLA.Backend)
Nx.default_backend(Nx.BinaryBackend)

后端选择:BinaryBackend 与 EXLA

EXLA 是 Google XLA(Accelerated Linear Algebra)编译器在 Elixir 侧的绑定,通过 NIF 加载 libexla 共享库。它带来三件事:算子融合、JIT 编译缓存、以及可选的 GPU 执行。

后端实现求值时机适用场景
Nx.BinaryBackend纯 Elixir立即单元测试、小数据、调试
EXLA.Backend (CPU)XLA NIF惰性/即时生产推理、训练
EXLA.Backend (CUDA)XLA GPU惰性大模型训练、批量推理
Torchx.Backendlibtorch惰性需要 PyTorch 算子覆盖

EXLA 的关键参数在应用配置里:

config :exla, :clients,
  cuda: [platform: :cuda, preallocate: true, memory_fraction: 0.9],
  host: [platform: :host]

preallocate: true 让 GPU 在启动时一次性占用指定比例显存,避免运行期碎片化;代价是同一张卡上跑多个 BEAM 节点会互相抢显存。

编译与执行分离是 EXLA 的性能来源。Nx.Defn.jit/2 会把一个 defn 函数编译成可复用的一元算子:

defmodule MyMath do
  import Nx.Defn

  defn softmax(x) do
    x
    |> Nx.subtract(Nx.reduce_max(x, axes: [-1], keep_axes: true))
    |> Nx.exp()
    |> Nx.divide(Nx.sum(Nx.exp(x), axes: [-1], keep_axes: true))
  end
end

# 首次调用触发编译,后续复用编译产物
MyMath.softmax(Nx.tensor([[1.0, 2.0, 3.0]]))

defn 是宏定义的受限 DSL:函数体内只能调用 Nx/defn 允许的算子,不能有任意副作用,条件分支要用 Nx.select/3 或 defn if 表达。这种限制换来的是可静态分析、可融合、可编译成 XLA HLO 图。数值计算与 Erlang 的 Port 与 NIF 互操作 在这里交汇:EXLA 本身就是一个加载进 VM 的 NIF,因此它的调度行为会直接影响 BEAM 的响应性。

张量运算、广播与归约

Nx 的算子命名与 NumPy 高度对齐,迁移成本低。常用族如下:

  • 逐元素(element-wise):Nx.add/2、Nx.multiply/2、Nx.exp/1、Nx.log/1、Nx.clip/2
  • 线性代数:Nx.dot/2、Nx.LinAlg.invert/1、Nx.LinAlg.svd/1
  • 归约(reduction):Nx.sum/2、Nx.mean/2、Nx.reduce_max/2、Nx.argmax/2
  • 变形:Nx.reshape/2、Nx.transpose/2、Nx.new_axis/2、Nx.squeeze/2
  • 索引:Nx.slice/4、Nx.take/2、Nx.gather/2、Nx.put_slice/4

广播(broadcasting) 规则与 NumPy 一致:从尾部对齐维度,长度相等或为 1 时可广播。理解广播能避免大量无谓的 Nx.tile 拷贝:

x = Nx.tensor([[1.0], [2.0], [3.0]])   # shape {3, 1}
y = Nx.tensor([10.0, 20.0])            # shape {2}
Nx.add(x, y)                            # shape {3, 2}

矩阵乘法与批量矩阵乘法要区分清楚:

a = Nx.tensor([[1, 2], [3, 4]], type: :f32)   # {2, 2}
b = Nx.tensor([[5, 6], [7, 8]], type: :f32)   # {2, 2}

Nx.dot(a, b)              # {2, 2},矩阵乘
Nx.dot(a, [0], b, [0])    # {2, 2},指定收缩轴

批量场景下常见的错误是用 Enum.map 逐样本调用 Nx.dot,这会退化成 N 次小算子调用,完全丢掉 XLA 的融合收益。正确做法是把样本堆叠成一个 {batch, ...} 张量一次性算完。

归约的 keep_axes 参数决定了形状是否保留,直接影响能否与后续运算广播:

x = Nx.tensor([[1.0, 2.0], [3.0, 4.0]])

Nx.sum(x, axes: [1])                    # {2}
Nx.sum(x, axes: [1], keep_axes: true)   # {2, 1},可直接与 x 广播相除

Axon:模型定义与训练

Axon 是 Elixir 的函数式神经网络库。模型是一个由 Axon.input 出发、经 Axon.dense 等层组合而成的数据流图(DAG),不是可变的 layer 对象。这与 PyTorch 的命令式风格不同,更接近 Keras Functional API。

defmodule MyModel do
  def build do
    Axon.input("features", shape: {nil, 784})
    |> Axon.dense(128, activation: :relu)
    |> Axon.dropout(rate: 0.2)
    |> Axon.dense(10, activation: :softmax)
  end
end

shape: {nil, 784} 里的 nil 是动态 batch 维,Axon 会把它标记为可变轴。模型初始化需要显式提供输入形状与随机种子:

model = MyModel.build()

template = Nx.template({32, 784}, :f32)
params = Axon.init(model, template, seed: 42)

训练循环是一个普通函数,可以用 Axon.Loop 的高层 API,也可以手写梯度步骤。手写能看清每一步在做什么:

model = Axon.build(MyModel.build(), mode: :train)
loss_fn = &Axon.Losses.categorical_cross_entropy(&1, &2, from_logits: false)

# 优化器由 init/update 两个函数组成
{init_opt, update_opt} = Axon.Optimizers.adam(learning_rate: 0.001)

step = fn {batch_x, batch_y}, {params, opt_state} ->
  grad_fn = Nx.Defn.grad(fn p -> loss_fn.(Axon.predict(model, p, batch_x), batch_y) end)
  grads = grad_fn.(params)
  update_opt.(grads, params, opt_state)
end

{params, opt_state} = step.({batch_x, batch_y}, {params, init_opt.(params)})

Nx.Defn.grad/1 做的是反向模式自动微分(reverse-mode autodiff),它把 defn 函数变换成返回梯度的函数。这一步同样走 XLA 编译,所以第一次调用会有一段编译延迟(通常几百毫秒到数秒),之后每次迭代都很快。

用 Axon.Loop 可以把训练、验证、指标、检查点串起来:

Axon.Loop.trainer(model, loss_fn, :adam)
|> Axon.Loop.metric(:accuracy, "accuracy")
|> Axon.Loop.validate(model, val_data)
|> Axon.Loop.checkpoint(event: :epoch_end)
|> Axon.Loop.run(train_data, params, epochs: 20)

常见陷阱:Axon.Loop 默认按 Axon.Loop.trainer 的损失函数类型推断输出形状;如果最后一层是 softmax,损失函数必须用 from_logits: false,否则会做两次 softmax,梯度被压平、训练几乎不收敛。

Nx.Serving:生产级批量推理

训练是离线任务,推理是线上服务,二者对延迟与吞吐的要求相反。Nx.Serving 是 Nx 提供的推理封装,核心能力是自动批处理(batching):把短时间窗口内到达的多个请求合并成一个大 batch 交给模型,显著提升 GPU 利用率。

serving =
  Nx.Serving.new(Nx.Defn.jit(&MyModel.predict/2), params)
  |> Nx.Serving.batch_size(64)
  |> Nx.Serving.batch_timeout(10)

result = Nx.Serving.run(serving, Nx.tensor([...]))

三个参数决定批处理行为:

  • batch_size:单批最大样本数,超过则拆批
  • batch_timeout:等待凑批的最长时间(毫秒),决定尾延迟下限
  • Nx.Serving.client_preprocessing / client_postprocessing:请求级的前后处理,在批处理之外执行

batch_timeout 是最需要调优的参数。设为 10ms 意味着 P99 延迟至少增加 10ms;设得过大则在低流量时白白等批。经验值是把 batch_timeout 设为「可接受的额外延迟」,并让 batch_size 与 GPU 显存匹配。

把 serving 挂进监督树后,它可以被多个进程并发调用:

children = [
  {Nx.Serving, serving: serving, name: MyInference, partitions: true}
]

Supervisor.start_link(children, strategy: :one_for_one)

partitions: true 会按调度器数量启动多个副本,每个副本独立批处理,避免单进程成为串行瓶颈。这与 Erlang 的并发模型天然契合——推理服务只是监督树里的一个子进程,崩溃后由监督者重启,不影响 HTTP 层。

在 Phoenix/Plug 层暴露推理接口时,通常把 Nx.Serving 放进 Application 的 children,控制器直接调用:

def predict(conn, %{"features" => features}) do
  tensor = Nx.tensor(features, type: :f32)
  %{predictions: preds} = Nx.Serving.batched_run(MyInference, tensor)
  json(conn, %{result: Nx.to_flat_list(preds)})
end

性能陷阱与调度器阻塞

Nx 与 EXLA 的性能问题几乎都来自同一类根因:NIF 调用阻塞调度器。BEAM 的调度器是协作式的,NIF 执行期间不会主动让出 CPU。EXLA 的默认执行模式是「dirty NIF」(在脏调度器上跑),但图编译、显存分配等操作仍可能占用正常调度器,表现为其他进程延迟抖动。

诊断方法与 BEAM 性能调优 中的流程一致:用 recon 观察调度器利用率与脏调度器队列。

%% 查看脏 CPU 调度器负载
recon:scheduler_usage(1000).

%% 找出长耗时 NIF 调用
recon:trace(5000, msacc).

其他常见问题:

数据搬运开销。BEAM term 与 XLA buffer 之间每次转换都要拷贝。如果每个请求都做 Nx.tensor(list),拷贝成本可能超过计算本身。缓解方式是让张量在服务边界内保持为 Nx.Tensor,只在最终输出时转回 Elixir 数据结构。

类型提升导致的隐式 f64。如前所述,整数输入混进 f32 权重会让整个计算图升格。用 Nx.as_type/2 在入口统一类型。

编译缓存失效。defn 的编译产物按输入形状缓存。如果每次请求的 batch 维不同(例如动态 batch),EXLA 会反复编译。解决方法是固定 batch 维度(用 padding 补齐),或在 Nx.Serving 里让 batch 维固定。

内存不释放。XLA 的 buffer 由 NIF 管理,不受 BEAM GC 控制。长时间运行的推理服务需要观察 :erlang.memory(:total) 之外的进程外内存,用 :erlang.system_info(:allocated_areas) 或操作系统层面的 RSS 监控。

对于需要跨节点分摊推理负载的场景,可以把多台机器组成 Erlang 集群,由 :pg 或 Registry 做请求路由,这与 分布式推理集群 的架构思路一致,区别只是通信层从 gRPC 换成了 Erlang distribution。

与 Python 生态的关系

Nx 并不试图重写 CUDA kernel 或训练框架。它的定位是在 BEAM 内做「够用」的数值计算:特征工程、轻量模型推理、embedding 检索、数值预处理。真正的大规模训练仍然交给 Python 侧的框架,通过 ONNX 或 safetensors 导出权重,再由 Axon 加载。

跨语言的关键是权重格式。Axon 提供了与 PyTorch 命名习惯的映射工具,可以在 Python 侧把 state_dict 转成 Nx 的 params map:

# 从 PyTorch 导出的 safetensors 加载权重
params = Axon.load_weights("model.safetensors", model)

如果模型本身是 Python 训练的,可以用 Python 侧框架生成 ONNX,再用 Ortex(ONNX Runtime 绑定)执行,避免手工重写网络结构。选型原则:结构简单、需要与业务数据同进程 → 用 Nx/Axon;结构复杂、算子依赖 PyTorch → 用 Ortex 或独立推理服务。

Embedding 场景是 Nx 最典型的落地:把文本向量化后,用 Nx.dot/2 计算余弦相似度,在 ETS 里存向量做近邻检索。这套组合的全部逻辑都能用 OTP 表达,无需引入外部服务,细节可参考 向量嵌入实践 。

数据管道:从 ETS/Stream 到张量

数值计算很少凭空产生,输入通常来自数据库、ETS 或消息队列。把 Elixir 的集合与 Nx 的张量接起来,关键是批(batch)的构造时机。

逐条 Nx.tensor/1 再堆叠是最差的写法,因为它为每个样本单独分配一次 buffer:

# 反例:N 次分配 + N 次拷贝
samples |> Enum.map(&Nx.tensor(&1, type: :f32)) |> Nx.stack()

更好的做法是先收集成普通列表,一次性转换:

batch =
  :ets.foldl(
    fn {_k, vec}, acc -> [vec | acc] end,
    [],
    :feature_cache
  )

tensor = Nx.tensor(batch, type: :f32)

如果数据量超过内存,用 Stream.chunk_every/2 做流式分批,让每批独立进入计算图,批与批之间释放:

"data.csv"
|> File.stream!()
|> Stream.drop(1)
|> Stream.map(&parse_row/1)
|> Stream.chunk_every(1024)
|> Enum.each(fn chunk ->
  chunk
  |> Nx.tensor(type: :f32)
  |> MyMath.normalize()
  |> write_back()
end)

这里 Enum.each 是必要的——Stream 是惰性的,没有终止操作就不会执行。分批大小要与后端对齐:CPU 后端上 10244096 行较合适,GPU 上则受显存限制,通常 2561024。

对于需要跨批次共享的统计量(均值、方差),可以先用一次流式归约算出全局值,再在第二遍里应用,避免把全量数据驻留在内存中:

{sum, count} =
  data
  |> Stream.map(&Nx.sum(&1, axes: [0]))
  |> Enum.reduce({Nx.tensor(0.0), 0}, fn s, {acc, n} -> {Nx.add(acc, s), n + 1} end)

mean = Nx.divide(sum, count)

与 OTP 组件的协作

把 Nx 接进既有系统时,几个模式反复出现:

场景承载组件注意点
在线单样本推理Nx.Serving依赖 batch_timeout 控尾延迟
离线批量打分Task.async_stream/3限制 max_concurrency 避免 NIF 抢占
特征缓存ETS + Nx.to_binary/1存二进制比存 Tensor 更省内存
周期性重训练GenServer + :timer训练期间用 dirty 调度器隔离
向量检索ETS + Nx.dot/2批量算相似度,避免逐条比较

Task.async_stream/3 里跑 Nx 运算要特别小心:每个 task 是一个普通 BEAM 进程,而 EXLA 调用会占用调度器。把 max_concurrency 设得过高会让所有调度器都被 NIF 占住,HTTP 层随之失去响应。经验值是 System.schedulers_online() 的一半,并配合 timeout: :infinity 与 ordered: false。

特征缓存用 Nx.to_binary/1 存二进制而不是存 Nx.Tensor 结构体,是因为后者在 ETS 里保存的是指向 XLA buffer 的引用,跨进程共享时生命周期难以管理;二进制则是普通 BEAM term,GC 可正常回收。

常见错误与排查

症状可能原因排查手段
训练 loss 不下降最后一层 softmax + from_logits: false 重复激活打印首层输出范围
显存持续增长每请求重新构造张量,buffer 未复用监控进程外 RSS
其他进程延迟抖动NIF 占用正常调度器recon:scheduler_usage/1
首次请求慢数秒defn 编译缓存未命中预热:启动时跑一次 dummy 输入
数值结果与 Python 不一致类型提升到 f64 或算子语义差异Nx.type/1 逐层校验
吞吐上不去动态 batch 维导致反复编译固定 batch 维 + padding

预热(warm-up)是生产部署的标准动作:在 Application.start/2 之后、接受流量之前,用一个与线上同形状的 dummy 张量跑一遍推理,把 XLA 编译成本从首个真实请求转移到启动阶段。

def warmup do
  dummy = Nx.broadcast(Nx.tensor(0.0, type: :f32), {64, 784})
  Nx.Serving.run(MyInference, dummy)
  :ok
end

实践建议

  1. 先 BinaryBackend 跑通正确性,再切 EXLA 优化性能。两者的数值结果应当一致(浮点误差范围内),不一致说明用了后端相关的非确定算子。
  2. 在 defn 里显式标注类型。入口处 Nx.as_type/2,避免隐式 f64 升格吃掉一半吞吐。
  3. 用 Nx.Serving 而不是手写批处理。批处理的超时、拆批、异常隔离都有现成实现,手写极易在高并发下出错。
  4. 把推理服务当作普通 OTP 子进程。这样限流、熔断、可观测性、热更新全部可以复用既有基础设施。
  5. 监控进程外内存。XLA buffer 不在 BEAM 堆上,:erlang.memory/0 看不到,必须结合 RSS 与 EXLA 自身的分配统计。
  6. 固定 batch 维度。动态形状会破坏编译缓存,用 padding 换取稳定的编译产物,通常比省那点填充计算更划算。

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「erlang」更多文章

  1. BEAM 内存剖析与泄漏排查:recon、observer 与堆分析
  2. 分布式一致性与网络分区:CRDT、libcluster 与脑裂治理
  3. 缓存、限流与熔断:Cachex、Hammer 与降级策略