ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

NetKet 并行计算完全指南:多 GPU 与多节点 HPC 集群实战

NetKet 并行计算完全指南:多 GPU 与多节点 HPC 集群实战 NetKet 并行计算完全指南多 GPU 与多节点 HPC 集群实战【免费下载链接】netketMachine learning algorithms for many-body quantum systems项目地址: https://gitcode.com/gh_mirrors/ne/netketNetKet 是一个面向多体量子系统的机器学习开源框架它把深度学习与量子物理紧密结合让研究者可以用神经网络求解海森堡模型、伊辛模型等复杂量子多体问题。当体系规模变大单张 GPU 已经力不从心NetKet 并行计算成为突破算力瓶颈的关键。本指南将带你快速掌握 NetKet 的多 GPU 与多节点 HPC 集群实战技巧从零开始跑通分布式量子变分蒙特卡洛VMC计算。NetKet 是什么量子计算研究的 AI 加速引擎 NetKetNetKet / Net-Ket是一款基于 JAX 构建的开源库核心能力包括能力模块典型功能模型库RBM、Jastrow、自动回归模型、Transformer 等神经网络波函数采样器MetropolisLocal、ExactSampler 等蒙特卡洛采样器优化器随机重配置SR、QGT、SGD 等 VMC 专用优化器驱动引擎nk.driver.VMC、nk.driver.VMC_SR等变分驱动它的分布式能力建立在 JAX 原生 sharding 机制之上旧版本中需要手动开启的NETKET_EXPERIMENTAL_SHARDING现在已默认启用上手门槛大大降低。相关模块可以在项目源码的netket/vqs/mc/、netket/sampler/、netket/models/目录中深入了解。为什么 NetKet 并行计算如此重要量子多体问题的计算量随体系尺寸指数增长单卡训练动辄数小时甚至数天。并行计算带来的收益是实实在在的✅更快的采样样本链可以均匀切分到多张 GPU 上并行采样✅更大的模型神经网络参数可以跨设备切分支撑更大规模网络✅更短的实验周期HPC 集群多节点协同让跑不完的任务变成睡一觉就好NetKet 并行计算的三种方式NetKet 官方提供了三套并行方案对应不同的硬件场景单机多 GPU零配置NetKet 自动使用所有可见 GPU最易上手多节点多 GPU面向 HPC 集群基于 JAX multi-controller 机制多 CPU MPICPU 集群专用实验特性通过 MPI 后端通信单机多 GPU开箱即用的默认并行方案如果你只在一台服务器上插了几张显卡那么最快配置方法是——什么都不用配置。NetKet 会自动检测并利用全部可见 GPUimport netket as nk import jax print(可用设备:, jax.devices()) print(本地设备:, jax.local_devices())想指定使用部分显卡在启动前设置环境变量即可# 只用 GPU 0 export CUDA_VISIBLE_DEVICES0 # 用 GPU 0 和 2 export CUDA_VISIBLE_DEVICES0,2更精细的数据放置如沿哪个轴切分数组可以通过jax.sharding.NamedSharding手动控制官方在netket/optimizer/qgt/qgt_onthefly.py等实现中展示了优秀的 sharding 实践。多节点多 GPUHPC 集群上的分布式计算方案多节点计算的关键一步是在脚本最开头、导入 netket 之前调用jax.distributed.initialize()import jax # 初始化分布式环境SLURM 下自动识别 jax.distributed.initialize() # 务必打印校验信息 print(f[{jax.process_index()}/{jax.process_count()}] devices:, jax.devices()) print(f[{jax.process_index()}/{jax.process_count()}] local devices:, jax.local_devices())校验要点jax.process_index()应输出[0/N]到[N-1/N]jax.local_devices()是每个进程可见的 GPUjax.devices()应为每节点 GPU 数 × 节点数。多 CPU MPICPU 集群的备选方案没有 GPU 的集群也无需放弃分布式。安装mpibackend4jax后在导入 JAX之前导入它import mpibackend4jax # 必须在 import jax 之前 import jax import netket as nk jax.distributed.initialize()启动方式使用传统 MPI 启动器mpirun -n 4 python your_script.py。注意 OpenMPI 5 不被支持macOS 建议使用 mpich。HPC 集群上的一键启动SLURM 作业脚本实战多节点计算在 SLURM 上最容易踩的坑是任务数与 GPU 数不匹配。JAX 假设每个 GPU 一个任务因此--ntasks-per-node必须等于每节点 GPU 数量同时不要使用--gpus-per-task与 JAX 自动检测冲突。一个典型的 2 节点 × 4 GPU 作业脚本#!/bin/bash #SBATCH --job-namenetket-distributed #SBATCH --nodes2 # 节点数 #SBATCH --ntasks-per-node4 # 必须等于每节点 GPU 数 #SBATCH --cpus-per-task5 #SBATCH --gresgpu:4 # 每节点 GPU 数 #SBATCH --time02:00:00 export JAX_PLATFORM_NAMEgpu srun uv run python your_netket_script.py 小贴士JAX 会自带安装配套 CUDA无需手动加载集群版 CUDA版本永远匹配且不会冲突。本地预测试技巧用 djaxrun 模拟多节点HPC 集群资源珍贵提交前先在本地验证脚本是明智之选。NetKet 自带的djaxrun工具源码见netket/tools/djaxrun.py可以模拟 SLURM 多进程环境# 用 2 个进程模拟 2 节点 djaxrun --simple -np 2 python3 your_script.py官方示例脚本Examples/Sharding/multi_process.py就是专门为这种预测试设计的包含完整的分布式 VMC 流程值得对照学习。并行计算中的常见陷阱与规避技巧 ⚠️分布式数组的安全打印与文件写入跨进程切分的数组sharded array不能直接print或随意保存。正确姿势是先复制到所有进程再操作replicated jax.lax.with_sharding_constraint( distributed_array, jax.sharding.NamedSharding(jax.sharding.get_abstract_mesh(), jax.P()), ) if jax.process_index() 0: print(replicated)文件写入同理NetKet 内置的 JsonLog 等日志器已自动处理分布式 I/O只有process_index() 0的进程真正落盘如果你自己写保存逻辑务必模仿这一模式避免多进程同时写文件造成冲突。死锁问题并行编程的第一大杀手JAX 的数组运算可能隐式触发全局通信类似 MPI_Allreduce。如果只在process_index() 0的分支里执行聚合操作其他进程永远不会执行对应通信就会永远卡死。规避法则很简单不要在jax.process_index() 0分支内对 JAX 数组做聚合运算分支内只操作已转成 NumPy 数组的数据发现jax.distributed.initialize()无输出时先检查节点间能否通信例如ping $(hostname)集群代理导致的通信故障某些集群的 HTTP 代理会导致 GRPC 通信异常no_proxy通配符如10.0.0.*GRPC 无法解析。最简单的处理是启动前清除代理环境变量import os for key in (http_proxy, https_proxy, no_proxy): os.environ.pop(key, None)快速上手的完整示例分布式 Ising 模型求解下面这段代码整合了前面所有要点可直接用于多节点集群import jax import netket as nk jax.distributed.initialize() print(f[{jax.process_index()}/{jax.process_count()}] 初始化完成, flushTrue) L 20 g nk.graph.Hypercube(lengthL, n_dim1, pbcTrue) hi nk.hilbert.Spin(s1/2, Ng.n_nodes) ha nk.operator.Ising(hilberthi, graphg, h1.0) model nk.models.RBM(alpha1) vs nk.vqs.MCState(samplernk.sampler.MetropolisLocal(hi), modelmodel, n_samples1024) gs nk.driver.VMC(ha, nk.optimizer.Sgd(learning_rate0.01), variational_statevs) gs.run(n_iter300, outdistributed_result) if jax.process_index() 0: print(优化完成最终能量:, gs.energy)运行后采样、能量估计和参数更新会自动在所有设备间协同完成你无需编写任何分布式逻辑。总结NetKet 并行计算的三大心法先本地后集群用djaxrun在本地模拟多进程验证无误再提交 SLURM先校验再训练启动后务必打印process_index/devices核对拓扑先防死锁再优化严格遵守进程 0 分支只处理 NumPy 数据的原则从单机多 GPU 的零配置体验到 HPC 集群的多节点协同NetKet 让量子多体计算真正走上了算力自由的道路。掌握这三板斧你的下一项大规模量子模拟实验就能稳定高效地跑起来 【免费下载链接】netketMachine learning algorithms for many-body quantum systems项目地址: https://gitcode.com/gh_mirrors/ne/netket创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表