TensorLayer 数据库任务分发实战:用 TensorHub 在 MongoDB 上编排分布式训练任务

发布时间:2026/9/28 6:43:05
TensorLayer 数据库任务分发实战:用 TensorHub 在 MongoDB 上编排分布式训练任务 人工智能深度学习机器学习强化学习【免费下载链接】TensorLayerDeep Learning and Reinforcement Learning Library for Scientists and Engineers项目地址https://gitcode.com/gh_mirrors/te/TensorLayer点击查看免费下载本指南以仓库 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 中明确给出了版本要求pymongo3.8.0另外gridfsGridFS 客户端也是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(iplocalhost, port27017, dbnametemp, project_nametutorial)TensorHub的完整构造函数签名及默认值如下见 tensorlayer/db.pyTensorHub(iplocalhost, port27017, dbnamedbname, usernameNone, passwordpassword, project_nameNone)ip/portMongoDB 地址与端口27017是 MongoDB 默认端口dbnameMongoDB 中的数据库名username/password认证信息不需要认证时可置为Noneproject_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, descriptionhandwriting digit)tl.files.load_mnist_dataset(shape(-1, 784))返回 6 元组训练 / 验证 / 测试集的X与y各 50000 / 10000 / 10000 条shape默认(-1, 784)也可传(-1, 28, 28, 1)保留图像形状见 mnist_dataset.pysave_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, protocolpickle.HIGHEST_PROTOCOL)读取时反向pickle.loads。这意味着存入的数据在分发端与运行端之间需要保持 Python 环境兼容。3.4 创建三个不同超参数的任务db.create_task( task_namemnist, scripttask_script.py, hyper_parametersdict(n_units1800, n_units2800), saved_result_keys[test_accuracy], description800-800 ) db.create_task( task_namemnist, scripttask_script.py, hyper_parametersdict(n_units1600, n_units2600), saved_result_keys[test_accuracy], description600-600 ) db.create_task( task_namemnist, scripttask_script.py, hyper_parametersdict(n_units1400, n_units2400), saved_result_keys[test_accuracy], description400-400 )create_task的参数语义见 tensorlayer/db.py参数类型含义task_namestr任务名运行端按它来匹配任务scriptstr任务脚本的文件名会被读入字节并随任务一起存入数据库hyper_parametersdict传入脚本的超参字典运行端执行脚本前会注入为全局变量saved_result_keyslist of str任务结束后需要从脚本全局作用域中回收的结果键名**kwargs-自定义附加信息如description、版本号每次create_task都会自动写入time时间戳并把任务状态置为pendingresult初始化为空字典。三份任务的差异只在n_units1/n_units2MLP 两个隐藏层的神经元数这正是典型的超参数搜索场景同一个脚本、同一份数据通过不同超参跑出多组结果最后统一比较。3.5 等待任务全部完成while db.check_unfinished_task(task_namemnist): 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_namemlp, 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_namemnist, sort[(time, -1)]) time.sleep(1)run_top_task是整套机制的执行核心源码实现值得细看tensorlayer/db.py用find_one_and_update原子地把一条pending任务置为runningsort[(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], nameinput) net tl.layers.Dropout(keep0.8, namedrop1)(ni) net tl.layers.Dense(n_unitsn_units1, acttf.nn.relu, namerelu1)(net) net tl.layers.Dropout(keep0.5, namedrop2)(net) net tl.layers.Dense(n_unitsn_units2, acttf.nn.relu, namerelu2)(net) net tl.layers.Dropout(keep0.5, namedrop3)(net) net tl.layers.Dense(n_units10, actNone, nameoutput)(net) M tl.models.Model(inputsni, outputsnet) return M注意n_units1、n_units2这两个变量并没有在脚本里赋值——它们正是由运行端从数据库任务中取出并注入全局作用域的。这也是「任务 脚本 超参数」这一设计能够成立的关键。训练部分使用tl.utils.fitbatch_size256、n_epoch20、Adam 学习率0.0001并以tl.utils.test计算测试准确率test_accuracy tl.utils.test(network, acc, X_test, y_test, batch_sizeNone, costcost) test_accuracy float(test_accuracy)5.3 把模型与结果存回数据库db.save_model(network, model_namemlp, namestr(n_units1) - str(n_units2), test_accuracytest_accuracy)save_model的参数见 tensorlayer/db.pynetworkTensorLayerModel实例model_name模型类别键此处为mlp与分发端find_top_model(model_namemlp, ...)对应**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(sesssess, model_namestr(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, sortNone, **kwargs)按条件取一条数据集find_datasets(dataset_name, **kwargs)取回所有匹配的数据集列表delete_datasets()清空当前项目数据集模型save_model(network, model_namemodel, **kwargs)保存架构 参数 自定义指标find_top_model(sortNone, model_namemodel, **kwargs)按条件与排序取回模型delete_model()清空当前项目模型任务create_task(task_name, script, hyper_parameters, saved_result_keys, **kwargs)推送任务run_top_task(task_name, sortNone)拉取并执行一条 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安装依赖pymongo3.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_nametutorial是实验隔离的关键不同项目使用不同project_name即可在同一 MongoDB 中互不干扰地并行管理多组实验分布式场景下分发端与运行端只需保证能访问同一个 MongoDB 实例即可ip参数从localhost改为服务器地址即可跨机使用。这套「数据库即任务队列」的方案把数据集共享、超参搜索、模型版本管理与自动评估全部收敛到 MongoDB 一层适合需要横向对比多个模型、或在多机间共享训练成果的工程化训练场景。更多细节可继续阅读 docs/modules/db.rst 以及 tensorlayer/db.py 中每个方法的 docstring。赞分享人工智能深度学习机器学习强化学习【免费下载链接】TensorLayerDeep Learning and Reinforcement Learning Library for Scientists and Engineers项目地址https://gitcode.com/gh_mirrors/te/TensorLayer点击查看免费下载相关推荐Windows 7 SP2让经典系统重获新生的终极解决方案Windows 7 SP2让经典系统重获新生的终极解决方案 你是否还在为Windows 7系统在新硬件上无法识别而烦恼想象一下这样的场景你刚买了一块高速N人工智能深度学习机器学习强化学习Ludwig 分布式训练实战使用 Ray Job Submission 在远程 Ray 集群上运行训练任务Ludwig 分布式训练实战使用 Ray Job Submission 在远程 Ray 集群上运行训练任务 本篇技术指南围绕 Ludwig 官方示例 exam人工智能深度学习机器学习大模型预训练微调LoRA多模态NLP计算机视觉模型推理服务10分钟上手FluvioFXUnity VFX流体模拟插件安装与配置全攻略10分钟上手FluvioFXUnity VFX流体模拟插件安装与配置全攻略 FluvioFX是一款专为Unity VFX Graph设计的流体动力学模拟插件上一篇Hutool工具库中BeanUtil.copyProperties方法处理Map类型时的类型转换问题解析下一篇终极tldraw链接编辑指南超链接创建与智能URL验证完整教程创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考