如何结合@pytest.mark.parametrize与skipif实现GPU数量参数化测试
解决方案:单测试用例适配不同GPU数量场景
完全可以通过@pytest.mark.parametrize结合@pytest.mark.skipif实现你的需求,不用写8个重复测试用例,这才是更优雅高效的实现方式,具体如下:
核心思路
- 先获取主机实际可用的GPU数量
- 参数化生成1到8的GPU数量测试参数
- 对每个参数判断:如果要求的GPU数超过主机实际数量,自动跳过该测试实例
代码实现(以PyTorch环境为例)
import pytest import torch # 获取主机实际GPU数量 actual_gpus = torch.cuda.device_count() @pytest.mark.parametrize("num_gpus", range(1, 9)) # 生成1-8的测试参数 @pytest.mark.skipif( lambda num_gpus: num_gpus > actual_gpus, reason=lambda num_gpus: f"主机仅拥有{actual_gpus}块GPU,无法执行{num_gpus}块GPU的测试" ) def test_target_function_with_gpus(num_gpus): # 这里编写你的目标函数测试逻辑,根据num_gpus配置GPU使用 print(f"正在用{num_gpus}块GPU执行测试") # 示例:调用目标函数 # target_function(num_gpus) assert True # 替换为实际断言逻辑
非PyTorch环境的GPU数量获取方式
如果你的环境不依赖PyTorch,可以通过解析nvidia-smi输出来获取GPU数量:
import subprocess def get_actual_gpus(): try: # 调用nvidia-smi获取GPU数量 output = subprocess.check_output( ["nvidia-smi", "--query-gpu=count", "--format=csv,noheader,nounits"], stderr=subprocess.STDOUT ) return int(output.strip().decode()) except (subprocess.CalledProcessError, ValueError, FileNotFoundError): # 处理无NVIDIA GPU或命令不存在的情况 return 0 # 替换前面的actual_gpus赋值 actual_gpus = get_actual_gpus()
效果说明
- 当主机有3块GPU时,参数1、2、3会正常执行,4-8会被自动跳过
- 所有测试逻辑只需要写一次,后续要调整测试范围(比如改成1-10),只需要修改
range(1,9)为range(1,11)即可,维护成本极低
内容的提问来源于stack exchange,提问作者malfonsoarquimea
相关产品推荐
相关产品推荐

