--- created: 2025-08-03 21:19:11 tags: - "Research" - "复现" - "MDMS" - "采样算法" - "代码解析" --- `sampling.py` 中的 `generalized_steps_overlapping` 函数对应于完整的 Algorithm 2。 **Algorithm 2: MDMS Diffusion Model Sampling** **输入 (Input):** * **`y` (Low-light image):** 低光照图像。在代码中对应 `x_cond` 的一部分(未经过 `data_transform` 的部分)。`x_cond` 包含了低光照图像和先验图像。 * **`yp` (prior image):** 先验图像。在代码中对应 `x_cond` 的一部分(未经过`data_transform` 的部分)。 * **`εθ(xt, y, yp, t)` (conditional diffusion model):** 条件扩散模型。在代码中对应 `model`。 * **`S` (sampling steps):** 采样步数。 在训练阶段通常是1000[^1], 代码中通过 `seq` 的长度来控制。 * **`D` (patch locations):** 图像块位置的数量。代码中通过 `corners`, `corners1`, `corners2` 列表的长度来表示。 **代码中 `generalized_steps_overlapping` 函数的参数:** * **`x`:** 初始噪声图像 (对应 Algorithm 2 的 `Xt`)。 * **`x_cond`:** 包含低光照图像 `y` 和先验图像 `yp` 的张量。 * **`seq`:** 采样步骤序列,决定了逆向过程中的时间步。 * **`model`:** 预训练的扩散模型 `εθ`。 * **`b`:** 噪声调度 (noise schedule) 的超参数。 * **`ii`, `jj`, `osize`**: 训练与测试时用于确定图像块位置(the location of the patches)和大小(size)的参数, 训练时 embedding 的一部分, 测试时无意义 * **`eta`:** 控制随机性的超参数(DDIM 中的 η)。 * **`corners`:** 64x64 图像块左上角的坐标列表。 * **`p_size`:** 64x64图像块的尺寸。 * **`corners1`:** 96x96 图像块左上角的坐标列表。 * **`corners2`:** 128x128 图像块左上角的坐标列表。 * **`manual_batching`:** 是否手动控制批处理大小。 **算法流程及代码逐行解释:** 1. **`1: Sample Xt ∼ N(0, I)`** * **算法描述:** 从标准正态分布(均值为0,方差为单位矩阵)中采样初始噪声图像 `Xt`。 * **代码实现:** `x` 作为函数的输入传入,它应该是一个从标准正态分布采样的张量。 2. **`2: for i = S, ..., 1 do`** * **算法描述:** 从 `S` 到 1 的逆向循环。这是 diffusion model 的逆向过程(denoising process),逐步去除噪声。 * **代码实现:** `for i, j in zip(reversed(seq), reversed(seq_next)):`。`seq` 是一个递减的时间步序列(例如 `[999, 998, ..., 0]`),`seq_next` 是 `seq` 的移位版本(例如 `[-1, 999, 998, ..., 1]`)。 3. **`3: t = (i - 1) * T/S + 1`** * **算法描述:** 计算当前时间步 `t`。`T` 是训练时的总时间步数(通常为 1000[^1])。 * **代码实现:** `t = (torch.ones(n) * i).to(x.device)`。这里 `i` 就是当前循环中的时间步。 4. **`4: t' = (i - 2) * T/S + 1`** * **算法描述:** 计算上一个时间步 `t'` * **代码实现:** `next_t = (torch.ones(n) * j).to(x.device)`。这里, `j`是`i`的前一个时间步 5. **`5: Φt = 0, W = 0`** * **算法描述:** 初始化累积变量 `Φt` 和权重 `W` 为 0。`Φt` 用于累积不同尺度图像块处理后的结果。 * **代码实现:** ```python et_output = torch.zeros(x_cond.size(0), 3, x_cond.size(2), x_cond.size(3), device=x.device) x_grid_mask = torch.zeros(x_cond.size(0), 3, x_cond.size(2), x_cond.size(3), device=x.device) # 后面还有 corners1 和 corners2 对应的 x_grid_mask1 和 x_grid_mask2 的初始化 ``` 这里 `et_output` 对应于 `Φt`,`x_grid_mask` 对应于 `W`(在累加 `Md` 之前)。 6. **`6: for ps = 64 × 64, 96 × 96, 128 × 128 do`** * **算法描述:** 遍历不同的图像块尺寸 (patch sizes)。这是 Multi-Scale Sampling (MSS) 的核心部分。 * **代码实现:** 代码中通过分别处理 `corners`, `corners1`, `corners2` 以及对应的 `p_size`, `p_size1`, `p_size2` 来实现多尺度循环。 7. **`7: for d = 1, ..., D do`** * **算法描述:** 循环处理每个图像块位置。 * **代码实现:** * 对于 `corners` (64x64): `for (hi, wi) in corners:`,其中 `(hi, wi)` 是图像块左上角的坐标。 * 对于 `corners1` (96x96) 和 `corners2` (128x128) 也是类似的循环。 8. **`8: x_t^d = Crop_ps(Md ◦ Xt), y^d = Crop_ps(Md ◦ y), and y_p^d = Crop_ps(Md ◦ yp)`** * **算法描述:** * `Md`:调整图像大小并进行掩码操作(详见之前的解释)。 * `Crop_ps()`: 从图像中裁剪出指定大小和位置的图像块。 * **代码实现:** ```python xt_patch = torch.cat([crop(xt, hi, wi, p_size, p_size) for (hi, wi) in corners], dim=0) x_cond_patch = torch.cat([data_transform(crop(x_cond, hi, wi, p_size, p_size)) for (hi, wi) in corners], dim=0) ``` * `crop(xt, hi, wi, p_size, p_size)`: 从 `xt` 中裁剪出以 `(hi, wi)` 为左上角,大小为 `p_size` x `p_size` 的图像块。 * `data_transform`: 将图像数据从 [0, 1] 范围转换到 [-1, 1] 范围。 * `xt_patch` 对应多个图像块的`x_t^d`的集合, `x_cond_patch`同理 * **`Md` 在哪里?** `crop` 函数实际上隐含了 `Md` 的操作。当你从 `xt` 中裁剪出一个图像块时,你已经考虑了它的位置(通过 `hi` 和 `wi`),这与 `Md` 的掩码操作是等效的。 调整大小的操作在训练流程中完成 9. **`9: Φps = Φps + Md · εθ(x_t^d, y^d, y_p^d, t)`** * **算法描述:** 将当前图像块及其条件输入到模型中,并将模型输出与 `Md` 相乘,累加到 `Φps`。注意, 这里的乘法是逐元素乘法. * **代码实现:** ```python x_input = torch.cat([x_cond_patch[i:i + manual_batching_size], xt_patch[i:i + manual_batching_size]], dim=1) outputs = model(x_input, t, ii_input, jj_input, osize_input) for idx, (hi, wi) in enumerate(corners[i:i+manual_batching_size]): et_output[0, :, hi:hi + p_size, wi:wi + p_size] += outputs[idx] ``` * `model(...)`: 这就是 `εθ`,它接收 `x_input`(包含 `xt_patch` 和 `x_cond_patch`)和 `t` 作为输入,输出预测的噪声。 * `et_output[0, :, hi:hi + p_size, wi:wi + p_size] += outputs[idx]`:将模型输出累加到 `et_output` 的相应位置。这里已经隐含了与“Md”逐像素相乘的步骤, 因为`et_output`的其他位置初始值为0 对于不在(hi, wi)位置上的`outputs[idx]`, 会被加到`et_output`上不属于该图像块的区域, 但在后续的步骤12会被`x_grid_mask`逐元素除法(对应于算法中的除以W)消除掉。 * 对`corners1`和`corners2`进行类似的操作,分别得到 `et_output1` 和 `et_output2`。 10. **`10: W = W + Md`** * **算法描述:** 累加 `Md` 到 `W`。 * **代码实现:** ```python for (hi, wi) in corners: x_grid_mask[:, :, hi:hi + p_size, wi:wi + p_size] += 1 # 对 corners1 和 corners2 执行类似操作 ``` `x_grid_mask` 在每个图像块位置 `(hi, wi)` 处加 1,记录该位置被多少个图像块覆盖。 11. **`11: end for`** (内层循环结束,遍历 `D`) 12. **`12: Φps = Φps ⊘ W, ⊘ means element-wise divide`** * **算法描述:** 对累加结果 `Φps` 进行归一化,除以累积权重 `W`。 * **代码实现:** `et0 = torch.div(et_output, x_grid_mask)` 将 `et_output` 除以 `x_grid_mask`,进行逐元素除法,实现归一化。 * 对 `et_output1` 和 `et_output2` 进行类似操作,得到 `et1` 和 `et2`。 13. **`13: Φt = (Φt + Φps)`** * **算法描述:** 将当前尺度处理后的结果 `Φps` 累加到总的累积变量 `Φt` 上。 * **代码实现:** ```python if corners1 != None: et1 = torch.div(et_output1, x_grid_mask1) et=(et0+et1)/2.0 if corners2 != None: et2 = torch.div(et_output2, x_grid_mask2) et = (et0 +et1+ et2) / 3.0 else: et=et0 ``` 根据是否存在`corners1`和`corners2`,将不同尺度的结果`et0`, `et1`, `et2`相加或取平均, 此处 `et`对应于算法描述中的`Φt` 14. **`14: end for`** (外层循环结束,遍历 `ps`) 15. **`15: Φt = Φt / 3.`** 算法描述: 对所有尺度求平均. 在上一步已经完成 16. **`16: Xt ← ...`** * **算法描述:** 使用 DDIM 的更新公式计算上一个时间步的噪声图像 `Xt`。 * **代码实现:** ```python x0_t = (xt - et * (1 - at).sqrt()) / at.sqrt() c1 = eta * ((1 - at / at_next) * (1 - at_next) / (1 - at)).sqrt() c2 = ((1 - at_next) - c1 ** 2).sqrt() xt_next = at_next.sqrt() * x0_t + c1 * torch.randn_like(x) + c2 * et xs.append(xt_next.to('cpu')) ``` * `at` 和 `at_next` 是根据噪声调度 `b` 计算的系数。 * `x0_t` 是根据当前噪声图像 `xt` 和预测的噪声 `et` 估计的“去噪”图像。 * `xt_next` 是根据 DDIM 公式计算的下一个时间步的噪声图像。 * `c1` 和 `c2` 是 DDIM 公式中的系数,`eta` 控制随机性。 17. **`17: end for`** (最外层循环结束,遍历 `i`) 18. **`18: return Xt`** * **算法描述:** 返回最终生成的图像 `Xt`。 * **代码实现:** 返回 `xs`(包含所有时间步的中间结果)和 `x0_preds`(每一步的去噪估计)。 **总结:** `sampling.py` 中的 `generalized_steps_overlapping` 函数完整地实现了 Algorithm 2 描述的 Multi-Scale Sampling DDIM 过程。代码通过以下方式与算法对应: * **循环结构:** 代码中的嵌套循环与算法中的 `for` 循环完全对应,实现了逆向采样、多尺度处理和图像块遍历。 * **变量映射:** 代码中的变量与算法中的变量有明确的对应关系,例如 `x` 对应 `Xt`,`x_cond` 对应 `y` 和 `yp`,`model` 对应 `εθ`,`et_output` 对应 `Φps`, `x_grid_mask` 对应`W`。 * **核心操作:** 代码中的 `crop` 函数、模型调用、`et_output` 的累加和归一化、DDIM 更新公式等,都与算法中的关键步骤一一对应。 * **多尺度处理:** 通过`corners` `corners1` `corners2`实现不同patch 希望这次的逐行解释能够帮助你更深入地理解 Algorithm 2 和 `sampling.py` 代码的实现细节。 [^1]: Shang 等. 训练时 T = 1000, 采样步数 S = 25。 [^6]: Shang 等. MDMS 扩散模型采样算法。