news 2026/9/20 18:00:59

深度学习驱动计算全息图生成:UNet实现实时相位恢复

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
深度学习驱动计算全息图生成:UNet实现实时相位恢复

简介:一份关于基于深度学习的计算全息图生成算法研究的论文复现与总结资料,内容聚焦如何用卷积神经网络替代传统迭代优化算法,在保证全息图质量的同时将生成速度提升一个数量级。资料面向全息显示、增强现实及计算成像领域的研究人员与工程师,尤其适合希望将深度学习引入全息图生成的初学者和进阶者。文档以docx格式呈现,共1个文件,压缩包大小约52KB,内容包含论文核心思想概括、带残差块的CNN模型完整代码、复数域损失函数设计、动态加权调整因子以及针对二维/三维/彩色图像的数据集生成与预处理方法,并配有详细注释和解释,便于对照复现。目前已有97人学习浏览。通过这份资料,读者可以快速理解模型架构与训练技巧,掌握使用TensorFlow构建全息图生成器的完整流程,同时了解迭代角谱法生成标签数据的具体实现,为进一步开展全息显示应用研究提供扎实基础。

1. 项目背景与整体思路拆解

1.1 全息图生成的痛点在哪里

计算全息图(Computer-Generated Holography,CGH)这个方向,我从入门到现在折腾了快一年。最早是被"用算法直接生成可以重建出三维场景的相位图"这个概念吸引的——你输入一张普通图片,经过计算得到一张看起来毫无规律的相位图,再用空间光调制器(SLM)照射,就能在空间中看到立体的光场重建。这种"无镜头成像"的感觉,确实有一种工程上的浪漫。

但传统CGH算法有个致命痛点:计算慢。经典的Gerchberg-Saxton(GS)算法和它的各种变体,需要在物平面和频谱平面之间反复迭代几十次甚至上百次,每次迭代都涉及两次傅里叶变换。一张512x512分辨率的相位图,在CPU上跑可能要几十秒到几分钟,即便上了GPU,实时生成也捉襟见肘。更麻烦的是,迭代算法对初始相位非常敏感,初始值选得不好,最后重建出的图像边缘会有明显的噪声和振铃。

这个问题的本质是什么?传统迭代算法是在"猜"一个相位分布,让这个相位经过光传播后能逼近目标振幅。而深度学习做的,是用大量样本去学习"目标图像到最优相位"之间的映射关系。一旦网络训练完成,前向推理只需要一次,几十毫秒就能输出相位图,这就把CGH从"离线计算"变成了"实时生成"。

1.2 为什么选择深度学习这条路

我在这篇文章里要复盘的就是我这一整套基于深度学习的计算全息图生成方案。整体思路并不复杂:训练一个卷积神经网络,输入是目标图像(也就是你想重建出来的画面),输出是对应的纯相位全息图,然后用角谱法(Angular Spectrum Method,ASM)在物理层面模拟光的传播,把相位图重建为目标图像,计算重建图和原图之间的误差,反向传播更新网络。

这个方案的核心优势在于:网络直接学习的是"图像到相位"的非线性映射,训练完成后完全不需要迭代,一次前向传播直接出结果。如果你愿意把一些参数(比如重建距离、波长)一起编码进网络,甚至可以做到动态调节重建距离。这让它在实时全息显示、微型投影、增强现实等方面非常有应用价值。

适合来读这篇文章的人,我默认是有一定深度学习基础、想快速上手计算全息方向的同学。你不需要懂复杂的光学理论,但至少要会PyTorch的基本操作,懂卷积网络的基本原理。我会把光学部分尽量讲得直白——只要会傅里叶变换,你就能理解计算全息的核心逻辑。

2. 核心原理与关键细节解析

2.1 计算全息的数学基础,其实就一个公式

要理解深度学习如何生成全息图,首先得理解光是如何传播的。计算全息的核心可以用一句话概括:光在自由空间里传播,本质上是一个线性系统,可以用傅里叶变换来描述。

光传播的角谱理论说,一个平面的复振幅分布 $U_0(x, y)$,在经过距离 $z$ 的传播后,其频谱 $U_z(f_x, f_y)$ 等于初始频谱乘以一个传播因子:

$$U_z(f_x, f_y) = U_0(f_x, f_y) \cdot H(f_x, f_y)$$

其中传播函数:

$$H(f_x, f_y) = \exp\left(j \frac{2\pi z}{\lambda} \sqrt{1 - (\lambda f_x)^2 - (\lambda f_y)^2}\right)$$

这里的 $\lambda$ 是波长,$f_x$、$f_y$ 是空间频率。整个过程相当于是把入射光场先傅里叶变换到频域,乘上一个相位因子,再逆傅里叶变换回到空间域。

这也解释了为什么计算全息图看起来是"花花的"——相位图本质上是在频域和空间域之间做变换的结果,它把目标的振幅信息编码进了相位里。而深度学习要做的就是学会这种编码方式,而不是像GS算法那样迭代地去逼近。

需要特别强调的是,这里的"距离" $z$ 直接影响传播函数的相位曲率。$z$ 越大,相位变化越快,高频分量的采样就越容易出问题。很多初学者在实现ASM时重建效果一团糟,八成就是空间频率的坐标没有对应好,或者距离过大导致混叠。我后面会给出具体代码和注意事项。

2.2 网络架构选型:UNet家族依然是主力

计算全息图生成的任务本质上是"图像到图像"的翻译任务,输入是目标亮度图,输出是相位图。这类任务天然适合编码器-解码器架构,而UNet是这一家族里最经典、最稳的选择。

UNet之所以好用,关键在它的跳跃连接(skip connection)。编码器逐渐下采样提取高层次特征,解码器逐步上采样恢复空间分辨率,跳跃连接把浅层的细节特征直接拼接到深层,保留了大量边缘和纹理信息。这对全息图生成非常关键——重建图像的边缘锐度直接取决于相位图中的高频细节,如果这些信息在编码过程中被丢失,重建结果就会发糊。

我在实际项目中,网络结构选择的是在UNet基础上略微调整后的变体。主干是4层编码器、4层解码器,每层两个3x3卷积加ReLU激活,通道数依次是64、128、256、512。解码器的最后一层用1x1卷积把通道数压到1,输出就是纯相位图的单通道。输入是灰度图的话就单通道输入,彩色图可以把RGB三个通道分别处理或堆叠输入。

关于相位输出,我在实践中试过两种方式。一种是直接回归相位值(0到2π),但这样会在相位跳变处引入难训练的不连续性;另一种是我最终采用的方案:输出两个通道,分别表示相位的正弦和余弦分量。这样做的好处是相位值是连续循环域,用 $\text{atan2}$ 从正弦余弦恢复出相位,避免模型被 0/2π 边界搞晕。这个细节很多人容易忽略,但对训练稳定性影响极大。

2.3 损失函数怎么设计才能重建得清晰

损失函数决定了网络优化的方向。计算全息的最终目标是让重建图像在人类视觉上接近原图,因此损失函数必须围绕"重建质量"来设计,而不是直接约束相位图的像素误差。

我采用的损失函数由两部分组成。第一部分是重建图的振幅与目标图之间的L1损失。为什么要用L1而不是L2?因为L2损失本质上是平方误差,对异常值(比如重建图上的孤立噪点)惩罚过重,容易导致重建图偏平滑;L1损失对边缘更友好,重建出的图像细节更锐利。

第二项是感知损失(Perceptual Loss),用VGG16的某个中间层输出做特征匹配。这里的思想是:像素级别的误差小,不代表人眼看着舒服;而深层特征的距离能更好地反映图像的语义和结构相似性。加入感知损失后,重建图像的噪点和斑块效应明显减少。

整个损失函数长这样:

$$\mathcal{L} = \lambda_1 \cdot \mathcal{L}{L1}\left(I{rec}, I_{target}\right) + \lambda_2 \cdot \mathcal{L}{perceptual}\left(I{rec}, I_{target}\right)$$

其中 $\lambda_1$ 取0.6,$\lambda_2$ 取0.4。需要注意,计算L1损失时,建议在复振幅域计算振幅 $|U|$,或者直接计算强度 $|U|^2$。我个人经验是振幅域的损失训练更稳定,强度域虽然在视觉上更贴近人眼感知,但在训练初期容易梯度爆炸。

3. 完整实现:从数据准备到训练推理

3.1 光传播模块代码实现

先上最核心的光传播模块,也就是角谱法(ASM)的PyTorch实现。这是整个项目的地基,所有损失计算和重建验证都建立在它上面:

import torch import torch.fft def asm_propagate(field: torch.Tensor, wavelength: float, pixel_size: float, distance: float): """ 角谱法光传播 Args: field: 复振幅场 (..., H, W) wavelength: 波长(米),比如 532e-9 pixel_size: SLM像素间距(米),比如 8e-6 distance: 传播距离(米),比如 0.1 Returns: 传播后的复振幅场 (..., H, W) """ H, W = field.shape[-2], field.shape[-1] # 生成频率坐标,注意是物理频率,不是像素索引! # 采样频率范围是 [-1/(2*dx), 1/(2*dx)],间隔是 1/(dx*N) fy = torch.fft.fftfreq(H, d=pixel_size, device=field.device) fx = torch.fft.fftfreq(W, d=pixel_size, device=field.device) fx, fy = torch.meshgrid(fx, fy, indexing='xy') # 传播函数 H = exp(j * k * z * sqrt(1 - (lambda*fx)^2 - (lambda*fy)^2)) k = 2 * torch.pi / wavelength temp = 1 - (wavelength * fx) ** 2 - (wavelength * fy) ** 2 # 为了避免负值开根号产生 NaN,把负区域置零(倏逝波部分直接忽略) temp = torch.clamp(temp, min=0.0) h = torch.exp(1j * k * distance * torch.sqrt(temp)) # 频域乘法 + 逆傅里叶 field_fft = torch.fft.fft2(field) out = torch.fft.ifft2(field_fft * h) return out

这段代码是整个项目的枢纽。我再解释几个关键点。

第一,空间频率的坐标生成。很多人直接用np.fft.fftfreq产生的像素频率,然后去乘波长和距离,单位混乱导致重建完全跑偏。这里的d=pixel_size是物理坐标系,$f_x$ 的单位是"1/米",波长乘频率后得到无量纲的量,这样才能正确计算传播函数的相位。这一步是我调试时排查最久的坑之一。

第二,torch.clamp的作用。当 $\lambda f_x > 1$ 时,开根号内变成负数,出现"倏逝波"(evanescent wave)。在数值模拟中这部分频率分量实际上不传播,直接置零处理即可,避免产生NaN。

第三,传播距离。在我的实验中发现,如果用8微米像素间距和532纳米波长的SLM,重建距离设置在0.05米到0.2米之间效果比较理想。距离太近,相位图变化过大;距离太远,高频信息丢失严重。

3.2 网络模型与相位生成

接下来是生成相位的网络结构。我用的是一个精简版UNet,输入为目标图像的灰度图,输出为相位编码(正弦+余弦双通道):

import torch import torch.nn as nn import torch.nn.functional as F 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), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): return self.conv(x) class UNetCGH(nn.Module): def __init__(self, in_ch=1, out_ch=2): super().__init__() # 编码器 self.e1 = DoubleConv(in_ch, 64) self.e2 = DoubleConv(64, 128) self.e3 = DoubleConv(128, 256) self.e4 = DoubleConv(256, 512) self.pool = nn.MaxPool2d(2) # 解码器 self.up3 = nn.ConvTranspose2d(512, 256, 2, stride=2) self.d3 = DoubleConv(512, 256) self.up2 = nn.ConvTranspose2d(256, 128, 2, stride=2) self.d2 = DoubleConv(256, 128) self.up1 = nn.ConvTranspose2d(128, 64, 2, stride=2) self.d1 = DoubleConv(128, 64) # 输出层 self.out = nn.Conv2d(64, out_ch, 1) def forward(self, x): # 下采样阶段 x1 = self.e1(x) x2 = self.e2(self.pool(x1)) x3 = self.e3(self.pool(x2)) x4 = self.e4(self.pool(x3)) # 上采样阶段,跳跃连接拼接 x = self.up3(x4) x = self.d3(torch.cat([x, x3], dim=1)) x = self.up2(x) x = self.d2(torch.cat([x, x2], dim=1)) x = self.up1(x) x = self.d1(torch.cat([x, x1], dim=1)) x = self.out(x) # 输出相位编码,最后一层用 tanh 把值压到 [-1, 1] return torch.tanh(x) def phase_from_encoding(encoding): sin_comp = encoding[:, 0, :, :] cos_comp = encoding[:, 1, :, :] phase = torch.atan2(sin_comp, cos_comp) # 将相位映射到 [0, 2*pi] phase = (phase + 2 * torch.pi) % (2 * torch.pi) return phase

这里有几个细节值得展开说。

编码器每层通道数从64翻倍到512,是考虑到全息图的高频信息丰富,需要足够的表示能力。但通道数也不是越大越好,我试过768甚至1024通道,训练速度几乎慢了一倍,重建精度提升却非常有限。512在这个任务里是一个性能和质量的平衡点。

输出层的激活函数用tanh而非直接线性输出,是因为tanh的输出范围天然匹配正弦余弦编码的范围 $[-1, 1]$,这相当于对网络的输出做了一次隐式的归一化,训练更稳定。后面通过atan2恢复相位的时候,无论sin_compcos_comp的比例被预测成什么样,都能得到一个循环一致的相位值。

3.3 数据集构建与训练循环

数据集方面,我用的是Flickr Faces数据集加上部分ImageNet的灰度图像,统一缩放到256x256的分辨率。每张图做随机裁剪、翻转、轻微旋转的数据增强。批量大小设为16,优化器用Adam,学习率初始为1e-4,每15个epoch衰减为原来的十分之一,总共训练60轮。

数据的标签不需要额外制作——输入图像本身就是标签。训练流程是这样的:

# 伪代码,完整逻辑已抽象为可复现步骤 phase_encoding = cgh_net(amplitude_target) # 网络前向,生成相位编码 phase = phase_from_encoding(phase_encoding) # 从编码恢复相位 # 构建复振幅场:振幅均匀为1,相位为预测值 field = torch.exp(1j * phase) # 纯相位调制 # 光传播到像平面 reconstructed = asm_propagate(field, wavelength=532e-9, pixel_size=8e-6, distance=0.1) # 计算重建振幅和损失 rec_amp = torch.abs(reconstructed) loss = l1_loss(rec_amp, amplitude_target) + 0.4 * perceptual_loss(rec_amp, amplitude_target) loss.backward() optimizer.step()

注意这里纯相位全息图意味着振幅恒为1,相位图中包含了全部信息。这是SLM的物理特性决定的——大多数空间光调制器只能调制相位,不能同时调制振幅。所以网络学习的本质是:如何在振幅受限的条件下,用纯相位分布逼近目标图像的光场。

训练过程中需要监控的指标有两个。一是重建图和目标图之间的PSNR,这是我判断收敛状态的主要指标;二是相位图的规范化梯度,如果梯度接近0,说明相位图已经饱和,可能出现了相位缠绕聚类的问题。

def compute_psnr(img1, img2, max_val=1.0): mse = torch.mean((img1 - img2) ** 2) psnr = 20 * torch.log10(torch.tensor(max_val) / torch.sqrt(mse)) return psnr.item()

3.4 推理阶段:从相位图到实际重建

训练完成后,推理阶段非常轻量。只需要把一张新图像输入网络,得到相位编码,然后用phase_from_encoding恢复相位图,把相位图送入asm_propagate即可重建出目标图像。

需要注意的是,实际部署到SLM上时,需要考虑SLM的相位调制范围。一般商用SLM的相位调制范围只有0到2π,刚好覆盖我们的相位输出范围,不需要额外处理。但有些SLM的调制范围不是完整的2π,这时候需要做相位校正,把预测相位映射到设备的可用范围内,否则重建画面会产生不可预测的畸变。

如果在CPU上推理,256x256的图大概需要60毫秒;在消费级GPU(比如RTX 3060)上只需要5到8毫秒,足够支撑实时全息显示的帧率需求。这个速度相较于传统GS算法动辄数秒的迭代,提升是质的飞跃。

4. 常见问题与排查技巧实录

4.1 重建图像模糊,细节丢失怎么办

这是训练初期最容易碰到的问题。网络输出看起来"差不多",但重建图的边缘始终是模糊的,高频细节出不来。

我排查这个问题的经验是,先确认光传播模块是否正确——用一个已知的相位图跑一遍前向传播和重建,看能不能还原。比如用球面波的相位图,重建后应该是聚焦的一个亮点,如果不是,说明ASM实现有bug。

排除物理模块后,问题大概率出在网络容量上。把解码器的通道数增加一倍,比如从64起步改成128起步,通常能明显改善。另外检查有没有用感知损失,L1损失往往倾向于低频正确而高频缺失,加入感知损失后高频细节会有显著提升。

最后提醒一个容易忽略的细节:训练数据和测试数据的振幅分布要统一。如果训练集里所有图像的灰度都集中在0到0.5,而测试图亮部接近1,网络会输出失真的相位。建议训练前把所有图像的像素值归一化到[0, 1]均匀分布,或者分位数归一化。

4.2 相位图出现"棋盘格"伪影

这个非常经典。表现为重建图像表面有一层细密的网格状纹理,像棋盘格一样。原因通常是网络在上采样过程中使用了转置卷积,转置卷积在高频处会引入周期性痕迹。

我的解决办法是:把上采样操作从ConvTranspose2d换成Upsample(scale_factor=2, mode='bilinear')加普通3x3卷积的组合。双线性插值不带来额外的可学习参数,也几乎不产生棋盘格伪影。这样改动之后,重建图清晰度反而有小幅提升,因为棋盘格伪影本身就是干扰信息。

另一种可能性是输入图像本身有摩尔纹或周期纹理。这种情况可以适当增大重建距离,让高频噪声在传播过程中自然衰减。

4.3 训练不收敛或loss震荡严重

训练初期的loss曲线如果剧烈震荡,先检查学习率。全息图生成任务的损失面比常规图像任务更崎岖,因为相位到振幅的映射涉及三角函数,梯度方向可能高度非线性。把初始学习率从1e-4降到3e-5,很多震荡问题就消失了。

另一个思路是使用梯度裁剪(gradient clipping),设置max_norm=0.5。这在相位回归任务里非常有用,因为相位在0到2π边界处的损失函数不平滑,容易产生过大的梯度。

网络对批大小的敏感度也值得注意。我试过batch size为4时,训练几乎无法收敛;提升到16后,loss下降曲线变得平滑。如果显存不够,可以尝试降低分辨率而不是降低batch size,256x256的batch 8、128x128的batch 32效果都不错。

4.4 迁移到其他数据集效果变差

如果训练好的网络在测试集上效果不错,但换成完全不同的数据集后重建效果退化明显,要考虑数据分布差异的问题。CGH网络学习的是光学映射规律,按理说不应该过拟合数据集,但实践中确实存在数据偏置。

我处理这种情况一般用的是微调策略:用原来的预训练权重做初始化,在新数据集上用小学习率(1e-5)微调几个epoch。因为网络已经学到了基本的相位生成规律,新数据只需要小幅调整特征分布即可。如果新数据集和原数据集差异极大(比如从人脸换成文字),建议把初始学习率适当加大,让网络更快适应新特征。

5. 个人经验与效果总结

我跑完整个项目后最深的体会是:计算全息和常规CV任务最大的区别在于,任何网络结果的评估都要经过"光传播"这一物理过程的"中转"。你在训练时觉得loss降得不错,可能重建出来就是不行;反之某一步loss看起来不好,重建效果却意外地好。这个过程中,物理模型的正确性比网络结构更关键。ASM模块一旦有细微的bug,后面所有努力都是白费。

另外有一个实操层面的心得:训练过程中的可视化很重要。建议每个epoch输出几张验证集的重建图,把重建图和目标图并排保存下来。loss曲线可能骗人,但图片不会。我靠这个方法发现了好几次问题——比如网络学到了相位图平均值为固定值,导致重建图虽然整体亮度对了但结构完全混乱,这种"假收敛"不通过可视化根本发现不了。

最后,我想分享一个后续可以扩展的方向。这套网络生成的是静态相位图,只对应固定的重建距离和固定的观察视角。如果想让全息图随着观察角度变化、或者实现动态变焦显示,可以在网络输入中额外加入观察角度或重建距离的编码,把条件信息喂给网络。我简单实验过把重建距离作为一维条件拼接到特征图上,效果很有限,估计需要更充分的条件注入方式,比如FiLM层。这个方向我觉得很有意思,后续会继续深挖。

对刚接触这个方向的朋友,我的建议是:先别急着调网络结构,严格按照文章里的顺序把物理模块和训练流程跑通,用标准的UNet拿到一波结果,再考虑改进。基础流程顺了,后面每一步优化你都心里有底。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/20 17:59:56

企业数字化转型中的IAM五大核心引擎解析

1. 项目概述在数字化转型浪潮中,企业IT架构正经历着从"以系统为中心"向"以身份为中心"的深刻变革。传统身份管理(IAM)系统往往被简单视为"账号管家",主要负责用户认证和权限分配这类基础工作。但现…

作者头像 李华
网站建设 2026/9/20 17:59:35

Flutter RSA加密在鸿蒙平台的适配与实践

1. 项目背景与核心价值在移动端开发领域,数据安全始终是重中之重。RSA算法作为非对称加密的经典实现,广泛应用于身份认证、数据传输等场景。但传统RSA实现往往存在几个痛点:密钥管理复杂、跨平台兼容性差、性能优化不足。simple_rsa库的出现为…

作者头像 李华
网站建设 2026/9/20 17:57:29

OpenClaw、Hermes、Claude Code、Codex CLI四大AI Agent对比与选型指南

最近后台和读者群里被问爆了一个问题:OpenClaw、Hermes Agent、Claude Code、Codex CLI这四个AI Agent到底有什么区别?到底该装哪个?我自己从春节后陆续把四个工具都装了一遍,有的在Mac上跑,有的丢到Linux服务器上&…

作者头像 李华
网站建设 2026/9/20 17:49:42

Windows 10声卡没声音?驱动重装全攻略:排查、卸载、安装与避坑

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华