49 lines
1.5 KiB
Python
49 lines
1.5 KiB
Python
import unittest
|
|
|
|
from main import ValidationError, analyze
|
|
|
|
|
|
class AnalyzeTests(unittest.TestCase):
|
|
def test_feasible_preflight(self):
|
|
result = analyze(
|
|
{
|
|
"model_name": "example/7b-model",
|
|
"architecture": "transformer",
|
|
"parameters_b": 7,
|
|
"dtype": "bf16",
|
|
"hardware": "国产算力卡",
|
|
"device_memory_gb": 24,
|
|
"cards": 1,
|
|
}
|
|
)
|
|
self.assertEqual(result["verdict"], "preflight_feasible")
|
|
self.assertGreater(result["estimates"]["headroom_gb"], 0)
|
|
|
|
def test_insufficient_memory(self):
|
|
result = analyze(
|
|
{
|
|
"model_name": "example/32b-model",
|
|
"architecture": "transformer",
|
|
"parameters_b": 32,
|
|
"dtype": "bf16",
|
|
"hardware": "国产算力卡",
|
|
"device_memory_gb": 24,
|
|
"cards": 2,
|
|
}
|
|
)
|
|
self.assertEqual(result["verdict"], "insufficient_memory")
|
|
|
|
def test_missing_information_is_explicit(self):
|
|
result = analyze({"model_name": "example/unknown-model"})
|
|
self.assertEqual(result["verdict"], "information_required")
|
|
self.assertIn("parameters_b", result["missing_fields"])
|
|
self.assertIn("device_memory_gb", result["missing_fields"])
|
|
|
|
def test_invalid_card_count_is_rejected(self):
|
|
with self.assertRaises(ValidationError):
|
|
analyze({"cards": 0})
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|