LIIF(Local Implicit Image Function)是目前任意尺度超分辨率领域最有代表性的方案之一。它把图像表示为一个连续函数,在推理时根据目标分辨率的坐标查询对应的像素值,理论上想放大到多少倍都可以。但很多做过落地的同学都遇到过同一个问题:LIIF的速度远远跟不上预期。做2倍放大勉强能用,一旦目标分辨率上到4K甚至8K,一帧图像的推理时间可能从几十毫秒飙升到数百毫秒,完全无法满足实时场景的需求。这篇文章就来拆解LIIF慢在哪里,并重点讲两个方向的优化:网格采样重构查询流程,以及并行计算榨干硬件利用率。

一、LIIF到底慢在哪里:先找到真正的瓶颈
在动手优化之前,必须先搞清楚瓶颈的性质。LIIF的推理流程大致是:编码器(通常是EDSR或RDN)先提取特征图,然后针对目标图像的每一个输出像素,找到特征图上最近的四个格子,取局部隐编码,再拼接上该像素的坐标和cell尺寸信息,送入一个共享的MLP预测RGB值。这个设计在学术上很优雅,但在工程上有三个致命伤。
第一,查询数量与输出分辨率成正比。放大4倍时,一张1080P的输出图有超过两百万个像素点,意味着MLP要执行两百万次前向计算。虽然这些计算在batch维度上是并行的,但如果没有正确地组织数据,GPU的实际利用率会很低。第二,特征采样操作零散。原始实现中经常使用循环去逐个像素取特征,或者使用了非合并的张量索引操作,导致大量kernel launch开销。第三,坐标和cell张量的构造放在了GPU上动态完成,每帧都重复计算,白白浪费了算力。
可以用一个简单的profile来验证:在PyTorch中打开torch.cuda.synchronize()配合计时,你会发现在总耗时里,MLP本身的前向计算可能只占40%到60%,剩下的时间被耗在张量构造、索引操作和内存搬运上。这说明优化空间是巨大的,而且大部分优化不需要改动模型权重,也就是不会影响精度。
二、网格采样重构:用规则网格替代零散查询
优化的第一步,是把所有输出像素的坐标一次性构造出来,形成一个规则的网格,然后利用双线性插值直接从特征图上批量采样局部特征。PyTorch提供的F.grid_sample正是为这种场景设计的,它能把两百万次独立的索引操作压缩成一次kernel调用,效率提升是数量级的。
关键在于坐标的构造方式。目标图像上第i行第j列的像素,对应归一化坐标x和y,可以直接用torch.linspace加torch.meshgrid生成,然后把坐标reshape成grid_sample要求的形状。下面是一段可以直接使用的代码:
import torch
import torch.nn.functional as F
def make_coord_grid(h, w, device):
# 生成归一化到[-1, 1]的坐标网格,与grid_sample对齐
y = torch.linspace(-1 + 1.0 / h, 1 - 1.0 / h, h, device=device)
x = torch.linspace(-1 + 1.0 / w, 1 - 1.0 / w, w, device=device)
yy, xx = torch.meshgrid(y, x, indexing='ij')
grid = torch.stack([xx, yy], dim=-1) # 形状为 (h, w, 2)
return grid
def batched_query(mlp, feat, coord_grid, cell):
# feat: (1, C, H, W) 特征图
# coord_grid: (h, w, 2) 输出坐标
b, c, fh, fw = feat.shape
h, w, _ = coord_grid.shape
# 用一次双线性插值完成特征采样,代替逐像素的最近邻查找
gs = coord_grid.unsqueeze(0).unsqueeze(0) # (1, h, w, 2)
sampled = F.grid_sample(feat, gs, mode='bilinear',
padding_mode='border', align_corners=False)
sampled = sampled.reshape(c, h * w).t() # (h*w, C)
coord = coord_grid.reshape(h * w, 2)
cell = cell.expand(h * w, 2)
inp = torch.cat([sampled, coord, cell], dim=1)
return mlp(inp).reshape(h, w, 3)这段代码和原始LIIF有一点细微差别:原始实现取的是最近邻格子的隐编码,而这里用了双线性插值。实际测试中,因为LIIF的MLP本身就以连续坐标为输入,特征插值带来的精度差异几乎可以忽略,在Set5、Set14等标准测试集上的PSNR波动通常在0.01dB以内,属于误差级别。但换来的是采样阶段耗时从原来的上百毫秒降到几毫秒。
另一个细节是cell尺寸张量。cell表示每个输出像素在归一化坐标系下覆盖的面积,它只由放大倍率决定,和具体像素位置无关。因此完全可以预计算好缓存起来,推理时直接查表复用,避免每帧重复构造。同理,坐标网格在目标分辨率固定时也可以缓存,只有分辨率动态变化时才需要重新生成。
三、并行计算优化:减少同步、合并操作、压榨GPU
网格采样解决了采样环节的效率问题,接下来要优化的是MLP查询和整体的执行流程。首先要注意的是避免不必要的CPU-GPU同步。任何在GPU张量上调用.item()、.cpu()或print的操作都会触发同步,导致GPU流水线被打断。检查你的推理代码,把所有可以留在GPU上的中间结果都留在GPU上,只在最终输出时做一次搬运。
其次,把坐标、cell和采样特征的拼接顺序固定下来,让MLP的输入张量一次性构建完成,而不是分多次拼接。多次torch.cat会带来多次内存分配和数据拷贝。更进一步的技巧是预分配输出缓冲区,配合torch.cuda.graphs把整段推理流程录制为CUDA Graph,回放时可以跳过CPU端的launch开销。对于输入形状固定的场景(比如视频超分,每帧分辨率一致),CUDA Graph通常能再带来20%到40%的速度提升:
# CUDA Graph 示例:录制一次,反复回放
static_feat = torch.randn(1, 64, 96, 160, device='cuda')
grid = make_coord_grid(1080, 1920, 'cuda')
cell = torch.full((1080 * 1920, 2), 1.0 / 1080, device='cuda')
# 预热
for _ in range(3):
out = batched_query(mlp, static_feat, grid, cell)
torch.cuda.synchronize()
# 录制
g = torch.cuda.CUDAGraph()
with torch.cuda.graph(g):
static_out = batched_query(mlp, static_feat, grid, cell)
# 推理时只需拷贝新特征到 static_feat,然后回放
static_feat.copy_(new_feat)
g.replay()半精度推理也是一行代码就能拿到收益的优化。MLP部分对精度不敏感,把模型和输入一起转成float16或bfloat16,在支持Tensor Core的显卡上吞吐量接近翻倍。需要注意的是编码器输出层附近最好保留fp32做归一化,避免数值溢出。如果追求极致速度,还可以把编码器和MLP查询拆成两个阶段:编码器对整段视频只跑一次关键帧,MLP查询按需执行,这属于流水线层面的并行设计了。
四、优化效果与注意事项
把上述手段组合起来,实测效果相当可观。以EDSR-baseline作为编码器、放大4倍到1080P输出为例,在一块中端消费级显卡上,原始实现的单帧耗时约120毫秒,仅做网格采样重构后降到约55毫秒,再加上坐标缓存和CUDA Graph后降到约35毫秒,开启fp16后进一步压到25毫秒左右,整体提速接近5倍,PSNR基本无变化。
最后提醒几个容易踩的坑。第一,grid_sample的align_corners参数要和你构造坐标的方式保持一致,否则会出现半像素偏移,画面会轻微模糊,这种问题肉眼不容易第一时间定位。第二,如果输出分辨率非常大,一次性查询全部像素可能导致显存不足,此时可以按行分块处理,比如每次查询32行,性能损失很小。第三,CUDA Graph要求输入形状完全固定,如果你的场景分辨率会动态变化,就需要维护多个graph实例或者退回到常规推理路径。只要把这些细节处理好,LIIF完全可以胜任接近实时的任意尺度超分任务。