add 寒武纪 MLU590 advisor agent
This commit is contained in:
50
test_main.py
Normal file
50
test_main.py
Normal file
@@ -0,0 +1,50 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user