You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何使用PyTorch加载FITS格式图像构建CNN?

当然可以轻松实现!

PyTorch的DataLoader本身只是负责批量加载、打乱数据这些流程控制,真正处理特殊文件格式的逻辑完全可以通过自定义Dataset类来实现——这也是PyTorch数据管道设计最灵活的地方,根本不需要去修改源码。

咱一步步来给你拆解解决方案:

1. 先搞定FITS文件的读取

首先需要一个能读取FITS格式的库,最常用的是astropy,先安装它:

pip install astropy

用astropy.io.fits就能轻松把FITS文件读成NumPy数组,这一步是核心。

2. 自定义Dataset类处理FITS

继承PyTorch的torch.utils.data.Dataset,在类里实现读取FITS、转成张量的逻辑就行。举个实际的例子:

import torch
from torch.utils.data import Dataset
from astropy.io import fits
import os

class FITSDataset(Dataset):
    def __init__(self, fits_dir, transform=None):
        # 遍历目录下所有FITS文件,存路径
        self.fits_paths = [os.path.join(fits_dir, f) for f in os.listdir(fits_dir) if f.endswith('.fits')]
        self.transform = transform

    def __len__(self):
        # 返回数据集大小
        return len(self.fits_paths)

    def __getitem__(self, idx):
        # 读取单个FITS文件
        fits_path = self.fits_paths[idx]
        with fits.open(fits_path) as hdul:
            # 一般FITS的图像数据在第一个HDU里,转成NumPy数组
            image_data = hdul[0].data.astype('float32')
        
        # 把NumPy数组转成PyTorch张量,注意调整维度(PyTorch是[通道, 高, 宽])
        # 如果是单通道FITS,要加一个通道维度
        if len(image_data.shape) == 2:
            image_tensor = torch.from_numpy(image_data).unsqueeze(0)
        else:
            # 多通道的话调整维度顺序(如果FITS是[高,宽,通道],转成[通道,高,宽])
            image_tensor = torch.from_numpy(image_data).permute(2, 0, 1)
        
        # 如果有transform(比如归一化、数据增强),就应用
        if self.transform:
            image_tensor = self.transform(image_tensor)
        
        # 这里可以根据你的任务返回标签,比如从文件名提取标签
        # label = ... (根据你的需求实现)
        return image_tensor # 或者 (image_tensor, label)

3. 用DataLoader加载自定义Dataset

这一步和加载普通PNG/JPG图片完全一样,直接把自定义的FITSDataset传给DataLoader就行:

from torch.utils.data import DataLoader

# 实例化Dataset
fits_dataset = FITSDataset(fits_dir='path/to/your/fits/files', transform=None)

# 用DataLoader批量加载
dataloader = DataLoader(
    fits_dataset,
    batch_size=32,
    shuffle=True,
    num_workers=4 # 根据你的CPU核心数调整
)

# 测试一下加载效果
for batch in dataloader:
    print(batch.shape) # 应该是 [batch_size, channels, height, width]
    break

关于你提到的ToPILImage转换器

如果你想用到PyTorch自带的一些基于PIL的图像增强(比如RandomResizedCrop、RandomHorizontalFlip),确实可以用ToPILImage先把张量转成PIL图像,处理完再转回张量。比如修改Dataset的__getitem__:

from torchvision.transforms import ToPILImage, ToTensor, Compose, RandomHorizontalFlip

# 定义transform链
transform = Compose([
    ToPILImage(),
    RandomHorizontalFlip(p=0.5),
    ToTensor()
])

# 实例化Dataset时传入transform
fits_dataset = FITSDataset(fits_dir='path/to/fits', transform=transform)

总结一下你的问题

  • 能不能用DataLoader轻松实现? 必须可以!只要自定义好处理FITS的Dataset,DataLoader就能无缝对接。
  • 需要修改源码吗? 完全不需要!PyTorch的Dataset类就是用来扩展自定义数据格式的,比改源码方便太多。

内容的提问来源于stack exchange,提问作者user8188120

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.27 04:11:05