cockpit-agent/scripts/run_ui_understanding.py

65 lines
2.3 KiB
Python
Raw Permalink Normal View History

2026-08-24 16:49:44 +08:00
import argparse
import json
from pathlib import Path
from cockpit_agent.common.model_registry import ModelRegistry
from cockpit_agent.config import load_config
from cockpit_agent.grounding.adapter import GroundingAdapter
from cockpit_agent.grounding.semantic_target import SemanticTargetBuilder
from cockpit_agent.intent.parser import ModelIntentParser
from cockpit_agent.perception.roi import crop_grounding_roi
from cockpit_agent.perception.ui_understanding import ModelUIUnderstanding
DEFAULT_CONFIG = Path(__file__).resolve().parents[1] / "configs" / "default.toml"
def main() -> None:
parser = argparse.ArgumentParser(description="Understand current cockpit UI state")
parser.add_argument("--image", required=True)
parser.add_argument("--instruction", required=True)
parser.add_argument("--output-dir", required=True)
parser.add_argument("--config", type=Path, default=DEFAULT_CONFIG)
args = parser.parse_args()
config = load_config(args.config)
registry = ModelRegistry()
model = registry.get_text_generator(config.intent)
intent = ModelIntentParser(model, config.intent.max_new_tokens).parse(
args.instruction
)
target = SemanticTargetBuilder().build_function_target(intent)
grounding_model = registry.get(config.grounding)
function_grounding = GroundingAdapter(
grounding_model,
config.grounding.max_new_tokens,
).ground(args.image, target)
output_dir = Path(args.output_dir).expanduser().resolve()
roi = crop_grounding_roi(
image_path=args.image,
function_grounding=function_grounding,
output_path=str(output_dir / "roi.jpg"),
padding_ratio=config.perception.roi_padding_ratio,
)
ui_model = registry.get(config.ui_understanding)
ui_state = ModelUIUnderstanding(
ui_model,
config.ui_understanding.max_new_tokens,
).understand(roi.image_path, intent, function_grounding)
print(
json.dumps(
{
"function_target": target,
"function_grounding": function_grounding.to_dict(),
"roi": roi.to_dict(),
"ui_state": ui_state.to_dict(),
},
indent=2,
ensure_ascii=False,
)
)
if __name__ == "__main__":
main()