58 lines
1.9 KiB
Python
58 lines
1.9 KiB
Python
|
|
import argparse
|
||
|
|
import json
|
||
|
|
from pathlib import Path
|
||
|
|
from typing import Any
|
||
|
|
|
||
|
|
from cockpit_agent.common.model_registry import ModelRegistry
|
||
|
|
from cockpit_agent.config import load_config
|
||
|
|
from cockpit_agent.pipeline import build_pipeline
|
||
|
|
|
||
|
|
|
||
|
|
DEFAULT_CONFIG = Path(__file__).resolve().parents[1] / "configs" / "default.toml"
|
||
|
|
|
||
|
|
|
||
|
|
def parse_args() -> argparse.Namespace:
|
||
|
|
parser = argparse.ArgumentParser(
|
||
|
|
description="Run the single-pass cockpit agent with mock execution",
|
||
|
|
)
|
||
|
|
parser.add_argument("--image", required=True, help="Local cockpit image")
|
||
|
|
parser.add_argument("--instruction", required=True, help="User instruction")
|
||
|
|
parser.add_argument("--output-dir", required=True, help="Artifact directory")
|
||
|
|
parser.add_argument("--config", type=Path, default=DEFAULT_CONFIG)
|
||
|
|
return parser.parse_args()
|
||
|
|
|
||
|
|
|
||
|
|
def main() -> None:
|
||
|
|
args = parse_args()
|
||
|
|
config = load_config(args.config)
|
||
|
|
registry = ModelRegistry()
|
||
|
|
pipeline = build_pipeline(config, registry)
|
||
|
|
print(f"Loaded model instances: {registry.loaded_model_count}")
|
||
|
|
run = pipeline.run(
|
||
|
|
image_path=args.image,
|
||
|
|
instruction=args.instruction,
|
||
|
|
output_dir=args.output_dir,
|
||
|
|
)
|
||
|
|
_print_section("INTENT", run["semantic_action"])
|
||
|
|
print(f"\nFunction target: {run['function_target']}")
|
||
|
|
_print_section("FUNCTION GROUNDING", run["function_grounding"])
|
||
|
|
_print_section("ROI", run["roi"])
|
||
|
|
_print_section("UI UNDERSTANDING", run["ui_state"])
|
||
|
|
_print_section("PLAN", run["action_plan"])
|
||
|
|
print(f"\nAction target: {run['action_target']}")
|
||
|
|
_print_section("ACTION GROUNDING", run["action_grounding"])
|
||
|
|
_print_section("MOCK ACTION", run["mock_action"])
|
||
|
|
print(f"\nResult image: {run['result_image']}")
|
||
|
|
|
||
|
|
|
||
|
|
def _print_section(title: str, value: dict[str, Any] | None) -> None:
|
||
|
|
print()
|
||
|
|
print("=" * 50)
|
||
|
|
print(title)
|
||
|
|
print("=" * 50)
|
||
|
|
print(json.dumps(value, indent=2, ensure_ascii=False))
|
||
|
|
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
main()
|