pi0.5 / tests /test_app.py
XiangpengYang's picture
feat: expose UR policy config in Gradio
df613f6
Raw
History Blame Contribute Delete
2.16 kB
from pathlib import Path
import unittest
from unittest import mock
class AppTests(unittest.TestCase):
def test_predict_ui_returns_table_file_and_status(self):
import app
result = type(
"Result", (),
{"actions": "table", "json_path": "/tmp/result.json", "status": "done"},
)()
with mock.patch.object(app.MODEL_MANAGER, "get", return_value=object()) as get_model, mock.patch.object(
app, "run_prediction", return_value=result
):
actual = app.predict_ui(
"owner/model", "checkpoint", "pi05_ur_demo_no_state",
object(), object(), "task",
1, 2, 3, 4, 5, 6, 0, 0,
)
self.assertEqual(
actual,
(
"table",
"/tmp/result.json",
"done Config=pi05_ur_demo_no_state (discrete_state_input=False).",
),
)
get_model.assert_called_once_with(
"owner/model", "checkpoint", "pi05_ur_demo_no_state"
)
def test_predict_ui_turns_exceptions_into_status(self):
import app
with mock.patch.object(app.MODEL_MANAGER, "get", side_effect=RuntimeError("load failed")):
table, output_file, status = app.predict_ui(
"model", "checkpoint", "pi05_ur_demo_state",
object(), object(), "task",
0, 0, 0, 0, 0, 0, 0, 0,
)
self.assertIsNone(table)
self.assertIsNone(output_file)
self.assertEqual(status, "Error: load failed")
def test_source_exposes_required_prediction_controls(self):
source = Path("app.py").read_text()
for label in (
"Fixed camera", "Wrist camera", "Task instruction", "Predict actions",
"TCP x", "TCP y", "TCP z", "TCP roll", "TCP pitch", "TCP yaw", "Gripper",
):
self.assertIn(label, source)
self.assertIn("default_concurrency_limit=1", source)
self.assertIn('label="Policy config"', source)
self.assertIn("value=DEFAULT_POLICY_CONFIG", source)
if __name__ == "__main__":
unittest.main()