简介:基于Unet网络实现天文图像降噪的源码与项目说明包,面向计算机、人工智能等专业学生及深度学习入门者,适合作为课程设计、毕业设计或图像去噪方向的项目起步参考。项目采用先为图像叠加噪声、再以Unet网络将带噪图像映射回原始图像的方式进行训练,数据生成、模型搭建与训练流程均完整提供。资源内含209个文件,以90个npy数据文件、91个png图像样本为主,配合3个Python脚本、2个ipynb教程/结果文件以及README说明文档,整体压缩包约35.94MB,目录结构清晰,便于按数据处理、模型训练和结果展示模块查阅。已有383人浏览学习。代码均在测试运行成功后上传,可在Google Colab平台直接运行,使用generate_data.py生成训练数据,并通过ipynb完成Unet模型的训练与效果对比;随附项目说明有助快速理解各文件作用,便于二次改进和迁移到其他图像降噪场景。
1. 天文图像的噪声模型,为什么降噪偏偏绕不开 Unet
一张叠加了 30 次曝光的仙女座星系照片,单帧读噪声可能在 5e- 左右,叠加后随机噪声确实会降低,但暗弱的旋臂结构往往依然淹没在背景方差里。天文降噪和普通手机拍照降噪最大的差异是:噪声不是简单的高斯白噪声,而是泊松噪声与读出噪声的混合体,而且信号本身的高动态范围让“一刀切”的滤波核完全失效。传统的均值滤波、中值滤波、小波阈值去噪,在平滑噪声的同时必然抹掉恒星点扩散函数(PSF)的细节,甚至把双星洗成单星。深度学习里的 Unet 因为其编码器-解码器结构和跳跃连接,天生适合做“保持形状的同时去掉噪声”这种任务,这就是本项目的核心思路。
如果你手头有一批 FITS 格式的天文照片,希望在不依赖专业软件的情况下自己训练一个针对性降噪模型,或者你只是想知道 Unet 作为图像分割经典网络怎么迁移到降噪任务上,这篇文章就是给你写的。整个方案只需要 Python 和 PyTorch,数据清洗、噪声模拟、模型训练、验证流程都会逐步展开。下面我会先讲清楚结构选型的原因,再给出一套可以直接跑的代码框架,最后聊几个真正影响降噪效果、但论文里很少写的细节。
2. Unet 做天文降噪的结构逻辑:编码器、跳跃连接与损失函数设计
2.1 天文降噪任务里 Unet 的编码器-解码器到底在学什么
Unet 的编码器是逐步下采样的卷积栈,每下采样一次,特征图分辨率减半、通道数翻倍。这个设计对天文图像非常友好,因为星云和星系的自相似结构可以用多尺度特征描述。编码器浅层学到的是高频边缘和点状源的锐利程度,深层学到的是大尺度背景梯度。但问题在于,如果只用编码器提取特征再直接上采样,很多暗弱细节会丢失,这正是当年全连接分割网络的问题。
解码器的工作是逐级恢复空间分辨率,它需要把深层的高层语义和浅层的边缘信息融合。值得注意的是,天文降噪和语义分割不同——分割任务是输出类别掩码,而降噪任务输出的是连续像素值。所以解码器最后一层的激活函数一般是线性或者不经过激活,直接输出估计的干净像素。在我的实现里,输入和输出尺寸保持一致,所以上采样路径不需要额外裁剪,这也是 Unet 做图像恢复比做分割更顺手的原因之一。
这里有一个关键点:Unet 并不是唯一的降噪架构,DnCNN、FFDNet 也常见。但 Unet 的优势在于它天然支持输入输出 shape 不变的端到端训练,而且跳跃连接让梯度可以更顺畅地流向浅层。对于天文图像这种信噪比空间分布不均匀的数据,浅层特征的保留尤为重要。我一般不会一开始就换 Attention U-Net 或 Res-UNet,除非基线效果不理想。
2.2 跳跃连接是保细节的关键,参数上怎么配合
跳跃连接把编码器第 i 层的输出直接拼接到解码器对应层。对于降噪任务,这一步的意义是提供原始分辨率的参考信息,相当于给解码器一条“旁路”,让它不需要从零重建细节。我在实践中发现,如果把跳跃连接去掉,模型虽然也能收敛,但星点轮廓会变模糊,PSF 的半峰全宽(FWHM)可能从 2.5 像素被抹成 4 像素。这就是为什么很多改进版本的 Unet 会引入注意力机制去控制跳跃连接的权重,而不是无脑拼接。
参数上,跳跃连接涉及一个容易忽略的点:输入图像的归一化范围。天文图像往往是 16 位 FITS,数值范围从 0 到 65535,如果直接输入网络,卷积权重初始化就崩了。我习惯先把数据线性拉伸到 [0,1] 范围。但有一个细节:如果只按全图最大最小拉伸,那么极亮的恒星会压制暗弱的细节,导致噪声模型不完整。更好的做法是先裁剪一定百分位的上下限,比如用 2% 和 98% 分位数做线性拉伸,然后再统一减均值、除方差。这样跳跃连接拼接的两层特征分布不会出现量级失衡。
2.3 损失函数选 L2 还是 SSIM?结合泊松噪声的加权方案
大部分降噪网络默认用 L2 损失,也就是均方误差。L2 损失在优化时梯度平滑,PSNR 指标也直接对应均方误差,所以用 L2 训练时 PSNR 通常涨得更快。但 L2 的缺点是容易产生过度平滑的结果,因为它在像素级对预测取平均,会抹掉细节中的高频波动。
另一种做法是组合损失:L_total = L_mse + λ * (1 - SSIM_loss)。SSIM 从亮度、对比度、结构三方面衡量感知质量,它对局部结构的保留比 L2 强。不过天文图像上直接使用 SSIM 有一个坑——SSIM 对局部窗口内的均值和方差敏感,而天文图像里星点所在区域的局部方差很大,SSIM 梯度会把这些区域当成重要区域来优化,反而可能让背景噪点被保留。我自己的方案是采用带噪声模型的损失加权:因为泊松噪声的信号方差等于信号强度,在暗弱区域方差小,在亮区方差大。因此可以使用变分近似,对每个像素的 MSE 按置信度加权。
下面是一个简化的带权损失代码:
import torch import torch.nn.functional as F def poisson_weighted_mse(pred, target, exposure=1.0): # 目标图像近似作为泊松噪声的期望值 variance = torch.clamp(target * exposure + 1e-6, min=1e-6) weights = 1.0 / torch.sqrt(variance) diff = (pred - target) ** 2 return (weights * diff).mean() def combined_loss(pred, target, lambda_ssim=0.1): mse = F.mse_loss(pred, target) ssim_val = ssim_loss(pred, target) # 这里用你选择的SSIM实现,返回 1-ssim return mse + lambda_ssim * ssim_val逻辑说明:poisson_weighted_mse中,我们把target当作真实信号的估计,用target的强度估算噪声方差。曝光时间exposure用于控制泊松噪声的相对大小。理论上亮区权重小,暗区权重大,这样模型不会只盯着亮核,而是把暗弱细节也拉起来。实际训练时我通常选用加权 MSE 加一个很小的 SSIM 项,而不是单用 L2。
参数建议:exposure根据你的数据决定。如果数据是单帧短曝光,取值 1.0 左右;如果是叠加后的图像,噪声方差已经降低,可把exposure设为 0.2~0.5,否则暗区权重会过大。lambda_ssim建议从 0.05 开始调,SSIM 项太大会导致梯度方向偏向结构相似,训练初期容易震荡。
3. 用 Python 搭一个可运行的天文图像降噪 Unet 最小项目
3.1 环境准备与依赖
这个项目只需要基础的深度学习库,不需要安装天文专业耗时依赖。Python 版本建议 3.8 或 3.10,PyTorch 2.x 都可以。如果你还没有配置环境,直接创建虚拟环境:
python -m venv astro_denoise_env source astro_denoise_env/bin/activate # Windows 下用 astro_denoise_env\Scripts\activate pip install torch torchvision numpy astropy tensorboard tqdm这里把astropy加进来了,它用来读取 FITS 文件。有些天文图像是 .fits 格式,里面包含 WCS 坐标信息和噪声参数,用astropy.io.fits读取最标准。注意torchvision可以帮助我们处理图像变换,但天文 FITS 不一定能被ImageFolder直接读取,所以数据加载一般用自定义 Dataset。
环境里最容易出问题的是 PyTorch 版本和 CUDA 版本不匹配。我用纯 CPU 训练小尺寸图像也可以,只是慢。建议先用 16x16 或 32x32 的裁剪块调试代码,确认无误后再上 GPU。
3.2 数据加载与噪声注入:泊松噪声模拟代码
真实天文图像配对数据很难获取,常见做法是用高信噪比图像加模拟噪声来构造训练对。高信噪比图可以来自哈勃数据、SDSS 数据或者你自己的长曝光叠加图。以下代码展示如何把干净图像变成带泊松噪声的退化图:
import numpy as np from astropy.io import fits def add_poisson_gaussian_noise(clean_img, gain=1.0, read_noise=10.0): """ clean_img: 归一化到 [0, 1] 的 float32 图像 gain: 电子增益 (e-/ADU) read_noise: 读噪声标准差 (e-) """ # 将归一化的强度转换回光子数:这里假设 clean_img 已经乘了一个参考满阱值 # 更合理的做法是直接用原始 e- 计数,这里为示例做线性映射 electrons = clean_img * 65535.0 / gain noisy_electrons = np.random.poisson(electrons).astype(np.float32) noisy_electrons += np.random.normal(0, read_noise, size=electrons.shape).astype(np.float32) # 转换回 [0,1] 范围并做截断 noisy_img = np.clip(noisy_electrons * gain / 65535.0, 0.0, 1.0) return noisy_img逻辑说明:先假设输入的clean_img代表归一化后的真实信号,乘以一个虚拟满阱电荷数(这里用 65535)映射到电子数目。泊松分布生成的随机数就是带信号依赖噪声的电子数,然后加一个高斯分布的读噪声,最后再除回去。这样生成的噪声模型与真实 CCD 接近。实际上真正的暗场噪声是非稳态的,但作为训练数据,这种模拟已经足够让网络学会区分散粒噪声和高斯噪声。
对训练数据,我建议不要每次都重新计算噪声,而是把干净图像保存成.npy数组,在训练循环里在线生成噪声。这样每个 epoch 噪声都会重新采样,等效于无限多的训练样本,对提升泛化性非常有帮助。
3.3 Unet 模型定义:面向图像恢复的紧凑版
经典的 Unet 最开始用于 256x256 输入,但天文图像可能非常大,为了避免显存爆炸,输入通常裁剪成 64x64 或 128x128。下面定义了一个适合降噪任务的轻型 Unet,重点是把通道数安排得比分割任务小一些,因为降噪任务不需要太深的语义信息:
import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv = nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1, bias=False), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True) ) def forward(self, x): return self.conv(x) class UNetDenoise(nn.Module): def __init__(self, in_channels=1, out_channels=1, features=[32, 64, 128, 256]): super().__init__() self.downs = nn.ModuleList() self.ups = nn.ModuleList() self.pool = nn.MaxPool2d(2) # 编码器 for f in features: self.downs.append(DoubleConv(in_channels, f)) in_channels = f self.bottleneck = DoubleConv(features[-1], features[-1]*2) # 解码器 for f in reversed(features): self.ups.append(nn.ConvTranspose2d(f*2, f, kernel_size=2, stride=2)) self.ups.append(DoubleConv(f*2, f)) self.final_conv = nn.Conv2d(features[0], out_channels, kernel_size=1) def forward(self, x): skip_connections = [] for down in self.downs: x = down(x) skip_connections.append(x) x = self.pool(x) x = self.bottleneck(x) skip_connections = skip_connections[::-1] for idx in range(0, len(self.ups), 2): x = self.ups[idx](x) skip = skip_connections[idx//2] if x.shape != skip.shape: x = nn.functional.interpolate(x, size=skip.shape[2:], mode="bilinear", align_corners=True) x = torch.cat((skip, x), dim=1) x = self.ups[idx+1](x) return self.final_conv(x)代码说明:DoubleConv是标准的两层 3x3 卷积,使用 BatchNorm 和 ReLU。features列表控制了各层通道数,从 32 逐步翻倍到 256。这个模型参数量约 8M,对小尺寸天文图块足够。解码器中的ConvTranspose2d是转置卷积,用来上采样。skip_connections保存编码器各层特征,并在解码器中用torch.cat拼接。
注意这个实现里有一个潜在的分辨率错位问题:当输入尺寸不是 2 的幂时,池化后的尺寸向下取整,上采样后的尺寸可能和跳跃连接不一致。我在代码里加入了interpolate来对齐,但这样会破坏一些空间对应关系。最好在数据加载时把图像统一裁剪成 64x64 或 128x128,从根本上避免该问题。
3.4 训练循环与超参数:batch size 怎么影响收敛
下面是一个最小训练循环,包含模型初始化、优化器和日志输出。训练时我通常使用 AdamW 而不是 Adam,因为配合权重衰减可以让网络更稳定。
device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = UNetDenoise(in_channels=1, out_channels=1).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-5) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=100, eta_min=1e-6) def train_one_epoch(loader, model, optimizer, criterion): model.train() total_loss = 0 for clean_batch in loader: clean_batch = clean_batch.to(device) # 在线加噪声 noisy_batch = add_noise_batch(clean_batch) # 你需要把噪声函数向量化或使用 numpy 处理 noisy_batch = torch.from_numpy(noisy_batch).float().to(device) pred = model(noisy_batch) loss = criterion(pred, clean_batch) optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() return total_loss / len(loader)超参数选择逻辑:batch size 小的时候 BatchNorm 的统计量不稳定,容易导致训练震荡。显存允许的情况下尽量用 32 以上。如果只能跑 8,那就把 BatchNorm 换成 GroupNorm,或者直接在DoubleConv里去掉 BatchNorm。我测试过,在天文图像上 GroupNorm 数量从 8 到 16 都有效,但 BatchNorm 在 batch size 小于 8 时确实会掉点。
学习率方面,1e-4 是起步值,配合 Cosine 衰减在 100 个 epoch 内逐渐降到 1e-6。如果发现 loss 不下降,先把学习率调小一个量级,而不是盲目增加模型容量。
3.5 训练中监控与断点保存
训练时不仅要看 loss,还要在每个 epoch 结束把模型权重保存下来。我习惯保存两份:一份是最新 epoch 的last.pth,一份是验证集 PSNR 最高的best.pth。代码如下:
best_psnr = 0 for epoch in range(epochs): train_loss = train_one_epoch(...) val_psnr = validate(model, val_loader) if val_psnr > best_psnr: best_psnr = val_psnr torch.save(model.state_dict(), "best_unet_denoise.pth") if epoch % 10 == 0: torch.save(model.state_dict(), f"checkpoint_epoch{epoch}.pth") scheduler.step()validate函数需要在无梯度模式下计算 PSNR,并且注意预测输出可能要截断到 0-1 再算指标。保存模型后,下次训练可以直接加载权重继续跑。
4. 训练调参实战:如何把 PSNR 和 SSIM 指标稳定提上去
4.1 预处理比网络结构更重要:减暗场、平场和归一化
实际天文图像拿到手之后,不能直接丢进网络。CCD 图像有偏置电平(bias)、暗电流(dark)和像素响应不均匀(flat)。这些系统性偏差如果不校准,Unet 学到的“噪声”里就会混入固定模式噪声,导致在真实数据集上泛化变差。标准的校准流程是:
from astropy.io import fits import numpy as np def calibrate(fits_path, bias_path, dark_path, flat_path): with fits.open(fits_path) as img_hdu, fits.open(bias_path) as bias_hdu, \ fits.open(dark_path) as dark_hdu, fits.open(flat_path) as flat_hdu: image = img_hdu[0].data.astype(np.float32) bias = bias_hdu[0].data.astype(np.float32) dark = dark_hdu[0].data.astype(np.float32) flat = flat_hdu[0].data.astype(np.float32) calibrated = (image - bias - dark) / np.maximum(flat, 1.0) # 去除异常像素做插值 bad_mask = ~np.isfinite(calibrated) calibrated[bad_mask] = np.nanmedian(calibrated) return calibrated这段代码里,flat要防止除零,所以用np.maximum把最小值限制到 1。暗场可能已经包含偏置,需要根据设备类型判断。做完校准后,再执行裁剪和归一化。如果跳过这步,网络会花大量容量去学偏置模式的周期条纹,而不是真正的噪声。
4.2 数据增强策略:裁剪、翻转与旋转
天文图像和自然图像不同,物体方向没有“上下”概念,所以旋转和翻转增强可以放心用。但注意,翻滚和旋转后图像中的星点 PSF 各向同性,所以不会引入伪影。我常用的增强如下:
import random import numpy as np def augment(image): # 随机90度旋转 if random.random() > 0.5: image = np.rot90(image, k=random.randint(1, 3)) if random.random() > 0.5: image = np.flip(image, axis=0) if random.random() > 0.5: image = np.flip(image, axis=1) return np.ascontiguousarray(image)注意增强是在干净图和噪声图上共用一套随机变换,否则配对被破坏。另外,不要用随机亮度抖动和色彩抖动,因为天文图像的亮度是有物理含义的,亮度扰动会破坏泊松噪声的比例关系。
4.3 训练阶段的三个关键参数:学习率调度、梯度裁剪和指数移动平均
除了学习率,梯度裁剪对 Unet 训练稳定性很有帮助。天文图像中亮恒星区域的梯度可能非常大,如果不裁剪,很容易导致 loss 炸成 nan。PyTorch 里一行代码:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)另一个实用技巧是使用指数移动平均(EMA)对模型参数做滑动平均。EMA 可以在训练末期提升 PSNR 0.2~0.5 dB。实现方式不复杂,维护一个影子权重,每次更新后按照衰减系数更新影子权重,最后用影子权重做验证。
学习率调度建议用 Cosine 或 ReduceLROnPlateau。如果使用 Cosine,初始学习率可以稍微大一点,因为后面会自动降下去。如果使用 ReduceLROnPlateau,它的参数patience建议设为 8~10,避免验证集指标微小波动导致误降。
4.4 容易过拟合的迹象与应对策略
天文图像数据量通常不大,一个深度 Unet 很容易在训练集上 PSNR 很高,验证集上效果普通。过拟合的典型表现是训练 loss 持续下降,验证 loss 在第 30 个 epoch 左右开始反弹。此时最优先的做法是降低模型复杂度:把features从[32,64,128,256]降为[16,32,64,128],参数量会缩小 4 倍,但降噪能力可能只下降 0.3dB。这比硬加 Dropout 更可靠。
另一种有效做法是 MixUp 或 CutMix 样式的数据混合,但这个在图像恢复任务里不太常用。我更推荐使用预训练权重迁移学习。如果实在没有预训练权重,就先在小尺寸裁剪块上训练,再利用上采样初始化大尺寸模型,这个过程称为渐进式训练。
另外需要注意的是,验证集应该选用和训练集不同的天体区域,而不是随机打乱。否则网络可能记住了固定星点的位置,导致验证指标虚高。一般按星表坐标切分,让训练集和验证集没有重叠的显著恒星。
5. 验证技巧:用噪声功率谱和残差图判断降噪质量,而不是只看 PSNR
5.1 PSNR 和 SSIM 的盲区
很多人在验证的时候只看 PSNR,但 PSNR 是全局像素误差的统计量,它无法告诉我们噪声是否在空间上被均匀地抑制。我在实践中遇到过这样的情况:模型把亮星周围的噪声消掉了,但背景区域出现规则的棋盘格伪影,PSNR 依然提升了 2dB。SSIM 对局部结构敏感,但对低频噪声反应比较迟缓。所以最终验证阶段,一定要回到天文图像本身的特点,用噪声功率谱来检查是否存在“过度平滑”或“伪纹理”。
噪声功率谱的计算方式:取一块没有明显星源的背景区域,对残差图(输出图像减原始干净图像)做二维 FFT,再对功率谱做径向平均。如果降噪后残差的功率谱在高频段仍接近原始噪声,说明模型没有有效抑制高频噪声;如果功率谱在某个频率出现尖峰,说明模型产生了周期伪影。
import numpy as np def radial_power_spectrum(img): # 输入二维图像,返回径向平均功率谱 fshift = np.fft.fftshift(np.fft.fft2(img - img.mean())) power = np.abs(fshift) ** 2 y, x = np.indices(power.shape) cx, cy = power.shape[1] // 2, power.shape[0] // 2 r = np.sqrt((x - cx) ** 2 + (y - cy) ** 2).astype(np.int) tbin = np.bincount(r.ravel(), weights=power.ravel()) nr = np.bincount(r.ravel()) radial = tbin / np.maximum(nr, 1) return radial这个函数返回一个一维数组,代表以中心频率为原点的径向平均功率。使用的时候,对原噪声图残差和降噪后残差分别画曲线,理想的结果是降噪后的曲线整体向下移动,而不是出现局部凸起。
5.2 残差图与径向亮度剖面
残差图的意义在于直观显示系统残留结构。把干净图和输出图相减,检查残差图中是否还有恒星的形状。如果残差在恒星位置呈现环形结构,说明 Unet 对 PSF 中心强度的估计有偏差。这时候可以考虑在损失函数中加入结构项,但更好的办法是检查预处理阶段是否做过星点对齐。如果输入的图像存在亚像素平移,Unet 很难同时保持位置精度和边缘锐度。
验证效果时,我最常做一个操作:画一条穿过亮星中心和暗弱背景的直线,对比原始噪声图、真实干净图和网络输出图沿这条线的亮度曲线。对于亮星,看 FWHM 是否保持;对于背景,看波动范围是否明显缩小。这个步骤比任何指标都直观,而且写报告时也更容易让人信服。
5.3 把模型打包进项目说明中的注意事项
如果你的最终交付物是python源代码+项目说明.zip,那么模型文件一般不要保存完整torch.save(model, ...),而是保存 state_dict。同时推荐导出为 TorchScript 或 ONNX,方便后续在 C++ 或推理库中使用。导出 TorchScript 的常见坑是 BatchNorm 层在某些操作符上不支持动态尺寸,所以固定输入尺寸最好。代码示例:
model.eval() example_input = torch.rand(1, 1, 128, 128) traced_model = torch.jit.trace(model, example_input) traced_model.save("unet_denoise_script.pt")注意trace方法会冻结控制流,如果模型中包含interpolate的动态尺寸分支,建议在导出前把输入固定为训练时尺寸。项目说明里要写清楚依赖版本、训练数据格式和最低显存要求,避免别人拿到 zip 后卡在环境配置上。我一般会在说明文件里放一个requirements.txt,并且写出 “CPU 环境下跑通完整训练流程至少需要 8GB 内存” 类似的话,这样使用者心里有底。
在最后还有一个容易被忽视的验证方法:把输出图像保存成 FITS 文件,在 DS9 里通过拉伸查看暗弱天区的噪声纹理。如果背景区域看起来像纯粹的高斯噪声,说明降噪成功;如果看到 coherent 的纹理或假结构,就要检查数据增强和归一化步骤。验证阶段多花 10 分钟看细节,胜过在指标排行榜上多追 0.1dB。
本文还有配套的精品资源,点击获取