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
相关产品推荐
相关产品推荐

