最新国产好看的视频,伊人天堂AV在线,国产Aaaaaa视频,蜜臀视频在线观看一区,人妻av色图,密臀久久久精品影片,青青视频免费观看毛片,久草在线观看视,国产三级精品色情在线

PyTorch分布式訓(xùn)練的實(shí)現(xiàn)

 更新時(shí)間:2026年01月30日 08:31:12   作者:盼小輝丶  
本文主要介紹了PyTorch分布式訓(xùn)練的實(shí)現(xiàn),文中通過示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來一起學(xué)習(xí)學(xué)習(xí)吧

0. 前言

在將預(yù)訓(xùn)練的機(jī)器學(xué)習(xí)模型投入生產(chǎn)環(huán)境之前,模型訓(xùn)練是不可或缺的關(guān)鍵環(huán)節(jié)。隨著深度學(xué)習(xí)的發(fā)展,大模型往往具有數(shù)百萬乃至數(shù)十億參數(shù)。使用反向傳播來調(diào)整這些參數(shù)需要大量的內(nèi)存和計(jì)算資源。即便如此,模型訓(xùn)練仍然可能需要數(shù)天甚至數(shù)月時(shí)間才能完成。
在本節(jié)中,我們將探討如何通過跨機(jī)器和機(jī)器內(nèi)多進(jìn)程的分布式訓(xùn)練來加速模型訓(xùn)練過程。我們將系統(tǒng)學(xué)習(xí) PyTorch 提供的三大分布式訓(xùn)練 API——torch.distributedtorch.multiprocessing 以及 torch.utils.data.distributed.DistributedSampler,使用這些 API 能夠極大的簡(jiǎn)化分布式訓(xùn)練,介紹如何使用 PyTorch 的分布式訓(xùn)練工具,在 CPUGPU 上加速訓(xùn)練。
通過本節(jié)學(xué)習(xí),將能夠充分釋放硬件設(shè)備的訓(xùn)練潛力。對(duì)于超大規(guī)模模型訓(xùn)練而言,本節(jié)所探討的工具不僅至關(guān)重要,在某些情況下甚至是不可或缺的。

1. 使用 PyTorch 進(jìn)行分布式訓(xùn)練

在本節(jié)中,我們將模型訓(xùn)練過程從常規(guī)訓(xùn)練轉(zhuǎn)換為分布式訓(xùn)練,探討 PyTorch 提供的分布式訓(xùn)練工具,這些工具能顯著提升訓(xùn)練速度并優(yōu)化硬件使用效率。

1.1 以常規(guī)方式訓(xùn)練模型

(1) 首先導(dǎo)入所需庫:

import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

import time
import argparse

device = torch.device("cpu")

(2) 接下來,定義卷積神經(jīng)網(wǎng)絡(luò) (Convolutional Neural Network, CNN) 模型架構(gòu):

class ConvNet(nn.Module):
    def __init__(self):
        super(ConvNet, self).__init__()
        self.cn1 = nn.Conv2d(1, 16, 3, 1)
        self.cn2 = nn.Conv2d(16, 32, 3, 1)
        self.dp1 = nn.Dropout(0.10)
        self.dp2 = nn.Dropout(0.25)
        self.fc1 = nn.Linear(4608, 64) # 4608 is basically 12 X 12 X 32
        self.fc2 = nn.Linear(64, 10)
 
    def forward(self, x):
        x = self.cn1(x)
        x = F.relu(x)
        x = self.cn2(x)
        x = F.relu(x)
        x = F.max_pool2d(x, 2)
        x = self.dp1(x)
        x = torch.flatten(x, 1)
        x = self.fc1(x)
        x = F.relu(x)
        x = self.dp2(x)
        x = self.fc2(x)
        op = F.log_softmax(x, dim=1)
        return op

(3) 然后,定義模型的訓(xùn)練過程:

def train(args):
    train_dataloader = torch.utils.data.DataLoader(
        datasets.MNIST('./data', train=True, download=True,
                       transform=transforms.Compose([
                           transforms.ToTensor(),
                           transforms.Normalize((0.1302,), (0.3069,))])),
        batch_size=128, shuffle=True)  
    model = ConvNet()
    optimizer = optim.Adadelta(model.parameters(), lr=0.5)
    model.train()

在函數(shù)的前半部分,使用 PyTorch 訓(xùn)練數(shù)據(jù)集定義了 PyTorch 的訓(xùn)練數(shù)據(jù)加載器。實(shí)例化卷積神經(jīng)網(wǎng)絡(luò) (ConvNet),并定義了優(yōu)化器。

    for epoch in range(args.epochs):
        for b_i, (X, y) in enumerate(train_dataloader):
            X, y = X.to(device), y.to(device)
            pred_prob = model(X)
            loss = F.nll_loss(pred_prob, y) # nll is the negative likelihood loss
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            if b_i % 10 == 0:
                print('epoch: {} [{}/{} ({:.0f}%)]\t training loss: {:.6f}'.format(
                    epoch, b_i, len(train_dataloader),
                    100. * b_i / len(train_dataloader), loss.item()))

在函數(shù)的后半部分,運(yùn)行訓(xùn)練循環(huán)預(yù)定義的 epoch 數(shù)。在循環(huán)內(nèi),通過批數(shù)據(jù)的方式遍歷整個(gè)訓(xùn)練數(shù)據(jù)集,本節(jié)中批大小為 128。對(duì)于每個(gè)包含 128 個(gè)訓(xùn)練數(shù)據(jù)點(diǎn)的批次,使用模型進(jìn)行前向傳播,以計(jì)算預(yù)測(cè)概率。然后,我們將預(yù)測(cè)結(jié)果結(jié)合真實(shí)標(biāo)簽計(jì)算批次損失,并通過反向傳播利用該損失梯度來調(diào)整模型參數(shù)。

(4) 將所有組件整合在 main() 函數(shù)中:

def main():
    parser = argparse.ArgumentParser()
    parser.add_argument('--epochs', default=1, type=int)
    args = parser.parse_args()
    start = time.time()
    train(args)
    print(f"Finished training in {time.time()-start} secs")

使用參數(shù)解析器,它可以幫助我們?cè)谶\(yùn)行 Python 訓(xùn)練程序時(shí)從命令行輸入超參數(shù),例如 epoch 數(shù)。我們還對(duì)訓(xùn)練過程進(jìn)行了計(jì)時(shí),以便可以將它與分布式訓(xùn)練過程進(jìn)行比較。

(5) 最后,確保通過命令行執(zhí)行腳本時(shí)能運(yùn)行 main() 函數(shù):

if __name__ == '__main__':
    main()

(6) 在命令行中執(zhí)行以下命令來運(yùn)行該 Python 腳本:

$ python convnet_undistributed.py -- epoches 1

本節(jié)中我們僅設(shè)置訓(xùn)練一個(gè) epoch,因?yàn)楫?dāng)前重點(diǎn)不在于模型精度,而在于模型訓(xùn)練耗時(shí)??梢钥吹捷敵鼋Y(jié)果如下所示:

訓(xùn)練 1 個(gè) epoch 大約花費(fèi)了 28 秒,一個(gè) epoch 共包含 469 個(gè)批次,實(shí)際訓(xùn)練時(shí)間會(huì)隨硬件配置差異而波動(dòng)。

1.2 分布式訓(xùn)練模型

通過使用 PyTorch 提供的分布式處理 API,即使存在跨進(jìn)程或跨機(jī)器重復(fù)傳遞數(shù)據(jù)的額外開銷,模型訓(xùn)練速度也能顯著提升。

(1) 首先,導(dǎo)入所需庫,增加幾個(gè)與分布式訓(xùn)練相關(guān)的模塊:

import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms

import torch.multiprocessing as mp
import torch.distributed as dist

import os
import time
import argparse

torch.multiprocessing 用于在單臺(tái)機(jī)器上生成多個(gè) Python 進(jìn)程(通常根據(jù) CPU 核心數(shù)生成對(duì)應(yīng)數(shù)量的進(jìn)程),而torch.distributed則實(shí)現(xiàn)不同機(jī)器間的通信協(xié)作,使它們能共同完成模型訓(xùn)練。執(zhí)行時(shí),我們需要在每臺(tái)參與訓(xùn)練的機(jī)器上顯式啟動(dòng)訓(xùn)練腳本。
PyTorch 內(nèi)置的通信后端(如 Gloo )會(huì)自動(dòng)處理機(jī)器間的通信協(xié)調(diào)。在每臺(tái)機(jī)器內(nèi)部,多進(jìn)程機(jī)制會(huì)進(jìn)一步將訓(xùn)練任務(wù)并行分配到各個(gè)進(jìn)程。

(2) 模型架構(gòu)定義部分保持不變:

class ConvNet(nn.Module):
    def __init__(self):
        super(ConvNet, self).__init__()
        self.cn1 = nn.Conv2d(1, 16, 3, 1)
        self.cn2 = nn.Conv2d(16, 32, 3, 1)
        self.dp1 = nn.Dropout2d(0.10)
        self.dp2 = nn.Dropout2d(0.25)
        self.fc1 = nn.Linear(4608, 64) # 4608 is basically 12 X 12 X 32
        self.fc2 = nn.Linear(64, 10)
 
    def forward(self, x):
        x = self.cn1(x)
        x = F.relu(x)
        x = self.cn2(x)
        x = F.relu(x)
        x = F.max_pool2d(x, 2)
        x = self.dp1(x)
        x = torch.flatten(x, 1)
        x = self.fc1(x)
        x = F.relu(x)
        x = self.dp2(x)
        x = self.fc2(x)
        op = F.log_softmax(x, dim=1)
        return op

(3) 定義 train() 函數(shù):

def train(cpu_num, args):
    rank = args.machine_id * args.num_processes + cpu_num                        
    dist.init_process_group(                                   
    backend='gloo',                                         
    init_method='env://',                                   
    world_size=args.world_size,                              
    rank=rank                                               
    ) 
    torch.manual_seed(0)
    device = torch.device("cpu")

可以看到,代碼開頭新增了兩條關(guān)鍵語句。首先是計(jì)算進(jìn)程的 rank 值——這本質(zhì)上是該進(jìn)程在整個(gè)分布式系統(tǒng)中的順序標(biāo)識(shí)符。舉例來說,若使用 2 臺(tái)各配備 4CPU 的機(jī)器進(jìn)行訓(xùn)練,為充分利用硬件資源可能需要啟動(dòng) 8 個(gè)進(jìn)程(每臺(tái)機(jī)器 4 個(gè))。此時(shí)就需要為這些進(jìn)程建立標(biāo)識(shí)體系:先為兩臺(tái)機(jī)器分配 ID 01,再為每臺(tái)機(jī)器內(nèi)的 4 個(gè)進(jìn)程分配子 ID 03。最終,第 n 臺(tái)機(jī)器上第 i 個(gè)進(jìn)程的全局 rank 值可通過以下公式確定:
r a n k = n × 4 + k rank=n\times 4+k rank=n×4+k
第二行代碼使用了 torch.distributed 模塊中的 init_process_group,該方法為每個(gè)啟動(dòng)的進(jìn)程配置以下關(guān)鍵參數(shù):

  • 用于機(jī)器間通信的后端(在節(jié)使用 Gloo)
  • 參與分布式訓(xùn)練的進(jìn)程總量(由 args.world_size 指定),亦稱 world_size
  • 當(dāng)前啟動(dòng)進(jìn)程的全局 rank

init_process_group 方法會(huì)阻塞所有進(jìn)程,直到跨機(jī)器的全部進(jìn)程都完成初始化才會(huì)繼續(xù)執(zhí)行。
PyTorch 提供了三種內(nèi)置的分布式訓(xùn)練后端:

  • Gloo
  • NCCL
  • MPI

簡(jiǎn)而言之,對(duì)于 CPU 上的分布式訓(xùn)練,使用 Gloo,對(duì)于 GPU,使用 NCCL。

    train_dataset = datasets.MNIST('./data', train=True, download=True,
                                   transform=transforms.Compose([
                                       transforms.ToTensor(),
                                       transforms.Normalize((0.1302,), (0.3069,))]))  
    train_sampler = torch.utils.data.distributed.DistributedSampler(
        train_dataset,
        num_replicas=args.world_size,
        rank=rank
    )
    train_dataloader = torch.utils.data.DataLoader(
       dataset=train_dataset,
       batch_size=args.batch_size,
       shuffle=False,            
       num_workers=0,
       sampler=train_sampler)
    model = ConvNet()
    optimizer = optim.Adadelta(model.parameters(), lr=0.5)
    model = nn.parallel.DistributedDataParallel(model)
    model.train()

與單機(jī)訓(xùn)練相比,分布式訓(xùn)練的關(guān)鍵改進(jìn)體現(xiàn)在數(shù)據(jù)加載與模型封裝兩個(gè)層面。我們將 MNIST 數(shù)據(jù)集實(shí)例化與數(shù)據(jù)加載器拆分為獨(dú)立步驟,其間插入 DistributedSampler 采樣器。該采樣器將訓(xùn)練數(shù)據(jù)均分為 world_size 個(gè)分區(qū),確保每個(gè)進(jìn)程處理等量數(shù)據(jù)。注意數(shù)據(jù)加載器的 shuffle 參數(shù)需設(shè)為 False,因?yàn)閿?shù)據(jù)分配已由采樣器控制。
代碼中的另一個(gè)新增部分是 nn.parallel.DistributedDataParallel 函數(shù),它應(yīng)用于模型對(duì)象。這部分可能是代碼中最重要的部分,因?yàn)?DistributedDataParallel 是實(shí)現(xiàn)分布式梯度下降算法的關(guān)鍵組件。其底層運(yùn)行機(jī)制如下:

  • 分布式環(huán)境中的每個(gè)派生進(jìn)程都會(huì)獲得獨(dú)立的模型副本
  • 每個(gè)進(jìn)程的模型都維護(hù)自己的優(yōu)化器,并與全局迭代保持同步的局部?jī)?yōu)化步驟
  • 在每次分布式訓(xùn)練迭代時(shí),各進(jìn)程獨(dú)立計(jì)算損失值及梯度,隨后跨進(jìn)程對(duì)這些梯度求取平均值
  • 平均后的梯度將通過全局反向傳播機(jī)制同步到所有模型副本,用于調(diào)整參數(shù)
  • 由于全局反向傳播步驟的存在,所有模型參數(shù)在每次迭代時(shí)都保持一致,從而實(shí)現(xiàn)自動(dòng)同步

DistributedDataParallel 通過讓每個(gè) Python 進(jìn)程運(yùn)行在獨(dú)立的解釋器上,有效規(guī)避了在單一解釋器下多線程實(shí)例化多個(gè)模型可能引發(fā)的全局解釋器鎖 (Global Interpreter Lock, GIL) 限制問題。這進(jìn)一步提升了性能表現(xiàn),特別是對(duì)于那些需要大量Python專屬運(yùn)算的模型而言。

    for epoch in range(args.epochs):
        for b_i, (X, y) in enumerate(train_dataloader):
            X, y = X.to(device), y.to(device)
            pred_prob = model(X)
            loss = F.nll_loss(pred_prob, y) # nll is the negative likelihood loss
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            if b_i % 10 == 0 and cpu_num==0:
                print('epoch: {} [{}/{} ({:.0f}%)]\t training loss: {:.6f}'.format(
                    epoch, b_i, len(train_dataloader),
                    100. * b_i / len(train_dataloader), loss.item()))

最后,訓(xùn)練循環(huán)幾乎和單機(jī)訓(xùn)練一樣。唯一的區(qū)別在于我們限制只有排名為0的進(jìn)程才能獲取日志信息。這是因?yàn)榕琶麨?0 的機(jī)器用于建立所有通信連接。因此,我們通常將排名為 0 的進(jìn)程作為參考來跟蹤模型訓(xùn)練性能。如果不加以限制,每個(gè)模型訓(xùn)練迭代都會(huì)產(chǎn)生與進(jìn)程數(shù)量相同的日志行數(shù)。

(4) 將所有組件整合在 main() 函數(shù)中:

def main():
    parser = argparse.ArgumentParser()
    parser.add_argument('--num-machines', default=1, type=int,)
    parser.add_argument('--num-processes', default=1, type=int)
    parser.add_argument('--machine-id', default=0, type=int)
    parser.add_argument('--epochs', default=1, type=int)
    parser.add_argument('--batch-size', default=128, type=int)
    args = parser.parse_args()
    
    args.world_size = args.num_processes * args.num_machines                
    os.environ['MASTER_ADDR'] = '127.0.0.1'              
    os.environ['MASTER_PORT'] = '8892'      
    start = time.time()
    mp.spawn(train, nprocs=args.num_processes, args=(args,))
    print(f"Finished training in {time.time()-start} secs")

首先,我們新增了以下參數(shù):

  • num_machines:機(jī)器總數(shù)量
  • num_processes:每臺(tái)機(jī)器上要啟動(dòng)的進(jìn)程數(shù)量
  • machine_id:當(dāng)前機(jī)器的序號(hào) ID。 需要注意的是,這個(gè) Python 腳本需要在每臺(tái)機(jī)器上單獨(dú)啟動(dòng)
  • batch_size:每個(gè)批次的數(shù)據(jù)樣數(shù)量。參數(shù)作用如下:
    • 所有進(jìn)程將各自計(jì)算梯度,這些梯度將在每次迭代中被平均,從而獲得整體梯度
    • 完整的訓(xùn)練數(shù)據(jù)集被分割為 world_size 個(gè)獨(dú)立的數(shù)據(jù)集

因此,在每次迭代時(shí),完整的批次數(shù)據(jù)需要被分割成 world_size 個(gè)子批次,每個(gè)進(jìn)程處理一個(gè)子批次。因?yàn)?batch_size 現(xiàn)在與 world_size 相關(guān)聯(lián),所以我們將其作為輸入?yún)?shù)提供,目的是為了簡(jiǎn)化訓(xùn)練接口。
參數(shù)定義后,計(jì)算 world_size 作為派生參數(shù)。接著,我們定義兩個(gè)重要的環(huán)境變量:

  • MASTER_ADDR:運(yùn)行 rank 0 進(jìn)程的主機(jī) IP 地址
  • MASTER_PORT:運(yùn)行 rank 0 進(jìn)程的主機(jī)上可用的端口號(hào)

rank 0 機(jī)器負(fù)責(zé)建立所有后端通信連接,因此整個(gè)系統(tǒng)必須能隨時(shí)定位到該主機(jī),這就是為什么需要提供其 IP 地址和端口號(hào)。本節(jié)中訓(xùn)練任務(wù)將在單臺(tái)本地機(jī)器上運(yùn)行,因此使用 localhost 地址即可,但在跨服務(wù)器的多機(jī)訓(xùn)練場(chǎng)景中,則需要提供 rank 0 服務(wù)器的真實(shí) IP 地址及空閑端口號(hào)。
最后一個(gè)變化是使用多進(jìn)程 (multiprocessing) 來在每臺(tái)機(jī)器上啟動(dòng) num_processes 個(gè)進(jìn)程,而非僅運(yùn)行單個(gè)訓(xùn)練進(jìn)程。分布式參數(shù)會(huì)傳遞給每個(gè)派生進(jìn)程,確保模型訓(xùn)練過程中各進(jìn)程與機(jī)器之間能自主協(xié)調(diào)。

(5) 分布式訓(xùn)練:

if __name__ == '__main__':
    main()

(6) 啟動(dòng)分布式訓(xùn)練腳本。首先使用分布式腳本進(jìn)行類非分布式運(yùn)行,將機(jī)器數(shù)量和進(jìn)程數(shù)量都設(shè)置為 1

$ python convnet_distributed.py --num-machines 1 --num-processes 1 --machine-id 0 --epochs 1 --batch-size 128

需要注意的是,由于本次訓(xùn)練只使用單個(gè)進(jìn)程,batch_size 與之前非分布式訓(xùn)練時(shí)保持一致(仍為 128)。運(yùn)行結(jié)果如下輸出:

若將此結(jié)果與上一節(jié)非分布式訓(xùn)練的輸出對(duì)比,可發(fā)現(xiàn)訓(xùn)練時(shí)間基本相當(dāng)(約 30 秒),損失值變化也十分相似。

(7) 接下來,運(yùn)行一個(gè)真正的分布式訓(xùn)練,使用 2 個(gè)進(jìn)程而不是 1 個(gè)進(jìn)程。相應(yīng)地,將 batch_size128 降至 64

$ python convnet_distributed.py --num-machines 1 --num-processes 2 --machine-id 0 --epochs 1 --batch-size 64

輸出結(jié)果如下所示:

可以看到,訓(xùn)練時(shí)間從 30 秒減少到了 20 秒。訓(xùn)練損失的變化趨勢(shì)沒有受到影響,這表明分布式訓(xùn)練可以加速訓(xùn)練過程,同時(shí)保持模型的準(zhǔn)確性。

(8) 接下來,使用 4 個(gè)進(jìn)程,并相應(yīng)地將批大小從 64 降至 32

$ python convnet_distributed.py --num-machines 1 --num-processes 4 --machine-id 0 --epochs 1 --batch-size 32

輸出結(jié)果如下所示:

可以看到,訓(xùn)練時(shí)間進(jìn)一步減少,從 20 秒降至 15 秒。訓(xùn)練損失的變化趨勢(shì)仍然與之前的訓(xùn)練相似。通過分布式訓(xùn)練,我們已經(jīng)將訓(xùn)練時(shí)間從 30 秒縮短到了 15 秒,減少了 2 倍。

(9) 進(jìn)一步增加進(jìn)程數(shù),使用 8 個(gè)進(jìn)程代替 4 個(gè)進(jìn)程,并相應(yīng)地將批次大小從 32 降至 16

$ python convnet_distributed.py --num-machines 1 --num-processes 8 --machine-id 0 --epochs 1 --batch-size 16

輸出結(jié)果如下所示:

與預(yù)期相反,訓(xùn)練時(shí)間不僅沒有進(jìn)一步縮短,反而從 15 秒略微增加至 18 秒。由于代碼在本地機(jī)器執(zhí)行,系統(tǒng)還存在其他進(jìn)程(如瀏覽器)會(huì)與部分分布式訓(xùn)練進(jìn)程爭(zhēng)奪資源。如果分布式訓(xùn)練模型是在遠(yuǎn)程機(jī)器上進(jìn)行的,同時(shí)這些機(jī)器的唯一任務(wù)就是進(jìn)行模型訓(xùn)練,在這樣的機(jī)器上,建議使用與 CPU 核心數(shù)相等甚至更多的進(jìn)程數(shù)。

(10) 最后需要指出的是,由于在本節(jié)中我們只使用了一臺(tái)機(jī)器,因此我們只需要啟動(dòng)一個(gè) Python 腳本來開始訓(xùn)練。然而,如果是在多臺(tái)機(jī)器上進(jìn)行訓(xùn)練,那么除了修改 MASTER_ADDRMASTER_PORT 外,還需要在每臺(tái)機(jī)器上啟動(dòng)一個(gè) Python 腳本。例如,如果有 2 臺(tái)機(jī)器,在機(jī)器 1 上執(zhí)行:

$ python distributed_script.py --num_machines=2 --num-processes 8 --machine_id=0 --epochs 1 --batch-size 16

在機(jī)器 2 上執(zhí)行:

$ python distributed_script.py --num_machines=2 --num-processes 8 --machine_id=1 --epochs 1 --batch-size 16

至此,我們完成了關(guān)于使用 PyTorchCPU 上實(shí)施分布式訓(xùn)練深度學(xué)習(xí)模型的實(shí)踐探討,這種方法能帶來顯著的加速效果。僅需添加少量代碼,就能將常規(guī) PyTorch 模型訓(xùn)練腳本升級(jí)為分布式訓(xùn)練模式。雖然上述實(shí)驗(yàn)基于簡(jiǎn)單的卷積網(wǎng)絡(luò),但由于我們完全無需修改模型架構(gòu)代碼,因此這套方案可直接擴(kuò)展到更復(fù)雜的模型訓(xùn)練場(chǎng)景。接下來,我們將簡(jiǎn)要討論如何應(yīng)用類似的代碼更改,實(shí)現(xiàn) GPU 環(huán)境下的分布式訓(xùn)練。

2. GPU 分布式訓(xùn)練

我們通常使用以下 PyTorch 代碼定義模型訓(xùn)練設(shè)備:

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

這行代碼的作用是自動(dòng)檢測(cè)可用計(jì)算設(shè)備,并優(yōu)先選擇 CUDA (GPU)。這種優(yōu)先選擇源于 GPU 通過并行化處理神經(jīng)網(wǎng)絡(luò)常規(guī)運(yùn)算(如矩陣乘法和加法)所能提供的顯著加速優(yōu)勢(shì)。本節(jié)我們將探討如何通過 GPU 分布式訓(xùn)練進(jìn)一步加速模型訓(xùn)練。

(1) 雖然導(dǎo)入語句和模型架構(gòu)定義代碼與使用 CPU 分布式訓(xùn)練一節(jié)完全一致,但 train() 函數(shù)中有幾處關(guān)鍵修改:

def train(gpu_num, args):
    rank = args.machine_id * args.num_gpu_processes + gpu_num                        
    dist.init_process_group(                                   
        backend='nccl',                                         
        init_method='env://',                                   
        world_size=args.world_size,                              
        rank=rank                      
    ) 
    model = ConvNet()
    torch.cuda.set_device(gpu_num)
    model.cuda(gpu_num)
    criterion = nn.NLLLoss().cuda(gpu_num)

在使用 GPU 時(shí),NCCL 是首選的通信后端。同時(shí),模型和損失函數(shù)都必須部署到 GPU 設(shè)備上,以確保充分利用 GPU 提供的并行矩陣運(yùn)算加速能力:

    train_dataset = datasets.MNIST('./data', train=True, download=True,
                                   transform=transforms.Compose([
                                       transforms.ToTensor(),
                                       transforms.Normalize((0.1302,), (0.3069,))]))  
    train_sampler = torch.utils.data.distributed.DistributedSampler(
        train_dataset,
        num_replicas=args.world_size,
        rank=rank
    )
    train_dataloader = torch.utils.data.DataLoader(
       dataset=train_dataset,
       batch_size=args.batch_size,
       shuffle=False,            
       num_workers=0,
       pin_memory=True,
       sampler=train_sampler)
    optimizer = optim.Adadelta(model.parameters(), lr=0.5)
    model = nn.parallel.DistributedDataParallel(model,
                                                device_ids=[gpu_num])
    model.train()

DistributedDataParallel API 包含一個(gè)關(guān)鍵參數(shù)——device_ids,用于指定調(diào)用該 APIGPU 進(jìn)程 ID。此外可以看到,數(shù)據(jù)加載器 (dataloader) 中新增了 pin_memory 參數(shù)并設(shè)為 True,該參數(shù)能顯著加速訓(xùn)練過程中從主機(jī)(此處指加載數(shù)據(jù)集的 CPU )到各設(shè)備 (GPU) 的數(shù)據(jù)傳輸。
pin_memory 機(jī)制的工作原理是將數(shù)據(jù)"鎖定" (pin) 在 CPU 內(nèi)存中,即把數(shù)據(jù)樣本分配到固定的頁鎖定內(nèi)存區(qū)域。訓(xùn)練時(shí),這些內(nèi)存區(qū)域的數(shù)據(jù)會(huì)被高效地拷貝到對(duì)應(yīng) GPU。該機(jī)制需與 non_blocking=True 參數(shù)配合使用:

    for epoch in range(args.epochs):
        for b_i, (X, y) in enumerate(train_dataloader):
            X, y = X.cuda(non_blocking=True), y.cuda(non_blocking=True)
            pred_prob = model(X)
            loss = criterion(pred_prob, y) 
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            if b_i % 10 == 0 and gpu_num==0:
                print('epoch: {} [{}/{} ({:.0f}%)]\t training loss: {:.6f}'.format(
                    epoch, b_i, len(train_dataloader),
                    100. * b_i / len(train_dataloader), loss.item()))

通過調(diào)用參數(shù) pin_memorynon_blocking,使得以下兩者之間的操作得以重疊:

  • CPUGPU 數(shù)據(jù)(真實(shí)標(biāo)簽)的傳輸
  • GPU 模型訓(xùn)練計(jì)算(或 GPU 內(nèi)核執(zhí)行)

這從根本上提升了整體 GPU 訓(xùn)練流程的效率。

(2) 除了 train() 函數(shù)的修改外,main() 函數(shù)同樣有所調(diào)整調(diào)整:

def main():
    parser = argparse.ArgumentParser()
    parser.add_argument('--num-machines', default=1, type=int,)
    parser.add_argument('--num-gpu-processes', default=1, type=int)
    parser.add_argument('--machine-id', default=0, type=int)
    parser.add_argument('--epochs', default=1, type=int)
    parser.add_argument('--batch-size', default=64, type=int)
    args = parser.parse_args()
    
    args.world_size = args.num_gpu_processes * args.num_machines                
    os.environ['MASTER_ADDR'] = '127.0.0.1'              
    os.environ['MASTER_PORT'] = '8892'      
    start = time.time()
    mp.spawn(train, nprocs=args.num_gpu_processes, args=(args,))
    print(f"Finished training in {time.time()-start} secs")

num_gpu_processes 替代了原來的 num_process 參數(shù),該參數(shù)值通過 torch.cuda.device_count() 自動(dòng)獲取可用 GPU 數(shù)量。其余代碼相應(yīng)調(diào)整,但 GPU 版本的核心邏輯與之前保持一致。執(zhí)行以下命令即可啟動(dòng) GPU 分布式訓(xùn)練:

$ python convnet_distributed_cuda.py --num-machines 1 --num-gpu-processes 1 --machine-id 0 --epochs 1 --batch-size 128

至此,我們已完成關(guān)于使用 PyTorch 進(jìn)行 GPU 分布式模型訓(xùn)練的簡(jiǎn)要探討。上述代碼同樣適用于其他深度學(xué)習(xí)模型,當(dāng)前深度學(xué)習(xí)模型大多采用 GPU 分布式訓(xùn)練方案。此外,Horovod、DeepSpeedPyTorch Lightning 等庫都提供了更簡(jiǎn)潔的 API 來簡(jiǎn)化 PyTorch 模型的分布式訓(xùn)練流程。

小結(jié)

在本節(jié)中,我們探討了機(jī)器學(xué)習(xí)中一個(gè)重要的實(shí)踐方面——如何優(yōu)化模型訓(xùn)練過程,介紹了使用 PyTorchCPUGPU 上進(jìn)行分布式訓(xùn)練的適用范圍與強(qiáng)大效能。

到此這篇關(guān)于PyTorch分布式訓(xùn)練的實(shí)現(xiàn)的文章就介紹到這了,更多相關(guān)PyTorch分布式訓(xùn)練內(nèi)容請(qǐng)搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!

相關(guān)文章

  • Python機(jī)器學(xué)習(xí)算法之決策樹算法的實(shí)現(xiàn)與優(yōu)缺點(diǎn)

    Python機(jī)器學(xué)習(xí)算法之決策樹算法的實(shí)現(xiàn)與優(yōu)缺點(diǎn)

    決策樹(Decision Tree)是一種基本的分類與回歸方法,這篇文章主要給大家介紹了關(guān)于Python機(jī)器學(xué)習(xí)算法之決策樹算法實(shí)現(xiàn)與優(yōu)缺點(diǎn)的相關(guān)資料,需要的朋友們下面隨著小編來一起學(xué)習(xí)學(xué)習(xí)吧
    2021-05-05
  • python數(shù)據(jù)庫開發(fā)之MongoDB安裝及Python3操作MongoDB數(shù)據(jù)庫詳細(xì)方法與實(shí)例

    python數(shù)據(jù)庫開發(fā)之MongoDB安裝及Python3操作MongoDB數(shù)據(jù)庫詳細(xì)方法與實(shí)例

    這篇文章主要介紹了python數(shù)據(jù)庫開發(fā)之MongoDB安裝及Python3操作MongoDB數(shù)據(jù)庫詳細(xì)方法與實(shí)例,需要的朋友可以參考下
    2020-03-03
  • Python中tkinter+MySQL實(shí)現(xiàn)增刪改查

    Python中tkinter+MySQL實(shí)現(xiàn)增刪改查

    這篇文章主要介紹了Python中tkinter+MySQL實(shí)現(xiàn)增刪改查,文中通過示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來一起學(xué)習(xí)學(xué)習(xí)吧
    2021-04-04
  • python進(jìn)階之JSON數(shù)據(jù)解析完整示例

    python進(jìn)階之JSON數(shù)據(jù)解析完整示例

    Python作為一種強(qiáng)大的編程語言,提供了多種方法來處理JSON數(shù)據(jù),使其在數(shù)據(jù)解析、處理和生成方面變得異常簡(jiǎn)便,這篇文章主要介紹了python進(jìn)階之JSON數(shù)據(jù)解析的相關(guān)資料,需要的朋友可以參考下
    2025-12-12
  • 使用Python生成個(gè)性化的電子郵件簽名

    使用Python生成個(gè)性化的電子郵件簽名

    在數(shù)字通信時(shí)代,電子郵件仍然是商務(wù)溝通和個(gè)人交流的重要工具,本文將詳細(xì)介紹如何使用Python構(gòu)建一個(gè)智能的個(gè)性化電子郵件簽名生成系統(tǒng),希望對(duì)大家有所幫助
    2025-11-11
  • 如何使用 Flask 做一個(gè)評(píng)論系統(tǒng)

    如何使用 Flask 做一個(gè)評(píng)論系統(tǒng)

    這篇文章主要介紹了如何使用 Flask 做一個(gè)評(píng)論系統(tǒng),幫助大家更好的理解和使用flask框架進(jìn)行python web開發(fā),感興趣的朋友可以了解下
    2020-11-11
  • python切片作為占位符使用實(shí)例講解

    python切片作為占位符使用實(shí)例講解

    在本篇內(nèi)容里小編給大家分享的是一篇關(guān)于python切片作為占位符使用實(shí)例講解內(nèi)容,有興趣的朋友們可以學(xué)習(xí)參考下。
    2021-02-02
  • opencv之為圖像添加邊界的方法示例

    opencv之為圖像添加邊界的方法示例

    這篇文章主要介紹了opencv之為圖像添加邊界的方法示例,文中通過示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來一起學(xué)習(xí)學(xué)習(xí)吧
    2019-12-12
  • django中F表達(dá)式和Q函數(shù)應(yīng)用與原理詳解

    django中F表達(dá)式和Q函數(shù)應(yīng)用與原理詳解

    F對(duì)象查詢與Q對(duì)象查詢,剛看到大家一定會(huì)感到很陌生,其實(shí)它們也是 Django 提供的查詢方法,而且非常的簡(jiǎn)單的高效,下面這篇文章主要給大家介紹了關(guān)于django中F表達(dá)式和Q函數(shù)應(yīng)用與原理的相關(guān)資料,需要的朋友可以參考下
    2023-05-05
  • python實(shí)現(xiàn)FTP循環(huán)上傳文件

    python實(shí)現(xiàn)FTP循環(huán)上傳文件

    這篇文章主要為大家詳細(xì)介紹了python實(shí)現(xiàn)FTP循環(huán)上傳文件,文中示例代碼介紹的非常詳細(xì),具有一定的參考價(jià)值,感興趣的小伙伴們可以參考一下
    2020-03-03

最新評(píng)論

民权县| 武城县| 桂林市| 江都市| 旺苍县| 利川市| 宝坻区| 镇巴县| 卓资县| 天台县| 交城县| 大厂| 肥乡县| 邹平县| 郴州市| 天等县| 陆川县| 鸡泽县| 伊宁市| 嵩明县| 阳西县| 苍梧县| 灵寿县| 泸州市| 随州市| 绥中县| 莱阳市| 阜新市| 岳普湖县| 南平市| 昆明市| 宣武区| 雅江县| 观塘区| 芜湖市| 孟连| 洞口县| 突泉县| 瑞昌市| 伊川县| 昭苏县|