찾다
백엔드 개발파이썬 튜토리얼신경망 훈련에 PyTorch를 사용하는 방법

신경망 훈련에 PyTorch를 사용하는 방법

소개:
PyTorch는 Python 기반의 오픈 소스 기계 학습 프레임워크로 유연성과 단순성으로 인해 많은 연구원과 엔지니어가 가장 먼저 선택합니다. 이 기사에서는 신경망 훈련에 PyTorch를 사용하는 방법을 소개하고 해당 코드 예제를 제공합니다.

1. PyTorch 설치
시작하기 전에 먼저 PyTorch를 설치해야 합니다. 공식 홈페이지(https://pytorch.org/)에서 제공하는 설치 가이드를 통해 운영체제 및 하드웨어에 적합한 버전을 선택하여 설치할 수 있습니다. 설치가 완료되면 Python으로 PyTorch 라이브러리를 가져오고 코드 작성을 시작할 수 있습니다.

2. 신경망 모델 구축
PyTorch를 사용하여 신경망을 훈련시키기 전에 먼저 적합한 모델을 구축해야 합니다. PyTorch는 자신만의 신경망 모델을 정의하기 위해 상속할 수 있는 torch.nn.Module이라는 클래스를 제공합니다. torch.nn.Module的类,您可以通过继承该类来定义自己的神经网络模型。

下面是一个简单的例子,展示了如何使用PyTorch构建一个包含两个全连接层的神经网络模型:

import torch
import torch.nn as nn

class Net(nn.Module):
    def __init__(self):
        super(Net, self).__init__()
        self.fc1 = nn.Linear(in_features=784, out_features=256)
        self.fc2 = nn.Linear(in_features=256, out_features=10)
    
    def forward(self, x):
        x = x.view(x.size(0), -1)
        x = self.fc1(x)
        x = torch.relu(x)
        x = self.fc2(x)
        return x

net = Net()

在上面的代码中,我们首先定义了一个名为Net的类,并继承了torch.nn.Module类。在__init__方法中,我们定义了两个全连接层fc1fc2。然后,我们通过forward方法定义了数据在模型中前向传播的过程。最后,我们创建了一个Net的实例。

三、定义损失函数和优化器
在进行训练之前,我们需要定义损失函数和优化器。PyTorch提供了丰富的损失函数和优化器的选择,可以根据具体情况进行选择。

下面是一个示例,展示了如何定义一个使用交叉熵损失函数和随机梯度下降优化器的训练过程:

loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.SGD(net.parameters(), lr=0.01)

在上面的代码中,我们将交叉熵损失函数和随机梯度下降优化器分别赋值给了loss_fnoptimizer变量。net.parameters()表示我们要优化神经网络模型中的所有可学习参数,lr参数表示学习率。

四、准备数据集
在进行神经网络训练之前,我们需要准备好训练数据集和测试数据集。PyTorch提供了一些实用的工具类,可以帮助我们加载和预处理数据集。

下面是一个示例,展示了如何加载MNIST手写数字数据集并进行预处理:

import torchvision
import torchvision.transforms as transforms

transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5,), (0.5,)),
])

train_set = torchvision.datasets.MNIST(root='./data', train=True, download=True, transform=transform)
train_loader = torch.utils.data.DataLoader(train_set, batch_size=32, shuffle=True)

test_set = torchvision.datasets.MNIST(root='./data', train=False, download=True, transform=transform)
test_loader = torch.utils.data.DataLoader(test_set, batch_size=32, shuffle=False)

在上面的代码中,我们首先定义了一个transform变量,用于对数据进行预处理。然后,我们使用torchvision.datasets.MNIST类加载MNIST数据集,并使用train=Truetrain=False参数指定了训练数据集和测试数据集。最后,我们使用torch.utils.data.DataLoader类将数据集转换成一个可以迭代的数据加载器。

五、开始训练
准备好数据集后,我们就可以开始进行神经网络的训练。在一个训练循环中,我们需要依次完成以下步骤:将输入数据输入到模型中,计算损失函数,反向传播更新梯度,优化模型。

下面是一个示例,展示了如何使用PyTorch进行神经网络训练:

for epoch in range(epochs):
    running_loss = 0.0
    for i, data in enumerate(train_loader):
        inputs, labels = data
        
        optimizer.zero_grad()
        
        outputs = net(inputs)
        loss = loss_fn(outputs, labels)
        
        loss.backward()
        optimizer.step()
        
        running_loss += loss.item()
        
        if (i+1) % 100 == 0:
            print('[%d, %5d] loss: %.3f' % (epoch+1, i+1, running_loss/100))
            running_loss = 0.0

在上面的代码中,我们首先使用enumerate函数遍历了训练数据加载器,得到了输入数据和标签。然后,我们将梯度清零,将输入数据输入到模型中,计算预测结果和损失函数。接着,我们通过backward方法计算梯度,再通过step方法更新模型参数。最后,我们累加损失,并根据需要进行打印。

六、测试模型
训练完成后,我们还需要测试模型的性能。我们可以通过计算模型在测试数据集上的准确率来评估模型的性能。

下面是一个示例,展示了如何使用PyTorch测试模型的准确率:

correct = 0
total = 0

with torch.no_grad():
    for data in test_loader:
        inputs, labels = data
        outputs = net(inputs)
        _, predicted = torch.max(outputs.data, 1)
        total += labels.size(0)
        correct += (predicted == labels).sum().item()

accuracy = 100 * correct / total
print('Accuracy: %.2f %%' % accuracy)

在上面的代码中,我们首先定义了两个变量correcttotal,用于计算正确分类的样本和总样本数。接着,我们使用torch.no_grad()

다음은 PyTorch를 사용하여 두 개의 완전히 연결된 레이어가 포함된 신경망 모델을 구축하는 방법을 보여주는 간단한 예입니다.

rrreee
위 코드에서는 먼저 Net이라는 클래스를 정의하고 torch.nn에서 상속합니다. 모듈 클래스. __init__ 메서드에서는 두 개의 완전히 연결된 레이어 fc1fc2를 정의합니다. 그런 다음 forward 메서드를 통해 모델에서 데이터의 순방향 전파 프로세스를 정의합니다. 마지막으로 Net 인스턴스를 만듭니다.

3. 손실 함수와 옵티마이저 정의

훈련 전에 손실 함수와 옵티마이저를 정의해야 합니다. PyTorch는 특정 상황에 따라 선택할 수 있는 다양한 손실 함수 및 최적화 프로그램을 제공합니다.
  1. 다음은 교차 엔트로피 손실 함수와 확률적 경사하강법 최적화 도구를 사용하여 훈련 과정을 정의하는 방법을 보여주는 예입니다.
  2. rrreee
  3. 위 코드에서는 교차 엔트로피 손실 함수와 확률적 경사하강법 최적화 도구를 할당합니다. 별도로 loss_fnoptimizer 변수가 제공됩니다. net.parameters()는 신경망 모델에서 학습 가능한 모든 매개변수를 최적화하려고 함을 나타내고 lr 매개변수는 학습 속도를 나타냅니다.
4. 데이터 세트 준비🎜 신경망을 훈련시키기 전에 훈련 데이터 세트와 테스트 데이터 세트를 준비해야 합니다. PyTorch는 데이터 세트를 로드하고 전처리하는 데 도움이 되는 몇 가지 실용적인 도구 클래스를 제공합니다. 🎜🎜다음은 MNIST 필기 숫자 데이터 세트를 로드하고 전처리하는 방법을 보여주는 예입니다. 🎜rrreee🎜위 코드에서는 먼저 transform 변수를 정의하여 데이터 전처리를 변환합니다. 그런 다음 torchvision.datasets.MNIST 클래스를 사용하여 MNIST 데이터세트를 로드하고 train=Truetrain=False 매개변수를 사용하여 훈련 데이터를 지정했습니다. 데이터 세트를 설정하고 테스트합니다. 마지막으로 torch.utils.data.DataLoader 클래스를 사용하여 데이터 세트를 반복 가능한 데이터 로더로 변환합니다. 🎜🎜5. 훈련 시작🎜 데이터 세트를 준비한 후 신경망 훈련을 시작할 수 있습니다. 훈련 루프에서는 입력 데이터를 모델에 입력하고, 손실 함수를 계산하고, 업데이트된 기울기를 역전파하고, 모델을 최적화하는 단계를 순서대로 완료해야 합니다. 🎜🎜다음은 신경망 훈련에 PyTorch를 사용하는 방법을 보여주는 예입니다. 🎜rrreee🎜위 코드에서는 먼저 enumerate 함수를 사용하여 훈련 데이터 로더를 탐색하여 입력 데이터와 레이블을 가져옵니다. 그런 다음 기울기를 0으로 만들고 입력 데이터를 모델에 공급한 다음 예측 및 손실 함수를 계산합니다. 다음으로 backward 메서드를 통해 기울기를 계산한 다음 step 메서드를 통해 모델 매개변수를 업데이트합니다. 마지막으로 손실을 누적하고 필요에 따라 인쇄합니다. 🎜🎜 6. 모델 테스트 🎜훈련이 완료된 후에도 모델의 성능을 테스트해야 합니다. 테스트 데이터 세트에 대한 정확도를 계산하여 모델의 성능을 평가할 수 있습니다. 🎜🎜다음은 PyTorch를 사용하여 모델의 정확성을 테스트하는 방법을 보여주는 예입니다. 🎜rrreee🎜위 코드에서는 먼저 corrighttotal 두 변수를 정의하고 사용합니다. 올바르게 분류된 샘플 수와 전체 샘플 수를 계산합니다. 다음으로 torch.no_grad() 컨텍스트 관리자를 사용하여 기울기 계산을 꺼서 메모리 소비를 줄입니다. 그런 다음 예측 결과를 순차적으로 계산하고 올바르게 분류된 샘플 수와 총 샘플 수를 업데이트합니다. 마지막으로 정확하게 분류된 샘플 수와 전체 샘플 수를 기준으로 정확도를 계산하여 인쇄합니다. 🎜🎜요약: 🎜이 글의 서론을 통해 신경망 훈련에 PyTorch를 사용하는 방법의 기본 단계를 이해했으며, 신경망 모델 구축, 손실 함수 및 옵티마이저 정의, 데이터 세트 준비, 훈련 시작 방법을 배웠습니다. 그리고 모델을 테스트합니다. 이 기사가 신경망 훈련에 PyTorch를 사용하는 작업과 연구에 도움이 되기를 바랍니다. 🎜🎜참고자료: 🎜🎜🎜PyTorch 공식 웹사이트: https://pytorch.org/🎜🎜PyTorch 문서: https://pytorch.org/docs/stable/index.html🎜🎜

위 내용은 신경망 훈련에 PyTorch를 사용하는 방법의 상세 내용입니다. 자세한 내용은 PHP 중국어 웹사이트의 기타 관련 기사를 참조하세요!

성명
본 글의 내용은 네티즌들의 자발적인 기여로 작성되었으며, 저작권은 원저작자에게 있습니다. 본 사이트는 이에 상응하는 법적 책임을 지지 않습니다. 표절이나 침해가 의심되는 콘텐츠를 발견한 경우 admin@php.cn으로 문의하세요.
파이썬 : 게임, Guis 등파이썬 : 게임, Guis 등Apr 13, 2025 am 12:14 AM

Python은 게임 및 GUI 개발에서 탁월합니다. 1) 게임 개발은 Pygame을 사용하여 드로잉, 오디오 및 기타 기능을 제공하며 2D 게임을 만드는 데 적합합니다. 2) GUI 개발은 Tkinter 또는 PYQT를 선택할 수 있습니다. Tkinter는 간단하고 사용하기 쉽고 PYQT는 풍부한 기능을 가지고 있으며 전문 개발에 적합합니다.

Python vs. C : 응용 및 사용 사례가 비교되었습니다Python vs. C : 응용 및 사용 사례가 비교되었습니다Apr 12, 2025 am 12:01 AM

Python은 데이터 과학, 웹 개발 및 자동화 작업에 적합한 반면 C는 시스템 프로그래밍, 게임 개발 및 임베디드 시스템에 적합합니다. Python은 단순성과 강력한 생태계로 유명하며 C는 고성능 및 기본 제어 기능으로 유명합니다.

2 시간의 파이썬 계획 : 현실적인 접근2 시간의 파이썬 계획 : 현실적인 접근Apr 11, 2025 am 12:04 AM

2 시간 이내에 Python의 기본 프로그래밍 개념과 기술을 배울 수 있습니다. 1. 변수 및 데이터 유형을 배우기, 2. 마스터 제어 흐름 (조건부 명세서 및 루프), 3. 기능의 정의 및 사용을 이해하십시오. 4. 간단한 예제 및 코드 스 니펫을 통해 Python 프로그래밍을 신속하게 시작하십시오.

파이썬 : 기본 응용 프로그램 탐색파이썬 : 기본 응용 프로그램 탐색Apr 10, 2025 am 09:41 AM

Python은 웹 개발, 데이터 과학, 기계 학습, 자동화 및 스크립팅 분야에서 널리 사용됩니다. 1) 웹 개발에서 Django 및 Flask 프레임 워크는 개발 프로세스를 단순화합니다. 2) 데이터 과학 및 기계 학습 분야에서 Numpy, Pandas, Scikit-Learn 및 Tensorflow 라이브러리는 강력한 지원을 제공합니다. 3) 자동화 및 스크립팅 측면에서 Python은 자동화 된 테스트 및 시스템 관리와 ​​같은 작업에 적합합니다.

2 시간 안에 얼마나 많은 파이썬을 배울 수 있습니까?2 시간 안에 얼마나 많은 파이썬을 배울 수 있습니까?Apr 09, 2025 pm 04:33 PM

2 시간 이내에 파이썬의 기본 사항을 배울 수 있습니다. 1. 변수 및 데이터 유형을 배우십시오. 이를 통해 간단한 파이썬 프로그램 작성을 시작하는 데 도움이됩니다.

10 시간 이내에 프로젝트 및 문제 중심 방법에서 컴퓨터 초보자 프로그래밍 기본 사항을 가르치는 방법?10 시간 이내에 프로젝트 및 문제 중심 방법에서 컴퓨터 초보자 프로그래밍 기본 사항을 가르치는 방법?Apr 02, 2025 am 07:18 AM

10 시간 이내에 컴퓨터 초보자 프로그래밍 기본 사항을 가르치는 방법은 무엇입니까? 컴퓨터 초보자에게 프로그래밍 지식을 가르치는 데 10 시간 밖에 걸리지 않는다면 무엇을 가르치기로 선택 하시겠습니까?

중간 독서를 위해 Fiddler를 사용할 때 브라우저에서 감지되는 것을 피하는 방법은 무엇입니까?중간 독서를 위해 Fiddler를 사용할 때 브라우저에서 감지되는 것을 피하는 방법은 무엇입니까?Apr 02, 2025 am 07:15 AM

Fiddlerevery Where를 사용할 때 Man-in-the-Middle Reading에 Fiddlereverywhere를 사용할 때 감지되는 방법 ...

Python 3.6에 피클 파일을로드 할 때 '__builtin__'모듈을 찾을 수없는 경우 어떻게해야합니까?Python 3.6에 피클 파일을로드 할 때 '__builtin__'모듈을 찾을 수없는 경우 어떻게해야합니까?Apr 02, 2025 am 07:12 AM

Python 3.6에 피클 파일로드 3.6 환경 보고서 오류 : modulenotfounderror : nomodulename ...

See all articles

핫 AI 도구

Undresser.AI Undress

Undresser.AI Undress

사실적인 누드 사진을 만들기 위한 AI 기반 앱

AI Clothes Remover

AI Clothes Remover

사진에서 옷을 제거하는 온라인 AI 도구입니다.

Undress AI Tool

Undress AI Tool

무료로 이미지를 벗다

Clothoff.io

Clothoff.io

AI 옷 제거제

AI Hentai Generator

AI Hentai Generator

AI Hentai를 무료로 생성하십시오.

인기 기사

R.E.P.O. 에너지 결정과 그들이하는 일 (노란색 크리스탈)
3 몇 주 전By尊渡假赌尊渡假赌尊渡假赌
R.E.P.O. 최고의 그래픽 설정
3 몇 주 전By尊渡假赌尊渡假赌尊渡假赌
R.E.P.O. 아무도들을 수없는 경우 오디오를 수정하는 방법
3 몇 주 전By尊渡假赌尊渡假赌尊渡假赌
WWE 2K25 : Myrise에서 모든 것을 잠금 해제하는 방법
4 몇 주 전By尊渡假赌尊渡假赌尊渡假赌

뜨거운 도구

mPDF

mPDF

mPDF는 UTF-8로 인코딩된 HTML에서 PDF 파일을 생성할 수 있는 PHP 라이브러리입니다. 원저자인 Ian Back은 자신의 웹 사이트에서 "즉시" PDF 파일을 출력하고 다양한 언어를 처리하기 위해 mPDF를 작성했습니다. HTML2FPDF와 같은 원본 스크립트보다 유니코드 글꼴을 사용할 때 속도가 느리고 더 큰 파일을 생성하지만 CSS 스타일 등을 지원하고 많은 개선 사항이 있습니다. RTL(아랍어, 히브리어), CJK(중국어, 일본어, 한국어)를 포함한 거의 모든 언어를 지원합니다. 중첩된 블록 수준 요소(예: P, DIV)를 지원합니다.

SecList

SecList

SecLists는 최고의 보안 테스터의 동반자입니다. 보안 평가 시 자주 사용되는 다양한 유형의 목록을 한 곳에 모아 놓은 것입니다. SecLists는 보안 테스터에게 필요할 수 있는 모든 목록을 편리하게 제공하여 보안 테스트를 더욱 효율적이고 생산적으로 만드는 데 도움이 됩니다. 목록 유형에는 사용자 이름, 비밀번호, URL, 퍼징 페이로드, 민감한 데이터 패턴, 웹 셸 등이 포함됩니다. 테스터는 이 저장소를 새로운 테스트 시스템으로 간단히 가져올 수 있으며 필요한 모든 유형의 목록에 액세스할 수 있습니다.

에디트플러스 중국어 크랙 버전

에디트플러스 중국어 크랙 버전

작은 크기, 구문 강조, 코드 프롬프트 기능을 지원하지 않음

SublimeText3 Linux 새 버전

SublimeText3 Linux 새 버전

SublimeText3 Linux 최신 버전

Dreamweaver Mac版

Dreamweaver Mac版

시각적 웹 개발 도구