- 人工智能
- 深度学习
- 机器学习
- 强化学习
【免费下载链接】TensorLayer
Deep Learning and Reinforcement Learning Library for Scientists and Engineers
本指南以仓库 examples/database 目录下的三个脚本为核心,系统讲解如何利用 TensorLayer 内置的
tl.db.TensorHub数据库模块,在 MongoDB 上完成「数据集共享 → 任务分发 → 多机训练 → 结果回收 → 最佳模型选取」的完整闭环。读完本文,你将掌握dispatch_tasks.py、run_tasks.py、task_script.py三端脚本的分工与写法,并理解其底层基于 GridFS + MongoDB 文档索引的存储原理,可直接在本地多终端或 GPU 服务器集群上复现这套训练任务编排方案。
一、示例总览:三阶段任务编排流程
examples/database/README.md 把整个数据库任务编排流程概括为三个阶段,与之对应的正是该目录下的三个 Python 脚本:
- 分发阶段:
dispatch_tasks.py(分发端)创建 3 个携带不同超参数的任务(均指向task_script.py),并把一份 MNIST 数据集推入数据库; - 执行阶段:在 GPU 服务器上(本地测试时可另开一个终端)运行
run_tasks.py(运行端),它会持续轮询数据库、拉取并执行待处理任务,最后把训练好的模型与结果回存数据库; - 汇总阶段:所有任务完成后,分发端按准确率(
test_accuracy)从数据库中选出最佳模型。
三个脚本的角色定位非常清晰:
| 脚本 | 角色 | 核心职责 |
|---|---|---|
| dispatch_tasks.py | 分发端(dispatcher) | 清理旧数据、保存数据集、创建任务、等待完成、选出最优模型 |
| run_tasks.py | 运行端(runner) | 常驻轮询,拉取并执行 pending 任务 |
| task_script.py | 任务脚本(task script) | 加载数据集、训练网络、评估并回存模型与结果 |
此外,README 还点明了另外两条使用主线,本文后续会分别展开:
- 模型的保存与加载:
task_script.py演示如何保存模型,dispatch_tasks.py演示如何按测试准确率查找并加载最优模型; - 数据集的保存与加载:
dispatch_tasks.py演示如何保存数据集,task_script.py演示如何从数据库取回数据集。
二、环境准备:MongoDB 与依赖安装
TensorHub 的现有实现基于 MongoDB。安装 MongoDB 后,请确保 Python 侧安装了pymongo,仓库在 requirements/requirements_db.txt 中明确给出了版本要求:
pymongo>=3.8.0另外gridfs(GridFS 客户端)也是db.py运行时必需的(见 tensorlayer/db.py 的导入)。结合 docs/modules/db.rst 的说明,TensorLayer 数据库的设计目标是解决大规模机器学习项目中的数据管理问题,包括:从企业数据仓库中检索训练数据、加载单机存储放不下的大数据集、对模型做版本化管理与横向比较、自动化训练/评估/部署流程。它的存储系统分为两层:
- 索引层:基于 MongoDB 这类 NoSQL 文档数据库,存储所有标签(tag)与指向 blob 的引用;
- Blob 层:基于 GridFS(文件系统),以大数据块存储视频、医学图像、模型参数等大对象。
这一点在源码中也有直接体现:TensorHub.__init__中建立了两个 GridFS 文件桶datasetFilesystem(存数据集)与modelfs(存模型参数),见 tensorlayer/db.py。
三、分发端 dispatch_tasks.py 逐段解析
dispatch_tasks.py完整演示了「推数据、推任务、等结果、取最优」四步,下面按代码顺序拆解。
3.1 连接数据库
db = tl.db.TensorHub(ip='localhost', port=27017, dbname='temp', project_name='tutorial')TensorHub的完整构造函数签名及默认值如下(见 tensorlayer/db.py):
TensorHub(ip='localhost', port=27017, dbname='dbname', username='None', password='password', project_name=None)ip/port:MongoDB 地址与端口,27017是 MongoDB 默认端口;dbname:MongoDB 中的数据库名;username/password:认证信息,不需要认证时可置为None;project_name:整个实验项目的标识(类似 GitHub 上的仓库名),用于隔离不同项目的数据。源码中如果未显式指定,会默认取当前脚本文件名(sys.argv[0]去掉扩展名,见 tensorlayer/db.py)。
3.2 清理旧数据
db.delete_tasks() db.delete_model() db.delete_datasets()这三个调用分别清空当前project_name下所有任务、模型与数据集(底层是delete_many,见 tensorlayer/db.py)。在重复实验前先清理,可以避免上次运行残留的脏数据干扰。
3.3 保存数据集,供其他服务器共享
X_train, y_train, X_val, y_val, X_test, y_test = tl.files.load_mnist_dataset(shape=(-1, 784)) db.save_dataset((X_train, y_train, X_val, y_val, X_test, y_test), 'mnist', description='handwriting digit')tl.files.load_mnist_dataset(shape=(-1, 784))返回 6 元组:训练 / 验证 / 测试集的X与y(各 50000 / 10000 / 10000 条),shape默认(-1, 784),也可传(-1, 28, 28, 1)保留图像形状,见 mnist_dataset.py;save_dataset(dataset, dataset_name, **kwargs)把任意 Python 对象序列化后存入 GridFS,同时自动记录时间戳(datetime.utcnow())并写入db.Dataset集合。除dataset_name外,还可以传任意自定义标签(如description、version、author)方便后续检索,见 tensorlayer/db.py。
注意:仓库源码(tensorlayer/db.py)显示数据序列化使用的是
pickle.dumps(ps, protocol=pickle.HIGHEST_PROTOCOL),读取时反向pickle.loads。这意味着存入的数据在分发端与运行端之间需要保持 Python 环境兼容。
3.4 创建三个不同超参数的任务
db.create_task( task_name='mnist', script='task_script.py', hyper_parameters=dict(n_units1=800, n_units2=800), saved_result_keys=['test_accuracy'], description='800-800' ) db.create_task( task_name='mnist', script='task_script.py', hyper_parameters=dict(n_units1=600, n_units2=600), saved_result_keys=['test_accuracy'], description='600-600' ) db.create_task( task_name='mnist', script='task_script.py', hyper_parameters=dict(n_units1=400, n_units2=400), saved_result_keys=['test_accuracy'], description='400-400' )create_task的参数语义(见 tensorlayer/db.py):
| 参数 | 类型 | 含义 |
|---|---|---|
task_name | str | 任务名,运行端按它来匹配任务 |
script | str | 任务脚本的文件名,会被读入字节并随任务一起存入数据库 |
hyper_parameters | dict | 传入脚本的超参字典,运行端执行脚本前会注入为全局变量 |
saved_result_keys | list of str | 任务结束后需要从脚本全局作用域中回收的结果键名 |
**kwargs | - | 自定义附加信息,如description、版本号 |
每次create_task都会自动写入time时间戳,并把任务状态置为pending,result初始化为空字典。三份任务的差异只在n_units1/n_units2(MLP 两个隐藏层的神经元数),这正是典型的超参数搜索场景:同一个脚本、同一份数据,通过不同超参跑出多组结果,最后统一比较。
3.5 等待任务全部完成
while db.check_unfinished_task(task_name='mnist'): print("waiting runners to finish the tasks") time.sleep(1)check_unfinished_task会查询当前项目下状态为pending或running的任务,只要还有未完成任务就返回True(见 tensorlayer/db.py)。分发端用这个轮询循环阻塞等待,直到运行端把全部任务执行完毕。
3.6 按准确率选出最佳模型
net = db.find_top_model(model_name='mlp', sort=[("test_accuracy", -1)]) print("the best accuracy {} is from model {}".format(net._test_accuracy, net._name))find_top_model的sort参数直接透传给 PyMongo 的find_one排序。这里[("test_accuracy", -1)]表示按测试准确率降序取第一条,即准确率最高的模型。它返回的是重建出的 TensorLayerModel对象,随后示例通过net._test_accuracy与net._name读取模型文档中记录的指标与名字(见 tensorlayer/db.py 的find_top_model与 dispatch_tasks.py)。
sort的其他常用写法在源码 docstring 中也有说明(tensorlayer/db.py):
# 最新模型(按入库时间降序) net = db.find_top_model(sort=[("time", -1)]) # 最旧模型(按入库时间升序) net = db.find_top_model(sort=[("time", 1)])四、运行端 run_tasks.py:常驻的任务消费者
run_tasks.py的运行端逻辑非常简短,核心是一个无限轮询循环:
while True: print("waiting task from distributor") db.run_top_task(task_name='mnist', sort=[("time", -1)]) time.sleep(1)run_top_task是整套机制的执行核心,源码实现值得细看(tensorlayer/db.py):
- 用
find_one_and_update原子地把一条pending任务置为running(sort=[("time", -1)]表示优先取最新推送的任务),避免多台服务器同时抢到同一条任务; - 取出任务中的
hyper_parameters,逐个注入到globals()中——这就是为什么task_script.py里可以直接引用n_units1、n_units2而不需要显式定义; - 在
tf.Graph().as_default()上下文中exec(_script, globals())执行任务脚本; - 执行完成后把状态更新为
finished,并按照saved_result_keys从脚本的全局作用域中收集结果(如test_accuracy)写入任务的result字段。
在真实部署中,你可以在多台 GPU 服务器上各跑一个run_tasks.py,它们会共享同一 MongoDB 实例,自动负载均衡地消费队列中的任务。本地测试时,只需在分发端之外再开一个终端执行该脚本即可模拟「第二台机器」。
五、任务脚本 task_script.py:训练、评估与回存
task_script.py是被运行端动态执行的工作负载脚本,它展示了两件事:从数据库取数据集、把模型与结果存回数据库。
5.1 从数据库加载数据集
X_train, y_train, X_val, y_val, X_test, y_test = db.find_top_dataset('mnist')find_top_dataset(dataset_name)从 GridFS 反序列化出之前保存的 MNIST 六元组(见 tensorlayer/db.py)。如果同一名字存了多份(例如不同版本),可用find_datasets('mnist')一次取回全部列表;也可以用sort参数指定取最新或最旧的一份。
5.2 定义 MLP 并训练
def mlp(): ni = tl.layers.Input([None, 784], name='input') net = tl.layers.Dropout(keep=0.8, name='drop1')(ni) net = tl.layers.Dense(n_units=n_units1, act=tf.nn.relu, name='relu1')(net) net = tl.layers.Dropout(keep=0.5, name='drop2')(net) net = tl.layers.Dense(n_units=n_units2, act=tf.nn.relu, name='relu2')(net) net = tl.layers.Dropout(keep=0.5, name='drop3')(net) net = tl.layers.Dense(n_units=10, act=None, name='output')(net) M = tl.models.Model(inputs=ni, outputs=net) return M注意n_units1、n_units2这两个变量并没有在脚本里赋值——它们正是由运行端从数据库任务中取出并注入全局作用域的。这也是「任务 = 脚本 + 超参数」这一设计能够成立的关键。训练部分使用tl.utils.fit(batch_size=256、n_epoch=20、Adam 学习率0.0001),并以tl.utils.test计算测试准确率:
test_accuracy = tl.utils.test(network, acc, X_test, y_test, batch_size=None, cost=cost) test_accuracy = float(test_accuracy)5.3 把模型与结果存回数据库
db.save_model(network, model_name='mlp', name=str(n_units1) + '-' + str(n_units2), test_accuracy=test_accuracy)save_model的参数(见 tensorlayer/db.py):
network:TensorLayerModel实例;model_name:模型类别键(此处为'mlp',与分发端find_top_model(model_name='mlp', ...)对应);**kwargs:自定义事件字段,如name、accuracy、loss、step 数等,全部会随模型文档一起入库,用于后续筛选与排序。
其内部实现为:把network.all_weights(参数列表)用 pickle 序列化后写入 GridFS 桶modelfs,把network.config(网络结构配置)与datetime.utcnow()时间戳一起作为文档插入db.Model集合。而find_top_model读取时则执行逆过程:先从 GridFS 取回参数,再用static_graph2net依据结构配置重建网络(见 tensorlayer/files/utils.py),最后通过assign_weights把参数加载进重建的网络。
这样,网络架构与参数都存进了数据库,任何一台连接同一 MongoDB 的机器都能取出并直接使用,无需传递权重文件。README 中注释的备用写法也说明了加载方式:
# net = db.find_model(sess=sess, model_name=str(n_units1)+'-'+str(n_units2))六、TensorHub 核心 API 速查表
结合 tensorlayer/db.py 与 docs/modules/db.rst 的文档,TensorHub的常用方法汇总如下:
| 类别 | 方法 | 说明 |
|---|---|---|
| 数据集 | save_dataset(dataset, dataset_name, **kwargs) | 保存任意对象,自动加时间戳 |
find_top_dataset(dataset_name, sort=None, **kwargs) | 按条件取一条数据集 | |
find_datasets(dataset_name, **kwargs) | 取回所有匹配的数据集(列表) | |
delete_datasets() | 清空当前项目数据集 | |
| 模型 | save_model(network, model_name='model', **kwargs) | 保存架构 + 参数 + 自定义指标 |
find_top_model(sort=None, model_name='model', **kwargs) | 按条件与排序取回模型 | |
delete_model() | 清空当前项目模型 | |
| 任务 | create_task(task_name, script, hyper_parameters, saved_result_keys, **kwargs) | 推送任务 |
run_top_task(task_name, sort=None) | 拉取并执行一条 pending 任务 | |
check_unfinished_task(task_name) | 判断是否还有未完成任务 | |
delete_tasks() | 清空当前项目任务 | |
| 日志 | save_training_log / save_validation_log / save_testing_log | 记录训练 / 验证 / 测试指标 |
delete_training_log / delete_validation_log / delete_testing_log | 按条件删除日志 |
数据库中的实体遵循「一切皆数据、一切皆可用查询标识」的设计原则(详见 docs/modules/db.rst):数据集、模型架构、模型参数、任务、日志五类实体都可以打标签(如description、version、accuracy),检索时用查询语句 + 排序即可定位目标对象,无需改动应用代码。
七、运行步骤与注意事项
完整复现本示例的步骤:
- 安装并启动 MongoDB(默认端口
27017); - 安装依赖:
pymongo>=3.8.0(见 requirements/requirements_db.txt)以及 TensorLayer 本体; - 运行分发端:
python dispatch_tasks.py——它会推入数据集与 3 个任务,然后进入等待循环; - 另开终端运行运行端:
python run_tasks.py——它会轮询并依次执行 3 个任务(各训练 20 个 epoch 的 MLP),把模型与test_accuracy回存数据库; - 分发端检测到任务全部完成(
check_unfinished_task返回False)后,自动按test_accuracy降序取出最佳模型并打印其准确率与名字。
需要留意的限制(基于源码结构推断):
- 任务脚本由运行端通过
exec()在globals()作用域中执行,超参数通过全局变量注入,因此task_script.py应避免定义与超参数同名的局部变量; - 数据与模型参数均以 pickle 序列化存入 GridFS,分发端与运行端的 Python / TensorFlow 版本差异可能影响反序列化兼容性;
- 示例中的
project_name='tutorial'是实验隔离的关键:不同项目使用不同project_name即可在同一 MongoDB 中互不干扰地并行管理多组实验; - 分布式场景下,分发端与运行端只需保证能访问同一个 MongoDB 实例即可,
ip参数从'localhost'改为服务器地址即可跨机使用。
这套「数据库即任务队列」的方案,把数据集共享、超参搜索、模型版本管理与自动评估全部收敛到 MongoDB 一层,适合需要横向对比多个模型、或在多机间共享训练成果的工程化训练场景。更多细节可继续阅读 docs/modules/db.rst 以及 tensorlayer/db.py 中每个方法的 docstring。
- 人工智能
- 深度学习
- 机器学习
- 强化学习
【免费下载链接】TensorLayer
Deep Learning and Reinforcement Learning Library for Scientists and Engineers
相关推荐
Windows 7 SP2:让经典系统重获新生的终极解决方案
Windows 7 SP2:让经典系统重获新生的终极解决方案 你是否还在为Windows 7系统在新硬件上无法识别而烦恼?想象一下这样的场景:你刚买了一块高速N
人工智能深度学习机器学习强化学习Ludwig 分布式训练实战:使用 Ray Job Submission 在远程 Ray 集群上运行训练任务
Ludwig 分布式训练实战:使用 Ray Job Submission 在远程 Ray 集群上运行训练任务 本篇技术指南围绕 Ludwig 官方示例 exam
人工智能深度学习机器学习大模型预训练微调LoRA多模态NLP计算机视觉模型推理服务10分钟上手FluvioFX:Unity VFX流体模拟插件安装与配置全攻略
10分钟上手FluvioFX:Unity VFX流体模拟插件安装与配置全攻略 FluvioFX是一款专为Unity VFX Graph设计的流体动力学模拟插件,
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考