用 Lisp 写 Jax 神经网络

老大鹏 专家 8小时前 更新于 2026年7月25日 497 浏览 8 点赞 约 2 分钟

把神经网络的数学表达用 Lisp 语法写出来,其实逻辑上极其顺畅。因为深度学习的数据流本质就是函数式编程:数据像流水一样经过一层层 Layer,每一层都是对前一层结果的组合。Jixp 就是基于这个逻辑搞的,它不是一个全新的框架,而是一个能把 Racket 风格的 DSL 编译成 Python 代码的编译器。

我研究了一下它的实现逻辑,最核心的价值在于它极大地简化了模型结构的描述。在 PyTorch 或原生 Jax 里,定义一个 Transformer Block 往往需要写一堆 __init__forward,而用 Jixp 只需要几行描述。

给个具体的代码片段,看看这种 DSL 是怎么定义 Transformer 结构的:

(let-dim
  (d 256)
  (heads 8)
  (define transformer-block (chain
    (residual (layernorm d)
              (attention d heads #:causal))
    (residual (layernorm d)
              (mlp d [4d] d #:bias)))))

这段代码被 Jixp 编译器处理后,会直接生成对应的 Python 代码,然后交给标准的 Python 训练栈去跑。这种做法把底层繁琐的实现细节给屏蔽掉了,你只需要关注模型本身的拓扑结构。

我在实际分析这个项目时发现几个关键点:

  • 编译链路: Lisp DSL → Jixp Compiler → Python → Jax。这意味着它不改变 Jax 的运行效率,因为最终跑的还是 Python/Jax。
  • 实测案例: 作者用这个工具在自己的 Obsidian 笔记库上练了一个 4M 参数的小模型做知识回溯。虽然规模极小,但证明了这套 DSL 定义模型并成功训练的可行性。
  • 潜在痛点: 这种方案目前最大的问题是调试链路太长。如果生成的 Python 代码报错,你得反推回 Lisp 定义里去找问题,对于习惯了直接在 Python 里打断点的人来说,学习成本较高。

从长远来看,这种用 DSL 描述模型的思路其实很前瞻。现在的模型权重文件(比如 .safetensors)只存了参数,没存结构。如果未来能有一种标准化的 DSL,一个模型定义文件可以同时被 vLLM、PyTorch 和各种自定义推理后端加载,那模型迁移的成本会低很多。

如果你想尝试部署,记得先配好 Jax 环境。一个简单的启动流程是:

1. 安装 Jixp 编译器环境。
2. 编写 .lisp 格式的模型定义文件。
3. 运行编译器生成 Python 脚本:

jixp compile model_def.lisp -o model_gen.py
4. 使用标准的 Python 训练循环加载 model_gen.py 中的类。

目前 Jixp 更像是一个极客的实验品,适合那些想快速原型化模型结构、或者本身就对函数式编程有执念的技术人。对于追求工业级稳定性的项目,直接写 Jax 依然是首选,但这种“解耦结构与实现”的思路绝对值得在设计 AI Agent 工作流或复杂模型时参考。

大模型LLM

全部回复 (2)

小阿伟的日常 初级 10小时前
以前试过用 Scheme 写简单的前向传播,逻辑确实比命令式语言清晰得多。
0 回复
数据分析师大山 中级 10小时前
这玩意儿编译到 Python 后,运行效率能跟原生 PyTorch 差不多吗?
0 回复

发表回复

支持 Markdown 格式