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

PyTorch自定义4分类模型全预测为left类问题排查求助

图像分类模型所有预测均为同一类别问题排查

我正在用自定义数据集构建4分类的图像分类模型(TingVGG),但模型对所有图像的预测结果均为left类。以下是我的模型结构、训练测试函数、模型保存加载及预测代码,恳请帮忙排查错误。这是我学习PyTorch的第2周,若表述有误还请见谅。

TingVGG模型结构

class TingVGG(nn.Module):
    def __init__(self, input_shape: int, hidden_units: int, output_shape: int) -> None:
        super().__init__()
        self.conv_block1 = nn.Sequential(nn.Conv2d(in_channels=input_shape,out_channels=hidden_units,kernel_size=3,stride=1,padding=1),
        nn.ReLU(),
        nn.Conv2d(in_channels=hidden_units, out_channels=hidden_units, kernel_size=3,stride=1,padding=1),
        nn.ReLU(),
        nn.MaxPool2d(kernel_size=2, stride=2)
        )
        
        self.conv_block2 = nn.Sequential(nn.Conv2d(in_channels=hidden_units,out_channels=hidden_units,kernel_size=3,stride=1,padding=1),
        nn.ReLU(),
        nn.Conv2d(in_channels=hidden_units, out_channels=hidden_units, kernel_size=3,stride=1,padding=1),
        nn.ReLU(),
        nn.MaxPool2d(kernel_size=2, stride=2)
        )
        
        self.classifier = nn.Sequential(nn.Flatten(), nn.Linear(in_features=hidden_units*32*32 ,out_features=output_shape))
        
        
    def forward(self, x: torch.Tensor):
        x = self.conv_block1(x)
        #print(x.shape)
        x = self.conv_block2(x)
        #print(x.shape)
        x = self.classifier(x)
        #print(x.shape)
        return x

训练与测试函数

def train_step(model: torch.nn.Module,
               dataloader: DataLoader,
               loss_fn: torch.nn.Module,
               optimizer: torch.optim.Optimizer,
               device=device
               ):
    # put the model into the train
    model.train()
    
    # Setup train loss and train accuracy values
    train_loss, train_acc = 0, 0
    
    # Loop through DataLoader and data batches
    for batch, (X, y) in enumerate(dataloader):
         # Send data to the target device
         X, y = X.to(device), y.to(device)
         
         # 1. forward pass
         y_pred = model(X) #output model logits
         
         # 2. Calculate the loss
         loss = loss_fn(y_pred, y)
         train_loss += loss.item()
         
         # 3. Optimize zero grad
         optimizer.zero_grad()
         
         # 4.Loss backward
         loss.backward()
         
         # 5.Optimzer step
         optimizer.step()
         
         # Calculate the accuracy metric
         y_pred_class = torch.argmax(torch.softmax(y_pred, dim=1), dim=1)
         train_acc += (y_pred_class ==y).sum().item()/len(y_pred)
    
    # Adjust metrics to get average loss and accuracy per batch
    train_loss = train_loss / len(dataloader)
    train_acc = train_acc / len(dataloader)
    return train_loss, train_acc


def test_step(model: torch.nn.Module,
               dataloader: DataLoader,
               loss_fn: torch.nn.Module,
               device=device
               ):
    # Put the model in eval mode
    model.eval()
    
    # Setup train loss and train accuracy values
    test_loss, test_acc = 0, 0
    
    # Turn on inference mode
    with torch.inference_mode():
        # Loop through DataLoader Batches
        for batch, (X, y) in enumerate(dataloader):
            # Send data to the target device
            X, y = X.to(device), y.to(device)
            
            # 1. forward pass
            test_pred_logits = model(X)
            
            # 2. Calculate the loss
            loss = loss_fn(test_pred_logits, y)
            test_loss += loss.item()
            
            # 3. Calculate the accuracy
            test_pred_labels = test_pred_logits.argmax(dim=1)
            test_acc += ((test_pred_labels == y).sum().item() / len(test_pred_labels))
            
    # Adjust metrics to get average loss and accuracy per batch
    test_loss = test_loss / len(dataloader)
    test_acc = test_acc / len(dataloader)
    return test_loss, test_acc


def train(model: torch.nn.Module,
          train_dataloader: DataLoader,
          test_dataloader: DataLoader,
          optimizer: torch.optim.Optimizer,
          loss_fn: torch.nn.Module = nn.CrossEntropyLoss(),
          epochs: int = 10,
          device = device):
    # 2. Create empty results dictionary
    results = {"train_loss": [],
               "train_acc": [],
               "test_loss": [],
               "test_acc": []}
    
    # 3. Loop through training and testing steps for a number of epochs
    for epoch in tqdm(range(epochs)):
        train_loss, train_acc = train_step(model= model, dataloader= train_dataloader,loss_fn=loss_fn, optimizer=optimizer, device=device)
        test_loss, test_acc = test_step(model= model, dataloader=test_dataloader,loss_fn=loss_fn, device=device)
        
        # 4. Print out what's happening
        print(f"Epoch: {epoch +1} | "
            f"train_loss: {train_loss:.4f} | "
            f"train_acc: {train_acc:.4f} | "
            f"test_loss: {test_loss:.4f} | "
            f"test_acc: {test_acc:.4f}")
        
        # 5. Update the results dictionary
        results["train_loss"].append(train_loss)
        results["train_acc"].append(train_acc)
        results["test_loss"].append(test_loss)
        results["test_acc"].append(test_acc)
        
    # 6. return the results at the end of the epoches
    return results


# Set random seed
torch.manual_seed(42)
torch.cuda.manual_seed(42)

# Set number of epoches
NUM_EPOCHS = 20

# Create and initialize of TinyVGG
model_0 = TingVGG(input_shape=1, # Number of channels in the input image (c, h, w) -> 3
                  hidden_units=20,
                  output_shape=len(train_data.classes)).to(device)

# Setup the loss function and optimizer
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(params= model_0.parameters(),
                             lr= 0.001)

# Start the timer
start_time = timer()

# Train model 0
model_0_results = train(model= model_0,
                        train_dataloader= train_dataloader_simple,
                        test_dataloader= test_dataloader_simple,
                        optimizer= optimizer,
                        loss_fn= loss_fn,
                        epochs= NUM_EPOCHS
                        )

# End the timer and print the results
end_time = timer()
print(f"Total training time: {end_time - start_time: .3f} seconds")

模型保存与加载

PATH='Model_3.pth'
torch.save(model_0, PATH)
model = torch.load('Model_3.pth', map_location='cpu')

预测代码

images_path = r'Rotation_Images\train\180'
class_names = ['180', 'left', 'right', 'zero']


for img in os.listdir(images_path):
    
    start_time =  timer()
    
    # Reads a file using pillow
    PIL_image = PIL.Image.open(images_path + "/" + img)

    # Convert to numpy array
    numpy_array = np.array(PIL_image)

    # Convert to PyTorch tensor
    tensor = torch.from_numpy(numpy_array)


    tensor = tensor.unsqueeze(0).type(torch.float32)

    tensor_image = tensor / 255. 
    
    # Create transform pipleine to resize image
    custom_image_transform = transforms.Compose([
        transforms.Resize((128, 128)),
    ])

    # Transform target image
    custom_image_transformed = custom_image_transform(tensor_image)
    
    model.eval()
    with torch.inference_mode():
        # Add an extra dimension to image
        custom_image_transformed_with_batch_size = custom_image_transformed.unsqueeze(dim=0)      
        
        # Make a prediction on image with an extra dimension
        custom_image_pred = model(custom_image_transformed_with_batch_size.to(device))
        
        # Convert logits -> prediction probabilities (using torch.softmax() for multi-class classification)
        custom_image_pred_probs = torch.softmax(custom_image_pred, dim=1)
        # Convert prediction probabilities -> prediction labels
        custom_image_pred_label = torch.argmax(custom_image_pred_probs, dim=1)
        # Find the predicted label
        
        custom_image_pred_class = class_names[custom_image_pred_label.cpu()]

问题排查方向

  • 数据集类别不平衡:统计训练集各类别样本数量,若left类占比过高,模型会倾向于预测该类。可通过过采样、欠采样或给损失函数添加类别权重(CrossEntropyLoss(weight=torch.tensor(class_weights).to(device)))解决。
  • 输入预处理不一致:检查训练时的图像预处理流程是否和预测完全一致。比如训练时是否做了resize、归一化,预测时是否遗漏或多做了步骤;确认输入张量维度是否匹配模型要求,单通道图像输入应为(batch_size, 1, H, W),预测时的张量转换是否正确。
  • 模型维度计算错误:验证模型线性层输入特征数是否正确。若训练时输入图像不是128x128,两次MaxPool后特征图尺寸会变化,导致hidden_units*32*32与实际Flatten后的特征数不匹配,引发训练异常。
  • 训练状态异常:查看训练过程中的准确率和损失值:
    • 若训练准确率接近25%(随机猜测水平),说明模型未学到有效特征,可能是学习率不合适、训练轮数不足、数据加载错误(如标签全为left类)。
    • 若训练准确率高但测试准确率低,需考虑过拟合,但当前所有预测为同一类更可能是训练数据本身存在问题。
  • 类别映射错误:确认预测时的class_names顺序是否与训练时train_data.classes完全一致,避免索引对应错误的类别名称。
  • 模型加载问题:验证加载的模型是训练完成后的版本,可在训练结束后对样本做一次预测,与加载模型后的预测结果对比。

内容的提问来源于stack exchange,提问作者Emad Younan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 14:54:52