使用HF Dataset类Transform输出指定尺寸pixel_values张量的问题
解决Hugging Face Dataset处理3D图像切片时维度丢失的问题
你的问题核心在于预处理函数逻辑和Dataset.map的工作方式不匹配。Dataset.map默认是逐样本(或指定批量)处理数据,而你当前的preprocess_data函数把整个数据集的所有样本切片平铺成了一个大列表,导致map将每个切片视为独立样本返回,自然丢失了每个样本对应的28个切片维度。
第一步:修正预处理函数,逐样本处理
将全局处理逻辑改为针对单个样本的处理,确保每个样本输出包含28个切片的pixel_values和对应标签:
def preprocess_sample(sample, num_slices=28): # 单个样本的3D数据:shape (28,28,28) img_3d = np.array(sample['image']).reshape(28, 28, 28) # 对每个切片做通道重复(适配ViT的3通道要求),生成28个(3,224,224)张量 slices = [ processor( np.repeat(img_3d[:, :, i][np.newaxis, ...], 3, axis=0), return_tensors='pt' )['pixel_values'][0] for i in range(num_slices) ] # 堆叠成(28,3,224,224)的张量 pixel_values = torch.stack(slices) return {'pixel_values': pixel_values, 'labels': sample['labels']}
然后用map函数逐样本处理数据集:
# 假设原始数据集为raw_ds processed_ds = raw_ds.map(preprocess_sample, remove_columns=['image'])
第二步:调整ViT模型的前向逻辑
自定义ViT模型,在MLP头前加入最大池化层聚合28个切片的特征:
from transformers import ViTForImageClassification import torch.nn as nn class ViTFor3DImage(ViTForImageClassification): def forward(self, pixel_values, labels=None): # 输入shape: (batch_size, 28, 3, 224, 224) batch_size, num_slices = pixel_values.shape[:2] # 合并batch和切片维度,适配ViT默认输入格式:(batch_size*28, 3, 224, 224) flattened_pixels = pixel_values.view(-1, 3, 224, 224) # 调用预训练ViT提取特征 outputs = self.vit(flattened_pixels) cls_token = outputs.last_hidden_state[:, 0, :] # 取CLS token,shape: (batch*28, hidden_size) # 还原维度并做最大池化聚合 cls_token = cls_token.view(batch_size, num_slices, -1) aggregated_features = torch.max(cls_token, dim=1)[0] # shape: (batch_size, hidden_size) # 过分类头 logits = self.classifier(aggregated_features) # 计算损失(兼容Trainer的损失逻辑) loss = None if labels is not None: loss_fct = nn.CrossEntropyLoss() loss = loss_fct(logits.view(-1, self.num_labels), labels.view(-1)) return {'loss': loss, 'logits': logits} if loss is not None else {'logits': logits}
第三步:用Trainer启动训练
此时处理后的数据集每个样本的pixel_values维度为(28,3,224,224),模型也能正确处理该输入,直接用Trainer训练即可:
from transformers import Trainer, TrainingArguments # 加载预训练ViT并初始化自定义模型 model = ViTFor3DImage.from_pretrained( 'google/vit-base-patch16-224-in21k', num_labels=你的类别数量 ) training_args = TrainingArguments( output_dir='./vit-3d-checkpoints', per_device_train_batch_size=8, num_train_epochs=10, logging_steps=50, evaluation_strategy='epoch', # 其他参数根据需求调整 ) trainer = Trainer( model=model, args=training_args, train_dataset=processed_ds['train'], eval_dataset=processed_ds['test'] ) trainer.train()
补充说明
- 原预处理函数将所有样本切片混在一起,导致Dataset把每个切片当作独立样本,所以取出的是单个切片维度。改成逐样本处理后,每个样本保留了28个切片的结构,符合预期。
- 模型中合并batch与切片维度,是为了复用预训练ViT的权重逻辑,之后再聚合特征,既利用了预训练成果,又适配了3D图像的处理需求。
- 之前自定义训练循环不收敛,大概率是数据处理或模型结构的问题,现在配合正确的预处理和模型,应该能解决收敛问题。
内容的提问来源于stack exchange,提问作者Killer Potato
相关产品推荐
相关产品推荐

