【问题标题】:Loading FITS images with PyTorch使用 PyTorch 加载 FITS 图像
【发布时间】:2018-10-18 06:27:08
【问题描述】:

我正在尝试使用 PyTorch 创建一个 CNN,但我的图像需要从 FITS 格式而不是传统的 .png 或 .jpeg 等格式导入。

有没有办法使用 torch.utils.data.DataLoader 轻松完成此任务,或者在源代码中是否有一个地方可以放入一个子句,该子句将在加载时处理 FITS 文件?

我查看了文档,发现最相关的是 ToPILImage 转换器,它将张量或 ndarray 转换为 PIL 图像。

目前我正在使用如下图像加载例程:

import torch
from torch.autograd import Variable
import torch.nn as nn
import torch.nn.functional as F
import torch.optim as optim
import torchvision.datasets as dset
import torchvision.transforms as transforms
import torchvision

batch_size = 4

transform = transforms.Compose(
                   [transforms.Resize((32,32)),
                    transforms.ToTensor(),
                    ])

trainset = dset.ImageFolder(root="Documents/Image_data",transform=transform)
train_loader = torch.utils.data.DataLoader(trainset, batch_size=batch_size,shuffle=True)

天文:http://www.astropy.org/

Pytorch:https://pytorch.org/

torch.utils:https://pytorch.org/docs/master/data.html

更新:也许使用 torchvision.datasets.DatasetFolder 而不是 DataLoader,插入到我自己的 FITS 处理程序中会起作用吗?

尝试使用此类时,我收到以下错误:

AttributeError: module 'torchvision.datasets' has no attribute 'DatasetFolder'

DatasetFolder 在这个时间点上是否真的得到了torchvision 的支持?

【问题讨论】:

    标签: python pytorch astropy fits


    【解决方案1】:

    您可以使用此方法将 FITS 图像导出为pyplot.imsave() 支持的任何格式:

    from astropy.io import fits
    import matplotlib.pyplot as plt
    
    image_data = fits.getdata(r"/path/to/image.fits")
    plt.imsave("/path/to/image.png", image_data, cmap="gray")
    

    【讨论】:

    • 一个好主意,不幸的是我需要在归档时将数据保持为 FITS 格式,以便我可以在天文管道中快速轻松地使用它。
    • 我不确定是什么问题。此答案演示了从 FITS 文件加载数据,然后将其写入单独的“.png”文件。您根本不会丢失 FITS 数据。否则,我不熟悉 PyTorch,但也许有一种方法可以扩展它以读取 FITS 文件。
    • 抱歉,我的意思是,与其将 FITS 转换为 png 并将图像保存以加载到 PyTorch 中,我更专注于将 FITS 直接读取到 PyTorch 中,而不需要中间阶段它将图像复制为 png 格式-我相信您最近的答案地址。
    【解决方案2】:

    通过阅读文档和代码的某些组合,我认为您不一定想使用 ImageFolder,因为它对 FITS 一无所知。

    您应该尝试使用更通用的DataSetFolder 类(实际上它是ImageFolder 的父类)。您将向它传递一个它应该处理的扩展列表(即['.fits'] 和一个接受 FITS 文件的“加载器”函数,并且似乎应该返回一个 PIL.Image

    您甚至可以按照ImageFolder 的示例创建自己的子类。例如

    class FitsFolder(DatasetFolder):
    
        EXTENSIONS = ['.fits']
    
        def __init__(self, root, transform=None, target_transform=None,
                     loader=None):
            if loader is None:
                loader = self.__fits_loader
    
            super(FitsFolder, self).__init__(root, loader, self.EXTENSIONS,
                                             transform=transform,
                                             target_transform=target_transform)
    
        @staticmethod
        def __fits_loader(filename):
            data = fits.getdata(filename)
            return Image.fromarray(data)
    

    __fits_loader 的确切详细信息可能取决于您的 FITS 文件的详细信息。这个基本示例只使用了高级fits.getdata() 函数,它返回FITS 文件中的第一个图像数组(一些FITS 文件可能有很多扩展名和很多图像,或者有表格等)。所以这部分取决于你。

    【讨论】:

    • 感谢您的回复。这看起来确实是一个很好的方法。然而,当试图实现这个想法时,我得到了以下错误:模块'torchvision.datasets'没有属性'DatasetFolder'。
    • 事实上,看起来这是最近才添加的:github.com/pytorch/vision/pull/444 所以如果你不能使用最新版本的包,你可能不得不重新发明轮子,不幸的是,但它可能看起来仍然基本相同(例如,您可以继承 ImageFolder,尽管您必须重新实现更多的 __init__ 方法)。
    • 嗯,有道理。我想如果我只是在本地复制源代码:github.com/pytorch/vision/blob/master/torchvision/datasets/… 那么我可以调用 DatasetFolder 并在上面实现你的方法吗?那会更简单。
    • 当然,您可以将其作为临时措施。也许会提醒您升级到 PyTorch 的未来版本时可以将其删除。
    【解决方案3】:

    几周前我遇到了与@user8188120 相同的问题。从文件夹结构中读取标签时,使用@Iguananaut 的答案非常有效。如果有人偶然发现这一点并需要从 csv 文件中读取,这也可能有效:

    labels = []
    transform = transforms.Compose([
        # here go your transforms
        ])
    
    
    class MyFitsDataset(data.Dataset):
        def __init__(self, csv_path):
            # Read the csv file
            self.data_info = pd.read_csv(csv_path, header=None)
            # First column contains the image paths
            self.image_arr = np.asarray(self.data_info.iloc[:, 0])
            # the rest contain the labels
            self.label_arr = np.asarray(self.data_info.iloc[:, 1:])  # for multi-label
            self.label_arr = np.asarray(self.data_info.iloc[:, 1])  # for single-label
            labels.append(self.label_arr)
            self.data_len = len(self.data_info.index)
    
        def __getitem__(self, index):
            single_image_name = self.image_arr[index]
    
            data = pyfits.open(single_image_name, axes=2)
            data = data[0].data.astype('float32')
            data = data.reshape(IMG_WIDTH, IMG_HEIGHT, CHANNELS)
    
            img = transform(data)
    
            # Get label(class) of the image based on the pandas column
            single_image_label = self.label_arr[index]
    
            return (img, single_image_label)
    
        def __len__(self):
            return self.data_len
    

    这也避免了使用 DatasetFolder 类,该类在最新版本的 PyTorch 中仍然不可用。我希望这对某人有所帮助。

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 1970-01-01
      • 2021-02-18
      • 2019-12-30
      • 2022-11-21
      • 2019-05-02
      • 2020-05-23
      • 2020-08-07
      • 2021-02-10
      相关资源
      最近更新 更多