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

PyTorch图像预测场景下向图像张量添加额外布尔信息的问题咨询

针对PyTorch图像序列预测任务补充序列属性的实现方案

基础问题修正

你提供的原有TrainFeeder代码中,PyTorch Dataset要求的魔法方法缺少双下划线,需要先将init改为__init__、getitem改为__getitem__、len改为__len__,否则DataLoader无法正常调用该数据集类。

新增JSON布尔属性读取逻辑

修改后的TrainFeeder代码如下,默认会读取序列文件夹下名为attr.json的文件,你可以根据实际的JSON文件名、布尔值对应的键名调整代码:

import os
import json
import glob
import cv2
import torch
from torch.utils.data import Dataset
from torchvision.transforms import ToTensor

class TrainFeeder(Dataset):
    def __init__(self, data_set):
        super(TrainFeeder, self).__init__()
        self.input_data = data_set
        if torch.cuda.current_device() == 0:
            print('There are total %d sequences in trainset' % len(self.input_data))

    def __getitem__(self, index):
        path = self.input_data[index]
        # 读取图像序列
        imgs_path = sorted(glob.glob(path + '/*.png'))
        imgs = []
        for img_path in imgs_path:
            img = cv2.imread(img_path)
            img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
            img = cv2.resize(img, (256,448))
            img = cv2.resize(img, (0, 0), fx=0.5, fy=0.5, interpolation=cv2.INTER_CUBIC)
            img_tensor = ToTensor()(img).float()
            imgs.append(img_tensor)
        imgs = torch.stack(imgs, dim=0) # 此时形状为 [6, 3, H, W],维度顺序:帧数、RGB通道数、高度、宽度
        
        # 读取JSON布尔属性
        json_path = os.path.join(path, 'attr.json') # 替换为你实际的JSON文件名
        with open(json_path, 'r', encoding='utf-8') as f:
            json_data = json.load(f)
        bool_val = json_data['你的布尔属性对应的键名'] # 替换为你实际的键名
        bool_tensor = torch.tensor(float(bool_val), dtype=torch.float32)

        # 可选方案1:将属性拼接到图像张量的通道维度
        # bool_tensor = bool_tensor.view(1, 1, 1, 1).expand(imgs.shape[0], 1, imgs.shape[2], imgs.shape[3])
        # imgs = torch.cat([imgs, bool_tensor], dim=1) # 拼接后形状为 [6, 4, H, W]
        # return imgs

        # 可选方案2:单独返回属性张量,不修改原有图像张量结构
        return imgs, bool_tensor

    def __len__(self):
        return len(self.input_data)

扩展维度选择建议

  • 如果你选择将布尔属性和图像张量拼接,推荐在*通道维度(dim=1)*扩展。该布尔属性是整个序列共享的属性,扩展为和每帧同尺寸的单通道张量,不会破坏序列的时间维度、空间维度结构,符合卷积层的输入逻辑。
  • 更推荐选择单独返回属性张量的方案,无需修改原有图像张量的结构,灵活性更高。

对原有预测系统的影响说明

  • 如果选择拼接通道的方案:原有模型的输入层接收的通道数是3,修改后输入通道数变为4,需要调整模型输入层的通道参数,否则会直接报错。如果你的预测任务不需要在标签中包含该属性,取标签的时候仅取前3通道即可,后续的损失计算、预测逻辑不需要做其他修改。
  • 如果选择单独返回属性张量的方案:原有模型的图像处理分支完全不需要修改,只需要新增一个小的特征分支处理该布尔属性,再和图像提取的特征做融合即可,对原有系统的侵入性最低,不会影响原有逻辑的正常运行。

内容的提问来源于stack exchange,提问作者Pia Lüdemann

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 12:54:02