如何通过插值调整1D向量尺寸?PyTorch interpolate函数报错求解
1D向量在PyTorch中的插值resize方法
PyTorch的torch.nn.functional.interpolate确实只支持3D及以上张量输入,但可以通过给1D向量增加维度的方式适配,或者用其他内置方法实现,下面是几种实用方案:
方法一:维度扩展+interpolate(推荐)
给原始1D张量增加batch和channel维度,将其转为(1, 1, N)格式的3D张量,使用interpolate完成插值后再压缩回1D:
import torch # 原始1D向量 x = torch.randn(5) # 扩展为(1, 1, 5)的3D张量 x_expanded = x.unsqueeze(0).unsqueeze(0) # 插值到目标长度10,mode选'linear'对应1D线性插值 x_resized = torch.nn.functional.interpolate( x_expanded, size=10, mode='linear', align_corners=False ) # 压缩维度回到1D x_resized = x_resized.squeeze()
mode参数:1D场景下用linear对应线性插值,nearest对应最近邻插值,按需选择align_corners:设为False是更通用的安全选项,避免边界值计算偏差
方法二:使用Upsample1d层
如果需要重复使用插值逻辑,可以直接用专门的Upsample1d层,本质和方法一一致:
import torch upsampler = torch.nn.Upsample1d(size=10, mode='linear', align_corners=False) x = torch.randn(5) x_resized = upsampler(x.unsqueeze(0).unsqueeze(0)).squeeze()
方法三:手动实现线性插值(适合理解原理)
如果不想依赖内置函数,也可以手动计算线性插值:
import torch x = torch.randn(5) target_len = 10 # 生成目标位置的索引(浮点型) target_indices = torch.linspace(0, len(x)-1, target_len) # 取索引的左右整数边界 left_idx = torch.floor(target_indices).long() right_idx = torch.ceil(target_indices).long() right_idx[right_idx >= len(x)] = len(x)-1 # 处理边界溢出 # 计算插值权重 weights = target_indices - left_idx.float() # 线性插值计算结果 x_resized = (1 - weights) * x[left_idx] + weights * x[right_idx]
内容的提问来源于stack exchange,提问作者sten
相关产品推荐
相关产品推荐

