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

如何结合@pytest.mark.parametrize与skipif实现GPU数量参数化测试

解决方案:单测试用例适配不同GPU数量场景

完全可以通过@pytest.mark.parametrize结合@pytest.mark.skipif实现你的需求,不用写8个重复测试用例,这才是更优雅高效的实现方式,具体如下:

核心思路

  1. 先获取主机实际可用的GPU数量
  2. 参数化生成1到8的GPU数量测试参数
  3. 对每个参数判断:如果要求的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 12:10:48