66 lines
2.1 KiB
Python
66 lines
2.1 KiB
Python
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
from PIL import Image
|
|
|
|
from cockpit_agent.grounding.adapter import GroundingAdapter
|
|
|
|
|
|
class FakeGroundingModel:
|
|
def __init__(self, output: str) -> None:
|
|
self.output = output
|
|
self.prompt: str | None = None
|
|
|
|
def generate(
|
|
self,
|
|
image_path: str,
|
|
prompt: str,
|
|
max_new_tokens: int = 128,
|
|
) -> str:
|
|
del image_path, max_new_tokens
|
|
self.prompt = prompt
|
|
return self.output
|
|
|
|
|
|
class GroundingAdapterTest(unittest.TestCase):
|
|
def test_uses_model_bbox_without_postprocessing(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
image_path = Path(directory) / "generic.png"
|
|
Image.new("RGB", (200, 100), color="white").save(image_path)
|
|
model = FakeGroundingModel('{"bbox_2d": [100, 200, 400, 600]}')
|
|
|
|
result = GroundingAdapter(model).ground(
|
|
str(image_path),
|
|
"车辆空调内循环控制",
|
|
)
|
|
|
|
self.assertEqual(result.bbox_relative, (100.0, 200.0, 400.0, 600.0))
|
|
self.assertEqual(result.bbox_pixel, (20, 20, 80, 60))
|
|
self.assertEqual(result.center_pixel, (50, 40))
|
|
self.assertIn("车辆空调内循环控制", model.prompt or "")
|
|
self.assertEqual(
|
|
set(result.to_dict()),
|
|
{
|
|
"semantic_query",
|
|
"bbox_relative",
|
|
"bbox_pixel",
|
|
"center_pixel",
|
|
"image_width",
|
|
"image_height",
|
|
"raw_output",
|
|
},
|
|
)
|
|
|
|
def test_invalid_model_output_is_not_repaired(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
image_path = Path(directory) / "generic.png"
|
|
Image.new("RGB", (32, 32), color="white").save(image_path)
|
|
adapter = GroundingAdapter(FakeGroundingModel("not a bbox"))
|
|
with self.assertRaises(ValueError):
|
|
adapter.ground(str(image_path), "任意语义控件")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|