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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 12:05:02