jax-llm-examples

纯JAX极简代码,跑通LLM训练与推理

jax-llm-examples 是一个用纯 JAX 实现的极简大语言模型示例集合,目标是让开发者以最少依赖理解并跑通 LLM 的训练与推理全流程。它不依赖 PyTorch 或 TensorFlow,只基于 JAX 及其生态(如 Flax/Optax 风格的手写模块),提供数据加载、模型定义、损失计算、优化器更新和文本生成等完整环节的清晰代码。项目强调性能与可读性的平衡:既展示 JAX 的 jit、vmap、pmap 等加速特性,又保持代码量小、结构直观,适合作为学习 JAX 做 LLM 的入门模板,也便于研究者快速改造实验。相比动辄数千行的训练框架,它把核心逻辑压缩到可通读的规模,帮助用户避开抽象层,直接观察张量流与梯度更新,是理解 LLM 底层实现和 JAX 并行能力的实用参考。

开源 free 模型训练
访问官网 ↗ GitHub ↗ 文档 ↗
GitHub 星标 ★ 281
维护状态 较活跃
是否开源
定价模式 free

项目数据

分类模型训练
开发团队jax-ml
所属国家
官网地址
定价模式free
价格说明开源免费,无在线服务
访问状态
是否开源
开源协议Apache-2.0
主要语言Python
技术栈/模型jax,llm,llm-inference
GitHub 星标★ 281
30天Star增速
HF 下载量
上线时间2025-02-06 00:00:00
最近更新2026-09-16 00:00:00
维护状态较活跃
中文支持
访问方式
移动端支持
综合评分
收录时间2026-09-16
浏览次数1

使用教程

核心亮点

  • 纯 JAX 实现,无 PyTorch/TensorFlow 依赖,代码量小可通读
  • 覆盖训练与推理全流程,含数据加载、模型、优化器和生成示例
  • 展示 jit/vmap/pmap 等 JAX 加速特性,便于学习并行与性能优化

不足之处

  • 项目处于早期,星标仅 278,生态与文档完善度有限
  • 示例规模偏小,缺少大规模分布式训练与生产部署的完整方案

适用场景

  • 学习 JAX 并想亲手实现 LLM 训练流程的开发者
  • 需要轻量模板快速验证 LLM 实验的研究者
  • 教学场景中演示 LLM 底层张量运算与梯度更新

替代项目

nanoGPT、minGPT

项目介绍

jax-llm-examples 是一个模型训练领域的开源项目,官方简介:Minimal yet performant LLM examples in pure JAX。项目使用 Python 开发,在 GitHub 上获得 278 星标。

上一篇:FATE-LLM

下一篇:ml-mdm

同类项目推荐

Long-RL 开源

让强化学习轻松驾驭超长序列,训练更稳更快

Long-RL: Scaling RL to Long Sequences (NeurIPS 2025)

★ 729 2026-08-09
TPA 开源

把注意力复杂度从平方降到线性,长序列不再卡顿

[NeurIPS 2025 Spotlight] TPA: Tensor ProducT ATTenTion Transformer (https://arxiv.or···

★ 463 2026-08-09
minimind 开源

两小时造出你的小模型,练手入门不求人

Train a 64M-parameter LLM from scratch in just 2h!

★ 61327 2026-08-09
pruna 开源

一键优化 AI 模型,推理提速又省资源。

Pruna is a model optimization framework built for developers, enabling you to delive···

★ 1279 2026-08-13