【问题标题】:Why does requires_grad turns from true to false when doing torch.nn.conv2d operation?为什么进行 torch.nn.conv2d 操作时 requires_grad 会从 true 变为 false?
【发布时间】:2021-01-10 00:56:50
【问题描述】:

我有 Unet 网络,它接收大脑的 MRI 图像,目标是分割大脑中的白色物质。图像的形状为 256x256x183(重新整形为 183x256x256)(FLAIR 和 T1 图像)。我遇到的问题是,在将输入发送到 Unet 网络之前,我的 pytorch 张量上有 requires_grad=True,但是在一次 torch.nn.conv2d 操作之后,requires_grad=False。这是一个大问题,因为梯度不会更新和学习。

from collections import OrderedDict

import torch
import torch.nn as nn


class UNet(nn.Module):
    
    def __init__(self, in_channels=3, out_channels=1, init_features=32):
        super(UNet, self).__init__()

        features = init_features
        self.encoder1 = UNet._block(in_channels, features, name="enc1")
        self.pool1 = nn.MaxPool2d(kernel_size=2, stride=2)
        self.encoder2 = UNet._block(features, features * 2, name="enc2")
        self.pool2 = nn.MaxPool2d(kernel_size=2, stride=2)
        self.encoder3 = UNet._block(features * 2, features * 4, name="enc3")
        self.pool3 = nn.MaxPool2d(kernel_size=2, stride=2)
        self.encoder4 = UNet._block(features * 4, features * 8, name="enc4")
        self.pool4 = nn.MaxPool2d(kernel_size=2, stride=2)

        self.bottleneck = UNet._block(features * 8, features * 16, name="bottleneck")

        self.upconv4 = nn.ConvTranspose2d(
            features * 16, features * 8, kernel_size=2, stride=2
        )
        self.decoder4 = UNet._block((features * 8) * 2, features * 8, name="dec4")
        self.upconv3 = nn.ConvTranspose2d(
            features * 8, features * 4, kernel_size=2, stride=2
        )
        self.decoder3 = UNet._block((features * 4) * 2, features * 4, name="dec3")
        self.upconv2 = nn.ConvTranspose2d(
            features * 4, features * 2, kernel_size=2, stride=2
        )
        self.decoder2 = UNet._block((features * 2) * 2, features * 2, name="dec2")
        self.upconv1 = nn.ConvTranspose2d(
            features * 2, features, kernel_size=2, stride=2
        )
        self.decoder1 = UNet._block(features * 2, features, name="dec1")

        self.conv = nn.Conv2d(
            in_channels=features, out_channels=out_channels, kernel_size=1
        )

    def forward(self, x):

        print(x.requires_grad) #<---- here it is true
        enc1 = self.encoder1(x)#<---- where the problem happens
       
        print(enc1.requires_grad) #<---- here it is false
        enc2 = self.encoder2(self.pool1(enc1))
        print(enc2.requires_grad)
        enc3 = self.encoder3(self.pool2(enc2))
        print(enc3.requires_grad)
        enc4 = self.encoder4(self.pool3(enc3))
        print(enc4.requires_grad)

        bottleneck = self.bottleneck(self.pool4(enc4))
        print(bottleneck.requires_grad)

        dec4 = self.upconv4(bottleneck)
        print(dec4.requires_grad)
        dec4 = torch.cat((dec4, enc4), dim=1)
        print(dec4.requires_grad)
        dec4 = self.decoder4(dec4)
        print(dec4.requires_grad)
        dec3 = self.upconv3(dec4)
        print(dec3.requires_grad)
        dec3 = torch.cat((dec3, enc3), dim=1)
        print(dec3.requires_grad)
        dec3 = self.decoder3(dec3)
        print(dec3.requires_grad)
        dec2 = self.upconv2(dec3)
        print(dec2.requires_grad)
        dec2 = torch.cat((dec2, enc2), dim=1)
        print(dec2.requires_grad)
        dec2 = self.decoder2(dec2)
        print(dec2.requires_grad)
        dec1 = self.upconv1(dec2)
        print(dec1.requires_grad)
        dec1 = torch.cat((dec1, enc1), dim=1)
        print(dec1.requires_grad)
        dec1 = self.decoder1(dec1)
        print(dec1.requires_grad)
        print("going out")
        return torch.sigmoid(self.conv(dec1))

    @staticmethod
    def _block(in_channels, features, name):
        return nn.Sequential(
            OrderedDict(
                [
                    (
                        name + "conv1",
                        nn.Conv2d(
                            in_channels=in_channels,
                            out_channels=features,
                            kernel_size=3,
                            padding=1,
                            bias=False,
                        ),
                    ),
                    (name + "norm1", nn.BatchNorm2d(num_features=features)),
                    (name + "relu1", nn.ReLU(inplace=True)),
                    (
                        name + "conv2",
                        nn.Conv2d(
                            in_channels=features,
                            out_channels=features,
                            kernel_size=3,
                            padding=1,
                            bias=False,
                        ),
                    ),
                    (name + "norm2", nn.BatchNorm2d(num_features=features)),
                    (name + "relu2", nn.ReLU(inplace=True)),
                ]
            )
        )

编辑: 这是训练代码

class run_network:
def __init__(self, eta, epoch, batch_size, train_file_path, validation_file_path, shuffle_after_epoch = True):
    self.eta = eta
    self.epoch = epoch
    self.batch_size = batch_size
    self.train_file_path = train_file_path
    self.validation_file_path = validation_file_path
    self.shuffle_after_epoch = shuffle_after_epoch

def __call__(self, is_train = False):
    
    device = torch.device("cpu" if not torch.cuda.is_available() else torch.cuda())
    unet = torch.hub.load('mateuszbuda/brain-segmentation-pytorch', 'unet',
    in_channels=3, out_channels=1, init_features=32, pretrained=True)
    unet.to(device)
    unet = unet.double()
    
    

    
    
    optimizer = optim.Adam(unet.parameters(), lr=self.eta)
    dsc_loss = DiceLoss()
    

    Load_training   = NiftiLoader(self.train_file_path)
    Load_validation = NiftiLoader(self.validation_file_path)
    
    mean_flair, mean_t1, std_flair, std_t1 = Load_training.average_mean_and_std(20, 79,99)

    total_mean = [mean_flair, mean_t1]
    total_std = [std_flair, std_t1]

    loss_train = []
    loss_validation = []


    

    for current_epoch in tqdm(range(self.epoch)):
        for phase in ["train", "validation"]:
            
            
            if phase == "train":
                mini_batch = Load_training.create_batch(self.batch_size, self.shuffle_after_epoch)
                unet.train()
                print("her22")

            if phase == "validation":
                print("her")
                mini_batch = Load_validation.create_batch(self.batch_size, self.shuffle_after_epoch)
                unet.eval()
            
            
            dim1, dim2, dim3 = mini_batch.shape
        
            for iteration in range(1):
                if phase == "train":
                    current_batch = Load_training.Load_Image_batch(mini_batch, iteration)
                    image_batch = Load_training.image_zero_mean_normalizer(current_batch)
                if phase == "validation":
                    current_batch = Load_validation.Load_Image_batch(mini_batch, iteration)
                    image_batch = Load_training.image_zero_mean_normalizer(current_batch, False, mean_list, std_list)


                image_dim0, image_dim1, image_dim2, image_dim3, image_dim4 = image_batch.shape
                image_batch = image_batch.reshape((
                                                    image_dim0, 
                                                    image_dim1*image_dim2, 
                                                    image_dim3, 
                                                    image_dim4
                                                    ))

                
                image_batch = np.swapaxes(image_batch, 0,1)
                image_batch = torch.as_tensor(image_batch)#.requires_grad_(True) #, requires_grad=True)
                image_batch = image_batch.to(device)
                print(image_batch.requires_grad)
                optimizer.zero_grad()
                
            
                with torch.set_grad_enabled(is_train == "train"):
                    for j in range(0, 10, 1): 
                        # [183*5, 3, 256, 256] -> [12, 3, 256, 256]  
                        # ANTALL ITERASJONER: (183*5/12) -> en chunk  
                            
                        input_image = image_batch[j:j+2,0:3,:,:]
                        print(input_image.requires_grad)
                        print("går inn")
                        y_predicted = unet(input_image)
                    
                        print(y_predicted.requires_grad)
                        print(image_batch[j:j+2,3,:,:].requires_grad)
                        loss = dsc_loss(y_predicted.squeeze(1), image_batch[j:j+2,3,:,:])
                        print(loss.requires_grad)
                       
                        if phase == "train":
                            loss_train.append(loss.item())
                            
                            loss.backward()
                            print(loss.item())
                            exit()
                            optimizer.step()
                            print(loss.item())
                            exit()
                        if phase == "validation":
                            loss_validation.append(loss.item())

迭代次数和打印语句用于试验可能的原因。

【问题讨论】:

    标签: python machine-learning pytorch conv-neural-network


    【解决方案1】:

    对我来说很好用。

    '''
    # I changed your code a little bit to catch up the problem.
    def forward(self, x):
    
            print("encoder1", x.requires_grad) #<---- here it is true
            enc1 = self.encoder1(x)#<---- where the problem happens
           
            print("encoder2", enc1.requires_grad) #<---- here it is false
    '''
    a = torch.randn(32, 3, 255, 255, requires_grad=True)
    # a.requires_grads = True
    print(a)
    UNet()(a)
    
    # This is the result:
    encoder1 True
    encoder2 True
    True
    True
    True
    True
    True
    

    你能告诉我你的训练来源吗?我想这就是问题所在。为什么需要更新输入数据?

    【讨论】:

    • 我不想更新输入数据,但问题是 loss.backwards() 不起作用。它引发错误 RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn.我对原始帖子进行了编辑并发布了培训脚本的源代码。
    【解决方案2】:

    训练代码很好,输入根本不需要梯度,如果你只想训练和更新权重。

    真正的问题是这里的这一行

     with torch.set_grad_enabled(is_train == "train"):
    

    因此,如果您不进行训练,则希望禁用渐变。问题是is_train 是一个布尔值(从这个判断:def __call__(self, is_train=False):),所以比较总是错误的,并且不会设置梯度。改成

    with torch.set_grad_enabled(is_train):
    

    你会没事的。

    【讨论】:

    • 啊,是的。我完全忘记了这一点。有时人们忘记了导致问题的最愚蠢和最明显的事情。无论如何感谢您指出哈哈
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2019-03-05
    • 2015-03-22
    • 2013-09-18
    • 2015-03-15
    • 2015-10-03
    • 1970-01-01
    相关资源
    最近更新 更多