修改3D-ResNets数据增强后遇RuntimeError:无法调整不可变存储
我基于3D-ResNets-PyTorch仓库开发数据增强脚本,原脚本包含center cropping、corner cropping和horizontalflip三种数据增强方式,现在只需要保留center cropping与horizontalflip。我注释掉了get_train_utils函数里随机裁剪(random)和角落裁剪(corner)的代码,仅保留中心裁剪逻辑,运行时却报错:runtime error trying to resize storage that is not resizable,怀疑该问题与DataLoader无法加载数据有关,相关代码如下:
训练数据处理函数
def get_train_utils(opt, model_parameters): assert opt.train_crop in ['random', 'corner', 'center'] spatial_transform = [] #if opt.train_crop == 'random': # spatial_transform.append( # RandomResizedCrop( # opt.sample_size, (opt.train_crop_min_scale, 1.0), #(opt.train_crop_min_ratio, 1.0 / opt.train_crop_min_ratio))) #elif opt.train_crop == 'corner': # scales = [1.0] # scale_step = 1 / (2**(1 / 4)) #for _ in range(1, 5): # scales.append(scales[-1] * scale_step) #spatial_transform.append(MultiScaleCornerCrop(opt.sample_size, scales)) if opt.train_crop == 'center': spatial_transform.append(Resize(opt.sample_size)) spatial_transform.append(CenterCrop(opt.sample_size)) normalize = get_normalize_method(opt.mean, opt.std, opt.no_mean_norm, opt.no_std_norm) if not opt.no_hflip: spatial_transform.append(RandomHorizontalFlip()) if opt.colorjitter: spatial_transform.append(ColorJitter()) spatial_transform.append(ToTensor()) if opt.input_type == 'flow': spatial_transform.append(PickFirstChannels(n=2)) spatial_transform.append(ScaleValue(opt.value_scale)) spatial_transform.append(normalize) spatial_transform = Compose(spatial_transform) assert opt.train_t_crop in ['random', 'center'] temporal_transform = [] if opt.sample_t_stride > 1: temporal_transform.append(TemporalSubsampling(opt.sample_t_stride)) # if opt.train_t_crop == 'random': # temporal_transform.append(TemporalRandomCrop(opt.sample_duration)) if opt.train_t_crop == 'center': temporal_transform.append(TemporalCenterCrop(opt.sample_duration)) temporal_transform = TemporalCompose(temporal_transform) train_data = get_training_data(opt.video_path, opt.annotation_path, opt.dataset, opt.input_type, opt.file_type, spatial_transform, temporal_transform) if opt.distributed: train_sampler = torch.utils.data.distributed.DistributedSampler( train_data) else: train_sampler = None train_loader = torch.utils.data.DataLoader(train_data, batch_size=opt.batch_size, shuffle=(train_sampler is None), num_workers=opt.n_threads, pin_memory=True, sampler=train_sampler, worker_init_fn=worker_init_fn)
验证数据处理函数
def get_val_utils(opt): normalize = get_normalize_method(opt.mean, opt.std, opt.no_mean_norm, opt.no_std_norm) spatial_transform = [ Resize(opt.sample_size), CenterCrop(opt.sample_size), ToTensor() ] if opt.input_type == 'flow': spatial_transform.append(PickFirstChannels(n=2)) spatial_transform.extend([ScaleValue(opt.value_scale), normalize]) spatial_transform = Compose(spatial_transform) temporal_transform = [] if opt.sample_t_stride > 1: temporal_transform.append(TemporalSubsampling(opt.sample_t_stride)) temporal_transform.append( TemporalEvenCrop(opt.sample_duration, opt.n_val_samples)) temporal_transform = TemporalCompose(temporal_transform) val_data, collate_fn = get_validation_data(opt.video_path, opt.annotation_path, opt.dataset, opt.input_type, opt.file_type, spatial_transform, temporal_transform) if opt.distributed: val_sampler = torch.utils.data.distributed.DistributedSampler( val_data, shuffle=False) else: val_sampler = None val_loader = torch.utils.data.DataLoader(val_data, batch_size=(opt.batch_size // opt.n_val_samples), shuffle=False, num_workers=opt.n_threads, pin_memory=True, sampler=val_sampler, worker_init_fn=worker_init_fn, collate_fn=collate_fn)
推理数据处理函数
def get_inference_utils(opt): assert opt.inference_crop in ['center', 'nocrop'] normalize = get_normalize_method(opt.mean, opt.std, opt.no_mean_norm, opt.no_std_norm) spatial_transform = [Resize(opt.sample_size)] if opt.inference_crop == 'center': spatial_transform.append(CenterCrop(opt.sample_size)) spatial_transform.append(ToTensor()) if opt.input_type == 'flow': spatial_transform.append(PickFirstChannels(n=2)) spatial_transform.extend([ScaleValue(opt.value_scale), normalize]) spatial_transform = Compose(spatial_transform) temporal_transform = [] if opt.sample_t_stride > 1: temporal_transform.append(TemporalSubsampling(opt.sample_t_stride)) temporal_transform.append( SlidingWindow(opt.sample_duration, opt.inference_stride)) temporal_transform = TemporalCompose(temporal_transform) inference_data, collate_fn = get_inference_data( opt.video_path, opt.annotation_path, opt.dataset, opt.input_type, opt.file_type, opt.inference_subset, spatial_transform, temporal_transform) inference_loader = torch.utils.data.DataLoader( inference_data, batch_size=opt.inference_batch_size, shuffle=False, num_workers=opt.n_threads, pin_memory=True, worker_init_fn=worker_init_fn, collate_fn=collate_fn) return inference_loader, inference_data.class_names
导入的模块
from spatial_transforms import (Compose, Normalize, Resize, CenterCrop, CornerCrop, MultiScaleCornerCrop, RandomResizedCrop, RandomHorizontalFlip, ToTensor, ScaleValue, ColorJitter, PickFirstChannels) from temporal_transforms import (LoopPadding, TemporalRandomCrop, TemporalCenterCrop, TemporalEvenCrop, SlidingWindow, TemporalSubsampling) from temporal_transforms import Compose as TemporalCompose
这个错误通常源于数据张量的存储不可调整大小,结合你的修改,可从以下几个方向排查修复:
检查时间裁剪参数匹配
你注释掉了随机时间裁剪的代码,但如果配置参数opt.train_t_crop设置的是random,当前代码没有对应的时间变换逻辑,会导致原始视频帧序列长度不匹配模型输入要求,进而引发张量尺寸不一致、存储调整失败的问题。要么把opt.train_t_crop改为center,要么恢复随机时间裁剪的代码,确保所有样本经过时间变换后长度统一为opt.sample_duration。确认数据预处理顺序与对象类型
训练的空间变换中,Resize和CenterCrop需要确保返回的是可修改的图像对象(比如PIL Image),而非提前转为不可变张量。如果这两个变换直接输出了固定存储的张量,后续的RandomHorizontalFlip等变换无法调整存储,就会触发错误。可以检查Resize和CenterCrop的实现代码,确认它们处理的是PIL图像或numpy数组,直到ToTensor步骤才转为张量。排查数据集样本尺寸一致性
即使做了变换,部分样本可能因原始尺寸问题导致变换后维度不统一。可以在创建train_data后,手动取出样本检查形状:sample = next(iter(train_data)) print(sample[0].shape)如果存在形状不一致的样本,需要在数据集中补充统一尺寸的处理逻辑,或调整变换逻辑覆盖所有边界情况。
临时关闭pin_memory测试
当pin_memory=True时,张量会被固定在内存中,可能导致部分变换无法调整存储。可以临时将train_loader的pin_memory改为False,如果错误消失,说明是固定内存导致的问题,需要确保所有尺寸调整操作在转成张量并固定内存前完成。检查Resize参数格式
确认opt.sample_size的格式是否符合Resize的要求:如果Resize期望传入元组(比如(224,224))但实际传入了单个整数,可能导致图像尺寸处理错误,进而引发后续存储问题。
内容的提问来源于stack exchange,提问作者anya

