mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-30 13:09:40 +00:00
feat(runtime): add grounded VQA command
This commit is contained in:
@@ -520,6 +520,7 @@ def _print_runtime_help() -> None:
|
|||||||
" /action resume the robot on the current task\n"
|
" /action resume the robot on the current task\n"
|
||||||
" /action <seconds> run the robot for N seconds, then auto-pause\n"
|
" /action <seconds> run the robot for N seconds, then auto-pause\n"
|
||||||
" /pause pause the action loop — robot holds position\n"
|
" /pause pause the action loop — robot holds position\n"
|
||||||
|
" /ask <question> pause and ask the policy about the current view\n"
|
||||||
" /help show this help\n"
|
" /help show this help\n"
|
||||||
" stop | quit | exit end the session",
|
" stop | quit | exit end the session",
|
||||||
flush=True,
|
flush=True,
|
||||||
@@ -552,6 +553,33 @@ def _clear_action_queue(runtime: Any) -> None:
|
|||||||
queue.clear()
|
queue.clear()
|
||||||
|
|
||||||
|
|
||||||
|
def _ask_runtime(runtime: Any, question: str) -> str:
|
||||||
|
"""Pause action dispatch and ask the adapter a grounded VQA question."""
|
||||||
|
question = question.strip()
|
||||||
|
if not question:
|
||||||
|
print("[runtime] usage: /ask <question>", flush=True)
|
||||||
|
return ""
|
||||||
|
runtime.state["mode"] = "paused"
|
||||||
|
runtime.state["action_deadline"] = None
|
||||||
|
_clear_action_queue(runtime)
|
||||||
|
generate_text = getattr(runtime.policy_adapter, "generate_text", None)
|
||||||
|
if not callable(generate_text):
|
||||||
|
print("[runtime] this policy adapter does not support text generation", flush=True)
|
||||||
|
return ""
|
||||||
|
observation = runtime._current_observation()
|
||||||
|
try:
|
||||||
|
answer = generate_text("vqa", observation, runtime.state, user_text=question)
|
||||||
|
except Exception as exc: # noqa: BLE001
|
||||||
|
logger.warning("VQA generation failed: %s", exc, exc_info=logger.isEnabledFor(logging.DEBUG))
|
||||||
|
print(f"[runtime] VQA failed: {type(exc).__name__}: {exc}", flush=True)
|
||||||
|
return ""
|
||||||
|
if not answer:
|
||||||
|
print("[runtime] the policy returned no answer", flush=True)
|
||||||
|
return ""
|
||||||
|
print(f"[policy] {answer}", flush=True)
|
||||||
|
return answer
|
||||||
|
|
||||||
|
|
||||||
def _handle_slash_command(runtime: Any, line: str) -> bool:
|
def _handle_slash_command(runtime: Any, line: str) -> bool:
|
||||||
"""Dispatch the runtime slash commands.
|
"""Dispatch the runtime slash commands.
|
||||||
|
|
||||||
@@ -560,6 +588,7 @@ def _handle_slash_command(runtime: Any, line: str) -> bool:
|
|||||||
(seconds), no argument resumes the current
|
(seconds), no argument resumes the current
|
||||||
task.
|
task.
|
||||||
``/pause`` pause the action loop — the robot holds.
|
``/pause`` pause the action loop — the robot holds.
|
||||||
|
``/ask <question>`` pause and ask about the current observation.
|
||||||
``/help`` print the command reference.
|
``/help`` print the command reference.
|
||||||
|
|
||||||
Returns ``True`` when ``line`` was a recognised command (consumed).
|
Returns ``True`` when ``line`` was a recognised command (consumed).
|
||||||
@@ -610,6 +639,10 @@ def _handle_slash_command(runtime: Any, line: str) -> bool:
|
|||||||
print("[runtime] paused — robot holding position", flush=True)
|
print("[runtime] paused — robot holding position", flush=True)
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
if cmd in {"/ask", "/vqa"}:
|
||||||
|
_ask_runtime(runtime, rest)
|
||||||
|
return True
|
||||||
|
|
||||||
if cmd in {"/help", "/?"}:
|
if cmd in {"/help", "/?"}:
|
||||||
_print_runtime_help()
|
_print_runtime_help()
|
||||||
return True
|
return True
|
||||||
@@ -994,6 +1027,8 @@ def _run_sim_interactive(
|
|||||||
if hasattr(runtime.policy, "reset"):
|
if hasattr(runtime.policy, "reset"):
|
||||||
runtime.policy.reset()
|
runtime.policy.reset()
|
||||||
print("[reset] new kitchen scene", flush=True)
|
print("[reset] new kitchen scene", flush=True)
|
||||||
|
elif low.startswith(("/ask ", "/vqa ")):
|
||||||
|
_ask_runtime(runtime, cmd.partition(" ")[2])
|
||||||
else:
|
else:
|
||||||
# Clear queued actions and rearm generation for a new command.
|
# Clear queued actions and rearm generation for a new command.
|
||||||
runtime.set_task(cmd)
|
runtime.set_task(cmd)
|
||||||
@@ -1094,6 +1129,8 @@ def _run_robot_interactive(
|
|||||||
elif low in {"/resume", "resume", "/run"}:
|
elif low in {"/resume", "resume", "/run"}:
|
||||||
runtime.state["mode"] = "action"
|
runtime.state["mode"] = "action"
|
||||||
print("[running]", flush=True)
|
print("[running]", flush=True)
|
||||||
|
elif low.startswith(("/ask ", "/vqa ")):
|
||||||
|
_ask_runtime(runtime, line.partition(" ")[2])
|
||||||
else:
|
else:
|
||||||
# New command: switch task/subtask immediately and regenerate.
|
# New command: switch task/subtask immediately and regenerate.
|
||||||
runtime.set_task(line)
|
runtime.set_task(line)
|
||||||
|
|||||||
@@ -18,7 +18,8 @@ from unittest.mock import MagicMock
|
|||||||
import pytest
|
import pytest
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from lerobot.runtime.cli import _build_rollout_runtime_io, _parse_args
|
from lerobot.runtime.cli import _ask_runtime, _build_rollout_runtime_io, _parse_args
|
||||||
|
from lerobot.runtime.language_runtime import RuntimeState
|
||||||
|
|
||||||
|
|
||||||
def test_parse_args_preserves_rollout_robot_overrides():
|
def test_parse_args_preserves_rollout_robot_overrides():
|
||||||
@@ -73,3 +74,27 @@ def test_rollout_runtime_io_uses_context_processors():
|
|||||||
|
|
||||||
assert observation["observation.state"].shape == (1, 1)
|
assert observation["observation.state"].shape == (1, 1)
|
||||||
robot.send_action.assert_called_once_with({"joint.pos": 2.0})
|
robot.send_action.assert_called_once_with({"joint.pos": 2.0})
|
||||||
|
|
||||||
|
|
||||||
|
def test_ask_runtime_pauses_and_routes_current_observation(capsys):
|
||||||
|
adapter = MagicMock()
|
||||||
|
adapter.generate_text.return_value = "The mug is beside the bowl."
|
||||||
|
runtime = SimpleNamespace(
|
||||||
|
state=RuntimeState(mode="action"),
|
||||||
|
policy_adapter=adapter,
|
||||||
|
_current_observation=lambda: {"image": "current"},
|
||||||
|
)
|
||||||
|
runtime.state.action_queue.extend([1, 2])
|
||||||
|
|
||||||
|
answer = _ask_runtime(runtime, "What is beside the bowl?")
|
||||||
|
|
||||||
|
assert answer == "The mug is beside the bowl."
|
||||||
|
assert runtime.state.mode == "paused"
|
||||||
|
assert not runtime.state.action_queue
|
||||||
|
adapter.generate_text.assert_called_once_with(
|
||||||
|
"vqa",
|
||||||
|
{"image": "current"},
|
||||||
|
runtime.state,
|
||||||
|
user_text="What is beside the bowl?",
|
||||||
|
)
|
||||||
|
assert "[policy] The mug is beside the bowl." in capsys.readouterr().out
|
||||||
|
|||||||
Reference in New Issue
Block a user