Python中张量Total Variation正则化的索引识别与多通道扩展实现
总变差损失代码维度说明及扩展实现
1 现有代码的维度对应规则
你提供的代码是PyTorch生态下的标准实现,默认处理的是批量多通道图像张量,维度顺序为(批量数N, 通道数C, 高度H, 宽度W),四个索引从左到右的对应关系如下:
- 第1位
::批量维度,遍历当前输入的所有图像 - 第2位
::通道维度C,即你提到的c维度,例如RGB图像对应3个通道 - 第3位:高度维度H,对应你提到的i维度,代表图像垂直方向的像素坐标
- 第4位:宽度维度W,对应你提到的j维度,代表图像水平方向的像素坐标
现有代码的计算逻辑:
tv_h为高度方向总变差:通过对高度维度切片取相邻像素做差,计算所有上下相邻像素差的平方和tv_w为宽度方向总变差:通过对宽度维度切片取相邻像素做差,计算所有左右相邻像素差的平方和
2 新增通道维度总变差的实现
如果需要额外计算通道方向的总变差,只需新增通道维度的相邻差计算逻辑即可,修改后的完整代码如下:
def compute_total_variation_loss(img, weight): # 高度(i维度)总变差 tv_h = ((img[:,:,1:,:] - img[:,:,:-1,:]).pow(2)).sum() # 宽度(j维度)总变差 tv_w = ((img[:,:,:,1:] - img[:,:,:,:-1]).pow(2)).sum() # 通道(c维度)总变差 tv_c = ((img[:,1:,:,:] - img[:,:-1,:,:]).pow(2)).sum() return weight * (tv_h + tv_w + tv_c)
如果你使用的是TensorFlow默认的
(N, H, W, C)维度顺序,只需对应调整各维度的索引位置即可。
内容的提问来源于stack exchange,提问作者Keivan
相关产品推荐
相关产品推荐

