Back to home page

EIC code displayed by LXR

 
 

    


File indexing completed on 2026-09-01 09:33:42

0001 # Description: Unit tests for the hardware matching functions
0002 import unittest
0003 
0004 from pandaserver.srvcore.hardware_matching import (
0005     compare_version_string,
0006     match_gpu_spec,
0007 )
0008 
0009 A100 = {
0010     "vendor": "NVIDIA",
0011     "model": "NVIDIA A100-SXM4-40GB",
0012     "vram": 40960,
0013     "architecture": "Ampere",
0014     "framework_version": "12.4",
0015     "driver_version": "575.57.08",
0016 }
0017 
0018 V100 = {
0019     "vendor": "NVIDIA",
0020     "model": "Tesla V100-SXM2-16GB",
0021     "vram": 16384,
0022     "architecture": "Volta",
0023     "framework_version": "12.2",
0024     "driver_version": "535.104.05",
0025 }
0026 
0027 ANY_GPU = {"vendor": "*", "model": "*"}
0028 
0029 
0030 class TestCompareVersionString(unittest.TestCase):
0031     def test_operators(self):
0032         self.assertTrue(compare_version_string("12.4", ">=12.0"))
0033         self.assertFalse(compare_version_string("11.8", ">=12.0"))
0034         self.assertTrue(compare_version_string("11.8", "<=12.0"))
0035         self.assertTrue(compare_version_string("12.4", ">12.0"))
0036         self.assertFalse(compare_version_string("12.0", ">12.0"))
0037         self.assertTrue(compare_version_string("11.8", "<12.0"))
0038         self.assertTrue(compare_version_string("12.0", "==12.0"))
0039         self.assertFalse(compare_version_string("12.4", "==12.0"))
0040 
0041     def test_single_equal_is_equality(self):
0042         self.assertTrue(compare_version_string("12.0", "=12.0"))
0043         self.assertFalse(compare_version_string("12.4", "=12.0"))
0044 
0045     def test_not_equal(self):
0046         # the != operator used to be unreachable since the operator pattern didn't accept !
0047         self.assertTrue(compare_version_string("12.4", "!=12.0"))
0048         self.assertFalse(compare_version_string("12.0", "!=12.0"))
0049 
0050     def test_multi_component_version(self):
0051         self.assertTrue(compare_version_string("575.57.08", ">=575.0"))
0052         self.assertFalse(compare_version_string("535.104.05", ">=575.0"))
0053 
0054     def test_invalid_input(self):
0055         # no operator
0056         self.assertIsNone(compare_version_string("12.0", "12.0"))
0057         # unparsable versions
0058         self.assertIsNone(compare_version_string("Ampere", ">=12.0"))
0059         self.assertIsNone(compare_version_string("12.0", ">=Ampere"))
0060 
0061 
0062 class TestMatchGpuSpec(unittest.TestCase):
0063     def test_wildcard_requirement(self):
0064         self.assertTrue(match_gpu_spec(ANY_GPU, [A100]))
0065         self.assertTrue(match_gpu_spec(ANY_GPU, [A100, V100]))
0066         # no GPU information is not a rejection for a wildcard requirement
0067         self.assertTrue(match_gpu_spec(ANY_GPU, []))
0068 
0069     def test_specific_requirement_without_gpus(self):
0070         self.assertFalse(match_gpu_spec({"vendor": "NVIDIA", "model": "*"}, []))
0071         self.assertFalse(match_gpu_spec({"vendor": "*", "model": ".*A100.*"}, []))
0072         self.assertFalse(match_gpu_spec({"vendor": "*", "model": "*", "vram": ">=40960"}, []))
0073         self.assertFalse(match_gpu_spec({"vendor": "*", "model": "*", "microarchitecture": "Ampere"}, []))
0074 
0075     def test_vendor(self):
0076         self.assertTrue(match_gpu_spec({"vendor": "NVIDIA", "model": "*"}, [A100]))
0077         self.assertTrue(match_gpu_spec({"vendor": "nvidia", "model": "*"}, [A100]))
0078         self.assertFalse(match_gpu_spec({"vendor": "AMD", "model": "*"}, [A100]))
0079         # any match, one of the GPUs is enough
0080         self.assertTrue(match_gpu_spec({"vendor": "NVIDIA", "model": "*"}, [{"vendor": "AMD", "model": "MI250"}, A100]))
0081 
0082     def test_model_inclusion(self):
0083         self.assertTrue(match_gpu_spec({"vendor": "*", "model": ".*A100.*"}, [A100]))
0084         # matching is case-insensitive
0085         self.assertTrue(match_gpu_spec({"vendor": "*", "model": ".*a100.*"}, [A100]))
0086         self.assertFalse(match_gpu_spec({"vendor": "*", "model": ".*A100.*"}, [V100]))
0087         # any match
0088         self.assertTrue(match_gpu_spec({"vendor": "*", "model": ".*A100.*"}, [V100, A100]))
0089 
0090     def test_model_exclusion(self):
0091         excl_p100 = {"vendor": "*", "model": {"pattern": ".*P100.*", "excl": True}}
0092         self.assertTrue(match_gpu_spec(excl_p100, [A100]))
0093         P100 = dict(A100, model="Tesla P100-PCIE-16GB")
0094         self.assertFalse(match_gpu_spec(excl_p100, [P100]))
0095         # excluded when any of the GPUs matches the pattern
0096         self.assertFalse(match_gpu_spec(excl_p100, [A100, P100]))
0097 
0098     def test_vram(self):
0099         self.assertTrue(match_gpu_spec(dict(ANY_GPU, vram=">=40960"), [A100]))
0100         self.assertFalse(match_gpu_spec(dict(ANY_GPU, vram=">=40960"), [V100]))
0101         # all match, every GPU has to meet the minimum
0102         self.assertFalse(match_gpu_spec(dict(ANY_GPU, vram=">=40960"), [A100, V100]))
0103         self.assertTrue(match_gpu_spec(dict(ANY_GPU, vram=">=16384"), [A100, V100]))
0104         self.assertTrue(match_gpu_spec(dict(ANY_GPU, vram="==40960"), [A100]))
0105 
0106     def test_microarchitecture(self):
0107         self.assertTrue(match_gpu_spec(dict(ANY_GPU, microarchitecture="Ampere"), [A100]))
0108         self.assertFalse(match_gpu_spec(dict(ANY_GPU, microarchitecture="Ampere"), [V100]))
0109         # a list of generations is accepted
0110         self.assertTrue(match_gpu_spec(dict(ANY_GPU, microarchitecture=["Ampere", "Hopper"]), [A100]))
0111         self.assertFalse(match_gpu_spec(dict(ANY_GPU, microarchitecture=["Ampere", "Hopper"]), [V100]))
0112         # any match, one of the GPUs is enough
0113         self.assertTrue(match_gpu_spec(dict(ANY_GPU, microarchitecture="Ampere"), [V100, A100]))
0114 
0115     def test_framework_version(self):
0116         self.assertTrue(match_gpu_spec(dict(ANY_GPU, version=">=12.0"), [A100, V100]))
0117         # all match, a single old GPU excludes the whole set
0118         self.assertFalse(match_gpu_spec(dict(ANY_GPU, version=">=12.3"), [A100, V100]))
0119         self.assertTrue(match_gpu_spec(dict(ANY_GPU, version=">=12.3"), [A100]))
0120 
0121     def test_driver_version(self):
0122         self.assertTrue(match_gpu_spec(dict(ANY_GPU, driver_version=">=575.0"), [A100]))
0123         self.assertFalse(match_gpu_spec(dict(ANY_GPU, driver_version=">=575.0"), [V100]))
0124         # all match
0125         self.assertFalse(match_gpu_spec(dict(ANY_GPU, driver_version=">=575.0"), [A100, V100]))
0126 
0127     def test_missing_attributes_in_gpus(self):
0128         bare = {"vendor": "NVIDIA", "model": "NVIDIA A100-SXM4-40GB"}
0129         self.assertTrue(match_gpu_spec({"vendor": "NVIDIA", "model": ".*A100.*"}, [bare]))
0130         # constraints on attributes which the GPU doesn't report are not satisfied
0131         self.assertFalse(match_gpu_spec(dict(ANY_GPU, vram=">=40960"), [bare]))
0132         self.assertFalse(match_gpu_spec(dict(ANY_GPU, version=">=12.0"), [bare]))
0133         self.assertFalse(match_gpu_spec(dict(ANY_GPU, driver_version=">=575.0"), [bare]))
0134         self.assertFalse(match_gpu_spec(dict(ANY_GPU, microarchitecture="Ampere"), [bare]))
0135 
0136     def test_combined_constraints(self):
0137         spec = {"vendor": "NVIDIA", "model": ".*A100.*", "vram": ">=40960", "microarchitecture": "Ampere", "version": ">=12.0", "driver_version": ">=575.0"}
0138         self.assertTrue(match_gpu_spec(spec, [A100]))
0139         self.assertFalse(match_gpu_spec(spec, [V100]))
0140         self.assertFalse(match_gpu_spec(spec, [A100, V100]))
0141 
0142 
0143 if __name__ == "__main__":
0144     unittest.main()