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

Python单元测试(unittest):sys与Torch相关测试问题咨询

Python中sys模块与Torch框架的单元测试(unittest)实现方案

一、sys模块相关单元测试实现

1. 测试命令行参数(sys.argv)

当代码依赖sys.argv获取命令行参数时,直接运行测试会受测试命令本身的参数干扰,可通过unittest.mock.patch模拟指定的参数列表:

import unittest
from unittest.mock import patch
import sys

def process_args():
    if len(sys.argv) < 2:
        return "no args"
    return f"hello {sys.argv[1]}"

class TestSysArgv(unittest.TestCase):
    @patch('sys.argv', ['script.py', 'Alice'])
    def test_process_args_with_arg(self):
        self.assertEqual(process_args(), "hello Alice")
    
    @patch('sys.argv', ['script.py'])
    def test_process_args_no_arg(self):
        self.assertEqual(process_args(), "no args")

2. 捕获sys.exit触发的退出行为

sys.exit()会抛出SystemExit异常,测试时需验证该异常是否被正确触发,以及退出码是否符合预期:

import unittest
import sys

def exit_on_error(flag):
    if flag:
        sys.exit(1)
    return "success"

class TestSysExit(unittest.TestCase):
    def test_exit_on_true(self):
        with self.assertRaises(SystemExit) as cm:
            exit_on_error(True)
        self.assertEqual(cm.exception.code, 1)
    
    def test_no_exit_on_false(self):
        self.assertEqual(exit_on_error(False), "success")

3. 验证sys.stdout/stderr输出

要测试代码向标准输出/错误流的输出内容,可使用io.StringIO临时替换流对象,捕获输出后断言:

import unittest
import sys
from io import StringIO

def print_message(msg):
    print(f"LOG: {msg}")

class TestSysOutput(unittest.TestCase):
    def test_print_message(self):
        # 替换stdout
        captured_output = StringIO()
        sys.stdout = captured_output
        print_message("test log")
        # 恢复stdout
        sys.stdout = sys.__stdout__
        self.assertEqual(captured_output.getvalue().strip(), "LOG: test log")

二、Torch框架相关单元测试实现

1. 张量结果的精确断言(处理浮点精度)

直接用==比较张量会因浮点精度问题出错,应使用torch.allclose()或torch.testing.assert_close():

import unittest
import torch

def add_tensors(a, b):
    return a + b

class TestTorchTensors(unittest.TestCase):
    def test_add_tensors(self):
        a = torch.tensor([1.0, 2.5])
        b = torch.tensor([3.0, 4.5])
        result = add_tensors(a, b)
        expected = torch.tensor([4.0, 7.0])
        # 容忍微小浮点误差
        self.assertTrue(torch.allclose(result, expected))
        # 更严格的断言(包含形状、数据类型)
        torch.testing.assert_close(result, expected)

2. 跨设备(CPU/GPU)兼容性测试

测试时可强制使用CPU,避免依赖GPU环境;也可通过skipIf跳过不支持CUDA的测试:

import unittest
import torch

def move_to_device(tensor, device):
    return tensor.to(device)

class TestTorchDevice(unittest.TestCase):
    def test_move_to_cpu(self):
        tensor = torch.tensor([1,2])
        result = move_to_device(tensor, torch.device('cpu'))
        self.assertEqual(result.device.type, 'cpu')
    
    @unittest.skipIf(not torch.cuda.is_available(), "CUDA not available")
    def test_move_to_cuda(self):
        tensor = torch.tensor([1,2])
        result = move_to_device(tensor, torch.device('cuda'))
        self.assertEqual(result.device.type, 'cuda')

3. 模型前向传播与输出验证

测试模型的输入输出形状、参数更新逻辑,需固定随机种子保证结果可复现:

import unittest
import torch
import torch.nn as nn

class SimpleModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.linear = nn.Linear(10, 2)
    
    def forward(self, x):
        return self.linear(x)

class TestTorchModel(unittest.TestCase):
    def setUp(self):
        # 固定随机种子
        torch.manual_seed(42)
        self.model = SimpleModel()
        self.input = torch.randn(3, 10)  # batch_size=3, feature_dim=10
    
    def test_forward_shape(self):
        output = self.model(self.input)
        self.assertEqual(output.shape, (3, 2))  # 验证输出形状
    
    def test_parameter_update(self):
        initial_weight = self.model.linear.weight.clone()
        optimizer = torch.optim.SGD(self.model.parameters(), lr=0.01)
        loss_fn = nn.CrossEntropyLoss()
        labels = torch.tensor([0,1,0])
        
        # 一次训练步骤
        optimizer.zero_grad()
        output = self.model(self.input)
        loss = loss_fn(output, labels)
        loss.backward()
        optimizer.step()
        
        # 验证权重已更新
        self.assertFalse(torch.allclose(self.model.linear.weight, initial_weight))

4. 模拟Torch依赖的外部资源

若代码依赖外部数据加载或第三方Torch工具,可使用unittest.mock.patch模拟返回值:

import unittest
from unittest.mock import patch
import torch

def load_data(path):
    # 实际中可能从磁盘加载张量
    return torch.load(path)

class TestTorchMock(unittest.TestCase):
    @patch('torch.load')
    def test_load_data(self, mock_load):
        # 模拟torch.load返回指定张量
        mock_load.return_value = torch.tensor([1,2,3])
        data = load_data("fake_path.pt")
        self.assertTrue(torch.allclose(data, torch.tensor([1,2,3])))
        mock_load.assert_called_once_with("fake_path.pt")

内容的提问来源于stack exchange,提问作者Jahidul Hasan Razib

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 04:50:07