Runnable code
完整项目见 GitHub 仓库,代码按职责组织如下:
上一篇我没有使用 PyTorch,而是用 NumPy + Autograd 手动搭了一遍 TorchOptics 论文 Figure 4 中的 inverse-design pipeline:
那次实现最大的价值,是把 differentiable optics 背后的计算过程拆开来看清楚。但实际做计算光学、holography、DOE optimization 或 optical neural network 时,更常见的工具显然是 PyTorch。原因并不只是“PyTorch 可以训练神经网络”,而是它同时提供了可微光学需要的几样东西:
这篇文章继续复现同一个 inverse-design problem,但完全使用原生 PyTorch。目标是理解:
一个 Fourier optics simulator 怎样自然地变成一个 PyTorch optimization model。
一、问题场景:三平面相位调制器的逆向设计
光学系统、目标与数值参数
我们继续使用 TorchOptics 论文 Figure 4 中的实验。输入是一束 Gaussian beam $\psi_{\rm in}(x,y)$,目标是在输出平面得到四束 Gaussian beam $\phi_{\rm target}(x,y)$。光学系统中放置三块 phase-only modulator:
三块调制器分别位于 $z=0$、$0.2,{\rm m}$ 和 $0.4,{\rm m}$,最终 target plane 位于 $z=0.6,{\rm m}$,所以整个系统是:
其中每一段传播距离都是 $0.2,{\rm m}$。
论文 Listing 2 使用 $250\times250$ sampling grid、$10,\mu{\rm m}$ spacing、$700,{\rm nm}$ wavelength、$150,\mu{\rm m}$ Gaussian waist,三块 trainable phase modulator 使用 Adam、learning rate $0.1$,共优化 400 iterations。
这一次依然不直接调用 TorchOptics,而是自己实现其中的物理传播和 optimization。
从 forward problem 到 inverse problem
普通 Fourier optics simulation 做的是
也就是给定 optical system,求输出光场。例如 Angular Spectrum Method:
这是 forward problem。现在则反过来:我们已经知道 $U_{\rm in}$ 和 $U_{\rm target}$,但不知道三块 phase modulator 应该具有怎样的相位 $\phi_1(x,y)$、$\phi_2(x,y)$ 和 $\phi_3(x,y)$。真正要求的是:
这就是 inverse design。
二、整体思路:把光学系统写成 PyTorch 模型
为什么 PyTorch 适合可微光学
表面上,这不是一个 neural network:没有 convolution layer、transformer 或 dataset。但从 optimization 的角度来看,它和训练神经网络实际上是一样的。神经网络的计算链是
现在的光学系统则是
唯一的区别是,神经网络里的参数是 $W$,这里的参数是 $\phi(x,y)$。所以在 PyTorch 看来,weight 和 phase 没有本质区别。只要它们是 torch.nn.Parameter,并且从 parameter 到 loss 的整个计算过程都由 differentiable operations 构成,PyTorch 就可以自动计算 gradient。
Complex autograd 如何连接 phase 与 loss
光学传播最大的特殊之处在于场是复数:
因此传播过程中会大量出现复指数、FFT、共轭和模长运算:
1 | torch.exp(1j * phi) |
PyTorch 原生支持 complex tensors,而且 torch.fft 的 FFT 运算支持 autograd。这里有一个非常重要的要求:最终 loss 必须是 real-valued scalar,比如 $L=1-\eta$。PyTorch 对 complex computation 使用 Wirtinger calculus 处理反向传播;对于这种“内部存在 complex operations,但最终优化 real-valued loss”的问题,可以直接进行 gradient descent。
因此,下面这条计算链是完全合法的:
三、准备输入场与目标场
建立 sampling grid
先建立空间坐标:
1 | import torch |
坐标满足
simulation window 大约为
这和 NumPy 版本没有任何物理区别。
构造 Gaussian input
Gaussian field 定义为
对应的 PyTorch 函数与输入场为:
1 | def gaussian_beam(X, Y, waist, x0=0.0, y0=0.0, amplitude=1.0): |
这里依然要区分 field 和 intensity:$U(x,y)\neq I(x,y)$,真正的 intensity 是
构造 four-Gaussian target
论文原始 Listing 2 中的 target 并不是程序生成的,而是通过类似下面的代码加载已有 spatial profile:
1 | target_field = torch.load("target.pt") |
这只是对论文实现的说明,不属于本文后续计算链。本文和上一篇一样,自行构造一个 four-Gaussian target:
四个中心为 $(-d,-d)$、$(-d,d)$、$(d,-d)$ 和 $(d,d)$,这里取 $d=0.5,{\rm mm}$。
1 | def four_gaussian_target(X, Y, offset, waist): |
所以这里并不是逐 pixel 恢复论文的 target.pt,而是复现相同的 inverse-design problem:
Field normalization
定义 power:
归一化函数和调用方式为:
1 | def normalize_field(U, dx, dy): |
此时 $\iint |U|^2dxdy=1$。
四、构建可微的传播算子
ASM transfer function
H 是传播距离为 $z$ 的 Angular Spectrum Method transfer function。它不是待优化参数,而是由采样网格、波长和传播距离一次性计算出的二维复数 tensor。对于输入场中的每个空间频率分量,H 给出它传播 $z$ 之后积累的相位与衰减:
构造过程分为三步。torch.fft.fftfreq() 根据像素数和采样间隔生成与离散 FFT 一一对应的空间频率坐标 $f_x$、$f_y$,fftshift() 再把零频移动到数组中心。接着将 cycles/m 转换为 rad/m,并依据自由空间色散关系计算纵向波矢:
最后,对每个 $(k_x,k_y)$ 计算 $e^{ik_z z}$,就得到与输入频谱尺寸相同的二维传输函数 H:
1 | def make_transfer_function( |
这里 H.shape == (Ny, Nx),并与 centered FFT 后的 spectrum 逐点相乘。由于系统的三段传播距离都为 $0.2,\mathrm{m}$,三次传播可以复用同一个 H;如果三段距离不同,则应分别构造 H1、H2 和 H3。
把 kz2 转换成 complex_dtype 是故意的。前面已经令 complex_dtype = torch.complex128,因此这里会以复数形式计算平方根。理论上,当 $k_x^2+k_y^2>k^2$ 时,$k_z$ 会变成 imaginary number,对应 evanescent component。虽然当前 sampling 参数实际上不会走到这一部分,把公式写成 complex 形式仍然更完整。
Centered FFT 与传播
和 NumPy 版本一样,先定义 centered FFT:
1 | def fft2c(U): |
传播函数直接把输入场变换到频域,与 H 相乘,再变换回空间域:
1 | def propagate(U, H): |
这些全部都是 torch.* operations,因此 PyTorch 会记录 $U\rightarrow\mathrm{FFT}\rightarrow H\rightarrow\mathrm{IFFT}$。torch.fft 本身支持 autograd,所以 gradient 可以穿过整个传播过程。
Phase modulation
phase-only SLM 满足
PyTorch 实现很直接:
1 | def phase_modulate(U, phi): |
这里 phi 是 real tensor,但 torch.exp(1j * phi) 会变成 complex tensor。
五、把物理系统封装为 nn.Module
Parameter 与 buffer:区分可训练参数和固定状态
前面已经定义了 phase_modulate()、propagate() 和传播传输函数 H。接下来把三块相位调制器及三段自由空间传播组合成一个完整模型。与 NumPy + Autograd 版本不同,这里不再分别维护 phi1、phi2、phi3 及其梯度,而是把它们注册为 torch.nn.Parameter,让 PyTorch 自动收集和更新。
传播传输函数 H 也属于模型状态,但它由波长、采样网格和传播距离决定,不应该被 optimizer 更新。因此用 register_buffer() 保存它,而不是把它定义成 Parameter。
1 | class ThreePlanePhaseSystem(torch.nn.Module): |
三张 phase map 都从零开始,因此模型初始状态对应三块透明的 phase-only modulator。model.parameters() 会返回 phi1、phi2 和 phi3,而 H 不会出现在 optimizer 的参数列表中。与此同时,H 仍会出现在 model.state_dict() 里,并且在调用 .to(device) 时与三张 phase map 一起移动到目标设备。
实例化模型并执行完整前向传播
类定义完成后,需要把前面构造的 H 传入模型,并将整个模型移动到与 U_in、U_target 相同的设备。由于三段传播距离相同,三个 buffer 都传入同一个 transfer function:
1 | model = ThreePlanePhaseSystem( |
model(U_in) 会自动调用 forward()。其内部计算链为:
因此,nn.Module 并不等于 neural network;它可以封装任意由 PyTorch operations 构成的 differentiable computational model。这里封装的不是神经网络层,而是一个具有三张可训练 phase map 的物理光学系统。至此,前面分别定义的调制函数、传播函数和 transfer function 已经连接成一条可以实际调用的 forward path,下一步只需要为 U_out 定义 loss 并执行反向传播。
六、定义损失并执行优化
Mode-overlap loss
论文使用:
这里写成更一般的 normalized mode overlap:
1 | def mode_overlap(U_out, U_target, dx, dy): |
损失写成 loss = 1 - eta。如果 $U_{\rm out}=U_{\rm target}$,那么 $\eta=1$,所以 $L=0$。
loss.backward() 如何计算 phase gradient
最核心的一行是:
1 | loss.backward() |
它实际上是在同时计算
反向传播会沿着整个 computational graph 逆向执行:
PyTorch 的 autograd 会在 forward execution 时动态记录这些 tensor operations,并在 backward 时利用 chain rule 计算 gradient。计算后,model.phi1.grad、model.phi2.grad 和 model.phi3.grad 分别保存
Adam 与 optimization loop
上一篇我们手写了 Adam,这次直接把模型参数交给 PyTorch:
1 | optimizer = torch.optim.Adam(model.parameters(), lr=0.1) |
model.parameters() 会自动返回 phi1、phi2 和 phi3,但不会包含 H,因为 H 是 buffer,不是 Parameter。完整训练循环为:
1 | for iteration in range(400): |
这几行就是整个 inverse design。
七、PyTorch 实现中的边界与注意点
为什么必须清除累计梯度
PyTorch 的 gradient 默认是累加的。第一次 loss.backward() 得到 $g_1$,第二次如果不清零,就会得到 $\mathrm{grad}=g_1+g_2$。但这里希望每个 iteration 使用当前 loss 对参数的 gradient,因此每轮训练前要调用:
1 | optimizer.zero_grad() |
更常见的写法是:
1 | optimizer.zero_grad(set_to_none=True) |
detach() 只能用于计算图之外
训练过程中不能随便写 U = U.detach(),因为 detach() 意味着从这里切断 computational graph。如果执行
1 | U = propagate(U, H) |
loss 就不能再通过 $U$ 反向传播到前面的 phase modulator。因此,detach() 只应该用于 plot、save 或 NumPy conversion 等计算图之外的操作。
优化结束后,可以先在不记录梯度的条件下重新计算输出,再转换成 NumPy array 用于绘图或保存:
1 | with torch.no_grad(): |
torch.no_grad() 已经阻止这次 forward 建立计算图,因此这里不再需要额外调用 detach()。如果处理的是训练过程中已经位于计算图内的 tensor,则应使用 tensor.detach().cpu().numpy()。
CPU、CUDA 与 Apple Silicon
如果有 NVIDIA GPU,可以使用 device = torch.device("cuda"),否则使用 device = torch.device("cpu")。PyTorch 的 FFT 可以运行在 CPU 或 CUDA accelerator 上,并且支持 autograd。
对于这篇 reference implementation,我建议优先保证
尤其在 Apple Silicon 上,MPS 的 complex FFT 后端目前仍存在一些特定 FFT shape / dimension 的兼容性问题,因此第一次验证算法时用 CPU 更稳妥。代码里因此使用:
1 | device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
八、从 NumPy + Autograd 到 PyTorch
可微光学训练循环的核心五步
把所有细节去掉以后,PyTorch 版本的核心就是:
1 | model = ThreePlanePhaseSystem( |
所以 differentiable optics 最核心的模式其实就是:
这里完全没有要求 physical model 必须是 neural network。它可以是 ASM、Fresnel propagation、Rayleigh-Sommerfeld propagation,也可以包含 lens、aperture、aberration、SLM、DOE、detector 或 PSF model。只要 $\text{parameter}\rightarrow\text{output}$ 这一整条计算链可微,就可以做 inverse design。
还能优化哪些光学参数
目前只优化 $\phi_1$、$\phi_2$ 和 $\phi_3$,但一旦 optical simulator 已经进入 PyTorch,其他参数同样可以成为 torch.nn.Parameter,例如 lens focal length $f$、传播距离 $z$、Zernike coefficients $a_1,a_2,\ldots,a_N$、DOE height profile $h(x,y)$、refractive-index distribution $n(x,y,z)$,甚至 wavelength-dependent optical parameters。
例如 phase DOE:
如果写成
1 | h = torch.nn.Parameter(...) |
PyTorch 就可以计算 $\partial L/\partial h$。这时优化的就不再只是数字 phase map,而可以逐渐靠近真正可制造的 optical element。
Differentiable optics 不等于 AI
下面两行代码看起来非常像 deep learning:
1 | loss.backward() |
但这个系统其实没有任何 AI。模型完全由 Maxwell / scalar diffraction physics 决定,不存在 training data、classification、neural network 或 feature extraction。它本质上是
PyTorch 只是计算工具,真正发生的是
这也是 differentiable optics 最有意思的地方:它把过去两个比较分离的世界连接起来。传统 computational optics 是
现代 optimization 则是
最终形成
而当这个 physical model 后面再连接 CNN、Transformer 或 reconstruction network 时,才真正进入所谓的 $\text{end-to-end optical + digital co-design}$。
TorchOptics 本身的设计目标也正是利用 PyTorch 实现 GPU-accelerated、fully differentiable Fourier-optics simulation,并支持 optical hardware 与 digital model 的联合优化。