如何获取PyTorch中FashionMNIST数据集的transform等超参数配置?
PyTorch FashionMNIST transform参数常见问题解答
为什么FashionMNIST官方文档页没有transform参数的详细说明
transform是torchvision所有内置数据集通用的继承参数,它的定义和规则写在父类VisionDataset的逻辑中,所以不会单独在FashionMNIST的子类文档页重复说明。
transform参数支持的传入值
该参数要求传入输入为PIL图像、输出为处理后数据的可调用对象,除了ToTensor()外,常用的可选值包括:
- 数据预处理类:
Normalize():对张量按预设的均值和标准差做逐通道归一化,加速模型训练收敛Resize():调整图像到指定尺寸Grayscale()/RGB():调整图像的通道数
- 数据增强类(仅建议训练集使用):
RandomHorizontalFlip():按设定概率随机水平翻转图像RandomCrop()/RandomResizedCrop():随机裁剪图像ColorJitter():随机调整图像的亮度、对比度、饱和度
- 自定义可调用对象:你自行实现的符合输入输出规则的转换逻辑,比如添加随机噪声、自定义掩膜等操作
如果需要依次执行多个转换操作,可以用transforms.Compose()将多个操作按顺序封装为一个组合对象,示例用法如下:
from torchvision import transforms # 训练集转换逻辑:先随机翻转,转张量,再做单通道归一化 train_transform = transforms.Compose([ transforms.RandomHorizontalFlip(p=0.5), transforms.ToTensor(), transforms.Normalize(mean=[0.286], std=[0.353]) ])
相关信息查询位置
- torchvision.transforms模块的官方文档:所有内置转换操作的参数、用法、适用场景都有完整说明
- VisionDataset基类的官方文档:
transform参数的定义、调用时机、传入规则的底层说明
内容的提问来源于stack exchange,提问作者JAEMTO
相关产品推荐
相关产品推荐

