You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

修改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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.10 19:25:56