如何使用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
相关产品推荐
相关产品推荐

