51 lines
1.8 KiB
Python
51 lines
1.8 KiB
Python
import unittest
|
|
|
|
from main import CARD_MODEL, CARD_VENDOR, ValidationError, analyze
|
|
|
|
|
|
class AnalyzeTests(unittest.TestCase):
|
|
def complete_payload(self):
|
|
return {
|
|
"model_name": "example/model",
|
|
"hardware": f"{CARD_VENDOR} {CARD_MODEL}",
|
|
"framework": "PyTorch",
|
|
"backend": "vendor-backend",
|
|
"sdk_version": "provided-by-user",
|
|
"driver_version": "provided-by-user",
|
|
}
|
|
|
|
def test_complete_preflight_targets_this_card(self):
|
|
result = analyze(self.complete_payload())
|
|
self.assertEqual(result["verdict"], "preflight_ready")
|
|
self.assertTrue(result["target_matches"])
|
|
self.assertEqual(result["official_target"]["model"], CARD_MODEL)
|
|
|
|
def test_mismatched_target_is_rejected(self):
|
|
payload = self.complete_payload()
|
|
payload["hardware"] = "其他厂商 其他卡型"
|
|
result = analyze(payload)
|
|
self.assertEqual(result["verdict"], "target_mismatch")
|
|
self.assertFalse(result["target_matches"])
|
|
|
|
def test_missing_information_is_explicit(self):
|
|
result = analyze({})
|
|
self.assertEqual(result["verdict"], "information_required")
|
|
self.assertIn("sdk_version", result["missing_fields"])
|
|
self.assertIn("driver_version", result["missing_fields"])
|
|
|
|
def test_multicard_and_quantization_risks_are_reported(self):
|
|
payload = self.complete_payload()
|
|
payload.update({"cards": 2, "precision": "int8"})
|
|
result = analyze(payload)
|
|
joined = " ".join(result["risks"])
|
|
self.assertIn("多卡", joined)
|
|
self.assertIn("量化", joined)
|
|
|
|
def test_invalid_card_count_is_rejected(self):
|
|
with self.assertRaises(ValidationError):
|
|
analyze({"cards": 0})
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|