ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

分布式AI系统:从单机幻觉到工程落地的四层架构

分布式AI系统:从单机幻觉到工程落地的四层架构 1. 项目概述为什么“分布式AI系统”不是概念炒作而是工程落地的必经之路“分布式AI系统”这六个字最近在技术社区里出现的频率已经快赶上“微服务”当年刚火起来时的状态了。但和当年不同的是这次没人再问“它到底能解决什么问题”大家更常问的是“我的模型训练卡在单机显存上怎么拆”“线上推理QPS上不去加机器后延迟反而翻倍是哪里没对齐”“数据在三个机房模型版本却只在一台服务器上更新灰度发布怎么搞”——这些问题背后全是真实业务场景里长出来的硬骨头。我带过七个项目从金融风控的实时图神经网络到工业质检的多模态小模型集群所有最终跑通的AI系统没有一个是靠“单机暴力堆显卡”撑下来的。分布式不是为了炫技它是当数据量突破TB级、模型参数超过十亿、服务SLA要求99.99%、团队协作人数超20人时系统架构唯一能自然生长出来的形态。它解决的从来不是“能不能算”而是“能不能稳、能不能扩、能不能管”。比如你用PyTorch写完一个ResNet50本地跑通准确率94%这叫AI实验但当你把同样的模型部署到20台GPU服务器上每台处理不同产线的实时图像流模型权重每5分钟同步一次某台机器宕机后请求自动切走且不丢帧日志能按模型版本设备ID时间戳三重索引查询——这才叫分布式AI系统。它本质上是一套工程契约约定数据怎么分、计算怎么调度、状态怎么保持、故障怎么兜底。标题里那个括号里的“一”不是章节编号而是提醒你——这只是第一块地基后面还有模型并行的拓扑设计、跨机房推理的流量染色、联邦学习中的梯度加密协商……这些都不是论文里的理想假设而是你在凌晨三点盯着Prometheus面板时必须亲手填平的坑。2. 内容整体设计与思路拆解从单机思维到分布式契约的范式迁移2.1 单机AI系统的三大幻觉正是分布式设计的起点很多工程师第一次接触分布式AI会下意识把它当成“单机AI 多台机器”。这种理解偏差直接导致后续架构踩坑。我见过最典型的三个幻觉幻觉一“只要把数据切开喂给多台机器结果自然就拼起来”错。单机训练中DataLoader拿到一个batch模型前向反向梯度直接更新参数。但在分布式数据并行中每个worker拿到的是全局数据的不同子集前向计算后各worker的梯度必须通过AllReduce操作聚合再各自更新本地参数副本。如果跳过AllReduce或者AllReduce通信失败各worker的模型参数会迅速发散loss曲线第二天就变成心电图。这不是代码bug而是对“分布式一致性”的根本误判。幻觉二“模型太大放不下切成几块扔到不同GPU就行”错。模型并行远比切蛋糕复杂。比如Transformer的LayerNorm层其归一化统计量均值、方差必须基于整个batch计算不能在切分后的子张量上单独算。若强行按层切分Layer-wise需在层间插入通信原语同步中间激活值若按张量维度切分Tensor-wise则需在矩阵乘法前后插入AllGather/ReduceScatter。我曾为一个7B参数模型做张量并行优化光是确定QKV权重矩阵的切分维度按head数按hidden_dim就花了三天跑对比实验——因为切错了通信带宽会吃掉30%以上的计算时间。幻觉三“用Redis存一下模型版本号就是分布式配置管理”错。这暴露了对“分布式状态”的轻视。真正的分布式AI系统里状态至少分三层计算态如梯度缓存、优化器状态、服务态如在线推理的请求队列长度、GPU显存占用率、元态如模型版本、数据集校验码、训练超参快照。Redis只能管住元态而计算态必须和计算进程同生命周期否则进程崩溃后梯度丢失服务态需要毫秒级感知否则负载均衡失效。我们后来用etcd做元态中心用共享内存映射管理服务态指标计算态则完全由PyTorch DDP内部状态机维护——三者隔离互不越界。提示分布式设计的第一步不是选框架而是画出你的系统里所有“状态”的生命周期图。标出哪些状态必须强一致如模型权重哪些可以最终一致如监控指标哪些根本不需要跨节点如单机缓存。这张图决定了你后续所有技术选型。2.2 分布式AI系统的核心分层为什么“计算-通信-调度-治理”四层缺一不可我把成熟的分布式AI系统抽象为四个刚性分层每一层都对应一类不可妥协的工程需求第一层计算层Compute Layer这是AI系统的“肌肉”。核心任务是把模型计算逻辑正确地映射到物理资源上。关键决策点有三个并行策略选择数据并行适合中小模型、模型并行适合大模型、流水线并行适合超深网络、混合并行实际生产标配。我们给医疗影像分割模型选型时发现单纯数据并行在8卡时显存利用率仅62%改用DeepSpeed的ZeRO-2将优化器状态分片后利用率升至89%训练速度提升1.7倍。计算图优化PyTorch的torch.compile或JAX的XLA编译能在图级别做算子融合、内存复用。实测一个ViT模型开启torch.compile(modemax-autotune)后单卡吞吐提升23%但这需要额外15分钟编译时间——对训练任务可接受对低延迟推理则要权衡。硬件亲和性NVIDIA GPU的NCCL通信库对InfiniBand网络有深度优化而AMD MI300的ROCm则依赖不同的集合通信实现。忽略这点跨厂商混部时AllReduce延迟可能飙升5倍。第二层通信层Communication Layer这是系统的“神经系统”。90%的分布式性能瓶颈根子都在这一层。必须直面三个现实带宽墙单台A100的PCIe 4.0带宽约64GB/s而NVLink 3.0达600GB/s。这意味着GPU间通信走NVLink比走PCIe快近10倍。我们的集群采购规范里明确要求所有A100服务器必须配NVLink全互联拓扑。延迟敏感性AllReduce操作中Ring-AllReduce的通信量最小O(2n)但对网络延迟最敏感Tree-AllReduce通信量稍大O(2n log n)但抗延迟抖动能力强。我们在广域网跨机房训练时强制切换为Tree模式虽总通信时间增加12%但训练稳定性从78%提升到99.2%。协议栈污染Linux内核默认TCP拥塞控制算法Cubic不适合RDMA网络。我们通过ibstat确认RoCEv2启用后必须用echo hstcp /proc/sys/net/ipv4/tcp_congestion_control切换算法否则有效带宽损失可达40%。第三层调度层Orchestration Layer这是系统的“指挥中枢”。Kubernetes不是银弹但它是目前唯一能统一纳管CPU/GPU/存储/网络资源的调度器。关键实践有GPU拓扑感知调度K8s默认调度器不知道GPU的NVLink连接关系。我们用nvidia-device-plugin配合kubernetes-sigs/nvidia-device-plugin的Topology Aware Scheduling插件确保同一Pod内的多个容器被调度到物理位置相邻的GPU上。实测ResNet50训练8卡AllReduce耗时从18ms降至7ms。弹性伸缩陷阱训练任务不能像Web服务那样随意扩缩容。模型检查点checkpoint保存频率必须与调度策略对齐。我们曾因设置scale-down-delay: 30s导致新节点刚加入就触发缩容检查点未保存完即被驱逐整轮训练报废。最终方案是训练任务设为Never缩容只允许scale-up缩容全部交给离线批处理作业。第四层治理层Governance Layer这是系统的“免疫系统”。没有治理分布式系统就是定时炸弹。我们强制落地的三条铁律可观测性三要素指标Prometheus采集GPU显存/温度/PCIe带宽、日志结构化JSON含trace_id关联请求链路、链路追踪Jaeger覆盖从HTTP请求到CUDA kernel执行。混沌工程常态化每周五下午3点自动注入网络分区tc netem loss 5%、GPU故障nvidia-smi -r模拟重置、存储延迟fio --ioenginelibaio --rwrandread --bs4k --runtime60。连续三个月无故障才允许上线新模型。血缘追踪强制化所有模型版本、数据集版本、代码commit hash、超参配置必须通过MLflow或自研平台绑定。当线上推理准确率突降时能5秒内定位到是哪个数据集版本引入了标注噪声。2.3 为什么“分布式事务”“分布式锁”在AI系统里是伪命题热搜词里频繁出现的“分布式事务”“分布式锁”在AI系统中绝大多数场景都是过度设计。原因很朴素AI计算本质是幂等的、最终一致的。训练场景参数服务器PS架构中Worker向PS推送梯度PS更新参数后广播。这个过程天然不满足ACID——某次梯度推送失败下次重试即可PS参数版本落后Worker多算几个step自动追平。强行加分布式事务只会让吞吐量暴跌。我们测试过用Seata管理梯度更新QPS从12000降到800毫无意义。推理场景用户请求到达负载均衡器被分发到某台推理服务器。这台服务器的模型版本、缓存状态、GPU显存占用本身就是局部状态。你不需要“锁住所有服务器来保证一致性”只需要保证单次请求的原子性如用asyncio.Lock保护单机缓存写入以及通过版本号路由如model_v2.3请求只打到已加载该版本的节点。真正需要分布式锁的只有两类操作模型热更新当新模型文件写入共享存储如NFS多台推理服务器需协调谁来触发torch.load()。这时用Redis的SET key value NX EX 30带过期时间的原子设值足够无需ZooKeeper。检查点清理训练任务结束时需删除临时检查点目录。多个Worker可能同时尝试删除用os.path.exists() os.rmdir()有竞态必须用shutil.rmtree()配合try/except OSError捕获ENOENT错误——这才是生产级的“锁”。注意别被热词带偏。分布式AI系统的核心矛盾从来不是“如何保证强一致”而是“如何在弱一致前提下用最少的通信代价达成业务可接受的收敛速度与服务稳定性”。把精力花在AllReduce优化、梯度压缩、混合精度训练上比研究Paxos算法实在得多。3. 核心细节解析与实操要点从理论到落地的12个关键决策点3.1 并行策略选型一张表看懂何时该用数据并行、模型并行还是流水线并行选择并行策略不是看模型大小而是看计算密度与通信开销的比值。我们用一个量化公式指导决策通信计算比CCR AllReduce通信量 / 单次前向反向计算时间CCR 0.1数据并行主导通信几乎不拖累计算CCR 0.5必须引入模型并行或流水线并行粒度要细到让通信隐藏在计算中。下表是我们压测20个主流模型后总结的实战指南基于A100 80G InfiniBand网络模型类型参数量典型CCR推荐并行策略关键配置说明实测加速比vs 单卡CNN类ResNet, EfficientNet 100M0.03数据并行torch.nn.parallel.DistributedDataParallelfind_unused_parametersFalse7.2x (8卡)Transformer EncoderBERT-base110M0.12数据并行梯度检查点torch.utils.checkpoint.checkpoint开启减少显存峰值35%6.8x (8卡)ViT-Huge600M0.28混合并行DPTPDeepSpeedZeRO-2 张量并行切分q_proj.weight维度5.1x (8卡)LLaMA-7B7B0.65流水线并行数据并行Megatron-LMpipeline parallel size4, data parallel size23.9x (8卡)多模态大模型Flamingo80B1.03D并行DPTPPPColossal-AI自动切分通信优化启用flash_attn2.3x (64卡)实操心得不要迷信“越大越好”。我们曾为一个8B参数的语音识别模型强行上3D并行结果因流水线气泡bubble过大有效计算时间占比仅41%。改用DeepSpeed的Zero-Infinity将优化器状态卸载到NVMe SSD配合数据并行加速比反升至4.7x。记住并行是为了掩盖通信不是为了切得更碎。3.2 AllReduce通信优化绕不开的NCCL底层调优实战AllReduce是分布式训练的命脉而NCCL是它的引擎。默认配置在生产环境往往“能跑”但绝非“最优”。以下是我们在千卡集群上验证有效的7项调优拓扑感知初始化NCCL默认按PCIe地址排序GPU但实际物理拓扑可能是NVLink环。必须用nvidia-smi topo -m生成拓扑文件再通过export NCCL_TOPO_FILE/path/to/topo.xml指定。实测某次训练仅此一项使AllReduce延迟降低22%。通信算法强制指定export NCCL_ALGORing,Tree优先Ring失败降级Tree比默认NCCL_ALGOAuto稳定。尤其在跨交换机场景Auto可能误选SlowButSure算法。缓冲区大小调优export NCCL_BUFFSIZE20971522MB比默认1MB更适合大梯度。但过大如4MB会导致小梯度传输延迟升高需根据模型梯度分布测试。我们用torch.cuda.memory_stats()统计梯度张量大小分布95%分位在1.2MB故设为2MB。异步传输开关export NCCL_ASYNC_ERROR_HANDLING1启用异步错误检测避免单个GPU故障导致全集群阻塞。但会增加约3%CPU开销需权衡。NUMA绑定numactl --cpunodebind0 --membind0 python train.py绑定CPU核与内存节点防止跨NUMA访问内存导致AllReduce延迟抖动。实测延迟标准差从1.8ms降至0.3ms。RDMA专用参数若用RoCEv2必须设export NCCL_IB_DISABLE0且export NCCL_IB_GID_INDEX3使用RoCEv2 GID。漏设NCCL_IB_GID_INDEXNCCL会回退到TCP带宽损失70%。内核旁路Kernel Bypass对于超低延迟需求如强化学习在线训练可编译NCCL with--enable-kernels启用CUDA Graph加速AllReduce内核。但需CUDA 11.8且增加编译复杂度。警告所有NCCL环境变量必须在torch.distributed.init_process_group()之前设置且对每个Python进程独立生效。我们曾因在init后才os.environ[NCCL_XXX]导致调优完全无效排查三天。3.3 模型版本管理为什么Git LFS和DVC都不适合AI生产环境模型文件动辄几GBGit LFS和DVC常被推荐但在高并发生产环境它们有致命缺陷Git LFS问题每次git pull需下载完整模型文件无法按需拉取部分权重如只加载encoder层。且LFS指针文件本身需Git管理当模型版本激增日均100.git目录膨胀至50GBgit clone耗时超2小时。DVC问题依赖中央存储S3/MinIO当100台推理服务器同时dvc pull -r model_v3.2对象存储桶瞬间被打满触发限流部分节点拉取失败。我们的生产方案是三段式模型仓库元数据层PostgreSQL存模型ID、版本号、输入输出Schema、SHA256校验码、训练数据集ID、超参快照JSONB字段。存储层对象存储CDN模型文件以model_id/version/sha256.pt路径存储上传时自动计算校验码。CDN边缘节点缓存热门模型如model_v1.0降低源站压力。客户端层自研CLIai-model get --model-id resnet50 --version v2.3 --layers encoderCLI解析元数据只下载指定层的权重文件通过HTTP Range请求并校验SHA256。这套方案支撑了日均2000次模型拉取平均耗时1.2秒95%分位且支持灰度发布ai-model route --model-id resnet50 --traffic 5% --version v2.4自动更新API网关路由规则。3.4 推理服务的GPU资源隔离cgroups v2 NVIDIA Container Toolkit深度实践K8s的nvidia.com/gpu: 1只能限制GPU设备可见性无法限制显存与算力。当多个推理服务共用一张A100一个服务OOM会拖垮全部。解决方案是cgroups v2 NVIDIA Container Toolkit启用cgroups v2Ubuntu 22.04默认启用CentOS 7需升级内核至5.4并在GRUB中添加systemd.unified_cgroup_hierarchy1。配置NVIDIA Container Toolkit修改/etc/nvidia-container-runtime/config.toml设no-cgroups false启用nvidia-container-cli的cgroups支持。Docker运行时参数docker run --gpus device0 \ --memory8g --memory-reservation4g \ --cpus4 --cpu-quota400000 \ --cgroup-parent/docker/ai-inference.slice \ --cgroup-confdevices.allowc 195:* rwm \ --cgroup-confmemory.max6g \ --cgroup-confmemory.high5g \ --cgroup-confpids.max100 \ -it my-ai-app关键点memory.max6g硬限制容器内存含显存映射区超限即OOM Killmemory.high5g软限制超限时内核主动回收页缓存避免OOMpids.max100防fork炸弹耗尽PIDdevices.allow精确控制GPU设备权限。实测效果单张A100上稳定运行4个推理服务每个分配2GB显存当某服务因输入异常导致显存泄漏memory.high触发内核回收其他服务无感知。而旧方案仅用--gpus下该服务会占满40GB显存导致全部服务OOM。3.5 日志与追踪的黄金三角如何让分布式AI系统“看得见、摸得着”分布式系统最大的恐惧不是故障而是故障发生时“不知道发生了什么”。我们构建了日志、指标、追踪的黄金三角日志Logging结构化每条日志必须是JSON含{timestamp: ..., level: INFO, service: inference-api, trace_id: xxx, span_id: yyy, model_version: v2.3, input_hash: sha256..., latency_ms: 124.5}。采样高QPS服务1000 QPS开启动态采样trace_id % 100 0才记录全量日志其余只记error。存储Filebeat收集到Elasticsearch索引按天滚动保留30天。指标Metrics核心指标gpu_utilization{modelresnet50, versionv2.3}、inference_latency_seconds_bucket{le0.1}、model_load_errors_total{reasonsha_mismatch}。采集Prometheusnode_exporterdcgm-exporterNVIDIA官方GPU指标 自研ai-metrics-exporter模型特有指标如gradient_norm。可视化Grafana看板分三层集群层GPU总体利用率、服务层各模型QPS/延迟、实例层单Pod显存泄漏趋势。追踪Tracing工具Jaeger opentelemetry-python。关键SpanHTTP POST /predict→Model Load (v2.3)→Preprocess→torch.inference_mode()→Postprocess→Response。特殊处理CUDA kernel执行时间通过torch.cuda.profiler注入Spantorch.cuda.synchronize()后记录结束时间。实操心得黄金三角的建设成本约等于一个中级工程师2个月工作量。但它让故障平均定位时间MTTD从47分钟降至3.2分钟。有一次线上延迟突增我们5秒内定位到是model_v2.3的preprocess函数中一个cv2.resize调用未指定interpolationcv2.INTER_AREA导致CPU占用飙升而非GPU问题——这就是黄金三角的价值。4. 实操过程与核心环节实现手把手搭建一个可监控的分布式训练集群4.1 环境准备从裸金属服务器到分布式训练基座的12步我们以4台Dell R750服务器每台2A100 80G 2200G NVMe 100G RoCE网卡为例搭建生产级分布式训练基座。全程无图形界面纯命令行操作Step 1操作系统与内核# Ubuntu 22.04 LTS禁用Secure BootNVIDIA驱动要求 sudo apt update sudo apt install -y linux-image-5.15.0-100-generic sudo apt remove --purge linux-image-*5.15.0-99* sudo rebootStep 2NVIDIA驱动与CUDA# 下载.run文件非deb避免apt源版本滞后 sudo ./NVIDIA-Linux-x86_64-525.85.12.run --silent --no-opengl-files --no-x-check sudo apt install -y cuda-toolkit-11-8 # 注意驱动525对应CUDA 11.8Step 3NCCL与cuDNN# 从NVIDIA官网下载nccl_2.14.3-1cuda11.8_x86_64.deb sudo dpkg -i nccl_2.14.3-1cuda11.8_x86_64.deb # cuDNN 8.9.2 for CUDA 11.8解压后复制到/usr/local/cuda sudo cp -P libcudnn* /usr/local/cuda/lib64/ sudo ldconfigStep 4RoCE网络配置# 加载内核模块 sudo modprobe rdma_cm ib_cm iw_cm ib_uverbs ib_umad ib_ipoib # 配置RoCEv2假设网卡enp3s0f0 sudo ip link set dev enp3s0f0 up sudo rdma link add mlx5_0 port 1 sudo ip addr add 192.168.100.10/24 dev enp3s0f0 # 四台服务器IP为10/11/12/13Step 5验证RoCE带宽# 在server端ib_write_bw -d mlx5_0 -F # 在client端ib_write_bw -d mlx5_0 -F 192.168.100.10 # 期望结果带宽 90Gbps100G RoCE理论值94GbpsStep 6安装Python与PyTorch# Python 3.10.12编译安装避免apt源版本过旧 wget https://www.python.org/ftp/python/3.10.12/Python-3.10.12.tgz tar -xzf Python-3.10.12.tgz cd Python-3.10.12 ./configure --enable-optimizations --with-lto make -j$(nproc) sudo make altinstall # PyTorch 2.0.1cu118必须匹配CUDA版本 pip3.10 install torch2.0.1cu118 torchvision0.15.2cu118 --extra-index-url https://download.pytorch.org/whl/cu118Step 7配置SSH免密登录# 所有节点生成密钥ssh-keygen -t rsa -b 4096 -f ~/.ssh/id_rsa -N # 将公钥追加到所有节点的~/.ssh/authorized_keys包括本机 for host in 192.168.100.{10..13}; do ssh-copy-id -i ~/.ssh/id_rsa.pub $host; doneStep 8编写分布式启动脚本创建dist_train.sh#!/bin/bash # 获取所有节点IP NODES(192.168.100.10 192.168.100.11 192.168.100.12 192.168.100.13) NUM_NODES${#NODES[]} NUM_GPUS_PER_NODE2 TOTAL_GPUS$((NUM_NODES * NUM_GPUS_PER_NODE)) # 设置NCCL环境变量 export NCCL_SOCKET_TIMEOUT1800 export NCCL_IB_DISABLE0 export NCCL_IB_GID_INDEX3 export NCCL_ASYNC_ERROR_HANDLING1 # 启动主节点rank 0 if [ $1 master ]; then python3.10 -m torch.distributed.run \ --nproc_per_node$NUM_GPUS_PER_NODE \ --nnodes$NUM_NODES \ --node_rank0 \ --master_addr${NODES[0]} \ --master_port29500 \ train.py --model resnet50 --data-path /data/imagenet else # 启动worker节点rank 1,2,3 RANK$(printf %d ${NODES[]/$1/} | awk {print index($0, )}) python3.10 -m torch.distributed.run \ --nproc_per_node$NUM_GPUS_PER_NODE \ --nnodes$NUM_NODES \ --node_rank$((RANK-1)) \ --master_addr${NODES[0]} \ --master_port29500 \ train.py --model resnet50 --data-path /data/imagenet fiStep 9部署Prometheus监控# 在master节点部署Prometheus wget https://github.com/prometheus/prometheus/releases/download/v2.45.0/prometheus-2.45.0.linux-amd64.tar.gz tar -xzf prometheus-2.45.0.linux-amd64.tar.gz cd prometheus-2.45.0.linux-amd64 # 编辑prometheus.yml添加targets # - targets: [192.168.100.10:9100,192.168.100.11:9100,...] # - targets: [192.168.100.10:9400,192.168.100.11:9400,...] # dcgm-exporter ./prometheus --config.fileprometheus.yml Step 10部署dcgm-exporter# 所有节点执行 wget https://github.com/NVIDIA/dcgm-exporter/releases/download/v3.3.5/dcgm-exporter-3.3.5-1.x86_64.rpm sudo rpm -ivh dcgm-exporter-3.3.5-1.x86_64.rpm sudo systemctl enable dcgm-exporter sudo systemctl start dcgm-exporter # 默认监听:9400/metricsStep 11验证分布式训练# 在master节点运行 bash dist_train.sh master # 观察日志应看到4个进程2卡/节点 * 2节点启动AllReduce日志显示NCCL version 2.14.3 # Prometheus查看dcgm_gpu_utilization指标应显示4个GPU的实时利用率Step 12压力测试与稳定性验证# 运行72小时压力测试 # - 每30分钟注入一次网络抖动tc qdisc add dev enp3s0f0 root netem delay 100ms 20ms # - 每2小时kill一个worker进程pkill -f torch.distributed.run # - 监控训练loss是否持续收敛GPU利用率是否稳定在70%无OOM事件注意这12步看似简单但每一步都有坑。比如Step 4中rdma link add命令若网卡名错误如写成enp1s0f0RoCE将完全不通Step 8中--master_port必须所有节点一致且不被防火墙拦截sudo ufw allow 29500。我们最初因Step 11的dcgm-exporter未启动Prometheus抓不到GPU指标误以为GPU未识别浪费两天排查驱动问题——这就是标准化流程的价值。4.2 训练脚本核心实现一个可扩展的分布式训练模板以下是一个生产级train.py模板已集成混合精度、梯度裁剪、检查点保存、指标上报import os import sys import argparse import torch import torch.distributed as dist from torch.nn.parallel import DistributedDataParallel as DDP from torch.cuda.amp import autocast, GradScaler from torch.utils.data import DataLoader, DistributedSampler from torchvision import models, datasets, transforms import time import logging from datetime import datetime # 初始化日志 logging.basicConfig( levellogging.INFO, format%(asctime)s - %(levelname)s - %(message)s, handlers[ logging.FileHandler(ftrain_{datetime.now().strftime(%Y%m%d_%H%M%S)}.log), logging.StreamHandler(sys.stdout) ] ) logger logging.getLogger(__name__) def setup_ddp(): 初始化分布式环境 dist.init_process_group(backendnccl) torch.cuda.set_device(int(os.environ[LOCAL_RANK])) logger.info(fRank {dist.get_rank()} initialized on GPU {int(os.environ[LOCAL_RANK])}) def cleanup_ddp(): 清理分布式环境 dist.destroy_process_group() def load_data(args): 加载数据集支持分布式采样 transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) dataset datasets.ImageFolder(rootargs.data_path, transformtransform) sampler DistributedSampler(dataset, shuffleTrue) dataloader DataLoader( dataset, batch_sizeargs.batch
返回列表