Ошибка видеокарты, питон

Пишу свою первую GAN-нейросеть на pytorch. Во время обучения вылетает ошибка CUDA error: an illegal memory access was encountered. Драйвера переустанавливал, размер батча менял, но ошибка продолжает вылетать и у меня на компьютере, и на google collab. Отсюда ощущение, что это какая-то ошибка в коде. Вот весь код:

import os
import torch
import torchvision
from torch import nn
from tqdm import tqdm
import torch.nn.functional as f
import matplotlib.pyplot as plt
import torchvision.transforms as transforms

torch.manual_seed(42)


class Discriminator(nn.Module):
    def __init__(self):
        super(Discriminator, self).__init__()

        self.conv1 = nn.Conv2d(in_channels=1, out_channels=8, kernel_size=6)

        self.max_pool = nn.MaxPool2d(kernel_size=2, stride=2)

        self.conv2 = nn.Conv2d(in_channels=8, out_channels=16, kernel_size=5)
        self.conv3 = nn.Conv2d(in_channels=16, out_channels=32, kernel_size=3)

        self.fc1 = nn.Linear(5*5*32, 256)
        self.fc2 = nn.Linear(256, 128)
        self.fc3 = nn.Linear(128, 1)

    def forward(self, x):
        x = f.relu(self.conv1(x))
        x = self.max_pool(x)
        x = f.relu(self.conv2(x))
        x = f.relu(self.conv3(x))

        x = x.view(-1, 5*5*32)

        x = f.relu(self.fc1(x))
        x = f.relu(self.fc2(x))
        x = self.fc3(x)

        x = f.softmax(x, dim=-1)

        return x


class Generator(nn.Module):
    def __init__(self):
        super(Generator, self).__init__()

        self.fc1 = nn.Linear(100, 256)
        self.fc2 = nn.Linear(256, 512)
        self.fc3 = nn.Linear(512, 1024)
        self.fc4 = nn.Linear(1024, 784)

    def forward(self, x):
        x = f.relu(self.fc1(x))
        x = f.relu(self.fc2(x))
        x = f.relu(self.fc3(x))
        x = f.tanh(self.fc4(x))

        x = x.view(x.size(0), 1, 28, 28)

        return x


class NetworkStuff:
    def __init__(self):

        self.batch_size = 16
        self.lr = 0.0001
        self.num_epoch = 50

        self.discriminator = Discriminator()
        self.generator = Generator()

        self.criterion = nn.BCELoss()

        self.optimizer_discriminator = torch.optim.Adam(
            self.discriminator.parameters(),
            self.lr
        )
        self.optimizer_generator = torch.optim.Adam(
            self.generator.parameters(),
            self.lr
        )

        if torch.cuda.is_available():
            self.device = torch.device("cuda")
        else:
            self.device = torch.device("cpu")

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

        self.train_set = torchvision.datasets.MNIST(
            root=os.path.abspath("data"), train=True, download=True,
            transform=self.transform
        )

        self.train_loader = torch.utils.data.DataLoader(
            self.train_set, batch_size=self.batch_size,
            shuffle=True
        )

    def show_samples(self, type_='generated'):

        if type_ == "dataset":

            real_samples, mnist_labels = next(iter(self.train_loader))
            plt.suptitle('dataset')

            for i in range(16):
                plt.subplot(4, 4, i + 1)
                plt.imshow(real_samples[i].reshape(28, 28), cmap="gray_r")
                plt.xticks([])
                plt.yticks([])

        elif type_ == "generated":

            latent_samples = torch.rand((self.batch_size, 100)).to(self.device)
            generated_samples = self.generator(latent_samples).detach()

            plt.suptitle('generated')
            for i in range(16):
                plt.subplot(4, 4, i + 1)
                plt.imshow(generated_samples[i].reshape(28, 28), cmap="gray_r")
                plt.xticks([])
                plt.yticks([])

        plt.show()

    def train(self):

        for epoch in tqdm(range(self.num_epoch)):

            for n, (real_samples, mnist_labels) in tqdm(enumerate(self.train_loader)):

                real_samples = real_samples.to(self.device)
                real_samples_labels = torch.ones((self.batch_size, 1)).to(self.device)

                latent_samples = torch.rand((self.batch_size, 100)).to(self.device)

                generated_samples = self.generator(latent_samples)
                generated_samples_labels = torch.zeros((self.batch_size, 1)).to(self.device)

                all_samples = torch.cat((real_samples, generated_samples))
                all_samples_labels = torch.cat((real_samples_labels, generated_samples_labels))

                # учим дискриминатор
                self.discriminator.zero_grad()
                output_discriminator = self.discriminator(all_samples)
                loss_discriminator = self.criterion(output_discriminator, all_samples_labels)
                loss_discriminator.backward()
                self.optimizer_discriminator.step()

                # учим генератор
                latent_samples = torch.rand((self.batch_size, 100)).to(self.device)
                real_samples_labels = torch.ones((self.batch_size, 1)).to(self.device)

                self.generator.zero_grad()
                generated_samples = self.generator(latent_samples)

                output_discriminator_generated = self.discriminator(generated_samples)

                generator_loss = self.criterion(output_discriminator_generated, real_samples_labels)
                generator_loss.backward()
                self.optimizer_generator.step()

                if n == self.batch_size - 1:
                    tqdm.write(f"Epoch: {epoch} Loss D.: {loss_discriminator}\n" +
                               f"Epoch: {epoch} Loss G.: {generator_loss}")


Ответы (1 шт):

Автор решения: DKay

Кажется ошибка была в том, что я забыл переместить нейросети на GPU (а все остальное переместил).

Исправление выглядит как-то так:

        self.discriminator = Discriminator().to(self.device)
        self.generator = Generator().to(self.device)
→ Ссылка