dist-keras
基于Apache Spark的分布式Keras训练框架,支持Parameter Server架构下的异步SGD优化,让深度学习模型训练突破单机GPU显存限制
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
基于Apache Spark的分布式Keras训练框架,支持Parameter Server架构下的异步SGD优化,让深度学习模型训练突破单机GPU显存限制
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
2016年,当大多数深度学习项目还在单机GPU上运行时,比利时鲁汶大学的研究生 Joeri Hermans 在 CERN(欧洲核子研究中心)完成了他的硕士论文,并开源了 Distributed Keras(即 dist-keras)这一框架。这个项目的诞生源于一个现实痛点:传统的深度学习训练受限于单张GPU的显存和算力,面对 ATLAS 实验产生的大量希格斯玻色子碰撞数据时显得力不从心。
作者在 README 中坦承,项目的一大灵感来自 Google 的 DistBelief 论文《Large-Scale Distributed Deep Networks》。他希望构建一个让研究人员能够专注于算法创新,而非被分布式系统复杂性所困扰的框架。如今,这一项目已被 CERN 官方收录(cerndb/dist-keras),拥有 622 颗 GitHub stars,在工业界和学术界都有实际应用案例。
图1:CERN(欧洲核子研究中心)是本项目的重要背景支撑
Distributed Keras 采用了经典的 Parameter Server(参数服务器) 架构,这是工业界最成熟的分布式深度学习通信模式之一。其核心思想是:训练集群中存在若干 Parameter Server 节点,负责汇总所有 Worker 的梯度更新并维护全局模型参数;各 Worker 节点持有模型副本,在本地数据分片上独立进行前向传播和反向传播,定期将梯度发送给 Parameter Server。
具体实现中,项目提供了四种参数服务器变体:
| 类名 | 算法 | 特点 |
|---|---|---|
DeltaParameterServer | 异步 SGD(带梯度差值压缩) | 只传输梯度变化量,节省带宽 |
ADAGParameterServer | Adaptive Delta Adaptive Gradient | 自适应学习率,每个worker独立 |
DynSGDParameterServer | Dynamic SGD | 动态调整worker数量 |
ExperimentalParameterServer | 实验性 | 支持 EASGD 等前沿算法 |
每种参数服务器都继承自统一的抽象基类,通过 get_model() 方法同步模型权重给 Worker。这种设计使得添加新的分布式优化器变得极为简单——只需实现 Worker 端的梯度计算逻辑和 Parameter Server 端的参数聚合逻辑即可。
项目实现了多种分布式优化算法,每种都封装为独立的 Worker 类:
核心代码逻辑上,每个 Worker 维护一个本地模型副本,通过 socket 与 Parameter Server 通信。关键实现包括:
class Worker(object):
def __init__(self, model, optimizer, loss, features_col="features",
label_col="label", batch_size=32, num_epoch=1, learning_rate=1.0):
# 序列化 Keras 模型以便网络传输
self.model = serialize_keras_model(model)
# 与 Parameter Server 建立连接
self.socket = connect(determine_host_address(), port)
trainers.py 中的 Trainer 类是更高层的封装,提供了统一的训练接口。它自动处理 Parameter Server 与 Worker 之间的协调,暴露 train() 方法,接受 Spark RDD 数据源。用户无需关心底层的 socket 通信细节,只需配置优化器类型和超参数即可启动分布式训练。
setup.py 中声明的安装依赖为:
install_requires=['theano', 'tensorflow', 'keras', 'flask']
这是一个典型的2016年前后技术栈——当时 Keras 的后端还可以选择 Theano,而 TensorFlow 刚刚发布不久(v0.8 左右)。从 trainers.py 的 import 语句可以看出,项目对两个后端都有适配:from keras import backend as K。然而这种双框架支持在 2017 年 TensorFlow 2.0 推出后逐渐失去了维护意义。
项目基于 Apache Spark 的 RDD API 构建,利用 Spark 的集群管理能力调度分布式训练任务。具体来说,Training 的输入是一个 Spark RDD,每个元素代表一个数据分片。Trainer 在每个 worker 上实例化一个 Keras 模型副本,通过 fit() 方法在本地 RDD 分片上训练。
然而,项目使用的 Spark API 较为底层(直接操作 RDD),而非后期更流行的 DataFrame 或 Dataset API。这也是为什么它需要手动设置 SPARK_HOME 和 PYTHONPATH 环境变量的原因之一。
从 examples 目录可以看到,项目支持的数据格式包括:
distkeras/transformers.py 提供了数据转换工具,如 CSV 解析、特征标准化等。
安装方式非常标准:
pip install --upgrade dist-keras
# 或开发模式
git clone https://github.com/JoeriHermans/dist-keras
cd dist-keras
pip install -e .
关键环境配置(必须):
export SPARK_HOME=/usr/lib/spark
export PYTHONPATH="$SPARK_HOME/python/:$SPARK_HOME/python/lib/py4j-0.9-src.zip:$PYTHONPATH"
这些配置告诉 Python 如何找到 Spark 的 PySpark 库和 py4j 网关。如果不配置,运行时会直接报错 No module named pyspark。
以 SingleTrainer(基线)为例:
from distkeras.trainers import SingleTrainer
from distkeras.evaluators import AccuracyEvaluator
# 从 Spark RDD 加载数据
train_rdd = sqlContext.read.format('libsvm').load('data/train.libsvm')
test_rdd = sqlContext.read.format('libsvm').load('data/test.libsvm')
# 配置训练器
trainer = SingleTrainer(model=model, num_epoch=10, batch_size=32)
# 训练
trained_model = trainer.train(train_rdd)
# 评估
evaluator = AccuracyEvaluator()
score = evaluator.evaluate(trained_model, test_rdd)
分布式训练的差别仅在于将 SingleTrainer 替换为 DOWNPOURTrainer 或 ADAGTrainer,其余代码保持不变。
schemes.py 提供了自动超参数调优功能,支持:
from distkeras.schemes import Scheme
scheme = Scheme(optimizer=trainer, num_epoch=100, evaluation_frequency=5)
本项目不适合快速部署,原因如下:
import Queue 等写法暗示对 Python 2 有特殊处理从部署评分来看,本项目的 deployment_support_score = 1(最低档),属于纯研究型工具。
| 场景 | 是否适用 |
|---|---|
| 学术研究分布式优化算法 | ✅ 极佳(代码结构清晰,易扩展) |
| 企业级深度学习生产系统 | ❌ 不推荐(框架老旧,无生产级工具) |
| GPU 集群上的大规模训练 | ❌ 不推荐(建议用 Horovod、Ray Train) |
| 快速原型验证 | ❌ 不推荐(Spark 集群搭建成本高) |
本项目最后更新于 2018年7月,距今已超过7年。在这段时间里,分布式深度学习领域发生了巨大变化:
然而,dist-keras 的核心设计思想——Parameter Server 模式、异步 SGD 变种、陈旧度(Staleness)控制——至今仍是分布式深度学习的重要理论基础。特别是 ADAG 和 EASGD 算法在参数同步策略上的创新,对理解当前主流框架的设计思路有很好的入门价值。
从可读性角度评价,代码有如下特点:
trainers、workers、parameter_servers、schemes 各司其职object 基类 + 文档说明),便于子类扩展dev-dependencies 可以推断有配套测试,但 examples 中的 .ipynb 文件(而非 .py 测试)说明测试粒度较粗整体评估:代码质量较高,但受限于时代(Python 2/3 过渡期),现代 Python 最佳实践(如 type hints、dataclass、async/await)均未采用。
项目作者 Joeri Hermans 的硕士论文详细阐述了分布式深度学习的理论基础,包括参数陈旧度(Parameter Staleness)的量化定义和 ADAG 算法的设计。这篇论文至今在分布式深度学习领域仍被引用。
总结:Distributed Keras 是一个具有历史意义的学术项目,代表了 2016 年前后分布式深度学习的最先进实践。它虽然已被更新的框架超越,但其架构设计理念和代码实现仍是理解现代分布式 AI 系统的重要参考。对于研究 Parameter Server 架构和异步 SGD 变体的开发者,这个项目依然值得一读。