使用PyTorch/torchvision批量裁剪缩放图像并获取缩放因子的方法
实现方案
你可以通过重写RandomResizedCrop的前向逻辑,或者手动分步调用裁剪、缩放接口两种方式获取缩放因子,两种方案均和原生变换的处理逻辑完全一致,不会改变原有裁剪缩放的行为。
方案1:自定义变换类(推荐)
继承原生RandomResizedCrop类,修改forward方法同时返回处理后的图像和缩放因子:
import torch import torchvision.transforms.functional as F from torchvision.transforms import RandomResizedCrop class RandomResizedCropWithScale(RandomResizedCrop): def forward(self, img): # 输入img为单张图像,形状为(C, H, W) top, left, crop_h, crop_w = self.get_params(img, self.scale, self.ratio) # 执行裁剪 cropped_img = F.crop(img, top, left, crop_h, crop_w) # 执行缩放 resized_img = F.resize( cropped_img, self.size, interpolation=self.interpolation, antialias=self.antialias ) # 计算缩放因子:输出尺寸 / 裁剪区域高度(固定宽高比1:1时crop_h=crop_w,用任意一个计算均可) scale_factor = self.size[0] / crop_h return resized_img, torch.tensor(scale_factor, dtype=torch.float32)
使用示例
# 初始化变换,参数和你原来的配置完全一致 scale_transform = RandomResizedCropWithScale(224, scale=(0.08, 1.0), ratio=(1.0, 1.0)) # 单张图像处理 scaled_img, scale = scale_transform(original_img) # 批量图像处理(输入形状为[B, C, H, W]) original_batch = torch.randn(8, 3, 512, 512) # 示例批量输入 scaled_batch = [] scale_list = [] for img in original_batch: scaled, s = scale_transform(img) scaled_batch.append(scaled) scale_list.append(s) scaled_batch = torch.stack(scaled_batch) scale_tensor = torch.stack(scale_list) # 最终得到形状为[B]的缩放因子张量
方案2:分步调用原生接口
如果你不想自定义类,也可以直接调用原生的参数生成、裁剪、缩放接口手动控制流程:
import torch import torchvision.transforms.functional as F from torchvision.transforms import RandomResizedCrop # 配置参数和原有变换对齐 output_size = 224 scale_range = (0.08, 1.0) ratio_range = (1.0, 1.0) # 单张图像处理示例 top, left, crop_h, crop_w = RandomResizedCrop.get_params(original_img, scale_range, ratio_range) cropped_img = F.crop(original_img, top, left, crop_h, crop_w) scaled_img = F.resize(cropped_img, output_size) scale_factor = output_size / crop_h
注意事项
- 你当前配置固定了宽高比为1:1,因此裁剪得到的高和宽相等,用任意一个计算缩放因子结果一致
- 如果后续需要适配可变宽高比的场景,可以分别计算高和宽的缩放因子:
height_scale = output_size[0]/crop_h、width_scale = output_size[1]/crop_w
内容的提问来源于stack exchange,提问作者mattroos
相关产品推荐
相关产品推荐

