```
├── .gitignore (700 tokens)
├── LICENSE (omitted)
├── README.md (4.3k tokens)
├── adb_utils.py (2.3k tokens)
├── api_task/
├── README.md (200 tokens)
├── api_logging.py (700 tokens)
├── api_task_service.py (3.8k tokens)
├── config_template.json (100 tokens)
├── assets/
├── audio/
├── voice_copy/
├── 1.mp3
├── 2.mp3
├── 3.mp3
├── 4.mp3
├── 5.mp3
├── 6.mp3
├── 7.mp3
├── 8.mp3
├── 9.mp3
├── voice_enterpage/
├── enterpage.mp3
├── voice_entersearch/
├── entersearch.mp3
├── voice_finish/
├── 1.mp3
├── 2.mp3
├── 3.mp3
├── 4.mp3
├── 5.mp3
├── finish6.mp3
├── finish7.mp3
├── finish8.mp3
├── finish9.mp3
├── voice_inputtext/
├── text.mp3
├── voice_openapp/
├── openapp.mp3
├── voice_point/
├── point.mp3
├── voice_press/
├── press.mp3
├── voice_self/
├── self.mp3
├── voice_swipe/
├── swipe.mp3
├── voice_temp/
├── final_page.mp3
├── voice_type/
├── type.mp3
├── voice_welcome/
├── 1.mp3
├── 2.mp3
├── 3.mp3
├── 4.mp3
├── 5.mp3
├── 6.mp3
├── 7.mp3
├── welcome10.mp3
├── welcome8.mp3
├── welcome9.mp3
├── audio/
├── audio_play.py (400 tokens)
├── tts.py (200 tokens)
├── cross_device/
├── function_call_utils.py (200 tokens)
├── instruction_mapper.py (400 tokens)
├── socket_utils.py (400 tokens)
├── vision_extractor.py (400 tokens)
├── cross_device_agent.py (1400 tokens)
├── eval/
├── eval_multi.sh (500 tokens)
├── run_eval_agent.py (1700 tokens)
├── run_predict_appcopilot_multi.py (1800 tokens)
├── run_predict_minicpm.py (1600 tokens)
├── utils/
├── SimHei.ttf
├── action_type.py (500 tokens)
├── action_utils.py (2.5k tokens)
├── convert_output.py (1100 tokens)
├── evaluator.py (4.4k tokens)
├── mv_vote.py (1300 tokens)
├── qwen_mobile_tool.py (4.3k tokens)
├── schema/
├── schema.json (300 tokens)
├── schema_for_extraction.json (800 tokens)
├── test_schema.py (3.8k tokens)
├── utils.py (1000 tokens)
├── utils_odyssey/
├── config.json (200 tokens)
├── configuration_qwen.py (400 tokens)
├── generation_config.json
├── his_index.json (6069.5k tokens)
├── model.safetensors.index.json (16.2k tokens)
├── modeling_qwen.py (10.1k tokens)
├── pytorch_model.bin.index.json (16.2k tokens)
├── qwen.tiktoken (512.2k tokens)
├── qwen_generation_utils.py (3k tokens)
├── special_tokens_map.json
├── tokenization_qwen.py (4.4k tokens)
├── tokenizer_config.json (100 tokens)
├── visual.py (2.9k tokens)
├── utils_qwen/
├── agent_function_call.py (2.3k tokens)
├── image_hash/
├── commands.txt (100 tokens)
├── hash_find.py (900 tokens)
├── image/
├── app1.jpg
├── app2.jpg
├── app3.jpg
├── app4.jpg
├── app5.png
├── page1.jpg
├── page2.jpg
├── page3.jpg
├── page4.jpg
├── page5.jpg
├── page6.jpg
├── page7.jpg
├── page8.jpg
├── search1.jpg
├── search10.jpg
├── search2.jpg
├── search3.jpg
├── search4.jpg
├── search5.jpg
├── search6.jpg
├── search7.jpg
├── search8.jpg
├── search9.jpg
├── images/
├── cuolecuolecuole.png
├── double_end.png
├── double_end_cn.png
├── emunew.png
├── logo.png
├── long_horizon.png
├── long_horizon_cn.png
├── triple_end.png
├── triple_end_cn.png
├── log/
├── experience_pool.py (700 tokens)
├── log_recorder.py (1100 tokens)
├── log_replay.py (400 tokens)
├── multi_step/
├── multi_step_execution.py (400 tokens)
├── multi_step_instruction.json (100 tokens)
├── ocr_model/
├── PP-OCRv5_server_det/
├── .gitattributes (300 tokens)
├── README.md (3.2k tokens)
├── config.json (600 tokens)
├── inference.json (80.5k tokens)
├── inference.pdiparams
├── inference.yml (200 tokens)
├── PP-OCRv5_server_rec/
├── .gitattributes (300 tokens)
├── README.md (3.2k tokens)
├── config.json (63.2k tokens)
├── inference.json (65k tokens)
├── inference.pdiparams
├── inference.yml (22.4k tokens)
├── omni_parser/
├── __init__.py
├── fix_log.jsonl
├── paser.py (1200 tokens)
├── payload.json (354.3k tokens)
├── readme/
├── README_chinese.md (2k tokens)
├── requirements.txt (100 tokens)
├── run_agent.py (3.1k tokens)
├── test/
├── api_test.py (400 tokens)
├── mv_test.py (1000 tokens)
├── user/
├── information.json (3k tokens)
├── ocr_service.py (1000 tokens)
├── user_manager.py (1700 tokens)
├── wrappers/
├── __init__.py
├── base_wrapper.py (200 tokens)
├── constants.py (1100 tokens)
├── cpm_wrapper.py (1800 tokens)
├── openai_wrapper.py (1000 tokens)
├── parallel_cpm_wrapper.py (1200 tokens)
├── qwenvl_wrapper.py (4k tokens)
├── schema_for_extraction.json (800 tokens)
├── schema_thought.json (300 tokens)
├── uitars_wrapper.py (1800 tokens)
├── utils.py (1500 tokens)
```
## /.gitignore
```gitignore path="/.gitignore"
# Created by https://www.toptal.com/developers/gitignore/api/python
# Edit at https://www.toptal.com/developers/gitignore?templates=python
### Python ###
# Byte-compiled / optimized / DLL files
__pycache__/
*.py[cod]
*$py.class
# C extensions
*.so
# Distribution / packaging
.Python
build/
develop-eggs/
dist/
downloads/
eggs/
.eggs/
lib/
lib64/
parts/
sdist/
var/
wheels/
share/python-wheels/
*.egg-info/
.installed.cfg
*.egg
MANIFEST
# PyInstaller
# Usually these files are written by a python script from a template
# before PyInstaller builds the exe, so as to inject date/other infos into it.
*.manifest
*.spec
# Installer logs
pip-log.txt
pip-delete-this-directory.txt
# Unit test / coverage reports
htmlcov/
.tox/
.nox/
.coverage
.coverage.*
.cache
nosetests.xml
coverage.xml
*.cover
*.py,cover
.hypothesis/
.pytest_cache/
cover/
# Translations
*.mo
*.pot
# Django stuff:
*.log
local_settings.py
db.sqlite3
db.sqlite3-journal
# Flask stuff:
instance/
.webassets-cache
# Scrapy stuff:
.scrapy
# Sphinx documentation
docs/_build/
# PyBuilder
.pybuilder/
target/
# Jupyter Notebook
.ipynb_checkpoints
# IPython
profile_default/
ipython_config.py
# pyenv
# For a library or package, you might want to ignore these files since the code is
# intended to run in multiple environments; otherwise, check them in:
# .python-version
# pipenv
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
# However, in case of collaboration, if having platform-specific dependencies or dependencies
# having no cross-platform support, pipenv may install dependencies that don't work, or not
# install all needed dependencies.
#Pipfile.lock
# poetry
# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
# This is especially recommended for binary packages to ensure reproducibility, and is more
# commonly ignored for libraries.
# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
#poetry.lock
# pdm
# Similar to Pipfile.lock, it is generally recommended to include pdm.lock in version control.
#pdm.lock
# pdm stores project-wide configurations in .pdm.toml, but it is recommended to not include it
# in version control.
# https://pdm.fming.dev/#use-with-ide
.pdm.toml
# PEP 582; used by e.g. github.com/David-OConnor/pyflow and github.com/pdm-project/pdm
__pypackages__/
# Celery stuff
celerybeat-schedule
celerybeat.pid
# SageMath parsed files
*.sage.py
# Environments
.env
.venv
env/
venv/
ENV/
env.bak/
venv.bak/
# Spyder project settings
.spyderproject
.spyproject
# Rope project settings
.ropeproject
# mkdocs documentation
/site
# mypy
.mypy_cache/
.dmypy.json
dmypy.json
# Pyre type checker
.pyre/
# pytype static type analyzer
.pytype/
# Cython debug symbols
cython_debug/
# PyCharm
# JetBrains specific template is maintained in a separate JetBrains.gitignore that can
# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore
# and can be added to the global gitignore or merged into this file. For a more nuclear
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
#.idea/
### Python Patch ###
# Poetry local configuration file - https://python-poetry.org/docs/configuration/#local-configuration
poetry.toml
# ruff
.ruff_cache/
# LSP config files
pyrightconfig.json
# End of https://www.toptal.com/developers/gitignore/api/python
# custom settings
*.iml
.idea/
.vscode/
.DS_Store
.gradle/
.pytest_cache/
**/__pycache__/
*.pyc
user/users_info/
user/ocr_output/
log/task_logs_new/
log/tasks_logs_new/
log/experience_pool.json
api_task/config.json
api_task/log/
```
## /README.md
# AppCopilot: Toward General, Accurate, Long‑Horizon, and Efficient Mobile Agent
<p align="center">
【English | <a href="readme/README_chinese.md">中文</a>】
</p>
<div align="center">
<img src="images/logo.png" alt="AppCopilot" width="600">
</div>
## 📖 Overview
With the rapid evolution of large language models and multimodal foundation models, the mobile-agent landscape has proliferated without converging on the fundamental challenges. This paper identifies four core problems that must be solved for mobile agents to deliver practical, scalable impact: (1) generalization across tasks, modalities, apps, and devices; (2) accuracy, specifically precise on-screen interaction and click targeting; (3) long-horizon capability for sustained, multi-step goals; and (4) efficiency, specifically high-performance runtime on resource-constrained devices.
We present AppCopilot, a multimodal, multi-agent, general-purpose on-device assistant that operates across applications and constitutes a full-stack, closed-loop system from data to deployment. AppCopilot operationalizes this position through an end-to-end autonomous pipeline spanning data collection, training, deployment, high-quality and efficient inference, and PC/mobile application development. At the model layer, it integrates multimodal foundation models with robust Chinese–English support. At the reasoning and control layer, it combines chain-of-thought reasoning, hierarchical task planning and decomposition, and multi-agent collaboration. At the execution layer, it enables user personalization and experiential adaptation, voice interaction, function/tool calling, cross-app and cross-device orchestration, and comprehensive mobile app support. The system design incorporates profiling-driven optimization for latency, memory, and energy across heterogeneous hardware.
Empirically, AppCopilot achieves significant improvements along all four dimensions: stronger generalization. higher-precision on-screen actions, more reliable long-horizon task completion, and faster, more resource-efficient runtime.
By articulating a cohesive position and a reference architecture that closes the loop from “data collection—training and deployment—high-quality, efficient inference—application development”, this paper offers a concrete roadmap for general-purpose digital assistants and provides actionable guidance for both academic research and industrial adoption.
## 🎉 News
At 2025-8-15, we are excited to announce the release of AppCopilot. AppCopilot is a general-purpose, on-device intelligent assistant that understands text and images, coordinates agents to complete complex tasks, works seamlessly across apps, and supports secure, real-time, cross-device collaboration.
## ⚡️ Quickstart
<details>
<summary>Click to expand</summary>
### AppCopilot Local Run
This section mainly introduces how to connect to the model trained on the server through the API and run AppCopilot locally.
#### Local Environment Basic Requirements
The following table shows the relevant dependency requirements for the local environment:
| **Dependency** | **Specific Requirements** |
|-----------------|---------------------------------------------------------------------|
| Operating System| An operating system that supports Android Studio |
| Software | Install Android Studio |
| Python Environment| Install Python environment, recommended Python version 3.12 |
| Network | Disable local VPN to ensure proper connection to the server's vllm API |
##### Install Android Studio
Android Studio is an integrated development environment (IDE) for Android platform development. It can be downloaded from the official [Android Studio website](https://developer.android.com/studio).
#### Server Environment Basic Requirements
The following table introduces the relevant dependency requirements for the server-side environment:
| **Dependency** | **Specific Requirements** |
|-----------------|---------------------------------------------------------------------|
| Operating System| An operating system that supports Conda and vLLM |
| Software | Install Conda, create a vLLM environment, and install vLLM dependencies|
##### Conda Installation
Conda is an open-source, cross-platform package manager and environment manager that helps users quickly install, run, and manage software packages and their dependencies. You can download it from the official [Conda website](https://anaconda.org/anaconda/conda).
After installing Conda, configure the Python virtual environment with the recommended Python version 3.12:
```bash
conda create --name vllm_env python=3.12
```
##### vLLM Installation
[vLLM](https://docs.vllm.ai/en/latest/) is an open-source high-performance library for large language model inference and services, providing faster responses for generative AI applications at a lower cost and higher efficiency. Here, configure the vLLM-related dependencies and install vLLM version 0.9.1 with the following command:
```bash
pip install vllm==0.9.1
```
##### Other Configuration
To connect to the server API and run AppCopilot, the other configuration requirements for the server environment are as follows::
```bash
pip install git+https://github.com/huggingface/transformers@f3f6c86582611976e72be054675e2bf0abb5f775
pip install accelerate
pip install qwen-vl-utils
pip install openai
git clone https://huggingface.co/Qwen/Qwen-VL-7B
```
#### Clone the Code
First, clone the folder from the remote repository to the local machine and add the necessary files:
```bash
mkdir AppCopilot
cd AppCopilot
git clone https://github.com/OpenBMB/AppCopilot.git .
```
To enhance the agent's ability to operate on Android phones, this project also requires the installation of the YADB tool to improve the native ADB functionality. It addresses the limitations of ADB in text input, screenshot capture, and UI layout extraction, providing more efficient and precise operations. Run the following command:
```bash
git clone https://github.com/ysbing/YADB.git ./YADB
```
#### Local System Environment Variable Configuration
##### Configure ADB Environment Variable
1.Windows System ADB Environment Variable Configuration:
On Windows, right-click on This PC, select Properties, and then click Advanced System Settings.
In the pop-up window, click Environment Variables, click New under System variables, enter the variable name: adb, and set the variable value to the directory path where adb is located (e.g., `C:\Android\Sdk\platform-tools`). Then find Path in System variables, and add the previously added ADB environment. Double-click Path, click New, and enter %adb%.
2.macOS/Linux System ADB Environment Variable Configuration:
On Linux or macOS, edit the ~/.bashrc or ~/.bash_profile file and add the ADB path at the end:
```bash
/Users/user/Android/Sdk/platform-tools
```
After saving the file, run `source ~/.bashrc` or `source ~/.bash_profile` to apply the configuration.
After completing the configuration, run `adb version` in the command line. If it correctly outputs the ADB version and related information, the configuration is successful.
##### Configure Emulator Environment Variable
The configuration method is similar to the ADB environment variable configuration.
1.Windows System Emulator Environment Variable Configuration:
On Windows, right-click on This PC, select Properties, and click Advanced System Settings.
In the pop-up window, click Environment Variables, click New under System variables, enter the variable name: emulator, and set the variable value to the directory path where emulator is located (e.g., `C:\Android\Sdk\emulator`). Then find Path in System variables, and add the previously added emulator environment. Double-click Path, click New, and enter %emulator%.
2.macOS/Linux System emulator Environment Variable Configuration:
On Linux or macOS, edit the ~/.bashrc or ~/.bash_profile file and add the emulator path at the end:
```bash
/Users/user/Library/Android/Sdk/emulator
```
After saving the file, run source ~/.bashrc or source ~/.bash_profile to apply the configuration. After completing the configuration, run `emulator version` in the command line. If it correctly outputs the emulator version and related information, the configuration is successful.
#### Configure the Android Device for Running
##### Configure Emulator Environment Variable
This project uses Android Studio to create and manage Android Virtual Devices (AVD). You can refer to the official Android Studio documentation to configure the emulator.
To view the list of available emulators and their names, run the following command:
```bash
emulator -avd <android> -dns-server <Local DNS Server>
```
Where `<android>` is the name of the emulator, and `<Local DNS Server>` is the local DNS address. You only need to specify the DNS Server the first time. After that, you can directly start the emulator with:`emulator -avd <android>`. If a snapshot corruption error occurs during debugging, you can add the `-no-snapshot-load` parameter when starting the emulator.
After completing the above configuration, the Android emulator should run locally, showing an interactive graphical interface, supporting mouse operations, and allowing network access through the host machine's network.
##### Configure Physical Device
In addition to using Android Virtual Devices (AVD), the agent can also control a physical Android phone through ADB. Below are the specific steps to use ADB to control a physical phone:
Enable Developer Mode on Physical Device:
On the phone, go to Settings -> About phone -> All parameters and information, and tap MIUI version 7 times to enable developer mode.
Enable USB Debugging Mode:
In Settings, find Developer options, and enable USB debugging.
Connect Physical Device via ADB:
Connect the phone to the computer via USB, then run adb devices in the command line. If you see the phone's serial number, the connection is successful.
##### Configure Python Environment Dependencies
It is recommended to install and use Python version 3.12. Enter the previously cloned GUI-Android directory and install the following dependencies:
```bash
pip install -r requirements.txt
```
##### Configure Model Keys
In the local code file `./wrappers/constants.py`, users need to manually configure the LLM key for future model calls:
```bash
# ----- model config -----
MODEL_EXTRACT = "AppCopilot"
ERROR_CALLING_LLM = "Error calling LLM"
MODEL_NOT_FOUND = "LLM not found"
# Replace with actual local endpoint port
END_POINT = "http://localhost:8001/v1/chat/completions"
PORTS = [8002, 8003, 8004]
# Replace with user-provided API key and Base URL
CLIENT_API_KEY = "switch to your own api key"
CLIENT_BASE_URL = "switch to your own base url"
CLIENT = OpenAI(api_key=CLIENT_API_KEY, base_url=CLIENT_BASE_URL)
```
##### Download AppCopilot Model
Download the pre-trained AppCopilot model from https://huggingface.co/ffcosmos/AppCopilot/tree/main and place it on the server to enable the next step of starting the vLLM inference service.
##### Start the vLLM Service on the Server
To enable AppCopilot to call the local large language model remotely, the vLLM inference service must be pre-deployed and started on the server side.
Start the server-side GUI model vLLM service:
```bash
#/your/model/path replace with actual GUI model path
vllm serve /your/model/path \
--served-model-name AppCopilot \
--tensor-parallel-size 1 \
--trust-remote-code \
--gpu-memory-utilization 0.9 \
--limit-mm-per-prompt image=10 \
--max_model_len 2048 \
--port 8001
```
Start the server-side Qwen2.5-VL-7B-Instruct model vLLM service:
```bash
#/your/model/path replace with actual Qwen2.5-VL-7B-Instruct model path
vllm serve /your/model/path \
--served-model-name Qwen2.5-VL-7B-Instruct \
--tensor-parallel-size 1 \
--trust-remote-code \
--gpu-memory-utilization 0.9 \
--port 8002
```
#### Local Run and Start AppCopilot
Before starting the program locally, you should first forward the port 8001 from the remote server to the local port 8001, and forward the port 8002 from the remote server to the local port 8002, to ensure that the local environment can access the model services on the server via the HTTP interface. This port forwarding operation can be executed via the terminal with the following commands:
```bash
ssh -L 8001:localhost:8001 username@model-server-ip
ssh -L 8002:localhost:8002 username@model-server-ip
```
##### Single-Device Run
Finally, to run AppCopilot on a single device, open the command-line interface in the terminal, navigate to the directory containing the `run_agent.py` file, and run the script with the required parameters. The following is an example command that enables voice input, audio feedback, and runs a custom task:
```bash
# Enable voice input, audio feedback, and run a custom task
python run_agent.py --custom-task
```
| Parameter | Type | Description |
| ---------------------------- | ------ | ---------------------------------------------- |
| `--predefined-task <TASK_NAME>` | str | Specify the name of a predefined task (task name must be in the built-in list). |
| `--custom-task` | flag | Enable custom task mode, skip predefined task selection. |
| `--enable-experience` | flag | Enable experience-based task matching mechanism. |
| `--enable-voice-input` | flag | Enable voice input (only valid in custom task mode). |
| `--enable-audio` | flag | Enable audio feedback. |
| `--show-tasks` | flag | Show all available predefined tasks and exit the program. |
| `--enable-vision-parser` | flag | Whether to call omniparser for coordinate calibration. |
| `--read-final-page` | flag | Whether to enable reading the final page. |
##### Multi-Device Cross-End Run
For multi-device cross-end scenarios, navigate to the directory containing the cross_device_agent.py file and run the script with the required parameters. The following table shows the available command-line arguments for cross-device running:
| Parameter | Type | Description |
| -------------------------- | ----- | ------------------------------------------------ |
| `--device1-serial` | str | ADB serial number for device 1 (optional) |
| `--device1-port` | int | Communication port for device 1 (default 11001). |
| `--device2-serial` | str | ADB serial number for device 2 (optional) |
| `--device2-port` | int | Communication port for device 2 (default 11002). |
| `--task` | str | Cross-device task instruction. |
### Model Inference Evaluation
#### Data Preparation
##### Android Control
Download [Android Control](https://github.com/google-research/google-research/tree/master/android_control) and save at ``eval/eval_data/tmp/android_control``
```
cd eval/eval_data
python process_ac.py
ln -s android_control_test android_control_high_test
ln -s android_control_test android_control_low_test
```
##### CAGUI
```
cd eval/eval_data
mkdir chinese_app_test && cd chinese_app_test
huggingface-cli download openbmb/CAGUI --repo-type dataset --include "CAGUI_agent/**" --local-dir ./ --local-dir-use-symlinks False --resume-download
mv CAGUI_agent test
```
##### aitz
Download [aitz](https://github.com/IMNearth/CoAT) and save at ``eval/eval_data/tmp/android_in_the_zoo``
```
cd eval/eval_data
mv tmp/android_in_the_zoo ./aitz_test
python process_aitz.py
```
##### gui-odyssey
Download [GUI-Odyssey](https://github.com/OpenGVLab/GUI-Odyssey?tab=readme-ov-file) and save at ``/eval/eval_data/tmp/GUI-Odyssey``. Copy [preprocessing.py](https://github.com/OpenGVLab/GUI-Odyssey/blob/master/data/preprocessing.py) and [format_converter.py](https://github.com/OpenGVLab/GUI-Odyssey/blob/master/data/format_converter.py) from the GUI-Odyssey repo to ``/eval/eval_data/tmp/GUI-Odyssey``
```
cd eval/eval_data/tmp/GUI-Odyssey
python preprocessing.py
python format_converter.py
python ../../process_odyssey.py
```
#### Running Inference
The scripts required for model evaluation are integrated into the `eval_multi.sh` script. Before running, please modify the path parameters in the script based on the actual data storage locations to ensure the correct loading and processing of files.
```bash
# Contents that need to be modified in eval.sh
# Configure basic parameters
data_name="evaluation dataset"
model_name="target model name"
base_output_dir="result directory"
# List of models to process
models_base_path=(
"models base path"
)
```
Before running model inference evaluation, ensure that the utils folder is correctly configured on the server. After configuring the path parameters correctly, you can execute the following command to start the model inference evaluation process:
```bash
bash eval_multi.sh
```
</details>
## ✨ **Demo Cases**
### Case 1: Single device control

This figure shows the execution using the "search bar filtering" path, demonstrating **active control** over requirement boundaries. The core logic is **precise requirement decomposition**: first search "restaurant" for full coverage, then set "Nearby" distance filter, finally switch to "Highest Rating" sorting. Each step directly corresponds to core instruction conditions.
The task is successfully completed, demonstrating long horizon capabilities in complex application scenarios.
### Case 2: Two-device coordinated control

As shown in this figure, Lili's device stores 5G Kuan Shijie viewing history data, while the user's device completes gift purchasing. This demonstrates **cross-device multi-agent collaboration**, **user preference extraction**, and **cross-application decision-making**.
After authorization verification, the Agent locates the history module and extracts key information from the most recent video. But raw video lists can't directly guide gift selection. Here, the Agent extracts IP keywords from "Crayon Shin-chan" to infer potential interests. This transcends simple data transfer by achieving a leap from data to preference to demand through content understanding.
On the user's device, the Agent receives "Crayon Shin-chan" keywords and launches Taobao. It locates relevant gifts through search, maintaining process coherence across multiple operations.
Critically, the task highlights the core value of cross-device service—breaking device barriers to achieve precise data-to-service docking. Traditional scenarios require manual preference inquiry and product search; the Agent automates the entire process from data collection to product recommendation.
### Case 3: Three-device coordinated control

As shown in this figure, Lili's and Fanfan's devices store viewing history data, while the user's device completes gift purchasing. This extends agent capabilities **from individual users to multi-user collaborative operations**. This isn't simple technical stacking but a paradigm shift from personal assistants to distributed collaboration networks.
Each agent focuses on parsing its own 5G Kuan Shijie history to **extract individual preferences**. The user's agent integrates preference data from both ends to drive targeted gift selection. In this architecture, each agent has independent computational space and decision boundaries, preserving core attributes representing individual intent while breaking limitations through collaboration.
The multi-device task confirms the feasibility of a **distributed intelligence system**. First, addressing reasoning and coordination challenges with incomplete information: Lili's and Fanfan's agents only know their own task status; the user's agent cannot directly access raw data on other devices. Through **data desensitization and intent recognition mechanisms**, agents collaborate accurately despite information gaps.
Second, addressing communication and negotiation mechanisms: agents achieve precise intent transmission through unified protocols despite heterogeneous systems. The successful execution validates that the **mobile agent system has upgraded from single-agent to a system-level architecture with multi-agent collaboration, distributed state modeling, and mechanism design capabilities**. This upgrade's core value enables intelligent services to break single-user boundaries and complete complex cross-domain long-horizon tasks through multiple autonomous agents collaborating—moving toward realizing the "theoretically expandable to massive terminals" vision of collective intelligence.
For more examples, please refer to the original paper.
## 🔎 Citation
```
@article{AppCopilot,
title = {AppCopilot: Toward General, Accurate, Long‑Horizon, and Efficient Mobile Agent},
author = {Jingru Fan and Yufan Dang and Jingyao Wu and Huatao Li and Runde Yang and Xiyuan Yang and Yuheng Wang and Chen Qian},
journal = {arXiv preprint arXiv:2509.02444},
url = {https://arxiv.org/abs/2509.02444},
year = {2025}
}
```
## 📬 Contact
If you have any questions, feedback, or would like to get in touch, please feel free to reach out to us via email at [qianc@sjtu.edu.cn](mailto:qianc@sjtu.edu.cn)
## /adb_utils.py
```py path="/adb_utils.py"
import subprocess
import datetime
import urllib.parse
import logging
import os
import io
import PIL.Image as Image
from typing import List, Dict, Any, Optional
from audio.audio_play import play_random_audio, VoiceType
from image_hash.hash_find import start_hash_find
from audio.tts import run_tts
logger = logging.getLogger(__name__)
logging.basicConfig(level=logging.INFO,
format="%(asctime)s %(levelname)s %(name)s: %(message)s")
# ---------------------------------------------------------------------------
# Low‑level helpers
# ---------------------------------------------------------------------------
def _run(cmd: List[str], timeout: int = 30) -> bytes:
"""Run a shell command and return raw stdout (raises on non‑zero exit)."""
logger.debug("$ %s", " ".join(cmd))
return subprocess.check_output(cmd, stderr=subprocess.STDOUT, timeout=timeout)
def _adb_prefix(serial: str | None) -> List[str]:
return ["adb", "-s", serial] if serial else ["adb"]
def _resize_pillow(origin_img, max_line_res: int = 1120):
"""Resize PIL image so that longest edge ≤ `max_line_res` using Lanczos."""
w, h = origin_img.size
if max_line_res is not None:
max_line = max_line_res
if h > max_line:
w = int(w * max_line / h)
h = max_line
if w > max_line:
h = int(h * max_line / w)
w = max_line
return origin_img.resize((w, h), resample=Image.Resampling.LANCZOS)
def _encode_text_for_adb(text: str) -> str:
"""Encode text for adb shell input. URL‑encode spaces as %s."""
def _esc(ch: str) -> str:
if ord(ch) < 128 and ch != " ":
return ch
if ch == " ":
return "%s"
return f"\\u{ord(ch):04x}"
return "".join(_esc(c) for c in text)
def _encode_ascii_for_adb(text: str) -> str:
"""Encode ASCII‑only string for `adb shell input text …` (spaces→%s)."""
return text.replace(" ", "%s")
# ---------------------------------------------------------------------------
# AndroidDevice class
# ---------------------------------------------------------------------------
class AndroidDevice:
"""Encapsulates a single, already‑connected Android handset."""
_yadb_pushed: bool = False
_yadb_local: str = os.path.join(os.path.dirname(__file__), "yadb/yadb")
def __init__(self, serial: str | None, audio_enable=False):
self.serial: str | None = serial
self.width: int = 0
self.height: int = 0
self.last_req_time: datetime.datetime = datetime.datetime.now()
self.audio_enable = audio_enable
# ---------- internal ----------
def _adb(self, *args: str, timeout: int = 30) -> bytes:
return _run(_adb_prefix(self.serial) + list(args), timeout)
def _ensure_yadb(self):
if AndroidDevice._yadb_pushed:
return
if not os.path.exists(AndroidDevice._yadb_local):
raise FileNotFoundError(f"yadb helper not found: {AndroidDevice._yadb_local}")
self._adb("push", AndroidDevice._yadb_local, "/data/local/tmp")
AndroidDevice._yadb_pushed = True
logger.info("yadb pushed to device for Unicode input support")
# ---------- public API ----------
def refresh_resolution(self) -> None:
"""Query and cache `wm size` (sets .width / .height)."""
raw = self._adb("shell", "wm", "size").decode()
try:
size_line = raw.split("Physical size: ")[1].splitlines()[0]
self.width, self.height = map(int, size_line.split("x"))
logger.info("Device %s resolution: %dx%d", self.serial or "<default>",
self.width, self.height)
except Exception as exc:
raise RuntimeError(f"Failed to parse wm size output: {raw}") from exc
# -------------------------------------------------------------------
# Step: execute user action
# -------------------------------------------------------------------
def step(self, data: Dict[str, Any]) -> None:
"""Execute a control step on the device (tap/swipe/key/text/clear)."""
logger.debug("Step: %s", data)
if "POINT" in data:
self._handle_point(data)
if "PRESS" in data:
self._handle_press(data["PRESS"])
if "TYPE" in data:
self._handle_type(data["TYPE"])
if "CLEAR" in data:
self._adb("shell", "input", "keyevent", "KEYCODE_CLEAR")
self.last_req_time = datetime.datetime.now()
if ("STATUS", "finish") in data.items() or ("STATUS", "impossible") in data.items():
logger.info("Task finished")
return True
return False
# -------------------------------------------------------------------
# State snapshot
# -------------------------------------------------------------------
def state(self) -> Dict[str, Any]:
return {
"width": self.width,
"height": self.height,
"last_req_time": self.last_req_time.isoformat(),
"screenshot": self.screenshot(),
}
# --- Device state ---------------------------------------------------
def screenshot(self, max_side: Optional[int] = None) -> Image.Image:
"""Grab screen; return Pillow Image. Optionally down‑scale with user rule."""
png_bytes = self._adb("exec-out", "screencap", "-p")
img = Image.open(io.BytesIO(png_bytes))
if max_side is not None:
img = _resize_pillow(img, max_side)
return img
# =================== private helpers ===================
def _handle_point(self, data: Dict[str, Any]) -> None:
search_thershold = 150
x1, y1 = data["POINT"]
print("point:",x1,y1)
x = int(x1 / 1000 * self.width)
y = int(y1 / 1000 * self.height)
if "to" in data:
if self.audio_enable:
play_random_audio(VoiceType.SWIPE)
if isinstance(data["to"], list):
x2, y2 = data["to"]
x2 = int(x2 / 1000 * self.width)
y2 = int(y2 / 1000 * self.height)
else: # directional swipe (up/down/left/right)
dirs = {
"up": (0, -0.15),
"down": (0, 0.15),
"left": (-0.15, 0),
"right": (0.15, 0),
}
if data["to"] not in dirs:
raise ValueError(f"Invalid swipe direction: {data['to']}")
dx_ratio, dy_ratio = dirs[data["to"]]
x2 = int(max(min(x + dx_ratio * self.width, self.width), 0))
y2 = int(max(min(y + dy_ratio * self.height, self.height), 0))
dur = str(data.get("duration", 150))
# Perform swipe
self._adb("shell", "input", "swipe", str(x), str(y), str(x2), str(y2), dur)
else: # simple tap
if self.audio_enable:
result = start_hash_find()
if result == "openapp":
play_random_audio(VoiceType.OPENAPP)
if result == "entersearch":
if y1< search_thershold:
play_random_audio(VoiceType.ENTERSEARCH)
else:
play_random_audio(VoiceType.ENTERPAGE)
if result == "enterpage":
if y1< search_thershold:
play_random_audio(VoiceType.ENTERSEARCH)
else:
play_random_audio(VoiceType.ENTERPAGE)
# Show tap location
#print("tap:", x,y)
self._adb("shell", "input", "tap", str(x), str(y))
def _handle_press(self, key: str) -> None:
if self.audio_enable:
play_random_audio(VoiceType.PRESS)
KEYS = {
"HOME": "KEYCODE_HOME",
"BACK": "KEYCODE_BACK",
"MENU": "KEYCODE_MENU",
"ENTER": "KEYCODE_ENTER",
"APPSELECT": "KEYCODE_APP_SWITCH",
"power": "KEYCODE_POWER",
"volume_up": "KEYCODE_VOLUME_UP",
"volume_down": "KEYCODE_VOLUME_DOWN",
"volume_mute": "KEYCODE_VOLUME_MUTE",
}
if key not in KEYS:
raise ValueError(f"Unknown PRESS value: {key}")
self._adb("shell", "input", "keyevent", KEYS[key])
# def _handle_type(self, raw):
# decoded = urllib.parse.unquote(raw)
# self._adb("shell", "am", "broadcast", '-a', 'ADB_INPUT_TEXT', '--es msg' , decoded)
# # self._adb("shell", "input", "text", decoded)
def _handle_type(self, raw):
text = urllib.parse.unquote(raw)
if self.audio_enable:
run_tts(text, output="assets/audio/voice_inputtext/text.mp3")
play_random_audio(VoiceType.TYPE)
play_random_audio(VoiceType.INPUTTEXT)
if all(ord(c) < 128 for c in text): # quick ASCII path
self._adb("shell", "input", "text", _encode_ascii_for_adb(text))
return
# Unicode → yadb
self._ensure_yadb()
safe = text.replace("'", "'\\''") # escape sigingle quotes for sh
cmd = (
"app_process -Djava.class.path=/data/local/tmp/yadb /data/local/tmp "
"com.ysbing.yadb.Main -keyboard '%s'" % safe
)
self._adb("shell", cmd)
# ---------------------------------------------------------------------------
# Public utility function
# ---------------------------------------------------------------------------
def list_connected_devices() -> List[str]:
"""列出所有已连接的ADB设备"""
lines = _run(["adb", "devices"]).decode().strip().splitlines()[1:]
return [l.split()[0] for l in lines if l.strip() and "device" in l]
# 支持指定设备序列号
def setup_device(serial: str | None = None, audio_enable=False) -> AndroidDevice:
"""创建AndroidDevice实例,可指定设备序列号"""
if serial is None:
devices = list_connected_devices()
if not devices:
raise RuntimeError("No authorised Android device found. Plug in & check adb.")
if len(devices) > 1:
logger.warning("Multiple devices detected; defaulting to the first (%s).", devices[0])
serial = devices[0]
dev = AndroidDevice(serial, audio_enable=audio_enable)
dev.refresh_resolution()
return dev
def change_ui_settings(mode: str = "open"):
mode = mode.strip()
if mode == "open":
subprocess.run("adb shell settings put system pointer_location 1", shell=True)
elif mode == "close":
subprocess.run("adb shell settings put system pointer_location 0", shell=True)
else:
raise ValueError("Invalid Openmode")
# ---------------------------------------------------------------------------
# Demo – run this file directly to test
# ---------------------------------------------------------------------------
if __name__ == "__main__":
device = setup_device()
logger.info("Device ready: serial=%s (%dx%d)", device.serial, device.width, device.height)
# Example: tap centre, take screenshot
x = 900
y = 800
device.step({"POINT": [x, y]})
png = device.screenshot()
target = os.path.join("screenshots", "screencap.png")
# logger.info("Screenshot saved → %s (%d bytes)", target, len(png))
```
## /api_task/README.md
# API Based Operations
- You need to create a `config.json` file under the specified `CONFIG_PATH` to store authentication information. The default value is `CONFIG_PATH = "./api_task/config.json"`.
- For the specific format, see [config template](config_template.json).
- For details about Bilibili authentication information, refer to the documentation: [Credential](https://nemo2011.github.io/bilibili-api/#/get-credential)
- All Bilibili-related code uses **asynchronous operations**. Since web scraping is involved, excessive high concurrency may lead to account bans.
## Usage
Currently supported API operations:
- Automatic email sending:
- Supports attaching local files
- Supports attaching files from Android phones using adb shell
- Supports binary files such as photos and videos
- Bilibili-related operations:
- Supports like/unlike, coin, and one-click triple actions
- Supports video search, crawling user information, and video information
## Demo
See [API_Demo](../test/api_test.py) for more information.
## /api_task/api_logging.py
```py path="/api_task/api_logging.py"
import logging
import os
import sys
from datetime import datetime
from pathlib import Path
# utils
def generate_timestamp():
return datetime.now().strftime("%Y%m%d-%H%M%S")
# Define a constant for the shared logger name
SHARED_LOGGER_NAME = "api_app_logger"
def setup_logging_config():
"""
Configures and returns a shared logger instance.
Ensures that the configuration is done only once to avoid duplicate handlers.
"""
# Attempt to get the existing logger instance
logger = logging.getLogger(SHARED_LOGGER_NAME)
# Check if the logger has already been configured with handlers; if so, return it directly.
# This prevents adding handlers multiple times if setup_logging_config() is called more than once.
if logger.handlers:
return logger
# If the logger is not configured, proceed with configuration.
logger.setLevel(logging.INFO) # Set the minimum logging level for the logger.
log_dir_home = Path.cwd() / "api_task" / "log"
log_file_path = None
try:
log_dir_home.mkdir(parents=True, exist_ok=True)
potential_log_file_path = log_dir_home / f"api_{generate_timestamp()}.log"
# Try to create or open the file to check for write permissions.
with potential_log_file_path.open("a", encoding="utf-8") as f:
f.write("") # Try to write an empty string to ensure writability.
log_file_path = str(potential_log_file_path)
except OSError as e:
print(
f"Warning: Could not create or write to log file at '{log_dir_home}'. Using /tmp instead. Error: {e}",
file=sys.stderr,
)
tmp_dir = Path("/tmp")
log_file_path = str(tmp_dir / "gpu_monitor.log")
try:
tmp_dir.mkdir(parents=True, exist_ok=True)
with (tmp_dir / "gpu_monitor.log").open("a", encoding="utf-8") as f:
f.write("") # Try again to ensure writability in /tmp.
except OSError as e:
print(
f"Critical Warning: Could not create or write to log file in /tmp. File logging will be disabled. Error: {e}",
file=sys.stderr,
)
log_file_path = None # Could not write to file, disabling file logging.
# Define the log format.
formatter = logging.Formatter("%(asctime)s %(levelname)s [%(name)s]: %(message)s")
# File Handler
if log_file_path:
file_handler = logging.FileHandler(log_file_path, encoding="utf-8")
file_handler.setLevel(
logging.INFO
) # The file handler records INFO level and above.
file_handler.setFormatter(formatter)
logger.addHandler(file_handler)
# Console Handler
console_handler = logging.StreamHandler(
sys.stdout
) # Explicitly direct output to stdout.
console_handler.setLevel(
logging.WARNING
) # The console handler only records WARNING level and above.
console_handler.setFormatter(formatter)
logger.addHandler(console_handler)
# Disable propagation to prevent log events from being passed to the root logger, which would cause duplicate output.
logger.propagate = False
return logger
if __name__ == "__main__":
setup_logging_config()
```
## /api_task/api_task_service.py
```py path="/api_task/api_task_service.py"
# Add two API modules: Bilibili-related video operations and automatic email sending assistant
import sys
import smtplib
import os
import json
import subprocess
import mimetypes
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from email.mime.multipart import MIMEMultipart
from email.mime.text import MIMEText
from email.mime.base import MIMEBase
from email import encoders
from email_validator import validate_email, EmailNotValidError
from typing import Optional, List, Tuple
from api_task.api_logging import setup_logging_config
from pathlib import Path
logger = setup_logging_config()
# pip3 install bilibili-api-python
from bilibili_api import Credential, user, sync, video, search
CONFIG_PATH = "./api_task/config.json"
def get_config_block(config_path, config_name: str) -> dict:
with open(config_path, "r", encoding="utf8") as file:
data = json.load(file)
for block in data:
if block.get("config_name") == config_name:
return block
raise ValueError(f"Config block with config_name '{config_name}' not found.")
# feat: add adb pull command for pulling files into local for email.
def pull_file_from_android(
device_file_path: str = "/sdcard/DCIM/Camera/",
local_destination_dir: Path | str | None = None,
) -> Path | None:
"""
Use the ADB pull command to pull files from an Android device to local.
Args:
device_file_path (str): Full path of the file on the Android device (e.g. /sdcard/Documents/my_file.txt).
local_destination_dir (Path | str | None): Target directory to store the file on the local computer.
If None, defaults to 'api_task/log/android_file' under the current working directory.
Returns:
Path | None: If successful, returns the full Path object of the pulled local file; otherwise returns None.
"""
# Handle default value and type for local_destination_dir
if local_destination_dir is None:
local_destination_dir = Path.cwd() / "api_task" / "log" / "android_file"
elif isinstance(local_destination_dir, str):
local_destination_dir = Path(local_destination_dir)
try:
os.makedirs(local_destination_dir, exist_ok=True)
logger.info(f"Created local destination directory: {local_destination_dir}")
except OSError as e:
logger.error(
f"Failed to create local destination directory {local_destination_dir}: {e}"
)
return None
file_name = Path(device_file_path).name
pulled_local_path = local_destination_dir / file_name # Compose the full local path
# Build the ADB pull command
# Note: If the target path of adb pull is a directory, it will create a file with the same name in that directory
command = [
"adb",
"pull",
device_file_path,
str(local_destination_dir),
] # Convert Path object to string for subprocess
logger.info(f"Attempting to pull file: {' '.join(command)}")
try:
result = subprocess.run(
command,
check=True, # Raises CalledProcessError if the command returns a non-zero exit code
capture_output=True, # Capture stdout and stderr
text=True, # Decode output as text
)
# Verify if the file was actually pulled to local
if pulled_local_path.exists():
logger.info(
f"Successfully pulled {device_file_path} to {pulled_local_path}"
)
logger.debug(f"ADB stdout: {result.stdout.strip()}")
logger.debug(f"ADB stderr: {result.stderr.strip()}")
return pulled_local_path
else:
logger.error(
f"Failed to pull file {device_file_path}. Local file not found at {pulled_local_path} after pull."
)
logger.debug(f"ADB stdout: {result.stdout.strip()}")
logger.debug(f"ADB stderr: {result.stderr.strip()}")
return None
except subprocess.CalledProcessError as e:
logger.error(
f"Failed to pull file {device_file_path}. ADB Error: {e.stderr.strip()}"
)
logger.debug(f"ADB stdout: {e.stdout.strip()}")
return None
except Exception as e:
logger.error(f"An unexpected error occurred during file pull: {e}")
return None
class EmailSender:
"""
A class to send emails with optional attachments using SMTP.
Attributes:
email_config_path (str): Path to the JSON config file containing sender and SMTP info.
sender_email (str): Sender's email address.
sender_password (str): Sender's email password.
smtp_server (str): SMTP server address.
smtp_port (int): SMTP server port.
"""
def __init__(self, email_config_path: str = CONFIG_PATH):
"""
Initialize the EmailSender by loading configuration from a JSON file.
Args:
email_config_path (str): Path to the JSON config file.
"""
self.email_config_path = email_config_path
self.sender_email, self.sender_password = self._get_usr_config()
self.smtp_server, self.smtp_port = self._get_smtp_config()
self.logger = setup_logging_config()
def _get_usr_config(self) -> Tuple[str, str]:
block = get_config_block(CONFIG_PATH, "email")
return str(block["sender_email"]).strip(), str(block["sender_password"]).strip()
def _get_smtp_config(self) -> Tuple[str, int]:
block = get_config_block(CONFIG_PATH, "email")
return str(block["smtp_server"]).strip(), block["smtp_port"]
def _attach_file(self, msg: MIMEMultipart, file_path: str):
"""
Attach a file to the email message if the file exists.
Args:
msg (MIMEMultipart): The email message object.
file_path (str): Path to the file to attach.
"""
if os.path.exists(file_path):
try:
# Try to get MIME type based on file extension
ctype, encoding = mimetypes.guess_type(file_path)
if ctype is None or encoding is not None:
# Fallback to generic type if unable to guess or encoding exists
ctype = "application/octet-stream"
maintype, subtype = ctype.split("/", 1)
with open(file_path, "rb") as attachment:
part2 = MIMEBase(maintype, subtype)
part2.set_payload(attachment.read())
encoders.encode_base64(part2)
filename = os.path.basename(file_path)
part2.add_header(
"Content-Disposition",
f"attachment; filename*=utf-8''{filename}",
)
part2.add_header(
"Content-Disposition", f"attachment; filename*=utf-8''{filename}"
)
# # Note: If the filename contains non-ASCII characters, there may be garbled text or replacements here
# part2.add_header(
# "Content-Disposition",
# f"attachment; filename=\"{filename}\"" # Use double quotes to enclose the filename
# )
msg.attach(part2)
except Exception as e:
self.logger.warning(f"Could not attach file {file_path}: {e}")
else:
self.logger.warning(
f"Warning: Attachment file not found at {file_path}. Skipping this attachment."
)
def send_mail(
self,
receiver_email: str,
subject: str = "hello world",
body: str = "hello world, just for fun!",
attach_local_file_path: Optional[str | List[str]] = None,
attach_android_file_path: Optional[str | List[str]] = None,
):
"""
Send an email with optional attachments.
Args:
receiver_email (str): Recipient's email address.
subject (str): Email subject.
body (str): Email body (plain text).
attach_file_path (Optional[str | List[str]]): Path(s) to files to attach.
sender_email (Optional[str]): Override sender email (default: from config).
sender_password (Optional[str]): Override sender password (default: from config).
Raises:
smtplib.SMTPAuthenticationError: If authentication fails.
Exception: For other errors during sending.
"""
# Validate sender email
try:
validate_email(self.sender_email, check_deliverability=False)
except EmailNotValidError as e:
self.logger.error(f"Sender email '{self.sender_email}' is not valid: {e}")
return
# --- Create the email ---
msg = MIMEMultipart()
msg["From"] = f"<{self.sender_email}>"
msg["To"] = receiver_email
msg["Subject"] = subject
msg.attach(MIMEText(body, "plain", "utf-8"))
msg.add_header("Content-Type", 'multipart/mixed; charset="utf-8"')
# --- Add attachments ---
# --- Add attachments on Android
if attach_android_file_path is not None:
if isinstance(attach_android_file_path, str):
attach_android_file_path = [attach_android_file_path]
# all transfer into local devices
new_path = [
str(pull_file_from_android(device_file_path=an_file_path))
for an_file_path in attach_android_file_path
]
for file_path_n in new_path:
self._attach_file(msg, file_path_n)
if attach_local_file_path is not None:
if isinstance(attach_local_file_path, str):
attach_local_file_path = [attach_local_file_path]
for file_path in attach_local_file_path:
self._attach_file(msg, file_path)
# --- Send the email ---
try:
with smtplib.SMTP_SSL(self.smtp_server, self.smtp_port) as server:
server.login(self.sender_email, self.sender_password)
server.send_message(msg)
self.logger.info("Email sent successfully!")
self.logger.info(
f"Message sent from {self.sender_email} to {receiver_email}"
)
except smtplib.SMTPAuthenticationError:
self.logger.error("Failed to send email!")
self.logger.error(
"Authentication failed: Please check if the sender email address and password/authorization code are correct."
)
except Exception as e:
self.logger.error("Failed to send email!")
self.logger.error(f"An error occurred: {e}")
class BilibiliOperator:
def __init__(self) -> None:
self._load_config()
self.credential = Credential(
sessdata=self.SESSDATA,
bili_jct=self.BILI_JCT,
buvid3=self.BUVID3,
dedeuserid=self.user_id,
)
self.logger = setup_logging_config()
def _load_config(self, config_name: str = "bilibili"):
block = get_config_block(CONFIG_PATH, config_name)
self.SESSDATA = block["SESSDATA"]
self.BILI_JCT = block["BILI_JCT"]
self.BUVID3 = block["BUVID3"]
self.user_id = block["dedeuserid"]
async def search_video(self, keyword: str):
# return await search.search_by_type(
# keyword=keyword,
# search_type=search.SearchObjectType.VIDEO,
# order_type=search.OrderVideo.SCORES,
# page=1,
# time_range=10
# )
return await search.search(keyword)
# Returns a relatively raw dictionary
async def get_info(self, video_bvid):
try:
v = video.Video(bvid=video_bvid, credential=self.credential)
info = await v.get_info()
self.logger.info(f"Video Title: {info.get('title', 'N/A')}")
self.logger.info(
f"Current like status of video {video_bvid}: {info.get('like', 'N/A')} (Bilibili API's 'like' field usually reflects your like status)"
)
except Exception as e:
self.logger.error(f"Error while fetching video info {video_bvid}: {e}")
async def like(self, video_bvid: str, mode: bool = True):
"""
Likes or unlikes a Bilibili video.
Args:
video_bvid (str): The BV ID of the video.
mode (bool): True for liking (default), False for unliking.
"""
# Determine the action for print statements
action_text = "liking" if mode else "unliking"
like_status_for_api = mode # v.like(True) for like, v.like(False) for unlike
self.logger.info(f"Attempting to {action_text} video: {video_bvid}")
try:
# Instantiate Video object with the BVID and credentials
v = video.Video(bvid=video_bvid, credential=self.credential)
# Await the asynchronous get_info() call
info = await v.get_info()
# Some basic information about the video, can be commented out for silent operation
self.logger.info(f"Video Title: {info.get('title', 'N/A')}")
self.logger.info(
f"Current like status of video {video_bvid}: {info.get('like', 'N/A')} (Bilibili API's 'like' field usually reflects your like status)"
)
# Await the asynchronous like/unlike call
await v.like(like_status_for_api) # Pass True for like, False for unlike
self.logger.info(
f"Task for {action_text} video {video_bvid} has completed successfully."
)
except Exception as e:
self.logger.error(f"Error while {action_text} video {video_bvid}: {e}")
async def coin(self, bvid: str, num_coins: int = 1, select_like: bool = True):
"""
Coins a Bilibili video.
Args:
bvid (str): The BV ID of the video to coin.
num_coins (int): Number of coins to give (1 or 2). Defaults to 1.
You usually have a daily limit of 2 coins.
select_like (bool): Whether to also like the video while coining. Defaults to True.
"""
if num_coins not in [1, 2]:
self.logger.error("Error: You can only give 1 or 2 coins per video.")
return
self.logger.info(f"Attempting to coin video: {bvid} with {num_coins} coin(s).")
if select_like:
self.logger.info("Simultaneously liking the video.")
try:
# Instantiate Video object with the BVID and credentials
v = video.Video(bvid=bvid, credential=self.credential)
# Perform the coin action
# The coin method takes num_coins and select_like (True/False to also like)
result = await v.pay_coin(num_coins, select_like)
if result:
self.logger.info(
f"Successfully coined video {bvid} with {num_coins} coin(s)."
)
if select_like:
self.logger.info(f"Video {bvid} also liked successfully.")
else:
self.logger.warning(
f"Failed to coin video {bvid}. Result: {result}"
) # Check API response if False
except Exception as e:
self.logger.error(f"An error occurred while coining video {bvid}: {e}")
async def triple_interation(self, bvid: str):
"""
Performs "One-click triple" (Like, Coin, Favorite) on a Bilibili video.
Args:
bvid (str): The BV ID of the video to interact with.
Returns:
dict: The result of the triple interaction API call, or None if an error occurred.
"""
self.logger.info(
f"--- Attempting operation: 'One-click triple' on video {bvid} ---"
)
try:
# Instantiate Video object
v = video.Video(bvid=bvid, credential=self.credential)
# Perform triple interaction
result = await v.triple()
# Check API return result
if result and result.get("code") == 0:
self.logger.info(
f" Task succeeded: 'One-click triple' on video {bvid} completed successfully."
)
else:
self.logger.warning(
f" Operation failed: 'One-click triple' on video {bvid} failed. API returned: {result}"
)
return result
except Exception as e:
self.logger.error(
f" Operation failed: Error occurred during 'One-click triple' on video {bvid}: {e}"
)
return None
async def get_user_info(self, uid: str):
"""
Retrieves detailed information for a Bilibili user.
Args:
uid (str): The User ID (UID) of the Bilibili user.
Returns:
dict: A dictionary containing the user's information.
"""
self.logger.info(f"--- Attempting operation: Get info for user UID: {uid} ---")
try:
u = user.User(uid=uid, credential=self.credential)
info = await u.get_user_info()
self.logger.info(
f" Task succeeded: Successfully retrieved info for user '{info.get('name', 'N/A')}' (UID: {uid})."
)
self.logger.info(
f" Gender: {info.get('sex', 'N/A')}, Level: LV{info.get('level', 'N/A')}"
)
self.logger.info(
f" Followers: {info.get('follower', 'N/A')}, Following: {info.get('following', 'N/A')}"
)
return info
except Exception as e:
self.logger.error(
f" Operation failed: Error occurred while getting info for user UID {uid}: {e}"
)
return None
async def follow(self, uid: str, mode: bool = True):
"""
Follows or unfollows a Bilibili user.
Args:
uid (str): The User ID (UID) of the Bilibili user.
mode (bool): True to follow (default), False to unfollow.
"""
action_text = "following" if mode else "unfollowing"
relation_type = (
user.RelationType.SUBSCRIBE if mode else user.RelationType.UNSUBSCRIBE
)
self.logger.info(f"--- Attempting operation: {action_text} user UID: {uid} ---")
try:
# Instantiate User object
u = user.User(uid=uid, credential=self.credential)
# Perform follow/unfollow operation
await u.modify_relation(relation_type)
self.logger.info(
f" Task succeeded: Successfully {action_text} user UID: {uid}."
)
except Exception as e:
self.logger.error(
f" Operation failed: Error occurred while {action_text} user UID {uid}: {e}"
)
```
## /api_task/config_template.json
```json path="/api_task/config_template.json"
[
{
"config_name": "email",
"sender_email": "todo",
"sender_password": "todo",
"smtp_server": "mail.sjtu.edu.cn",
"smtp_port": 465
},
{
"config_name": "bilibili",
"SESSDATA": "todo",
"BILI_JCT": "todo",
"BUVID3": "todo",
"dedeuserid": "todo"
}
]
```
## /assets/audio/voice_copy/1.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_copy/1.mp3
## /assets/audio/voice_copy/2.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_copy/2.mp3
## /assets/audio/voice_copy/3.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_copy/3.mp3
## /assets/audio/voice_copy/4.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_copy/4.mp3
## /assets/audio/voice_copy/5.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_copy/5.mp3
## /assets/audio/voice_copy/6.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_copy/6.mp3
## /assets/audio/voice_copy/7.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_copy/7.mp3
## /assets/audio/voice_copy/8.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_copy/8.mp3
## /assets/audio/voice_copy/9.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_copy/9.mp3
## /assets/audio/voice_enterpage/enterpage.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_enterpage/enterpage.mp3
## /assets/audio/voice_entersearch/entersearch.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_entersearch/entersearch.mp3
## /assets/audio/voice_finish/1.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_finish/1.mp3
## /assets/audio/voice_finish/2.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_finish/2.mp3
## /assets/audio/voice_finish/3.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_finish/3.mp3
## /assets/audio/voice_finish/4.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_finish/4.mp3
## /assets/audio/voice_finish/5.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_finish/5.mp3
## /assets/audio/voice_finish/finish6.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_finish/finish6.mp3
## /assets/audio/voice_finish/finish7.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_finish/finish7.mp3
## /assets/audio/voice_finish/finish8.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_finish/finish8.mp3
## /assets/audio/voice_finish/finish9.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_finish/finish9.mp3
## /assets/audio/voice_inputtext/text.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_inputtext/text.mp3
## /assets/audio/voice_openapp/openapp.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_openapp/openapp.mp3
## /assets/audio/voice_point/point.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_point/point.mp3
## /assets/audio/voice_press/press.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_press/press.mp3
## /assets/audio/voice_self/self.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_self/self.mp3
## /assets/audio/voice_swipe/swipe.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_swipe/swipe.mp3
## /assets/audio/voice_temp/final_page.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_temp/final_page.mp3
## /assets/audio/voice_type/type.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_type/type.mp3
## /assets/audio/voice_welcome/1.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_welcome/1.mp3
## /assets/audio/voice_welcome/2.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_welcome/2.mp3
## /assets/audio/voice_welcome/3.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_welcome/3.mp3
## /assets/audio/voice_welcome/4.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_welcome/4.mp3
## /assets/audio/voice_welcome/5.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_welcome/5.mp3
## /assets/audio/voice_welcome/6.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_welcome/6.mp3
## /assets/audio/voice_welcome/7.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_welcome/7.mp3
## /assets/audio/voice_welcome/welcome10.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_welcome/welcome10.mp3
## /assets/audio/voice_welcome/welcome8.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_welcome/welcome8.mp3
## /assets/audio/voice_welcome/welcome9.mp3
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/assets/audio/voice_welcome/welcome9.mp3
## /audio/audio_play.py
```py path="/audio/audio_play.py"
import pygame
import time
import random
import os
from enum import Enum
AUDIO_PATH = "assets/audio"
class VoiceType(Enum):
COPY = "voice_copy"
FINISH = "voice_finish"
POINT = "voice_point"
PRESS = "voice_press"
SWIPE = "voice_swipe"
TYPE = "voice_type"
WELCOME = "voice_welcome"
OPENAPP = "voice_openapp"
ENTERSEARCH = "voice_entersearch"
ENTERPAGE = "voice_enterpage"
SELF = "voice_self"
TEMP = "voice_temp"
INPUTTEXT = "voice_inputtext"
def play_random_audio(voice_type: VoiceType):
"""
Plays a random MP3 audio file from the specified voice type folder.
Args:
voice_type (VoiceType): The type of voice/audio to play. Must be a member of the VoiceType enum.
Behavior:
- Selects a random MP3 file from the corresponding folder under AUDIO_PATH.
- Plays the selected audio file using pygame.
- Waits until playback is finished before quitting the mixer.
- Prints an error message if the folder or MP3 files are not found, or if the voice type is invalid.
"""
if voice_type not in VoiceType:
print(f"无效的音频类型: {voice_type}")
return
voice_folder = voice_type.value
folder_path = os.path.join(AUDIO_PATH, voice_folder)
if not os.path.exists(folder_path):
print(f"音频文件夹不存在: {folder_path}")
return
mp3_files = [f for f in os.listdir(folder_path) if f.lower().endswith(".mp3")]
if not mp3_files:
print("未找到任何 MP3 文件!")
return
selected_file = random.choice(mp3_files)
full_path = os.path.join(folder_path, selected_file)
try:
pygame.mixer.init()
pygame.mixer.music.load(full_path)
pygame.mixer.music.play()
# wait for the audio to be finished
while pygame.mixer.music.get_busy():
time.sleep(0.1)
finally:
# ! important
pygame.mixer.quit()
```
## /audio/tts.py
```py path="/audio/tts.py"
import edge_tts
import asyncio
async def text_to_speech(
text, voice="zh-CN-XiaoxiaoNeural", rate="+0%", output="output.mp3"
):
"""
Convert Chinese text to a speech MP3 file.
Args:
text: The Chinese text to synthesize.
voice: Voice type (default is "Xiaoxiao", female).
rate: Speech rate (e.g. '+20%' means 20% faster; '-20%' means slower).
output: Output MP3 file name.
"""
communicate = edge_tts.Communicate(text=text, voice=voice, rate=rate)
await communicate.save(output)
print(f"Speech has been saved as {output}")
def run_tts(text, output="assets/audio/voice_temp/output.mp3"):
"""
Run text-to-speech conversion.
"""
voice = "zh-CN-XiaoxiaoNeural" # Female voice, recommended
rate = "+20%" # 20% faster than standard speed
asyncio.run(text_to_speech(text, voice=voice, rate=rate, output=output))
return output
```
## /cross_device/function_call_utils.py
```py path="/cross_device/function_call_utils.py"
import json
def generate_function_call(function_name: str, params: dict) -> str:
"""
生成符合格式的函数调用指令(设备1调用设备2时使用)
"""
call = {"action": "function_call", "function": function_name, "parameters": params}
return json.dumps(call, ensure_ascii=False)
def parse_function_call(call_str: str) -> dict:
"""
解析设备1发送的函数调用指令(设备2接收时使用)
"""
try:
return json.loads(call_str)
except json.JSONDecodeError:
return {"error": "invalid format"}
def generate_function_response(result: str, status: str = "success") -> str:
"""
生成函数调用的响应结果(设备2返回结果给设备1时使用)
"""
response = {"action": "function_response", "status": status, "result": result}
return json.dumps(response, ensure_ascii=False)
def parse_function_response(response_str: str) -> dict:
"""
解析设备2返回的函数响应结果(设备1接收时使用)
"""
try:
return json.loads(response_str)
except json.JSONDecodeError:
return {"error": "invalid response"}
```
## /cross_device/instruction_mapper.py
```py path="/cross_device/instruction_mapper.py"
import json
import numpy as np
import os
import sys
from wrappers.cpm_wrapper import MiniCPMWrapper
from wrappers.constants import CLIENT, SUPPORTED_FUNCTIONS, TASK_SPLIT_PROMPT
current_dir = os.path.dirname(os.path.abspath(__file__))
outer_dir = os.path.dirname(current_dir)
sys.path.append(outer_dir)
class InstructionMapper:
"""
负责将用户指令拆分为两台安卓设备的协同操作任务,通过调用llm,生成符合格式要求的设备任务描述、依赖关系及函数调用信息
"""
def __init__(self):
self.model = MiniCPMWrapper(
model_name="AgentCPM-GUI", temperature=0.6, use_history=False
)
def split_task(self, user_instruction: str) -> dict:
prompt = TASK_SPLIT_PROMPT.format(user_instruction=user_instruction)
dummy_image = np.zeros((100, 100, 3), dtype=np.uint8) # 空图像
messages = [
{"role": "system", "content": TASK_SPLIT_PROMPT},
{"role": "user", "content": user_instruction},
]
response = CLIENT.chat.completions.create(
messages=messages,
model="gpt-4o-mini",
temperature=0,
top_p=1.0,
n=1,
).model_dump()
response = response["choices"][0]["message"]
# response = self.model.predict_mm(prompt, [dummy_image])
json_str = self._extract_json(response["content"])
task_info = json.loads(json_str)
cleaned_task_info = {}
for key, value in task_info.items():
cleaned_key = key.strip('\n "')
cleaned_task_info[cleaned_key] = value
required_fields = [
"device1_tasks",
"device2_tasks",
"dependency",
"function_call",
]
for field in required_fields:
if field not in cleaned_task_info:
cleaned_task_info[field] = (
[] if "tasks" in field else "" if field == "dependency" else {}
)
return cleaned_task_info
def _extract_json(self, text: str) -> str:
start = text.find("{")
end = text.rfind("}") + 1
if start == -1 or end == 0:
raise ValueError("模型输出不包含有效JSON")
return text[start:end]
```
## /cross_device/socket_utils.py
```py path="/cross_device/socket_utils.py"
import socket
import threading
import logging
class DualSocket:
"""
双设备Socket通信管理类
该类用于实现两台设备(如安卓手机)之间的双向通信,
支持同时作为服务器监听端口和作为客户端发送消息,通过回调函数处理接收到的消息,
适用于需要设备间实时指令交互和数据传输的场景(如多设备协同完成任务)。
"""
def __init__(self, listen_port: int, peer_port: int, device_name: str = "设备"):
self.listen_port = listen_port # 监听端口
self.peer_port = peer_port # 对方设备端口
self.device_name = device_name # 设备名称
self.is_running = False
self.server_thread = None
self.message_handler = None
def start_server(self, handler):
self.message_handler = handler
self.is_running = True
self.server_thread = threading.Thread(target=self._server_loop, daemon=True)
self.server_thread.start()
def _server_loop(self):
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.bind(("localhost", self.listen_port))
s.listen(1)
while self.is_running:
conn, addr = s.accept()
with conn:
data = conn.recv(1024).decode(encoding="utf-8")
if data and self.message_handler:
response = self.message_handler(data)
conn.sendall(response.encode(encoding="utf-8"))
def send_message(self, message: str) -> str:
try:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
s.connect(("localhost", self.peer_port))
s.sendall(message.encode(encoding="utf-8"))
response = s.recv(1024).decode(encoding="utf-8")
return response
except Exception as e:
return ""
def stop(self):
self.is_running = False
if self.server_thread:
self.server_thread.join(timeout=1)
```
## /cross_device/vision_extractor.py
```py path="/cross_device/vision_extractor.py"
import numpy as np
import os
import sys
import base64
import time
current_dir = os.path.dirname(os.path.abspath(__file__))
outer_dir = os.path.dirname(current_dir)
sys.path.append(outer_dir)
from wrappers.constants import CLIENT
from PIL import Image
from io import BytesIO
class VisionExtractor:
def __init__(self, model_name: str = "gpt-4o"):
self.model = model_name
def query_image(self, image: np.ndarray, prompt: str) -> str:
"""
向gpt-4o发送图像和提示词,获取视觉分析结果
:param image: 输入图像(numpy数组格式)
:param prompt: 提示词(描述需要分析的任务)
:return: 模型返回的文本结果
"""
try:
img = Image.fromarray(image)
if img.mode in ("RGBA", "P", "LA"):
img = img.convert("RGB")
elif img.mode not in ("RGB", "L"):
img = img.convert("RGB")
buffered = BytesIO()
img.save(buffered, format="JPEG", quality=90)
buffered.seek(0)
img_base64 = base64.b64encode(buffered.getvalue()).decode("utf-8")
# 调用vlm
time.sleep(3) # 避免429错误,增加请求间隔
response = CLIENT.chat.completions.create(
model=self.model,
messages=[
{
"role": "user",
"content": [
{"type": "text", "text": prompt},
{
"type": "image_url",
"image_url": {
"url": f"data:image/jpeg;base64,{img_base64}"
},
},
],
}
],
max_tokens=300,
temperature=0.2,
)
return response.choices[0].message.content
except Exception as e:
return f"处理图像时出错:{str(e)}"
```
## /cross_device_agent.py
```py path="/cross_device_agent.py"
import sys
import os
import time
import argparse
sys.path.append(os.getcwd())
from run_agent import GUITaskExecutor
from cross_device.socket_utils import DualSocket
from cross_device.instruction_mapper import InstructionMapper
from cross_device.function_call_utils import (
generate_function_call,
parse_function_call,
generate_function_response,
parse_function_response,
)
from adb_utils import setup_device, list_connected_devices, change_ui_settings
from PIL import Image
from cross_device.vision_extractor import VisionExtractor
from user.ocr_service import OCRService
class CrossDeviceCoordinator:
"""跨设备协同控制器
用于协调两台安卓设备通过Socket通信完成用户指令,
实现设备间的任务分配、函数调用和结果交互。
"""
def __init__(
self,
device1_serial: str,
device1_port: int,
device2_serial: str,
device2_port: int,
):
self.device1 = setup_device(device1_serial)
self.device2 = setup_device(device2_serial)
self.socket1 = DualSocket(
listen_port=device1_port,
peer_port=device2_port,
device_name=f"设备1({self.device1.serial})",
)
self.socket2 = DualSocket(
listen_port=device2_port,
peer_port=device1_port,
device_name=f"设备2({self.device2.serial})",
)
self.mapper = InstructionMapper()
self.ocr_service = OCRService()
self.common_run_params = {
"ocr_service": self.ocr_service,
"enable_audio": False,
"enable_vision_parser": False,
"return_result": False,
}
self.enable_experience = False
def _device2_function_handler(self, call_str: str) -> str:
"""设备2的函数调用处理函数
接收设备1的调用指令,执行对应操作(如提取关键词),返回处理结果
"""
call = parse_function_call(call_str)
if "error" in call:
return generate_function_response("无效调用格式", "error")
if call["function"] == "extract_keyword":
task = call["parameters"]["task"]
task_executor = GUITaskExecutor(serial=self.device2.serial, **self.common_run_params)
task_executor.run_task(query=task, enable_experience=self.enable_experience)
screenshot = self.device2.screenshot()
keyword = self._extract_info_from_screenshot(
screenshot, call["parameters"]["extract_prompt"]
)
if keyword:
return generate_function_response(keyword)
return generate_function_response("提取关键词失败", "error")
def _extract_info_from_screenshot(
self, screenshot: Image.Image, prompt: str
) -> str:
"""从截图中提取信息
使用视觉语言模型,根据提示词从截图中提取所需内容
"""
from cross_device.vision_extractor import VisionExtractor
extractor = VisionExtractor(model_name="gpt-4o")
import numpy as np
img_array = np.array(screenshot)
result = extractor.query_image(img_array, prompt)
return self._postprocess_extraction(result)
def _postprocess_extraction(self, text: str) -> str:
"""提取结果后处理
对模型返回的提取结果进行简单清洗(如正则匹配提取关键部分)
"""
import re
match = re.search(r"标题[::]\s*(.+)", text)
return match.group(1).strip() if match else text.strip()
def _agent1_call_agent2(self, function_name: str, params: dict) -> str:
"""设备1调用设备2的函数
生成函数调用指令,通过Socket发送给设备2,接收并返回处理结果
"""
call_str = generate_function_call(function_name, params)
self.socket2.start_server(handler=self._device2_function_handler)
response_str = self.socket1.send_message(call_str)
response = parse_function_response(response_str)
if response["status"] == "success":
return response["result"]
else:
return ""
def start_workflow(self, user_instruction: str) -> None:
"""启动跨设备工作流程
解析用户指令,分配任务给设备1和设备2,协调执行并处理结果
"""
try:
task_info = self.mapper.split_task(user_instruction)
device1_task = task_info["device1_task"]
function_call = task_info["function_call"]
keyword = self._agent1_call_agent2(
function_name=function_call["name"], params=function_call["parameters"]
)
if not keyword:
print("获取关键词失败,终止流程")
return
print(f"从设备2获取到关键词: {keyword}")
self._run_device1(device1_task, keyword)
except Exception as e:
print(f"流程执行出错:{e}")
finally:
self.socket1.stop()
self.socket2.stop()
def _run_device1(self, task: str, keyword: str) -> None:
"""执行设备1的任务
将任务中的"等待关键词"替换为实际关键词,然后执行任务
"""
new_task = task.replace("等待设备2发送的关键词", f"输入关键词:{keyword}")
print(f"在设备1上执行: {new_task}")
task_executor = GUITaskExecutor(serial=self.device1.serial, **self.common_run_params)
task_executor.run_task(query=new_task, enable_experience=self.enable_experience)
time.sleep(2)
def main():
"""命令行入口函数
解析命令行参数,初始化设备协调器,启动工作流程
"""
# 解析命令行参数
parser = argparse.ArgumentParser(description="跨设备协同控制工具")
parser.add_argument(
"--device1-serial", type=str, default=None, help="设备1的ADB序列号(可选)"
)
parser.add_argument(
"--device1-port", type=int, default=11001, help="设备1的通信端口(默认11001)"
)
parser.add_argument(
"--device2-serial", type=str, default=None, help="设备2的ADB序列号(可选)"
)
parser.add_argument(
"--device2-port", type=int, default=11002, help="设备2的通信端口(默认11002)"
)
parser.add_argument("--task", type=str, help="跨设备任务指令")
parser.add_argument(
"--list-devices", action="store_true", help="列出所有已连接的ADB设备"
)
args = parser.parse_args()
if args.list_devices:
devices = list_connected_devices()
if not devices:
print("没有找到已连接的ADB设备")
else:
print("已连接的ADB设备:")
for i, device in enumerate(devices):
print(f"{i+1}. {device}")
return
if not args.task:
parser.error("请提供 --task 参数指定跨设备任务指令")
devices = list_connected_devices()
if not devices:
raise RuntimeError("没有找到已连接的 ADB 设备")
if args.device1_serial is None:
if len(devices) >= 1:
args.device1_serial = devices[0]
else:
raise RuntimeError("至少需要1台设备")
if args.device2_serial is None:
if len(devices) >= 2:
args.device2_serial = devices[1]
else:
raise RuntimeError("至少需要2台设备")
try:
# 打开UI设置(如无障碍模式等)
change_ui_settings(mode="open")
coordinator = CrossDeviceCoordinator(
device1_serial=args.device1_serial,
device1_port=args.device1_port,
device2_serial=args.device2_serial,
device2_port=args.device2_port,
)
coordinator.start_workflow(args.task)
finally:
change_ui_settings(mode="close")
if __name__ == "__main__":
main()
```
## /eval/eval_multi.sh
```sh path="/eval/eval_multi.sh"
#!/bin/bash
if [ ! -f "run_predict_appcopilot_multi.py" ]; then
echo "Error: run_predict_appcopilot_multi.py script not found"
exit 1
fi
if [ ! -f "run_eval_agent.py" ]; then
echo "Error: run_eval_agent.py script not found"
exit 1
fi
# Configure basic parameters
data_name="DATASET NAME"
model_name="AppCopilot"
base_output_dir="./eval_results/AppCopilot/${data_name}/${model_name}"
models_base_path=(
"MODEL DIR"
)
start_time=$(date +%s)
echo "===== Start model inference and evaluation ====="
echo "Start time: $(date)"
# Iterate over models to run inference and evaluation
for model_path in "${models_base_path[@]}"; do
model_short_name=$(basename "${model_path}")
echo "===== Processing model: ${model_short_name} (Original Path: ${model_path}) ====="
model_root_dir="${base_output_dir}/${model_short_name}"
mkdir -p "${model_root_dir}"
inference_output_dir="${model_root_dir}/inference"
mkdir -p "${inference_output_dir}"
echo "[Step 1/2] Running inference for model ${model_short_name}..."
python run_predict_appcopilot_multi.py \
--model_path "${model_path}" \
--output_dir "${inference_output_dir}" \
--data_name "${data_name}"
inference_exit_code=$?
if [ $inference_exit_code -ne 0 ]; then
echo "Error: Inference for model ${model_short_name} failed, exit code: ${inference_exit_code}"
continue
fi
all_jsonl_path="${inference_output_dir}/all.jsonl"
if [ ! -f "${all_jsonl_path}" ]; then
echo "Error: Inference result file ${all_jsonl_path} not found"
continue
fi
eval_output_dir="${model_root_dir}/results"
mkdir -p "${eval_output_dir}"
echo "[Step 2/2] Running evaluation for model ${model_short_name}..."
python run_eval_agent.py \
--input_path "${all_jsonl_path}" \
--output_dir "${eval_output_dir}" \
--data_name "${data_name}"
eval_exit_code=$?
if [ $eval_exit_code -ne 0 ]; then
echo "Error: Evaluation for model ${model_short_name} failed, exit code: ${eval_exit_code}"
else
echo "Inference and evaluation for model ${model_short_name} completed!"
echo "Result path for model ${model_short_name}: ${model_root_dir}"
fi
done
# Calculate and print total time and end info
end_time=$(date +%s)
total_time=$((end_time - start_time))
echo "===== Model inference and evaluation completed ====="
echo "End time: $(date)"
echo "Total time: ${total_time} seconds (approximately $(echo "scale=2; $total_time/60" | bc) minutes)"
echo "Root directory for all model results: ${base_output_dir}"
```
## /eval/run_eval_agent.py
```py path="/eval/run_eval_agent.py"
import os
from shutil import ExecError
import json
import random
from collections import defaultdict
from tqdm import tqdm
from utils.convert_output import convert2aitz
import argparse
from utils.evaluator import ActionEvaluator
from utils.utils import get_dataset_dir
import logging
# set logging
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s %(levelname)s %(processName)s %(message)s',
handlers=[
logging.FileHandler("eval_gui_agent.log"),
logging.StreamHandler()
]
)
class EvalDataset(object):
# subset of dataset to eval
DATASET_DIR = {
'general': '{}/general',
'google_apps': '{}/google_apps',
'install': '{}/install',
'single': '{}/single',
'web_shopping': '{}/web_shopping',
'domestic': '{}/domestic',
'odyssey': '{}/odyssey',
'android_control': '{}/android_control',
}
def __init__(self, data_dir, split="test", ratio=1.0) -> None:
self.ratio = ratio
self.data_dir = os.path.join(data_dir, split)
self.episode_data = self._load_data_()
self.data = self._split_to_steps_(self.episode_data)
def _load_data_(self):
valid_paths = defaultdict(list)
for subset in self.DATASET_DIR:
subdata_dir = self.DATASET_DIR[subset].format(self.data_dir)
if os.path.exists(subdata_dir):
sequence_names = os.listdir(subdata_dir)
for seq_name in sequence_names:
seq_dir = os.path.join(subdata_dir, seq_name)
if not os.path.isdir(seq_dir): continue
episode_path = os.path.join(seq_dir, f"{seq_name}.json")
valid_paths[subset].append(episode_path)
sampled_paths = []
for subset, v_paths in valid_paths.items():
N = len(v_paths)
k = int(self.ratio * N)
sampled_paths += random.sample(v_paths, k) if self.ratio < 1.0 else v_paths
ep_data = []
for episode_path in sampled_paths:
try:
with open(episode_path, "r") as f:
episode_data = json.load(f)
ep_data.append(episode_data)
except json.JSONDecodeError as e:
logging.error(f"JSON decoding failed, file: {episode_path}, error: {e}")
except Exception as e:
logging.error(f"Error occurred when loading file {episode_path}: {e}")
return ep_data
def _split_to_steps_(self, episode_data):
data = []
for edx, episode in enumerate(episode_data):
for idx, step in enumerate(episode):
try:
if step.get('subset') is None:
step['subset'] = step['image_path'].split('/')[0]
step['image_full_path'] = os.path.join(self.data_dir, step['image_path'])
data.append(step)
except KeyError as e:
logging.error(f"Missing key {e}, at episode {edx}, step {idx}")
except Exception as e:
logging.error(f"Error processing episode {edx}, step {idx}: {e}")
return data
def __len__(self):
return len(self.data)
def __getitem__(self, index):
return self.data[index]
def process_step_data(step_data, evaluator, save_dir):
"""
Process a single step of data, load prediction results and evaluate.
Args:
step_data (dict): Data of a single step.
Returns:
dict or None: None means the step hasn't been predicted yet. Empty JSON indicates JSON parsing failed.
"""
subset = step_data.get('subset')
episode_id = step_data.get('episode_id')
step_id = step_data.get('step_id')
if subset is None or episode_id is None or step_id is None:
raise ValueError(f"Missing subset/episode_id/step_id in test step data: {step_data}")
save_dir_ep = os.path.join(save_dir, f"{subset}-{episode_id}")
cur_save_path = os.path.join(save_dir_ep, f"{subset}-{episode_id}_{step_id}.json")
try:
# Ensure directory exists
os.makedirs(os.path.dirname(cur_save_path), exist_ok=True)
# Check if the file exists; if not, the step hasn't been predicted yet
if not os.path.exists(cur_save_path):
return None
# Load the existing file
with open(cur_save_path, "r") as file:
try:
pred = json.load(file)
except json.JSONDecodeError as e:
logging.error(f"JSON decoding failed, file: {cur_save_path}, error: {e}")
pred = {'action_predict': {'COA': {'txt': {'ACTION': None, 'ARGS': None, 'STATUS': None}}}}
assert pred is not None
# Use global evaluator to evaluate
result = evaluator(step_data, pred)
return result
except Exception as e:
raise ExecError(f"An error occurred, indicating unhandled edge case.")
def evaluate(args):
# Get dataset path
args.data_dir, args.split, _ = get_dataset_dir(args.data_name)
# Convert to aitz format
convert2aitz(os.path.abspath(args.input_path), os.path.abspath(args.output_dir), max_workers=16)
save_dir = os.path.abspath(args.output_dir)
results_save_file = os.path.join(save_dir, "result.json")
# Initialize dataset
eval_data = EvalDataset(data_dir=args.data_dir)
logging.info(f"Total steps: {len(eval_data)}, total episodes: {len(eval_data.episode_data)}.")
evaluator = ActionEvaluator(save_dir, args.eval_android_control)
results = list(tqdm(map(process_step_data, eval_data.data, [evaluator]*len(eval_data.data), [save_dir]*len(eval_data.data)),total=len(eval_data.data), desc="Processing steps", ncols=100))
results = list(filter(lambda x: x is not None, results))
# Save results
try:
os.makedirs(os.path.dirname(results_save_file), exist_ok=True)
with open(results_save_file, "w") as f:
json.dump(results, f, indent=4, ensure_ascii=False)
logging.info(f"Evaluation results saved to {results_save_file}")
except Exception as e:
logging.error(f"Error saving results to {results_save_file}: {e}")
# Aggregate episode results
episode_results = defaultdict(list)
for result in results:
subset = result.get("subset")
episode_id = result.get("episode_id")
if subset is None or episode_id is None:
logging.warning(f"Result missing subset/episode_id: {result}")
continue
episode_key = f"{subset}-{episode_id}"
episode_results[episode_key].append(result)
# Compute final evaluation metrics
try:
episode_metrics = ActionEvaluator.compute_episode_metrics(episode_results)
atomic_metrics = ActionEvaluator.compute_atomic_metrics(results)
logging.basicConfig(level=logging.INFO, format='%(message)s')
logging.info(f"episode_metrics: {episode_metrics}")
logging.info(f"atomic_metrics: {atomic_metrics}")
logging.info(
f"success_rate: {episode_metrics.get('success_rate')}, "
f"goal_progress: {episode_metrics.get('goal_progress')}, "
f"type_acc: {atomic_metrics.get('total', {}).get('type_acc')}, "
f"exact_acc: {atomic_metrics.get('total', {}).get('exact_acc')}"
)
summary_save_file = os.path.join(args.output_dir, 'summary.json')
logging.info(f"Evaluation summary saved to {summary_save_file}")
with open(summary_save_file, 'w') as f:
json.dump(episode_metrics|atomic_metrics, f, ensure_ascii=False)
except Exception as e:
logging.error(f"Error computing evaluation metrics: {e}")
# ==========================================================================================
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="GUI Agent Eval")
parser.add_argument("--seed", type=int, default=2020, help="Random seed")
parser.add_argument("--input_path", type=str, required=True, help="Path to input prediction JSONL file")
parser.add_argument("--output_dir", type=str, required=True, help="Directory to save results")
parser.add_argument("--data_name", type=str, required=True, choices=['gui_odyssey_test', 'chinese_app_test', 'aitz_test', 'android_control_high_test', 'android_control_low_test'], help="Eval dataset name")
parser.add_argument("--eval_android_control", action="store_true", help="For evaluating android control, which is different from other datasets according to qwen's scripts")
args = parser.parse_args()
logging.info(f"Received arguments: {args}")
random.seed(args.seed)
evaluate(args)
```
## /eval/run_predict_appcopilot_multi.py
```py path="/eval/run_predict_appcopilot_multi.py"
import sys
import multiprocessing
import os
os.chdir(os.path.dirname(os.path.abspath(__file__)))
import json
import torch
import random
import jsonschema
import requests
from tqdm import tqdm
from transformers import AutoTokenizer,AutoModelForCausalLM
from concurrent.futures import ProcessPoolExecutor,as_completed,ThreadPoolExecutor
from PIL import Image
from utils.utils import get_dataset_dir
from utils.mv_vote import action_majority_vote
import argparse
import logging
import time
DEVICES = [
"cuda:0", "cuda:1", "cuda:2",
"cuda:3", "cuda:5", "cuda:6",
"cuda:7"
]
current_file_path = os.path.abspath(__file__)
current_dir = os.path.dirname(current_file_path)
if current_dir not in sys.path:
sys.path.append(current_dir)
def compact_json_dumps(obj):
return json.dumps(obj, indent=None, separators=(",", ":"), ensure_ascii=False)
ACTION_SCHEMA = json.load(open(os.path.join(current_dir, 'utils/schema', 'schema.json'), encoding="utf-8"))
items = list(ACTION_SCHEMA.items())
insert_index = 3
items.insert(insert_index, ("required", ["thought"])) # enable/disable thought by setting it to "required"/"optional"
ACTION_SCHEMA = dict(items)
SYSTEM_PROMPT = f'''# Role
你是一名熟悉安卓系统触屏GUI操作的智能体,将根据用户的问题,分析当前界面的GUI元素和布局,生成相应的操作。
# Task
针对用户问题,根据输入的当前屏幕截图,输出下一步的操作。
# Rule
- 以紧凑JSON格式输出
- 输出操作必须遵循Schema约束
# Schema
{json.dumps(ACTION_SCHEMA, indent=None, ensure_ascii=False, separators=(',', ':'))}'''
EXTRACT_SCHEMA = json.load(open(os.path.join(current_dir, 'utils/schema', 'schema_for_extraction.json'), encoding="utf-8"))
_llm = None
_tokenizer = None
def _init_llm(model_name):
global _llm,_tokenizer
if _llm is None:
_llm = AutoModelForCausalLM.from_pretrained(model_name, trust_remote_code=True,torch_dtype=torch.bfloat16)
if _tokenizer is None:
_tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
def move_to(device):
global _llm,_tokenizer
if _llm is None:
raise ValueError("Error, LLM is not initialized.")
_llm = _llm.to(device)
if _tokenizer is None:
raise ValueError("Error, Tokenizer is not initialized.")
return f"Moved to {device}"
def run_single_agent(msg):
outputs = _llm.chat(
image=None, msgs=msg, system_prompt=SYSTEM_PROMPT,
tokenizer=_tokenizer, temperature=0.1, top_p=0.3, n=1,
)
return extract_and_validate_json(outputs)
def run_episode(episode, msg,):
global _llm,_tokenizer
with ThreadPoolExecutor(max_workers=3) as executor:
futures = [executor.submit(run_single_agent, msg) for _ in range(3)]
results = []
for future in as_completed(futures):
try:
res = future.result(timeout=15)
if res is not None:
results.append(res)
except Exception as e:
print(f"[Error] Single agent failed: {e}")
if results:
episode["pred"] = action_majority_vote(results)
else:
episode["pred"] = {"FALLBACK": True, "reason": "All agents failed"}
print("Aggegated result:")
print(episode["pred"] )
return episode
def extract_and_validate_json(input_string):
try:
json_obj = json.loads(input_string)
jsonschema.validate(json_obj, EXTRACT_SCHEMA)
return json_obj
except json.JSONDecodeError as e:
print("Error, JSON is NOT valid.")
return input_string
except Exception as e:
print(f"Error, JSON is NOT valid according to the schema.{input_string}", e)
return input_string
def load_image(episode, image_path, data_name):
# resize the image proportionally so that the longer side is at most 1120
def __resize__(origin_img):
resolution = origin_img.size
w,h = resolution
max_line_res = 1120
if max_line_res is not None:
max_line = max_line_res
if h > max_line:
w = int(w * max_line / h)
h = max_line
if w > max_line:
h = int(h * max_line / w)
w = max_line
img = origin_img.resize((w,h),resample=Image.Resampling.LANCZOS)
return img
image = Image.open(image_path).convert("RGB")
image = __resize__(image)
if data_name == 'android_control_low_test':
query = episode['low_instruction']
else:
query = episode['instruction']
messages = []
messages.append(
{
"role": "user",
"content": [
f"<Question>{query}</Question>\n当前屏幕截图:",
image
]
}
)
return (episode,messages)
def predict(args):
args.data_dir, args.split, data_subset = get_dataset_dir(args.data_name)
print(f"Predicting on: {args.data_dir}/{args.split}")
print(f"Data subset: {data_subset}")
if multiprocessing.get_start_method(allow_none=True) != "spawn":
multiprocessing.set_start_method("spawn", force=True)
with ProcessPoolExecutor(max_workers=len(DEVICES),initializer=_init_llm,initargs=(args.model_path,)) as poolexec:
tasks = []
print("Moving model to devices")
futures = [poolexec.submit(move_to, dev) for dev in DEVICES]
for fut in futures:
print(fut.result())
for dataset in data_subset:
save_dir = os.path.join(args.output_dir, dataset)
if not os.path.exists(save_dir):
os.makedirs(save_dir)
episode_dir = os.path.join(args.data_dir, args.split, dataset)
output_file = os.path.join(save_dir, "predict.jsonl")
# Get the list of all episodes files
if os.path.exists(episode_dir):
episodes_files = os.listdir(episode_dir)
else:
continue
future = []
all_tasks = []
print("Loading episodes")
with ThreadPoolExecutor(max_workers=16) as executor:
for episodes_file in episodes_files:
episodes_path = os.path.join(episode_dir, episodes_file, f"{episodes_file}.json")
try:
with open(episodes_path, 'r', encoding='utf-8') as f:
episodes = json.load(f)
except Exception as e:
print(f"Failed to load {episodes_path}: {e}")
continue
# Skip this file on error
for episode in episodes:
episode["category"] = dataset
image_path = os.path.join(episode_dir, episodes_file, f"{episodes_file}_{episode['step_id']}.jpeg")
if not os.path.exists(image_path):
image_path = image_path.replace(".jpeg", ".png")
if not os.path.exists(image_path):
image_path = episode['image_path']
future.append(executor.submit(load_image, episode, image_path, args.data_name))
for f in as_completed(future):
all_tasks.append(f.result())
with open(output_file, "w", encoding="utf-8") as f_out:
print("Predicting")
tasks = []
for task_value in all_tasks:
tasks.append(poolexec.submit(run_episode, *task_value))
for task in tqdm(as_completed(tasks), total=len(tasks), dynamic_ncols=True):
try:
episode = task.result()
episode_json = json.dumps(episode, ensure_ascii=False)
f_out.write(episode_json + "\n")
f_out.flush()
except Exception as e:
print(f"Error: {e}")
continue
print(f"Prediction saved at: {output_file}.")
os.system(f"cat {args.output_dir}/*/predict.jsonl > {args.output_dir}/all.jsonl")
print(f"Merged prediction saved at: {args.output_dir}/all.jsonl.")
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="GUI Agent Inference")
parser.add_argument("--seed", type=int, default=2020, help="Random seed")
parser.add_argument("--model_path", type=str, required=True, help="Model path")
parser.add_argument("--output_dir", type=str, required=True, help="Directory to save results")
parser.add_argument("--data_name", type=str, required=True, choices=['gui_odyssey_test', 'chinese_app_test', 'aitz_test', 'android_control_high_test', 'android_control_low_test'], help="Eval dataset name")
args = parser.parse_args()
random.seed(args.seed)
print(f'Loading model at : {args.model_path}')
print(f'Saving results at: {args.output_dir}')
predict(args)
```
## /eval/run_predict_minicpm.py
```py path="/eval/run_predict_minicpm.py"
import sys
import multiprocessing
import os
os.chdir(os.path.dirname(os.path.abspath(__file__)))
import json
import torch
import random
import jsonschema
from tqdm import tqdm
from transformers import AutoTokenizer,AutoModelForCausalLM
from concurrent.futures import ProcessPoolExecutor,as_completed,ThreadPoolExecutor
from PIL import Image
from utils.utils import get_dataset_dir
import argparse
import logging
import time
DEVICES = [
"cuda:0", "cuda:1", "cuda:2", "cuda:3",
"cuda:4","cuda:5", "cuda:6", "cuda:7",
]
current_file_path = os.path.abspath(__file__)
current_dir = os.path.dirname(current_file_path)
if current_dir not in sys.path:
sys.path.append(current_dir)
def compact_json_dumps(obj):
return json.dumps(obj, indent=None, separators=(",", ":"), ensure_ascii=False)
ACTION_SCHEMA = json.load(open(os.path.join(current_dir, 'utils/schema', 'schema.json'), encoding="utf-8"))
items = list(ACTION_SCHEMA.items())
insert_index = 3
items.insert(insert_index, ("required", ["thought"])) # enable/disable thought by setting it to "required"/"optional"
ACTION_SCHEMA = dict(items)
SYSTEM_PROMPT = f'''# Role
你是一名熟悉安卓系统触屏GUI操作的智能体,将根据用户的问题,分析当前界面的GUI元素和布局,生成相应的操作。
# Task
针对用户问题,根据输入的当前屏幕截图,输出下一步的操作。
# Rule
- 以紧凑JSON格式输出
- 输出操作必须遵循Schema约束
# Schema
{json.dumps(ACTION_SCHEMA, indent=None, ensure_ascii=False, separators=(',', ':'))}'''
EXTRACT_SCHEMA = json.load(open(os.path.join(current_dir, 'utils/schema', 'schema_for_extraction.json'), encoding="utf-8"))
_llm = None
_tokenizer = None
def _init_llm(model_name):
global _llm,_tokenizer
if _llm is None:
_llm = AutoModelForCausalLM.from_pretrained(model_name, trust_remote_code=True,torch_dtype=torch.bfloat16)
if _tokenizer is None:
_tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
def move_to(device):
global _llm,_tokenizer
if _llm is None:
raise ValueError("Error, LLM is not initialized.")
_llm = _llm.to(device)
if _tokenizer is None:
raise ValueError("Error, Tokenizer is not initialized.")
return f"Moved to {device}"
def run_episode(episode, msg,):
global _llm,_tokenizer
outputs = _llm.chat(image=None, msgs=msg, system_prompt=SYSTEM_PROMPT, tokenizer=_tokenizer, temperature=0.1,top_p=0.3,n=1,)
episode["pred"] = extract_and_validate_json(outputs)
return episode
def extract_and_validate_json(input_string):
try:
json_obj = json.loads(input_string)
jsonschema.validate(json_obj, EXTRACT_SCHEMA)
return json_obj
except json.JSONDecodeError as e:
print("Error, JSON is NOT valid.")
return input_string
except Exception as e:
print(f"Error, JSON is NOT valid according to the schema.{input_string}", e)
return input_string
def load_image(episode, image_path, data_name):
# resize the image proportionally so that the longer side is at most 1120
def __resize__(origin_img):
resolution = origin_img.size
w,h = resolution
max_line_res = 1120
if max_line_res is not None:
max_line = max_line_res
if h > max_line:
w = int(w * max_line / h)
h = max_line
if w > max_line:
h = int(h * max_line / w)
w = max_line
img = origin_img.resize((w,h),resample=Image.Resampling.LANCZOS)
return img
image = Image.open(image_path).convert("RGB")
image = __resize__(image)
if data_name == 'android_control_low_test':
query = episode['low_instruction']
else:
query = episode['instruction']
messages = []
messages.append(
{
"role": "user",
"content": [
f"<Question>{query}</Question>\n当前屏幕截图:",
image
]
}
)
return (episode,messages)
def predict(args):
args.data_dir, args.split, data_subset = get_dataset_dir(args.data_name)
print(f"Predicting on: {args.data_dir}/{args.split}")
print(f"Data subset: {data_subset}")
if multiprocessing.get_start_method(allow_none=True) != "spawn":
multiprocessing.set_start_method("spawn", force=True)
with ProcessPoolExecutor(max_workers=len(DEVICES),initializer=_init_llm,initargs=(args.model_path,)) as poolexec:
tasks = []
print("Moving model to devices")
futures = [poolexec.submit(move_to, dev) for dev in DEVICES]
for fut in futures:
print(fut.result())
for dataset in data_subset:
save_dir = os.path.join(args.output_dir, dataset)
if not os.path.exists(save_dir):
os.makedirs(save_dir)
episode_dir = os.path.join(args.data_dir, args.split, dataset)
output_file = os.path.join(save_dir, "predict.jsonl")
# Get the list of all episodes files
if os.path.exists(episode_dir):
episodes_files = os.listdir(episode_dir)
else:
continue
future = []
all_tasks = []
print("Loading episodes")
with ThreadPoolExecutor(max_workers=16) as executor:
for episodes_file in episodes_files:
episodes_path = os.path.join(episode_dir, episodes_file, f"{episodes_file}.json")
try:
with open(episodes_path, 'r', encoding='utf-8') as f:
episodes = json.load(f)
except Exception as e:
print(f"Failed to load {episodes_path}: {e}")
continue
# Skip this file on error
for episode in episodes:
episode["category"] = dataset
image_path = os.path.join(episode_dir, episodes_file, f"{episodes_file}_{episode['step_id']}.jpeg")
if not os.path.exists(image_path):
image_path = image_path.replace(".jpeg", ".png")
if not os.path.exists(image_path):
image_path = episode['image_path']
future.append(executor.submit(load_image, episode, image_path, args.data_name))
for f in as_completed(future):
all_tasks.append(f.result())
with open(output_file, "w", encoding="utf-8") as f_out:
print("Predicting")
tasks = []
for task_value in all_tasks:
tasks.append(poolexec.submit(run_episode, *task_value))
for task in tqdm(as_completed(tasks), total=len(tasks), dynamic_ncols=True):
try:
episode = task.result()
episode_json = json.dumps(episode, ensure_ascii=False)
f_out.write(episode_json + "\n")
f_out.flush()
except Exception as e:
print(f"Error: {e}")
continue
print(f"Prediction saved at: {output_file}.")
os.system(f"cat {args.output_dir}/*/predict.jsonl > {args.output_dir}/all.jsonl")
print(f"Merged prediction saved at: {args.output_dir}/all.jsonl.")
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="GUI Agent Inference")
parser.add_argument("--seed", type=int, default=2020, help="Random seed")
parser.add_argument("--model_path", type=str, required=True, help="Model path")
parser.add_argument("--output_dir", type=str, required=True, help="Directory to save results")
parser.add_argument("--data_name", type=str, required=True, choices=['gui_odyssey_test', 'chinese_app_test', 'aitz_test', 'android_control_high_test', 'android_control_low_test'], help="Eval dataset name")
args = parser.parse_args()
random.seed(args.seed)
print(f'Loading model at : {args.model_path}')
print(f'Saving results at: {args.output_dir}')
predict(args)
```
## /eval/utils/SimHei.ttf
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/eval/utils/SimHei.ttf
## /eval/utils/action_type.py
```py path="/eval/utils/action_type.py"
# coding=utf-8
# Copyright 2024 The Google Research Authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""AndroidInTheWild action types."""
# https://github.com/google-research/google-research/blob/master/android_in_the_wild/action_type.py
import enum
class ActionType(enum.IntEnum):
"""Integer values for each supported action type in AndroidInTheWild."""
# Placeholders for unused enum values
# UNUSED_0 = 0 # used for long point
# UNUSED_1 = 1 # used for no action
UNUSED_2 = 2
UNUSED_8 = 8
UNUSED_9 = 9
########### Agent actions ###########
LONG_POINT = 0 # long ponint
NO_ACTION = 1 # no action
# A type action that sends text to the emulator. Note that this simply sends
# text and does not perform any clicks for element focus or enter presses for
# submitting text.
TYPE = 3
# The dual point action used to represent all gestures.
DUAL_POINT = 4
# These actions differentiate pressing the home and back button from touches.
# They represent explicit presses of back and home performed using ADB.
PRESS_BACK = 5
PRESS_HOME = 6
# An action representing that ADB command for hitting enter was performed.
PRESS_ENTER = 7
########### Episode status actions ###########
# An action used to indicate the desired task has been completed and resets
# the environment. This action should also be used in the case that the task
# has already been completed and there is nothing to do.
# e.g. The task is to turn on the Wi-Fi when it is already on
STATUS_TASK_COMPLETE = 10
# An action used to indicate that desired task is impossible to complete and
# resets the environment. This can be a result of many different things
# including UI changes, Android version differences, etc.
STATUS_TASK_IMPOSSIBLE = 11
```
## /eval/utils/action_utils.py
```py path="/eval/utils/action_utils.py"
# coding=utf-8
# Copyright 2023 The Google Research Authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import json
import jax.numpy as jnp
import numpy as np
from .action_type import ActionType
def extract_gt_action(example):
ex_action_type = example['result_action_type']
if ex_action_type == ActionType.DUAL_POINT:
lift_yx = json.loads(example['result_lift_yx'])
touch_yx = json.loads(example['result_touch_yx'])
if is_tap_action(np.array(touch_yx), np.array(lift_yx)):
action_type = 'CLICK'
w, h = example['image_width'], example['image_height']
assert w and h, "Invalid image size"
click_y, click_x = round(lift_yx[0] * h), round(lift_yx[1] * w) # Note: image size range, not 0-1000
action = (click_y, click_x)
else:
action_type = 'SCROLL'
v_change = abs(touch_yx[0] - lift_yx[0])
h_change = abs(lift_yx[1] - touch_yx[1])
is_scroll_up = lift_yx[0] < touch_yx[0] # touch is lower
is_scroll_left = lift_yx[1] < touch_yx[1] # touch is bigger
if v_change >= 0.9*h_change: # vertical
action = "scroll up" if is_scroll_up else "scroll down"
else: # horizonal
action = "scroll left" if is_scroll_left else "scroll right"
elif ex_action_type in (
ActionType.PRESS_BACK,
ActionType.PRESS_HOME,
):
button = ActionType(ex_action_type).name.split('_')[1].lower()
action = f'press the {button} button'
action_type = f'PRESS {button}'.upper()
elif ex_action_type == ActionType.PRESS_ENTER:
action = "press enter"
action_type = 'PRESS ENTER'
elif ex_action_type == ActionType.TYPE:
action_text = example['result_action_text']
action = f'input text "{action_text}"'
action_type = 'INPUT'
elif ex_action_type == ActionType.STATUS_TASK_COMPLETE:
action = 'stop and set the query as completed'
action_type = 'STOP'
elif ex_action_type == ActionType.STATUS_TASK_IMPOSSIBLE:
action = 'stop and set the query as impossible'
action_type = 'STOP'
elif ex_action_type == ActionType.LONG_POINT:
lift_yx = json.loads(example['result_lift_yx'])
w, h = example['image_width'], example['image_height']
assert w and h, "Invalid image size"
click_y, click_x = round(lift_yx[0] * h), round(lift_yx[1] * w) # Note: image size range, not 0-1000
action = (click_y, click_x)
action_type = "LONG_POINT"
elif ex_action_type == ActionType.NO_ACTION:
duration_time = example['duration']
action = f'no action for {duration_time} ms'
action_type = 'NO_ACTION'
else:
raise NotImplementedError
return action, action_type
'''======================================
Global Args
======================================'''
_TAP_DISTANCE_THRESHOLD = 0.14 # Fraction of the screen
ANNOTATION_WIDTH_AUGMENT_FRACTION = 1.4
ANNOTATION_HEIGHT_AUGMENT_FRACTION = 1.4
# Interval determining if an action is a tap or a swipe.
_SWIPE_DISTANCE_THRESHOLD = 0.04
def is_tap_action(normalized_start_yx, normalized_end_yx):
distance = jnp.linalg.norm(
jnp.array(normalized_start_yx) - jnp.array(normalized_end_yx))
return distance <= _SWIPE_DISTANCE_THRESHOLD
def _is_non_dual_point_action(action_type):
return jnp.not_equal(action_type, ActionType.DUAL_POINT)
'''===============================================================================
=== Utilites for performing action matching on AndroidInTheWild data. ===
=== Note: this code is implemented using JAX so it can be 'vmap'ed and ===
=== efficiently appied over a batch of data. ===
==============================================================================='''
def _yx_in_bounding_boxes(
yx, bounding_boxes
):
"""Check if the (y,x) point is contained in each bounding box.
Args:
yx: The (y, x) coordinate in pixels of the point.
bounding_boxes: A 2D int array of shape (num_bboxes, 4), where each row
represents a bounding box: (y_top_left, x_top_left, box_height,
box_width). Note: containment is inclusive of the bounding box edges.
Returns:
is_inside: A 1D bool array where each element specifies if the point is
contained within the respective box.
"""
y, x = yx
# `bounding_boxes` has shape (n_elements, 4); we extract each array along the
# last axis into shape (n_elements, 1), then squeeze unneeded dimension.
top, left, height, width = [
jnp.squeeze(v, axis=-1) for v in jnp.split(bounding_boxes, 4, axis=-1)
]
# The y-axis is inverted for AndroidEnv, so bottom = top + height.
bottom, right = top + height, left + width
return jnp.logical_and(y >= top, y <= bottom) & jnp.logical_and(x >= left, x <= right)
def _resize_annotation_bounding_boxes(
annotation_positions, annotation_width_augment_fraction,
annotation_height_augment_fraction):
"""Resize the bounding boxes by the given fractions.
Args:
annotation_positions: Array of shape (N, 4), where each row represents the
(y, x, height, width) of the bounding boxes.
annotation_width_augment_fraction: The fraction to augment the box widths,
E.g., 1.4 == 240% total increase.
annotation_height_augment_fraction: Same as described for width, but for box
height.
Returns:
Resized bounding box.
"""
height_change = (
annotation_height_augment_fraction * annotation_positions[:, 2])
width_change = (
annotation_width_augment_fraction * annotation_positions[:, 3])
# Limit bounding box positions to the screen.
resized_annotations = jnp.stack([
jnp.maximum(0, annotation_positions[:, 0] - (height_change / 2)),
jnp.maximum(0, annotation_positions[:, 1] - (width_change / 2)),
jnp.minimum(1, annotation_positions[:, 2] + height_change),
jnp.minimum(1, annotation_positions[:, 3] + width_change),
],axis=1)
return resized_annotations
def is_tap_action(normalized_start_yx, normalized_end_yx):
distance = jnp.linalg.norm(
jnp.array(normalized_start_yx) - jnp.array(normalized_end_yx))
return distance <= _SWIPE_DISTANCE_THRESHOLD
def _is_non_dual_point_action(action_type):
return jnp.not_equal(action_type, ActionType.DUAL_POINT)
def _check_tap_actions_match(
tap_1_yx,
tap_2_yx,
annotation_positions,
matching_tap_distance_threshold_screen_percentage,
annotation_width_augment_fraction,
annotation_height_augment_fraction,
):
"""Determines if two tap actions are the same."""
resized_annotation_positions = _resize_annotation_bounding_boxes(
annotation_positions,
annotation_width_augment_fraction,
annotation_height_augment_fraction,
)
# Check if the ground truth tap action falls in an annotation's bounding box.
tap1_in_box = _yx_in_bounding_boxes(tap_1_yx, resized_annotation_positions)
tap2_in_box = _yx_in_bounding_boxes(tap_2_yx, resized_annotation_positions)
both_in_box = jnp.max(tap1_in_box & tap2_in_box)
# If the ground-truth tap action falls outside any of the annotation
# bounding boxes or one of the actions is inside a bounding box and the other
# is outside bounding box or vice versa, compare the points using Euclidean
# distance.
within_threshold = (
jnp.linalg.norm(jnp.array(tap_1_yx) - jnp.array(tap_2_yx))
<= matching_tap_distance_threshold_screen_percentage
)
return jnp.logical_or(both_in_box, within_threshold)
def _check_drag_actions_match(
drag_1_touch_yx,
drag_1_lift_yx,
drag_2_touch_yx,
drag_2_lift_yx,
):
"""Determines if two drag actions are the same."""
# Store drag deltas (the change in the y and x coordinates from touch to
# lift), magnitudes, and the index of the main axis, which is the axis with
# the greatest change in coordinate value (e.g. a drag starting at (0, 0) and
# ending at (0.3, 0.5) has a main axis index of 1).
drag_1_deltas = drag_1_lift_yx - drag_1_touch_yx
drag_1_magnitudes = jnp.abs(drag_1_deltas)
drag_1_main_axis = np.argmax(drag_1_magnitudes)
drag_2_deltas = drag_2_lift_yx - drag_2_touch_yx
drag_2_magnitudes = jnp.abs(drag_2_deltas)
drag_2_main_axis = np.argmax(drag_2_magnitudes)
return jnp.equal(drag_1_main_axis, drag_2_main_axis)
def check_actions_match(
action_1_touch_yx,
action_1_lift_yx,
action_1_action_type,
action_2_touch_yx,
action_2_lift_yx,
action_2_action_type,
annotation_positions,
tap_distance_threshold = _TAP_DISTANCE_THRESHOLD,
annotation_width_augment_fraction = ANNOTATION_WIDTH_AUGMENT_FRACTION,
annotation_height_augment_fraction = ANNOTATION_HEIGHT_AUGMENT_FRACTION,
):
"""Determines if two actions are considered to be the same.
Two actions being "the same" is defined here as two actions that would result
in a similar screen state.
Args:
action_1_touch_yx: The (y, x) coordinates of the first action's touch.
action_1_lift_yx: The (y, x) coordinates of the first action's lift.
action_1_action_type: The action type of the first action.
action_2_touch_yx: The (y, x) coordinates of the second action's touch.
action_2_lift_yx: The (y, x) coordinates of the second action's lift.
action_2_action_type: The action type of the second action.
annotation_positions: The positions of the UI annotations for the screen. It
is A 2D int array of shape (num_bboxes, 4), where each row represents a
bounding box: (y_top_left, x_top_left, box_height, box_width). Note that
containment is inclusive of the bounding box edges.
tap_distance_threshold: The threshold that determines if two taps result in
a matching screen state if they don't fall the same bounding boxes.
annotation_width_augment_fraction: The fraction to increase the width of the
bounding box by.
annotation_height_augment_fraction: The fraction to increase the height of
of the bounding box by.
Returns:
A boolean representing whether the two given actions are the same or not.
"""
action_1_touch_yx = jnp.asarray(action_1_touch_yx)
action_1_lift_yx = jnp.asarray(action_1_lift_yx)
action_2_touch_yx = jnp.asarray(action_2_touch_yx)
action_2_lift_yx = jnp.asarray(action_2_lift_yx)
# Checks if at least one of the actions is global (i.e. not DUAL_POINT),
# because if that is the case, only the actions' types need to be compared.
has_non_dual_point_action = jnp.logical_or(
_is_non_dual_point_action(action_1_action_type),
_is_non_dual_point_action(action_2_action_type),
)
different_dual_point_types = jnp.logical_xor(
is_tap_action(action_1_touch_yx, action_1_lift_yx),
is_tap_action(action_2_touch_yx, action_2_lift_yx),
)
is_tap = jnp.logical_and(
is_tap_action(action_1_touch_yx, action_1_lift_yx),
is_tap_action(action_2_touch_yx, action_2_lift_yx),
)
taps_match = _check_tap_actions_match(
action_1_touch_yx,
action_2_touch_yx,
annotation_positions,
tap_distance_threshold,
annotation_width_augment_fraction,
annotation_height_augment_fraction,
)
taps_match = jnp.logical_and(is_tap, taps_match)
drags_match = _check_drag_actions_match(
action_1_touch_yx, action_1_lift_yx, action_2_touch_yx, action_2_lift_yx
)
drags_match = jnp.where(is_tap, False, drags_match)
return jnp.where(
has_non_dual_point_action,
jnp.equal(action_1_action_type, action_2_action_type),
jnp.where(
different_dual_point_types,
False,
jnp.logical_or(taps_match, drags_match),
),
)
```
## /eval/utils/convert_output.py
```py path="/eval/utils/convert_output.py"
import json
import os
import jsonschema
from concurrent.futures import ProcessPoolExecutor, as_completed
from tqdm import tqdm
# Get the absolute path of the current file
current_file_path = os.path.abspath(__file__)
schema_dir = os.path.dirname(os.path.dirname(current_file_path))
EXTRACT_SCHEMA = json.load(open(os.path.join(schema_dir, 'utils/schema', 'schema_for_extraction.json'), encoding="utf-8"))
def load_json_data(file_path):
data = []
# Determine file type, support both JSON and JSONL
if file_path.endswith('.json'):
# Handle JSON file
with open(file_path, 'r') as file:
data = json.load(file)
elif file_path.endswith('.jsonl'):
# Handle JSONL file
with open(file_path, 'r') as file:
first_line = file.readline().strip()
try:
json.loads(first_line)
data.append(json.loads(first_line))
except json.JSONDecodeError:
pass
for line in file:
line = line.strip()
if line:
data.append(json.loads(line))
return data
def parse_action(data):
try:
jsonschema.validate(data, EXTRACT_SCHEMA)
actions = {}
parameters = {}
status = data.get("STATUS", "continue") # Default value
# Define actions
action_keys = ["POINT", "to", "PRESS", "TYPE"]
# Extract actions
for key in action_keys:
if key in data:
actions[key] = data[key]
# Extract global parameters
parameters["duration"] = data.get("duration", EXTRACT_SCHEMA["properties"]["duration"]["default"])
# Handle "to" parameter, if present
if "to" in data:
parameters["to"] = data["to"]
return actions, parameters, status
except Exception as e:
print('Error, JSON is NOT valid according to the schema.')
return None, None, None
# Use multiprocessing to speed up processing
def process_step(args):
task, episode_id, step_id, pred, base_path = args
try:
actions, parameters, status = parse_action(pred)
transformed_entry = {
"action_predict": {
"COA": {
"txt": {
"ACTION": actions,
"ARGS": parameters,
"STATUS": status
},
}
}
}
folder = f"{task}-{episode_id}"
file_name = f"{folder}_{step_id}.json"
output_file_path = os.path.join(base_path, folder, file_name)
with open(output_file_path, 'w', encoding='utf-8') as output_file:
json.dump(transformed_entry, output_file, indent=4, ensure_ascii=False)
return f"Saved transformed entry to: {output_file_path}"
except Exception as e:
return f"Error processing step {step_id} in episode {episode_id}: {e}"
# # Multi-threaded version
def convert2aitz(input_path, output_path, max_workers=None):
data = load_json_data(input_path)
base_path = os.path.join(output_path)
folders = set()
tasks = []
for item in data:
task = item.get("category", item.get("subset", "unknown"))
episode_id = item.get("episode_id", "unknown")
steps = item.get("steps", [item])
for index, each_step in enumerate(steps):
step_id = index if "steps" in item else each_step.get("step_id", index)
folder = f"{task}-{episode_id}"
folders.add(folder)
pred = each_step.get("pred", {})
tasks.append((task, episode_id, step_id, pred, base_path))
for folder in folders:
folder_path = os.path.join(base_path, folder)
os.makedirs(folder_path, exist_ok=True)
with ProcessPoolExecutor(max_workers=max_workers) as executor:
futures = [executor.submit(process_step, task_args) for task_args in tasks]
for future in tqdm(as_completed(futures), total=len(futures), desc="Processing steps"):
result = future.result()
print(result)
# # Single-threaded version
def convert2aitz_single_thread(input_path, output_path):
data = load_json_data(input_path)
base_path = os.path.join(output_path)
for item in data:
task = item.get("category", "unknown")
episode_id = item.get("episode_id", "unknown")
steps = item.get("steps", [item])
for index, each_step in enumerate(steps):
step_id = index if "steps" in item else each_step.get("step_id", index)
actions, parameters, status = parse_action(each_step["pred"])
transformed_entry = {
"action_predict": {
"COA": {
"txt": {
"ACTION": actions,
"ARGS": parameters,
"STATUS": status
},
}
}
}
folder = f"{task}-{episode_id}"
file_name = f"{folder}_{step_id}.json"
folder_path = os.path.join(base_path, folder)
output_path = os.path.join(folder_path, file_name)
os.makedirs(folder_path, exist_ok=True)
with open(output_path, 'w', encoding='utf-8') as output_file:
json.dump(transformed_entry, output_file, indent=4, ensure_ascii=False)
print(f"Saved transformed entry to: {output_path}")
```
## /eval/utils/evaluator.py
```py path="/eval/utils/evaluator.py"
import os
import json
import numpy as np
import Levenshtein
import math
from PIL import Image, ImageDraw, ImageFont
from utils.action_type import ActionType
from utils.utils import annotate_and_save_image
from typing import List, Union
# Based on evaluator of Qwen 2.5 VL
# https://github.com/QwenLM/Qwen2.5-VL/issues/904
# https://gist.github.com/LukeForeverYoung/274a073ca77c9dc46022cb8cc5382223
# https://gist.github.com/LukeForeverYoung/1f5d19495788de0d905c5ac6341153f5
# Get the absolute path of the current file
current_file_path = os.path.abspath(__file__)
schema_dir = os.path.dirname(os.path.dirname(current_file_path))
EXTRACT_SCHEMA = json.load(open(os.path.join(schema_dir, 'utils/schema', 'schema_for_extraction.json'), encoding="utf-8"))
# CONSTANTS
_TAP_DISTANCE_THRESHOLD = 0.14 # Fraction of the screen
_TAP_DISTANCE_THRESHOLD_AC = 0.04 # for android control, align with qwen's code.
_SWIPE_DISTANCE_THRESHOLD = 0.04 # Interval determining if an action is a tap or a swipe.
ANNOTATION_WIDTH_AUGMENT_FRACTION= 1.2 # aitw set it to 1.4, aitz and qwen 2.5 vl set it to 1.2.
ANNOTATION_HEIGHT_AUGMENT_FRACTION= 1.2 # We follow qwen setting.
default_duration = EXTRACT_SCHEMA["properties"]["duration"]["default"] # default 200
def _resize_annotation_bounding_boxes(
annotation_position: Union[List[float], List[List[float]]],
width_factor: float = 1.2,
height_factor: float = 1.2,
):
"""Uniformly enlarge bbox(es) by the given factors."""
def _resize(box: List[float]):
y, x, h, w = box
h_delta = (height_factor - 1) * h
w_delta = (width_factor - 1) * w
y = max(0, y - h_delta / 2)
x = max(0, x - w_delta / 2)
h = min(1, h + h_delta)
w = min(1, w + w_delta)
return [y, x, h, w]
if not annotation_position:
return []
if isinstance(annotation_position[0], list):
return [_resize(b) for b in annotation_position]
return _resize(annotation_position)
def is_tap_action(normalized_start_yx, normalized_end_yx):
distance = np.linalg.norm(np.array(normalized_start_yx) - np.array(normalized_end_yx))
return distance <= _SWIPE_DISTANCE_THRESHOLD
def check_inside(x, y, bbox_list):
bbox_array = np.array(bbox_list)
y_min, x_min, height, width = bbox_array[:, 0], bbox_array[:, 1], bbox_array[:, 2], bbox_array[:, 3]
y_max, x_max = y_min + height, x_min + width
# Check whether (x, y) is inside any of the bounding boxes
within_x = (x_min <= x) & (x <= x_max)
within_y = (y_min <= y) & (y <= y_max)
within_bbox = within_x & within_y
if np.any(within_bbox):
within_bbox_coords = bbox_array[within_bbox]
return True, within_bbox_coords
else:
return False, None
def obtain_gt_bbox(coordinate, bbox_list, eval_android_control=False):
x, y = coordinate['x'], coordinate['y']
if len(bbox_list) == 0:
return []
if not eval_android_control:
is_inside, bbox_inside = check_inside(x, y, bbox_list)
if is_inside:
return bbox_inside.tolist()
else:
return []
else:
def get_center_distance(box):
ymin, xmin, h, w = box
center_y = ymin + h/2
center_x = xmin + w/2
return ((center_y - y) ** 2 + (center_x - x) ** 2) ** 0.5
sorted_boxes = sorted(bbox_list, key=get_center_distance)
# return the 5 nearest bboxes
return sorted_boxes[:5]
def _get_direction(point1, point2):
try:
x1, y1 = point1["x"], point1["y"]
x2, y2 = point2["x"], point2["y"]
assert x1 is not None
assert x2 is not None
assert y1 is not None
assert y2 is not None
vector = (x2 - x1, y2 - y1)
vx, vy = vector
except Exception as e:
return "no direction"
directions = {
"up": (0, -1),
"down": (0, 1),
"left": (-1, 0),
"right": (1, 0)
}
vector_length = math.sqrt(vx ** 2 + vy ** 2)
if vector_length == 0:
return "no direction"
unit_vector = (vx / vector_length, vy / vector_length)
max_cosine = -float('inf')
closest_direction = None
for direction, dir_vector in directions.items():
dx, dy = dir_vector
dir_length = math.sqrt(dx ** 2 + dy ** 2)
cos_theta = (unit_vector[0] * dx + unit_vector[1] * dy) / dir_length
if cos_theta > max_cosine:
max_cosine = cos_theta
closest_direction = direction
return closest_direction
def get_direction(point, to):
if isinstance(to, str):
if to in ["up", "down", "left", "right"]:
return to
else:
return "no direction"
elif isinstance(to, list):
try:
point1 = {"x": point[0], "y": point[1]}
point2 = {"x": to[0], "y": to[1]}
return _get_direction(point1, point2)
except Exception as e:
return "no direction"
class ActionEvaluator(object):
def __init__(self, save_dir, eval_android_control=False) -> None:
self.save_dir = save_dir
# compatible with aitz evaluator
self.demo_mode = "COA"
self.screen_mode = "txt"
self._aitz_action_type_ = ActionType
self._stop_status = [
"finish",
"satisfied",
"impossible",
"interrupt",
"need_feedback"
]
self.eval_android_control = eval_android_control
def action_map(self, action_api: dict):
action = action_api.get('ACTION', None)
args = action_api.get('ARGS', None)
status = action_api.get('STATUS', None)
duration = args.get('duration', default_duration) if args else None
if action is None and args is None and status is None:
print('Schema error.')
return None, {}
elif status in self._stop_status:
return "stop", {}
elif "TYPE" in action:
return "type", action['TYPE']
elif "POINT" in action and "to" not in args and duration == default_duration: # click
return "click", action['POINT']
elif "POINT" in action and "to" in args and duration == default_duration: # swipte
return "scroll", {"start": action['POINT'], "end": args['to']}
elif "POINT" in action and "duration" in args and duration > default_duration: # long press
return "long_point", {"coordinate": action['POINT'], "duration": args['duration']}
elif "PRESS" in action:
return "press", action['PRESS']
elif "duration" in args: # pause and wait
return "stop", args['duration']
else:
raise ValueError("Unknown action type.")
def _parse_action_(self, pred, image_width=None, image_height=None):
pd_action_type, pd_action_yx, pd_action_idx, pd_action_direction, pd_action_text, pd_action_button, pd_duration = (None, ) * 7
pr = pred.get('action_predict', {})
if self.demo_mode not in pr: return (None, ) * 7
action = pr[self.demo_mode].get(self.screen_mode, {})
if not action: return (None, ) * 7
pd_action_type, pd_action_args = self.action_map(action)
if pd_action_type is None: print('Unknown action: ', action)
# scale factors
scale_x = 1000
scale_y = 1000
if pd_action_type == "click":
try:
pd_action_yx = {"x": pd_action_args[0] / scale_x, "y": pd_action_args[1] / scale_y}
except Exception as e:
pd_action_yx = {"x": 0.0, "y": 0.0}
elif pd_action_type == "long_point":
try:
pd_action_yx = {"x": pd_action_args["coordinate"][0] / scale_x, "y": pd_action_args["coordinate"][1] / scale_y}
except Exception as e:
pd_action_yx = {"x": 0.0, "y": 0.0}
else:
pd_action_yx = None
# Not supporting click by id
pd_action_idx = None
# Process swipe
pd_action_direction = get_direction(pd_action_args["start"], pd_action_args["end"]) if pd_action_type == "scroll" else None
# Process text input
pd_action_text = pd_action_args if pd_action_type == "type" else None
# Process button press
pd_action_button = pd_action_args.lower() if pd_action_type == "press" else None
# Process long press
pd_duration = pd_action_args["duration"] if pd_action_type == "long_point" else None
# Treat pause and wait as normal stop
return pd_action_type, pd_action_yx, pd_action_idx, pd_action_text, pd_action_button, pd_action_direction, pd_duration
def _parse_answer_(self, gt):
gt_cand_nodes=None
gt_action_text=None
gt_action_type=None
gt_action_yx=None
gt_action_direction=None
gt_action_button=None
gt_duration=None
if gt['result_action_type'] == self._aitz_action_type_.TYPE:
gt_action_type = "type"
gt_action_text = gt['result_action_text']
elif gt['result_action_type'] == self._aitz_action_type_.DUAL_POINT: # Might be swipe or click
normalized_start_yx = gt['result_touch_yx']
normalized_start_yx = json.loads(normalized_start_yx)
normalized_end_yx = gt['result_lift_yx']
normalized_end_yx =json.loads(normalized_end_yx)
if is_tap_action(normalized_start_yx, normalized_end_yx):
gt_cand_nodes = json.loads(gt['ui_positions'])
gt_action_type = "click"
gt_action_yx = {"y": normalized_start_yx[0], "x": normalized_start_yx[1]}
else:
point1 = {"y": normalized_start_yx[0], "x": normalized_start_yx[1]}
point2 = {"y": normalized_end_yx[0], "x": normalized_end_yx[1]}
gt_action_type = "scroll"
gt_action_direction = _get_direction(point1, point2)
elif gt['result_action_type'] == self._aitz_action_type_.LONG_POINT:
normalized_start_yx = gt['result_touch_yx']
normalized_start_yx = json.loads(normalized_start_yx)
normalized_end_yx = gt['result_lift_yx']
normalized_end_yx =json.loads(normalized_end_yx)
gt_cand_nodes = json.loads(gt['ui_positions'])
gt_action_type = "long_point"
gt_action_yx = {"y": normalized_start_yx[0], "x": normalized_start_yx[1]}
gt_duration = gt['duration']
elif gt['result_action_type'] == self._aitz_action_type_.PRESS_BACK:
gt_action_type = "press"
gt_action_button = "back"
elif gt['result_action_type'] == self._aitz_action_type_.PRESS_HOME:
gt_action_type = "press"
gt_action_button = "home"
elif gt['result_action_type'] == self._aitz_action_type_.PRESS_ENTER:
gt_action_type = "press"
gt_action_button = "enter"
elif gt['result_action_type'] == self._aitz_action_type_.STATUS_TASK_COMPLETE or gt['result_action_type'] == self._aitz_action_type_.STATUS_TASK_IMPOSSIBLE:
gt_action_type = "stop"
gt_action_text = gt['result_action_text']
elif gt['result_action_type'] == self._aitz_action_type_.NO_ACTION:
gt_action_type = "stop"
gt_duration = gt['duration']
else:
raise ValueError("Unknow action type.")
return gt_action_type, gt_action_yx, gt_cand_nodes, \
gt_action_text, gt_action_button, gt_action_direction, gt_duration
def __call__(self, gt, pred, annotate_image=False):
""" eval_single_step """
pd_action_detail = None
pixel_distance = None
image_width, image_height = gt['image_width'], gt['image_height']
subset, episode_id, step_id, task_desc = gt['subset'], gt['episode_id'], gt['step_id'], gt['instruction']
# get ground truth information
gt_action_type, gt_action_yx, gt_cand_nodes, \
gt_action_text, gt_action_button, gt_action_direction, gt_duration = self._parse_answer_(gt)
if not gt_action_type: print(gt['result_action_type'])
gt_action_detail = {
"click": gt_action_yx,
"scroll": gt_action_direction,
"type": gt_action_text,
"press": gt_action_button,
"long_point": gt_action_yx,
"stop": "stop"
}.get(gt_action_type, None)
# get predict action information
pd_action_type, pd_action_yx, pd_action_idx, \
pd_action_text, pd_action_button, pd_action_direction, pd_duration = self._parse_action_(pred, image_width, image_height)
pd_action_detail={
"click": pd_action_yx,
"scroll": pd_action_direction,
"type": pd_action_text,
"press": pd_action_button,
"long_point": pd_action_yx,
"stop": "stop"
}.get(pd_action_type, None)
# compute metrics
hit_format = True if pd_action_type is not None else False # invalid actions are set to None when converting format
type_match = (pd_action_type is not None and gt_action_type == pd_action_type)
exact_match = False
text_dist = None
if type_match and (pd_action_type == "click" or pd_action_type == "long_point"):
gt_cand_nodes = _resize_annotation_bounding_boxes(gt_cand_nodes, ANNOTATION_WIDTH_AUGMENT_FRACTION, ANNOTATION_HEIGHT_AUGMENT_FRACTION)
gt_bbox = obtain_gt_bbox(gt_action_yx, gt_cand_nodes, self.eval_android_control)
if gt_bbox == []:
y_gt, x_gt = gt_action_yx["y"], gt_action_yx["x"]
y_pd, x_pd = pd_action_yx["y"], pd_action_yx["x"]
distance = np.linalg.norm(np.array([x_gt, y_gt]) - np.array([x_pd, y_pd]))
exact_match = bool(distance <= (_TAP_DISTANCE_THRESHOLD_AC if self.eval_android_control else _TAP_DISTANCE_THRESHOLD))
reference_point = gt_action_yx["x"], gt_action_yx["y"]
else:
reference_point = gt_action_yx["x"], gt_action_yx["y"]
for bbox in gt_bbox:
ymin, xmin, height, width = bbox
ymax, xmax = ymin + height, xmin + width
exact_match = ((ymin <= pd_action_yx["y"] <= ymax) and (xmin <= pd_action_yx["x"] <= xmax))
if exact_match:
reference_point = (xmax + xmin) / 2, (ymax + ymin) / 2
break
if not exact_match:
y_gt, x_gt = gt_action_yx["y"], gt_action_yx["x"]
y_pd, x_pd = pd_action_yx["y"], pd_action_yx["x"]
distance = np.linalg.norm(np.array([x_gt, y_gt]) - np.array([x_pd, y_pd]))
exact_match = bool(distance <= (_TAP_DISTANCE_THRESHOLD_AC if self.eval_android_control else _TAP_DISTANCE_THRESHOLD))
# Calculate pixel mse, here the distance is calculated in the normalized space [0, 1000]
pixel_distance = np.linalg.norm(np.array([pd_action_yx["x"], pd_action_yx["y"]])*1000 - np.array(reference_point)*1000)
if type_match and pd_action_type == "scroll":
exact_match = (pd_action_direction == gt_action_direction)
if type_match and pd_action_type == "type":
pd_text_norm = pd_action_text.lower().strip()
gt_text_norm = gt_action_text.lower().strip()
text_dist = Levenshtein.ratio(pd_text_norm, gt_text_norm)
# align with Qwen‑2.5‑VL eval
exact_match = (pd_text_norm in gt_text_norm or \
gt_text_norm in pd_text_norm)
if type_match and pd_action_type == "press":
exact_match = (pd_action_button == gt_action_button)
if type_match and pd_action_type == "stop":
exact_match = True
# for visualization
if annotate_image:
output_folder = os.path.join(self.save_dir, "image_output")
annotate_and_save_image(gt['image_full_path'], output_folder,
gt_action_type, gt_action_detail,
pd_action_type, pd_action_detail, type_match,
exact_match,subset, episode_id, step_id, task_desc)
if not type_match or (type_match and not exact_match):
match_type = "No Type Match" if not type_match else "No Exact Match"
print(f"\n{match_type}, pd action: {pd_action_type}, detail: {pd_action_detail}; gt action: {gt_action_type}, detail: {gt_action_detail}, {subset}_{episode_id}_{step_id}")
return {
"subset": subset,
"episode_id": episode_id,
"step_id": step_id,
"answer": {
"action_type": gt_action_type,
"action_detail": gt_action_detail
},
"pred": {
"action_type": pd_action_type,
"action_detail": pd_action_detail
},
"type_match": type_match,
"exact_match": exact_match,
"text_dist": text_dist,
"format_hit": hit_format,
"pixel_distance": pixel_distance,
}
@staticmethod
def compute_episode_metrics(episode_results):
success, progress = [], []
total_exact_matches = 0
total_steps = 0
for __, eplist in episode_results.items():
ep_success, ep_progress = True, 0
for ex in eplist:
if ex['exact_match'] is True:
ep_progress += 1
total_exact_matches += 1
else:
ep_success = False
if not ep_success:
break
success.append(ep_success)
progress.append(ep_progress/len(eplist)*1.0)
total_steps = 0
for __, eplist in episode_results.items():
for ex in eplist:
total_steps += 1
num_episodes = len(success)
num_successes = sum(success)
return {
"total_episodes": num_episodes,
"total_steps": total_steps,
"num_successes": num_successes,
"total_exact_matches": total_exact_matches,
"success_rate": round(sum(success) / len(success), 4),
"goal_progress": round(sum(progress) / len(progress), 4)}
@staticmethod
def compute_atomic_metrics(step_results):
recorder = {
'total': {'count':0, 'type_match':0, 'exact_match':0, "hit": 0},
# -------------------------------------------
'CLICK': {'count':0, 'type_match':0, 'exact_match':0},
'TYPE': {'count':0, 'type_match':0, 'exact_match':0, 'text_dist': []},
'SCROLL': {'count':0, 'type_match':0, 'exact_match':0},
'PRESS': {'count':0, 'type_match':0, 'exact_match':0},
'STOP': {'count':0, 'type_match':0, 'exact_match':0},
'LONG_POINT': {'count':0, 'type_match':0, 'exact_match':0},
}
for step in step_results:
recorder['total']['count'] += 1
recorder['total']['hit'] += step.get('format_hit', 0)
# Get action_type and ensure it is a string
action_type = step.get('answer', {}).get('action_type')
if isinstance(action_type, str):
action_type = action_type.upper()
else:
action_type = ''
if action_type in recorder:
recorder[action_type]['count'] += 1
recorder[action_type]['type_match'] += step.get('type_match', 0)
recorder['total']['type_match'] += step.get('type_match', 0)
recorder[action_type]['exact_match'] += step.get('exact_match', 0)
recorder['total']['exact_match'] += step.get('exact_match', 0)
if 'text_dist' in recorder[action_type] and step.get('text_dist') is not None:
recorder[action_type]['text_dist'].append(step['text_dist'])
# Initialize scores dictionary, including counts and ratios
scores = {
metric_key: {
'count': recorder[metric_key]['count'],
'type_acc': round(
recorder[metric_key]['type_match'] / recorder[metric_key]['count'],
4
) if recorder[metric_key]['count'] > 0 else 0,
'exact_acc': round(
recorder[metric_key]['exact_match'] / recorder[metric_key]['count'],
4
) if recorder[metric_key]['count'] > 0 else 0
}
for metric_key in ['total', 'CLICK', 'LONG_POINT', 'SCROLL', 'PRESS', 'STOP', 'TYPE']
}
# Calculate hit_rate
scores['total']['hit_rate'] = round(
recorder['total']['hit'] / recorder['total']['count'], 4
) if recorder['total']['count'] > 0 else 0
# Calculate average text_dist for TYPE
if recorder['TYPE']['text_dist']:
scores['TYPE']['text_dist_avg'] = round(
sum(recorder['TYPE']['text_dist']) / len(recorder['TYPE']['text_dist']), 4
)
else:
scores['TYPE']['text_dist_avg'] = 0
# Calculate pixel distance
pixel_distances = [
step['pixel_distance'] for step in step_results
if step.get('pixel_distance') is not None
]
median_pixel_distance = round(
float(np.median(pixel_distances)), 4
) if pixel_distances else -1
mean_pixel_distance = -1
if pixel_distances:
pixel_distances = np.array(pixel_distances)
filtered_distances = pixel_distances[pixel_distances < 1e15]
if len(filtered_distances) > 0:
mean_pixel_distance = round(
float(np.mean(filtered_distances)), 4
)
scores['mean_pixel_distance'] = mean_pixel_distance
scores['median_pixel_distance'] = median_pixel_distance
return scores
```
## /eval/utils/mv_vote.py
```py path="/eval/utils/mv_vote.py"
import numpy as np
import io
from PIL import Image
import requests
import aiohttp
import asyncio
from collections import Counter
from typing import List, Dict, Any
# Function to perform majority voting on a list of actions
def clear_mv(actions: List[Dict[str, Any]]) -> Dict[str, Any]:
"""对 CLEAR 操作进行多数投票处理"""
vote_result = {}
clear_thoughts = []
clear_actions = []
for act in actions:
if "CLEAR" in act:
thought = act.get("thought", "")
clear_action = act["CLEAR"]
clear_actions.append(clear_action)
clear_thoughts.append(thought)
if clear_thoughts:
if_clear = Counter(clear_actions)
most_common_clear, _ = if_clear.most_common(1)[0]
most_common_thought = next(thought for thought, action in zip(clear_thoughts, clear_actions) if action == most_common_clear)
vote_result["CLEAR"] = most_common_clear
if most_common_clear:
vote_result["CLEAR"] = True
else:
vote_result["CLEAR"] = False
if "thought" not in vote_result:
vote_result["thought"] = most_common_thought
return vote_result
def status_mv(actions: List[Dict[str, Any]]) -> Dict[str, Any]:
"""对 STATUS 操作进行多数投票处理"""
vote_result = {}
status_thoughts = []
for act in actions:
if "STATUS" in act:
thought = act.get("thought", "")
status_action = act["STATUS"]
status_thoughts.append((thought, status_action))
if status_thoughts:
# 统计 STATUS 操作的出现次数
status_counter = Counter([pt[1] for pt in status_thoughts])
most_common_status, _ = status_counter.most_common(1)[0]
# 获取对应的 thought
thought = next(pt[0] for pt in status_thoughts if pt[1] == most_common_status)
vote_result["STATUS"] = most_common_status
vote_result["thought"] = thought
return vote_result
def type_mv(actions: str | None) -> list[str]:
"""对 TYPE 操作进行多数投票处理"""
vote_result = {}
type_thoughts = []
for act in actions:
if "TYPE" in act:
thought = act.get("thought", "")
type_action = act["TYPE"]
type_thoughts.append((thought, type_action))
if type_thoughts:
# 统计 TYPE 操作的出现次数
type_counter = Counter([pt[1] for pt in type_thoughts])
most_common_type, _ = type_counter.most_common(1)[0]
# 获取对应的 thought
thought = next(pt[0] for pt in type_thoughts if pt[1] == most_common_type)
vote_result["TYPE"] = most_common_type
vote_result["thought"] = thought
return vote_result
def press_mv(actions: List[Dict[str, Any]]) -> Dict[str, Any]:
"""对 PRESS 操作进行多数投票处理"""
vote_result = {}
press_thoughts = []
for act in actions:
if "PRESS" in act:
thought = act.get("thought", "")
press_action = act["PRESS"]
press_thoughts.append((thought, press_action))
if press_thoughts:
# 统计 PRESS 操作的出现次数
press_counter = Counter([pt[1] for pt in press_thoughts])
most_common_press, _ = press_counter.most_common(1)[0]
# 获取对应的 thought
thought = next(pt[0] for pt in press_thoughts if pt[1] == most_common_press)
vote_result["PRESS"] = most_common_press
vote_result["thought"] = thought
return vote_result
def point_mv(actions: List[Dict[str, Any]]) -> Dict[str, Any]:
"""对 POINT 操作进行多数投票处理"""
vote_result = {}
point_thoughts = []
for act in actions:
if "POINT" in act:
point = tuple(act["POINT"])
to = act["to"] if "to" in act else None
duration = act["duration"] if "duration" in act else None
thought = act.get("thought", "")
point_thoughts.append((point, thought, to, duration))
if point_thoughts:
# 分别统计x和y坐标的出现次数
x_coords = [pt[0][0] for pt in point_thoughts] # 获取所有x坐标
y_coords = [pt[0][1] for pt in point_thoughts] # 获取所有y坐标
x_avg = sum(x_coords) / len(x_coords)
y_avg = sum(y_coords) / len(y_coords)
min_distance = float('inf')
closest_thought = ""
closest_to = None
for point, thought, to, duration in point_thoughts:
x, y = point
distance = ((x - x_avg) ** 2 + (y - y_avg) ** 2) ** 0.5
if distance < min_distance:
min_distance = distance
closet_point = point
closest_thought = thought
closest_to = to
closest_duration = duration
vote_result["POINT"] = list(closet_point)
if closest_to is not None:
vote_result["to"] = closest_to
if closest_duration is not None:
vote_result["duration"] = closest_duration
vote_result["thought"] = closest_thought
return vote_result
def duration_mv(actions: List[Dict[str, Any]]) -> Dict[str, Any]:
"""对 duration 操作进行多数投票处理"""
vote_result = {}
durations = []
duration_thoughts = []
for act in actions:
if "duration" in act:
duration = act["duration"]
thought = act.get("thought", "")
durations.append(duration)
duration_thoughts.append((duration, thought))
if duration_thoughts:
avg_duration = sum(durations) / len(durations)
closest_duration = min(durations, key=lambda x: abs(x - avg_duration))
thought = next(pt[1] for pt in duration_thoughts if pt[0] == closest_duration)
vote_result["duration"] = int(closest_duration)
vote_result["thought"] = thought
return vote_result
def get_action_type(action: Dict[str, Any]) -> str:
"""提取单条 action 的主类型"""
for key in ["POINT", "PRESS", "TYPE", "CLEAR","STATUS","duration"]:
if key in action:
return key
return "UNKNOWN"
def action_majority_vote(actions: List[Dict[str, Any]]) -> Dict[str, Any]:
vote_result = {}
if not actions:
return {}
# 1. 找出每个 action 的主类型
types = [get_action_type(act) for act in actions if get_action_type(act) != "UNKNOWN"]
if not types:
return {}
# 2. 对类型做多数投票
type_counter = Counter(types)
main_type, _ = type_counter.most_common(1)[0]
# 3. 调用对应的 Majority Vote 函数
dispatch_table = {
"POINT": point_mv,
"PRESS": press_mv,
"TYPE": type_mv,
"CLEAR": clear_mv,
"STATUS": status_mv,
"duration": duration_mv
}
return dispatch_table[main_type](actions)
```
## /eval/utils/qwen_mobile_tool.py
```py path="/eval/utils/qwen_mobile_tool.py"
import os.path as osp
from matplotlib import pyplot as plt
from PIL import Image
import numpy as np
from icecream import ic
import math
import argparse
from PIL import Image, ImageDraw, ImageFont, ImageColor
import torch
from transformers import Qwen2_5_VLForConditionalGeneration, AutoProcessor
from qwen_agent.llm.fncall_prompts.nous_fncall_prompt import (
NousFnCallPrompt,
Message,
ContentItem,
)
from utils.evaluator import get_direction
from qwen_vl_utils import smart_resize
import json
from utils.action_utils import *
from utils.utils_qwen.agent_function_call import MobileUse
from IPython.display import display
import os
#os.environ["CUDA_VISIBLE_DEVICES"] = "4" # Only use the 4th GPU
args = type('Args', (), {})
import torch
torch.manual_seed(1)
def aitw_2_uitars(aitw_action: dict):
"""
Convert AITW action to UITARS action format
"""
ex_action_type = aitw_action['result_action_type']
if ex_action_type == ActionType.DUAL_POINT:
lift_yx = json.loads(aitw_action['result_lift_yx'])
touch_yx = json.loads(aitw_action['result_touch_yx'])
if is_tap_action(np.array(touch_yx), np.array(lift_yx)):
# Click action
click_y, click_x = lift_yx[0], lift_yx[1]
click_x = int(click_x* 1000)
click_y = int(click_y* 1000)
return f"click(start_box=\'<|box_start|>({click_x},{click_y})<|box_end|>\')"
else:
# Swipe action
touch_yx_new = {
"x": touch_yx[1],
"y": touch_yx[0]
}
lift_yx_new = {
"x": lift_yx[1],
"y": lift_yx[0]
}
direction = get_direction(touch_yx_new, lift_yx_new)
return f"scroll(direction='{direction}')"
elif ex_action_type == ActionType.PRESS_BACK:
return f"press_back()"
elif ex_action_type == ActionType.PRESS_HOME:
return f"press_home()"
elif ex_action_type == ActionType.PRESS_ENTER:
return f"press_enter()"
elif ex_action_type == ActionType.TYPE:
return f"type(content='{aitw_action['result_action_text']}')"
elif ex_action_type == ActionType.STATUS_TASK_COMPLETE:
return f"finished()"
elif ex_action_type == ActionType.STATUS_TASK_IMPOSSIBLE:
return f"finished()"
elif ex_action_type == ActionType.LONG_POINT:
lift_yx = json.loads(aitw_action['result_lift_yx'])
touch_yx = json.loads(aitw_action['result_touch_yx'])
click_y, click_x = lift_yx[0], lift_yx[1]
click_x = int(click_x* 1000)
click_y = int(click_y* 1000)
return f"long_press(start_box=\'<|box_start|>({click_x},{click_y})<|box_end|>\')"
elif ex_action_type == ActionType.NO_ACTION:
return f"wait()"
elif ex_action_type == ActionType.OPEN_APP:
return f"open(app_name='{aitw_action['result_action_app_name']}')"
else:
print('aitw_action:',aitw_action)
raise NotImplementedError
# Return formatted JSON string
return json.dumps(qwen_action)
def aitz_2_qwen2_5(aitz_action: dict, resized_height: int, resized_width: int) -> str:
"""
Convert AITZ action to Qwen2.5 action format
Args:
aitz_action (dict): AITZ format action, contains ACTION and ARGS
resized_height (int): Screen height
resized_width (int): Screen width
Returns:
str: Qwen2.5 format action string
"""
aitz_action = json.loads(aitz_action)
print(aitz_action)
action_type = aitz_action["ACTION"]
args = aitz_action["ARGS"]
qwen_action = {}
# Handle click action
if action_type == "CLICK_ELEMENT":
bbox = args["bbox"]
# Calculate center point from bbox [x1, y1, x2, y2]
center_x = (bbox[0] + bbox[2]) / 2
center_y = (bbox[1] + bbox[3]) / 2
# Convert coordinates to screen coordinates
center_x = int(center_x * resized_width)
center_y = int(center_y * resized_height)
qwen_action = {
"action": "click",
"coordinate": [center_x, center_y]
}
# Handle swipe action
elif action_type == "SCROLL":
direction = args["direction"]
# Set start and end points according to direction
mid_x = resized_width // 2
mid_y = resized_height // 2
if direction == "up":
# Swipe up from the middle of the screen (start at bottom, end at top)
qwen_action = {
"action": "swipe",
"coordinate": [mid_x, mid_y + 300],
"coordinate2": [mid_x, mid_y - 300]
}
elif direction == "down":
# Swipe down from the middle of the screen
qwen_action = {
"action": "swipe",
"coordinate": [mid_x, mid_y - 300],
"coordinate2": [mid_x, mid_y + 300]
}
elif direction == "left":
# Swipe left from the middle of the screen
qwen_action = {
"action": "swipe",
"coordinate": [mid_x + 300, mid_y],
"coordinate2": [mid_x - 300, mid_y]
}
elif direction == "right":
# Swipe right from the middle of the screen
qwen_action = {
"action": "swipe",
"coordinate": [mid_x - 300, mid_y],
"coordinate2": [mid_x + 300, mid_y],
}
# Handle text input
elif action_type == "INPUT":
qwen_action = {
"action": "type",
"text": args["text"]
}
# Handle system buttons
elif action_type == "PRESS BACK":
qwen_action = {
"action": "system_button",
"button": "Back"
}
elif action_type == "PRESS HOME":
qwen_action = {
"action": "system_button",
"button": "Home"
}
elif action_type == "PRESS ENTER":
qwen_action = {
"action": "system_button",
"button": "Enter"
}
# Handle terminate action
elif action_type == "STOP":
qwen_action = {
"action": "terminate",
"status": args.get("task_status", "success")
}
# Build complete Qwen2.5 format output
if qwen_action:
return f'{{"name":"mobile_use","arguments":{json.dumps(qwen_action)}}}'
else:
return ""
def qwen2_5_2_aitz(output_text: str, resized_height: int, resized_width: int) -> str:
"""
Convert Qwen2.5 output to AITZ output
"""
action = json.loads(output_text.split('<tool_call>\n')[1].split('\n</tool_call>')[0])
qwen_action = action['arguments']
action_name = qwen_action['action']
# Handle click action, treat long_press as click because there is no corresponding action
if action_name == "click" or action_name == "long_press":
x, y = qwen_action["coordinate"]
# Convert coordinates to bbox format [x1, y1, x2, y2]
# Use 0.1 times the screen width and height as the click area
# Normalize
x = x/ resized_width
y = y/ resized_height
return {"ACTION": "CLICK_ELEMENT", "ARGS": {"bbox": [int((x-0.1)*999), int((y-0.1)*999), int((x+0.1)*999), int((y+0.1)*999)]}}
# Handle swipe action
elif action_name == "swipe":
x1, y1 = qwen_action["coordinate"]
x2, y2 = qwen_action["coordinate2"]
# hack short swipe and should be click (copied from Qwen's evaluation logic)
if np.linalg.norm([x2 - x1, y2 - y1]) <= 0.04:
action_name = "click"
x1=x1/ resized_width
y1=y1/ resized_height
x2=x2/ resized_width
y2=y2/ resized_height
return {"ACTION": "CLICK_ELEMENT", "ARGS": {"bbox": [int((x1-0.1)*999), int((y1-0.1)*999), int((x1+0.1)*999), int((y1+0.1)*999)]}}
# Determine swipe direction based on start and end points
if abs(x2 - x1) > abs(y2 - y1): # Horizontal swipe
direction = "right" if x2 > x1 else "left"
else: # Vertical swipe
direction = "down" if y2 > y1 else "up"
return {"ACTION": "SCROLL", "ARGS": {"direction": direction}}
# Handle text input
elif action_name == "type":
return {"ACTION": "INPUT", "ARGS": {"text": qwen_action["text"]}}
# Handle system buttons
elif action_name == "system_button":
button = qwen_action["button"]
if button == "Back":
return {"ACTION": "PRESS_BACK", "ARGS": {}}
elif button == "Home":
return {"ACTION": "PRESS_HOME", "ARGS": {}}
elif button == "Enter":
return {"ACTION": "PRESS_ENTER", "ARGS": {}}
# Handle terminate action
elif action_name == "terminate":
return {"ACTION": "STOP", "ARGS": {"task_status": qwen_action["status"]}}
# For other actions (such as key, wait, open, long_press, etc.), may need to ignore or handle specially
# key, open, wait cannot find corresponding action, long_press is treated as click here
return {"ACTION": "", "ARGS": {}}
model_path = "Qwen/Qwen2.5-VL-7B-Instruct"
model_path = "/home/test/test03/models/Qwen2.5-VL-7B-Instruct"
user_query_template = 'The user query:{user_request} (You have done the following operation on the current device):'
#model = Qwen2_5_VLForConditionalGeneration.from_pretrained(model_path, torch_dtype=torch.bfloat16, attn_implementation="flash_attention_2",device_map="auto")
#processor = AutoProcessor.from_pretrained(model_path)
def get_qwen_response(user_query: str, screenshot: str, args=None, model_path: str = "/home/test/test03/models/Qwen2.5-VL-7B-Instruct") -> tuple:
"""
Get response from Qwen model
Args:
user_query: User query text
screenshot: Screenshot path
model_path: Model path, default is the official model
Returns:
tuple: (response_text, status_code)
"""
try:
# Set default args
if args is None:
args = type('Args', (), {
'greedy': False,
'top_p': 0.01,
'top_k': 1,
'temperature': 0.01,
'repetition_penalty': 1.0,
'presence_penalty': 0.0,
'out_seq_length': 1024,
'seed': 1
})
# Build parameters using args
generation_params = {
'do_sample': not getattr(args, 'greedy', False),
'top_p': getattr(args, 'top_p', 0.01),
'top_k': getattr(args, 'top_k', 1),
'temperature': getattr(args, 'temperature', 0.01),
'repetition_penalty': getattr(args, 'repetition_penalty', 1.0),
'presence_penalty': getattr(args, 'presence_penalty', 0.0),
'max_new_tokens': getattr(args, 'out_seq_length', 1024),
'seed': getattr(args, 'seed', 1)
}
# Handle image size
dummy_image = Image.open(screenshot)
#print(dummy_image.size)
resized_height, resized_width = smart_resize(
dummy_image.height,
dummy_image.width,
factor=processor.image_processor.patch_size * processor.image_processor.merge_size,
min_pixels=processor.image_processor.min_pixels,
max_pixels=processor.image_processor.max_pixels,
)
#print(resized_height, resized_width)
# Initialize mobile device interface
mobile_use = MobileUse(
cfg={"display_width_px": resized_width, "display_height_px": resized_height}
)
# Build message
message = NousFnCallPrompt.preprocess_fncall_messages(
messages=[
Message(role="system", content=[ContentItem(text="You are a helpful assistant.")]),
Message(role="user", content=[
ContentItem(text=user_query_template.format(user_request=user_query)),
ContentItem(image=f"file://{screenshot}")
]),
],
functions=[mobile_use.function],
lang=None,
)
message = [msg.model_dump() for msg in message]
# Handle input
text = processor.apply_chat_template(
message,
tokenize=False,
add_generation_prompt=True
)
print('text:',text)
inputs = processor(
text=[text],
images=[dummy_image],
padding=True,
return_tensors="pt"
).to('cuda')
# If you need to set a random seed, set it before generate
if hasattr(args, 'seed'):
import torch
torch.manual_seed(args.seed)
# Call generate with correct parameters
output_ids = model.generate(
**inputs,
**generation_params
)
generated_ids = [output_ids[len(input_ids):] for input_ids, output_ids in zip(inputs.input_ids, output_ids)]
output_text = processor.batch_decode(
generated_ids,
skip_special_tokens=True,
clean_up_tokenization_spaces=True
)[0]
aitz_answer=qwen2_5_2_aitz(output_text,resized_height, resized_width)
return json.dumps(aitz_answer), 200
except Exception as e:
print(f"Error: {str(e)}")
return str(e), 500
user_query_template_history = '''The user query: {user_request}
Before answering, explain your reasoning step-by-step in <thinking></thinking> tags, and insert them before the <tool_call></tool_call> XML tags.
After answering, summarize your action in <conclusion></conclusion> tags, and insert them after the <tool_call></tool_call> XML tags.
Task progress (You have done the following operation on the current device):
{history_actions}'''
def aitw_2_qwen2_5_action(aitw_action: dict, resized_height: int, resized_width: int) -> str:
"""
Convert AITW action to Qwen2.5 action format
"""
ex_action_type = aitw_action['result_action_type']
qwen_action = {"name": "mobile_use", "arguments": {}}
if ex_action_type == ActionType.DUAL_POINT:
lift_yx = json.loads(aitw_action['result_lift_yx'])
touch_yx = json.loads(aitw_action['result_touch_yx'])
if is_tap_action(np.array(touch_yx), np.array(lift_yx)):
# Click action
click_y, click_x = lift_yx[0], lift_yx[1]
click_x = int(click_x* resized_width)
click_y = int(click_y* resized_height)
qwen_action["arguments"] = {
"action": "click",
"coordinate": [click_x, click_y]
}
else:
# Swipe action
qwen_action["arguments"] = {
"action": "swipe",
"coordinate": [int(touch_yx[1]* resized_width), int(touch_yx[0]* resized_height)], # Start point
"coordinate2": [int(lift_yx[1]* resized_width), int(lift_yx[0]* resized_height)] # End point
}
elif ex_action_type == ActionType.PRESS_BACK:
button = "Back"
qwen_action["arguments"] = {
"action": "system_button",
"button": button
}
elif ex_action_type == ActionType.PRESS_HOME:
button = "Home"
qwen_action["arguments"] = {
"action": "system_button",
"button": button
}
elif ex_action_type == ActionType.PRESS_ENTER:
button = "Enter"
qwen_action["arguments"] = {
"action": "system_button",
"button": button
}
elif ex_action_type == ActionType.TYPE:
qwen_action["arguments"] = {
"action": "type",
"text": aitw_action['result_action_text']
}
elif ex_action_type == ActionType.STATUS_TASK_COMPLETE:
qwen_action["arguments"] = {
"action": "terminate",
"status": "success"
}
elif ex_action_type == ActionType.STATUS_TASK_IMPOSSIBLE:
qwen_action["arguments"] = {
"action": "terminate",
"status": "failure"
}
elif ex_action_type == ActionType.LONG_POINT:
qwen_action["arguments"] = {
"action": "long_press",
"coordinate": [int(aitw_action['result_touch_yx'][1]* resized_width), int(aitw_action['result_touch_yx'][0]* resized_height)],
"time": 2
}
elif ex_action_type == ActionType.NO_ACTION:
qwen_action["arguments"] = {
"action": "wait",
"time": 2
}
else:
print('aitw_action:',aitw_action)
raise NotImplementedError
# Return formatted JSON string
return json.dumps(qwen_action)
def aitw_2_qwen2_5(aitw_action: dict, resized_height: int, resized_width: int) -> str:
"""
Convert AITW action to Qwen2.5 prompt
"""
aitw_action = json.loads(aitw_action)
action=aitw_2_qwen2_5_action(aitw_action,resized_height, resized_width)
thinking = f"<thinking>\n{aitw_action['coat_action_think']}\n</thinking>\n"
action = f"<tool_call>\n{action}\n</tool_call>\n"
result = f'<conclusion>\n"{aitw_action["coat_action_desc"]}"\n</conclusion>'
return thinking + action + result
def get_qwen_response_history(user_query: str, screenshot: str, history_actions: list, model_path: str = "/home/test/test03/models/Qwen2.5-VL-7B-Instruct") -> tuple:
"""
Get response from Qwen model
Args:
user_query: User query text
screenshot: Screenshot path
history_actions: History actions
model_path: Model path, default is the official model
Returns:
tuple: (response_text, status_code)
"""
#try:
print('history_actions:',history_actions)
# Handle image size
dummy_image = Image.open(screenshot)
#print(dummy_image.size)
resized_height, resized_width = smart_resize(
dummy_image.height,
dummy_image.width,
factor=processor.image_processor.patch_size * processor.image_processor.merge_size,
min_pixels=processor.image_processor.min_pixels,
max_pixels=processor.image_processor.max_pixels,
)
#print('max_pixels:',processor.image_processor.max_pixels)
# 12845056, exceeds 4096*3112 (4k resolution), should be enough for mobile resolution
# Convert history_actions to Qwen2.5 format
if history_actions:
history_actions_str = "".join([f"Step {i+1}: {aitw_2_qwen2_5(action,resized_height, resized_width).replace('<tool_call>','').replace('</tool_call>','').strip()}; " for i, action in enumerate(history_actions)])
else:
history_actions_str = ""
#print(resized_height, resized_width)
# Initialize mobile device interface
mobile_use = MobileUse(
cfg={"display_width_px": resized_width, "display_height_px": resized_height}
)
# Build message
message = NousFnCallPrompt.preprocess_fncall_messages(
messages=[
Message(role="system", content=[ContentItem(text="You are a helpful assistant.")]),
Message(role="user", content=[
ContentItem(text=user_query_template_history.format(user_request=user_query,history_actions=history_actions_str)),
ContentItem(image=f"file://{screenshot}")
]),
],
functions=[mobile_use.function],
lang=None,
)
message = [msg.model_dump() for msg in message]
# Handle input
text = processor.apply_chat_template(
message,
tokenize=False,
add_generation_prompt=True
)
print('text:',text)
inputs = processor(
text=[text],
images=[dummy_image],
padding=True,
return_tensors="pt"
).to('cuda')
# Modify generation_params definition
generation_params = {
# Replace 'greedy' with 'do_sample'
'do_sample': not getattr(args, 'greedy', False),
'top_p': getattr(args, 'top_p', 0.01),
'top_k': getattr(args, 'top_k', 1),
'temperature': getattr(args, 'temperature', 0.01),
'repetition_penalty': getattr(args, 'repetition_penalty', 1.0),
# 'presence_penalty' is not supported, can be removed
# Replace 'out_seq_length' with 'max_new_tokens'
# 'seed' is not directly supported, needs to be set externally
}
# If you need to set a random seed, set it before generate
# Call generate with correct parameters
output_ids = model.generate(
**inputs,
max_new_tokens=getattr(args, 'out_seq_length', 2048),
**generation_params
)
generated_ids = [output_ids[len(input_ids):] for input_ids, output_ids in zip(inputs.input_ids, output_ids)]
output_text = processor.batch_decode(
generated_ids,
skip_special_tokens=True,
clean_up_tokenization_spaces=True
)[0]
aitz_answer=qwen2_5_2_aitz(output_text,resized_height, resized_width)
return json.dumps(aitz_answer), 200
#except Exception as e:
# print(f"Error: {str(e)}")
# return str(e), 500
# Example usage
if __name__ == "__main__":
user_query = 'Open the file manager app and view the au_uu_SzH3yR2.mp3 file in MUSIC Folder'
screenshot = "/home/test/test03/fuyikun/CoAT/data-example/GOOGLE_APPS-523638528775825151/GOOGLE_APPS-523638528775825151_0.png"
response, state = get_qwen_response(user_query, screenshot)
print(f"Response: {response}")
print(f"State: {state}")
```
## /eval/utils/schema/schema.json
```json path="/eval/utils/schema/schema.json"
{
"type": "object",
"description": "执行操作并决定当前任务状态",
"additionalProperties": false,
"properties": {
"thought": {
"type": "string",
"description": "智能体的思维过程"
},
"POINT": {
"$ref": "#/$defs/Location",
"description": "点击屏幕上的指定位置"
},
"to": {
"description": "移动,组合手势参数",
"oneOf": [
{
"enum": [
"up",
"down",
"left",
"right"
],
"description": "从当前点(POINT)出发,执行滑动手势操作,方向包括向上、向下、向左、向右"
},
{
"$ref": "#/$defs/Location",
"description": "移动到某个位置"
}
]
},
"duration": {
"type": "integer",
"description": "动作执行的时间或等待时间,毫秒",
"minimum": 0,
"default": 200
},
"PRESS": {
"type": "string",
"description": "触发特殊按键,HOME为回到主页按钮,BACK为返回按钮,ENTER为回车按钮",
"enum": [
"HOME",
"BACK",
"ENTER"
]
},
"TYPE": {
"type": "string",
"description": "输入文本"
},
"STATUS": {
"type": "string",
"description": "当前任务的状态。特殊情况:satisfied,无需操作;impossible,任务无法完成;interrupt,任务中断;need_feedback,需要用户反馈;",
"enum": [
"continue",
"finish",
"satisfied",
"impossible",
"interrupt",
"need_feedback"
],
"default": "continue"
}
},
"$defs": {
"Location": {
"type": "array",
"description": "坐标为相对于屏幕左上角位原点的相对位置,并且按照宽高比例缩放到0~1000,数组第一个元素为横坐标x,第二个元素为纵坐标y",
"items": {
"type": "integer",
"minimum": 0,
"maximum": 1000
},
"minItems": 2,
"maxItems": 2
}
}
}
```
## /eval/utils/schema/schema_for_extraction.json
```json path="/eval/utils/schema/schema_for_extraction.json"
{
"type": "object",
"description": "执行操作并决定当前任务状态",
"additionalProperties": false,
"properties": {
"thought": {
"type": "string"
},
"POINT": {
"description": "点击屏幕上的指定位置",
"$ref": "#/$defs/Location"
},
"to": {
"description": "移动,组合手势参数",
"oneOf": [
{
"enum": [
"up",
"down",
"left",
"right"
],
"description": "结合 POINT 操作,实现向上下左右滑动"
},
{
"$ref": "#/$defs/Location",
"description": "移动到某个位置"
}
]
},
"duration": {
"type": "integer",
"description": "动作执行的时间或等待时间,毫秒",
"minimum": 0,
"default": 200
},
"PRESS": {
"type": "string",
"description": "触发特殊按键,HOME为回到主页按钮,BACK为返回按钮,ENTER为回车按钮,APPSELECT为查看已打开APP列表按钮",
"enum": [
"HOME",
"BACK",
"ENTER",
"APPSELECT"
]
},
"TYPE": {
"type": "string",
"description": "输入文本"
},
"DEEP_LINK": {
"type": "null",
"description": "跳转到最近打开的 APP"
},
"CLEAR": {
"type": "null",
"description": "清空输入框的内容"
},
"STATUS": {
"type": "string",
"description": "当前任务的状态。特殊情况:satisfied,无需操作;impossible,任务无法完成;interrupt,任务中断;need_feedback,需要用户反馈;",
"enum": [
"continue",
"start",
"finish",
"satisfied",
"impossible",
"interrupt",
"need_feedback"
],
"default": "continue"
}
},
"$defs": {
"Location": {
"type": "array",
"description": "坐标为相对于屏幕左上角位原点的相对位置,并且按照宽高比例缩放到 0~1000,数组第一个元素为横坐标 x,第二个元素为纵坐标 y",
"items": {
"type": "integer",
"minimum": 0,
"maximum": 1000
},
"minItems": 2,
"maxItems": 2
}
},
"allOf": [
{
"if": {
"required": ["to"]
},
"then": {
"required": ["POINT"]
}
},
{
"if": {
"anyOf": [
{ "not": { "required": ["STATUS"] } },
{ "properties": { "STATUS": { "enum": ["continue", "start"] } } }
]
},
"then": {
"anyOf": [
{ "required": ["POINT"] },
{ "required": ["PRESS"] },
{ "required": ["TYPE"] },
{ "required": ["DEEP_LINK"] },
{ "required": ["CLEAR"] },
{ "required": ["duration"] }
]
}
},
{
"oneOf": [
{
"required": ["POINT"],
"not": {
"anyOf": [
{ "required": ["PRESS"] },
{ "required": ["TYPE"] },
{ "required": ["DEEP_LINK"] },
{ "required": ["CLEAR"] }
]
}
},
{
"required": ["PRESS"],
"not": {
"anyOf": [
{ "required": ["POINT"] },
{ "required": ["TYPE"] },
{ "required": ["DEEP_LINK"] },
{ "required": ["CLEAR"] }
]
}
},
{
"required": ["TYPE"],
"not": {
"anyOf": [
{ "required": ["POINT"] },
{ "required": ["PRESS"] },
{ "required": ["DEEP_LINK"] },
{ "required": ["CLEAR"] }
]
}
},
{
"required": ["DEEP_LINK"],
"not": {
"anyOf": [
{ "required": ["POINT"] },
{ "required": ["PRESS"] },
{ "required": ["TYPE"] },
{ "required": ["CLEAR"] }
]
}
},
{
"required": ["CLEAR"],
"not": {
"anyOf": [
{ "required": ["POINT"] },
{ "required": ["PRESS"] },
{ "required": ["TYPE"] },
{ "required": ["DEEP_LINK"] }
]
}
},
{
"not": {
"anyOf": [
{ "required": ["POINT"] },
{ "required": ["PRESS"] },
{ "required": ["TYPE"] },
{ "required": ["DEEP_LINK"] },
{ "required": ["CLEAR"] }
]
}
}
]
}
]
}
```
## /eval/utils/schema/test_schema.py
```py path="/eval/utils/schema/test_schema.py"
import json
from jsonschema import validate, ValidationError
import os
# Get the absolute path of the current file
current_file_path = os.path.abspath(__file__)
schema_dir = os.path.dirname(current_file_path)
# This is the schema for extracting actions, not the schema for training LLM — the latter is a subset of the former.
schema = json.load(open(os.path.join(schema_dir, 'schema_for_extraction.json'), encoding="utf-8"))
# test cases
test_cases = [
{
"name": "Valid Case: Only POINT",
"data": {
"POINT": [500, 300],
"duration": 300,
"STATUS": "continue"
},
"expected": True
},
{
"name": "Valid Case: POINT with to (direction)",
"data": {
"POINT": [500, 300],
"to": "left",
"duration": 300,
"STATUS": "continue"
},
"expected": True
},
{
"name": "Valid Case: POINT with to (Location)",
"data": {
"POINT": [500, 300],
"to": [600, 400],
"duration": 300,
"STATUS": "start"
},
"expected": True
},
{
"name": "Valid Case: Only PRESS",
"data": {
"PRESS": "HOME",
"duration": 200,
"STATUS": "finish"
},
"expected": True
},
{
"name": "Valid Case: Only TYPE",
"data": {
"TYPE": "Hello, World!",
"duration": 250,
"STATUS": "satisfied"
},
"expected": True
},
{
"name": "Invalid Case: to without POINT",
"data": {
"to": "up",
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: POINT with PRESS",
"data": {
"POINT": [500, 300],
"PRESS": "HOME",
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: PRESS with TYPE",
"data": {
"PRESS": "BACK",
"TYPE": "Some text",
"duration": 200,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: Multiple Actions",
"data": {
"POINT": [500, 300],
"PRESS": "HOME",
"TYPE": "Hello",
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: POINT with invalid to value",
"data": {
"POINT": [500, 300],
"to": "invalid_direction",
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: PRESS with invalid enum",
"data": {
"PRESS": "INVALID_KEY",
"duration": 200,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: LOCATION with out of range coordinates",
"data": {
"POINT": [1500, 300],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: Missing STATUS (should default to continue)",
"data": {
"PRESS": "ENTER",
"duration": 200
},
"expected": True # STATUS has a default, so it's valid
},
{
"name": "Invalid Case: Additional Property",
"data": {
"PRESS": "HOME",
"extra_property": "not_allowed",
"duration": 200,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Valid Case: POINT with to and default STATUS",
"data": {
"POINT": [400, 200],
"to": "down",
"duration": 300
},
"expected": True
},
{
"name": "Valid Case: POINT with default STATUS",
"data": {
"POINT": [400, 200],
},
"expected": True
},
{
"name": "Invalid Case: Negative duration",
"data": {
"PRESS": "HOME",
"duration": -100,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Valid Case: Only STATUS 'finish'",
"data": {
"STATUS": "finish"
},
"expected": True
},
{
"name": "Valid Case: Only STATUS 'satisfied'",
"data": {
"STATUS": "finish"
},
"expected": True
},
{
"name": "Valid Case: Only STATUS 'impossible'",
"data": {
"STATUS": "finish"
},
"expected": True
},
{
"name": "Valid Case: Only STATUS 'interrupt'",
"data": {
"STATUS": "finish"
},
"expected": True
},
{
"name": "Valid Case: Only STATUS 'need_feedback'",
"data": {
"STATUS": "finish"
},
"expected": True
},
{
"name": "Valid Case: POINT at boundary (0,0)",
"data": {
"POINT": [0, 0],
"duration": 200,
"STATUS": "continue"
},
"expected": True
},
{
"name": "Valid Case: POINT at boundary (1000,1000)",
"data": {
"POINT": [1000, 1000],
"duration": 200,
"STATUS": "start"
},
"expected": True
},
{
"name": "Valid Case: duration missing (should default to 200)",
"data": {
"PRESS": "BACK",
"STATUS": "finish"
},
"expected": True
},
{
"name": "Valid Case: to as Location with boundary coordinates",
"data": {
"POINT": [500, 500],
"to": [0, 1000],
"duration": 250,
"STATUS": "satisfied"
},
"expected": True
},
{
"name": "Invalid Case: duration as string",
"data": {
"PRESS": "ENTER",
"duration": "300",
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: POINT with one coordinate missing",
"data": {
"POINT": [500],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: to as array with one element",
"data": {
"POINT": [500, 300],
"to": [600],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: PRESS as null",
"data": {
"PRESS": None,
"duration": 200,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: TYPE as null",
"data": {
"TYPE": None,
"duration": 200,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: 'continue' STATUS missing other action",
"data": {
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: to as empty array",
"data": {
"POINT": [500, 300],
"to": [],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: POINT as non-array",
"data": {
"POINT": "500,300",
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Valid Case: POINT and to with minimum duration (0)",
"data": {
"POINT": [500, 300],
"to": "up",
"duration": 0,
"STATUS": "continue"
},
"expected": True
},
{
"name": "Valid Case: POINT and to with maximum duration (10000)",
"data": {
"POINT": [500, 300],
"to": "down",
"duration": 10000,
"STATUS": "finish"
},
"expected": True
},
{
"name": "Invalid Case: POINT with non-integer values",
"data": {
"POINT": [500.5, "300"],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: to as direction and POINT missing",
"data": {
"to": "left",
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: to as invalid Location array size (3 elements)",
"data": {
"POINT": [500, 300],
"to": [600, 400, 200],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Valid Case: Only STATUS with duration",
"data": {
"STATUS": "finish",
"duration": 500
},
"expected": True
},
{
"name": "Invalid Case: STATUS with invalid enum value",
"data": {
"STATUS": "unknown_status",
"duration": 200
},
"expected": False
},
{
"name": "Valid Case: POINT with to as direction and missing duration (should default)",
"data": {
"POINT": [400, 200],
"to": "right",
"STATUS": "start"
},
"expected": True
},
{
"name": "Valid Case: PRESS with missing duration (should default)",
"data": {
"PRESS": "APPSELECT",
"STATUS": "need_feedback"
},
"expected": True
},
{
"name": "Valid Case: TYPE with missing duration (should default)",
"data": {
"TYPE": "Sample Text",
"STATUS": "impossible"
},
"expected": True
},
{
"name": "Invalid Case: POINT as empty array",
"data": {
"POINT": [],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: POINT with negative coordinates",
"data": {
"POINT": [-100, 300],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: PRESS with lowercase value",
"data": {
"PRESS": "home",
"duration": 200,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: TYPE with empty string",
"data": {
"TYPE": "",
"duration": 200,
"STATUS": "continue"
},
"expected": True # Empty string is valid as per schema
},
{
"name": "Valid Case: POINT with to as Location and missing duration (should default)",
"data": {
"POINT": [500, 500],
"to": [600, 600],
"STATUS": "start"
},
"expected": True
},
{
"name": "Valid Case: POINT at boundary (0,0)",
"data": {
"POINT": [0, 0],
"duration": 200,
"STATUS": "continue"
},
"expected": True
},
{
"name": "Valid Case: POINT at boundary (1000,1000)",
"data": {
"POINT": [1000, 1000],
"duration": 200,
"STATUS": "start"
},
"expected": True
},
{
"name": "Valid Case: duration missing (should default to 200)",
"data": {
"PRESS": "BACK",
"STATUS": "finish"
},
"expected": True
},
{
"name": "Valid Case: to as Location with boundary coordinates",
"data": {
"POINT": [500, 500],
"to": [0, 1000],
"duration": 250,
"STATUS": "satisfied"
},
"expected": True
},
{
"name": "Invalid Case: duration as string",
"data": {
"PRESS": "ENTER",
"duration": "300",
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: POINT with one coordinate missing",
"data": {
"POINT": [500],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: to as array with one element",
"data": {
"POINT": [500, 300],
"to": [600],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: PRESS as null",
"data": {
"PRESS": None,
"duration": 200,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: TYPE as null",
"data": {
"TYPE": None,
"duration": 200,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: 'continue' STATUS missing other action",
"data": {
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: 'start' STATUS missing other action",
"data": {
"STATUS": "start"
},
"expected": False
},
{
"name": "Invalid Case: Empty object",
"data": {},
"expected": False
},
{
"name": "Invalid Case: to as empty array",
"data": {
"POINT": [500, 300],
"to": [],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: POINT as non-array",
"data": {
"POINT": "500,300",
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Valid Case: POINT and to with minimum duration (0)",
"data": {
"POINT": [500, 300],
"to": "up",
"duration": 0,
"STATUS": "continue"
},
"expected": True
},
{
"name": "Valid Case: POINT and to with maximum duration (10000)",
"data": {
"POINT": [500, 300],
"to": "down",
"duration": 10000,
"STATUS": "finish"
},
"expected": True
},
{
"name": "Invalid Case: POINT with non-integer values",
"data": {
"POINT": [500.5, "300"],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: to as direction and POINT missing",
"data": {
"to": "left",
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: to as invalid Location array size (3 elements)",
"data": {
"POINT": [500, 300],
"to": [600, 400, 200],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Valid Case: Only STATUS with duration",
"data": {
"STATUS": "finish",
"duration": 500
},
"expected": True
},
{
"name": "Invalid Case: STATUS with invalid enum value",
"data": {
"STATUS": "unknown_status",
"duration": 200
},
"expected": False
},
{
"name": "Valid Case: POINT with to as direction and missing duration (should default)",
"data": {
"POINT": [400, 200],
"to": "right",
"STATUS": "start"
},
"expected": True
},
{
"name": "Valid Case: PRESS with missing duration (should default)",
"data": {
"PRESS": "APPSELECT",
"STATUS": "need_feedback"
},
"expected": True
},
{
"name": "Valid Case: TYPE with missing duration (should default)",
"data": {
"TYPE": "Sample Text",
"STATUS": "impossible"
},
"expected": True
},
{
"name": "Invalid Case: POINT as empty array",
"data": {
"POINT": [],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: POINT with negative coordinates",
"data": {
"POINT": [-100, 300],
"duration": 300,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: PRESS with lowercase value",
"data": {
"PRESS": "home",
"duration": 200,
"STATUS": "continue"
},
"expected": False
},
{
"name": "Invalid Case: TYPE with empty string",
"data": {
"TYPE": "",
"duration": 200,
"STATUS": "continue"
},
"expected": True # Empty string is valid as per schema
},
{
"name": "Valid Case: POINT with to as Location and missing duration (should default)",
"data": {
"POINT": [500, 500],
"to": [600, 600],
"STATUS": "start"
},
"expected": True
},
{
"name": "Valid Case: only duration (just wait)",
"data": {
"duration": 200,
},
"expected": True
},
{
"name": "Valid Case: POINT with to as Location and missing duration (should default), with thought",
"data": {
"POINT": [500, 500],
"to": [600, 600],
"STATUS": "start",
"thought": "I am thinking"
},
"expected": True
},
]
def run_tests(schema, test_cases):
print("Starting test cases...\n")
for idx, test in enumerate(test_cases, 1):
data = test["data"]
expected = test["expected"]
name = test["name"]
try:
validate(instance=data, schema=schema)
result = True
error_message = ""
except ValidationError as e:
result = False
error_message = e.message
status = "PASS" if result == expected else "FAIL"
print(f"Test Case {idx}: {name}")
print(f" Expected Result: {'Valid' if expected else 'Invalid'}")
print(f" Actual Result: {'Valid' if result else 'Invalid'}")
if status == "FAIL":
print(f" Error Message: {error_message}")
print(f" Test Status: {status}\n")
print("Testing completed.")
if __name__ == "__main__":
run_tests(schema, test_cases)
```
## /eval/utils/utils.py
```py path="/eval/utils/utils.py"
from colorama import init, Fore, Style
import os
from utils.action_type import ActionType
def annotate_and_save_image(img_path, output_folder, gt_action_type, gt_action_detail, pd_action_type, pd_action_detail, type_match, exact_match, subset, episode_id, step_id, task_desc):
"""Save an annotated image with action details to the specified folder."""
# Load the image and get its dimensions
image = Image.open(img_path)
draw = ImageDraw.Draw(image)
# Dynamically compute font size based on image height
base_height = 1080 # Reference height, e.g., 1080p
font_size = max(12, int(image.height / base_height * 20)) # Ensure minimum font size is 12
current_file_path = os.path.abspath(__file__)
current_dir = os.path.dirname(current_file_path)
try:
font = ImageFont.truetype(os.path.join(current_dir, './SimHei.ttf'), font_size)
except IOError:
# If the specified font file does not exist, use the default font
font = ImageFont.load_default()
w, h = image.width, image.height
# Create annotation text
annotation_text = (
f"taskDesc: {task_desc}\n"
f"taskID: {subset}{episode_id}_{step_id}\n"
f"GT action: {gt_action_type}\n"
f"GT detail: {gt_action_detail}\n"
f"PD action: {pd_action_type}\n"
f"PD detail: {pd_action_detail}\n"
f"type_match: {'Yes' if type_match else 'No'}\n"
f"exac_match: {'Yes' if exact_match else 'No'}"
)
# Calculate text size and wrap lines if necessary
max_width = w - 20 # Max width for the text
lines = []
for line in annotation_text.split('\n'):
# Split line by words to check width
words = line.split()
current_line = ""
for word in words:
# Check width after adding a word
test_line = current_line + " " + word if current_line else word
if draw.textlength(test_line, font=font) > max_width:
# If line is too long, start a new line
lines.append(current_line)
current_line = word
else:
current_line = test_line
lines.append(current_line) # Add the final line
# Draw each line on the image
y_text = 10
line_spacing = int(font_size * 1.2) # Line spacing is 1.2 times the font size
for line in lines:
draw.text((10, y_text), line, font=font, fill='red')
y_text += line_spacing # Move to next line position
# Draw rectangle and point based on conditions
if pd_action_type == 'click' and type_match:
if isinstance(gt_action_detail, (list, tuple)) and len(gt_action_detail) == 4:
ymin, xmin, height, width = gt_action_detail # Parse GT action details
pd_x = pd_action_detail.get("x", 0) * w
pd_y = pd_action_detail.get("y", 0) * h
gt_box = [xmin * w, ymin * h, (xmin + width) * w, (ymin + height) * h]
draw.rectangle(gt_box, outline="red", width=max(1, int(font_size / 5))) # Adjust line width dynamically
point_radius = max(5, int(font_size / 2)) # Adjust point radius dynamically
draw.ellipse(
(pd_x - point_radius, pd_y - point_radius, pd_x + point_radius, pd_y + point_radius),
fill="red",
outline="blue",
width=max(1, int(font_size / 10))
)
# Save the annotated image to the output folder
if not os.path.exists(output_folder):
os.makedirs(output_folder, exist_ok=True)
output_file_name = os.path.basename(img_path).replace('.png', '_annotated.png')
output_path = os.path.join(output_folder, output_file_name)
image.save(output_path)
return output_path
def get_dataset_dir(data_name):
data_list = ['aitz_test', 'chinese_app_test', 'gui_odyssey_test', 'android_control_high_test', 'android_control_low_test']
assert data_name in data_list, "Error, unkonw eval dataset."
data_split = None
data_dir = None
data_subset = None
current_file_path = os.path.abspath(__file__)
data_dir = os.path.dirname(os.path.dirname(current_file_path))
match data_name:
case 'aitz_test':
data_dir = os.path.join(data_dir, "eval_data", "aitz_test")
data_split = "test"
data_subset = ["general", "install", "web_shopping", "google_apps"]
case 'chinese_app_test':
data_dir = os.path.join(data_dir, "eval_data", "chinese_app_test")
data_split = "test"
data_subset = ["domestic"]
case 'gui_odyssey_test':
data_dir = os.path.join(data_dir, "eval_data", "odyssey")
data_split = "test"
data_subset = ["odyssey"]
case 'android_control_high_test':
data_dir = os.path.join(data_dir, "eval_data", "android_control_high_test")
data_split = "test"
data_subset = ["android_control"]
case 'android_control_low_test':
data_dir = os.path.join(data_dir, "eval_data", "android_control_low_test")
data_split = "test"
data_subset = ["android_control"]
return data_dir, data_split, data_subset
```
## /eval/utils/utils_odyssey/config.json
```json path="/eval/utils/utils_odyssey/config.json"
{
"architectures": [
"QWenLMHeadModel"
],
"attn_dropout_prob": 0.0,
"auto_map": {
"AutoConfig": "configuration_qwen.QWenConfig",
"AutoModelForCausalLM": "modeling_qwen.QWenLMHeadModel"
},
"bf16": true,
"emb_dropout_prob": 0.0,
"fp16": false,
"fp32": false,
"hidden_size": 4096,
"his_len": 4,
"initializer_range": 0.02,
"intermediate_size": 22016,
"kv_channels": 128,
"layer_norm_epsilon": 1e-06,
"max_position_embeddings": 8192,
"model_type": "qwen",
"no_bias": true,
"num_attention_heads": 32,
"num_hidden_layers": 32,
"onnx_safe": null,
"rotary_emb_base": 10000,
"rotary_pct": 1.0,
"scale_attn_weights": true,
"seq_length": 2048,
"tie_word_embeddings": false,
"tokenizer_type": "QWenTokenizer",
"torch_dtype": "bfloat16",
"transformers_version": "4.50.0",
"use_cache": true,
"use_dynamic_ntk": true,
"use_flash_attn": false,
"use_logn_attn": true,
"visual": {
"heads": 16,
"image_size": 448,
"image_start_id": 151857,
"layers": 48,
"mlp_ratio": 4.9231,
"output_dim": 4096,
"patch_size": 14,
"width": 1664
},
"vocab_size": 151936
}
```
## /eval/utils/utils_odyssey/configuration_qwen.py
```py path="/eval/utils/utils_odyssey/configuration_qwen.py"
# Copyright (c) Alibaba Cloud.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
from transformers import PretrainedConfig
class QWenConfig(PretrainedConfig):
model_type = "qwen"
keys_to_ignore_at_inference = ["past_key_values"]
def __init__(
self,
vocab_size=151936,
hidden_size=4096,
num_hidden_layers=32,
num_attention_heads=32,
emb_dropout_prob=0.0,
attn_dropout_prob=0.0,
layer_norm_epsilon=1e-6,
initializer_range=0.02,
max_position_embeddings=8192,
scale_attn_weights=True,
use_cache=True,
bf16=False,
fp16=False,
fp32=False,
kv_channels=128,
rotary_pct=1.0,
rotary_emb_base=10000,
use_dynamic_ntk=True,
use_logn_attn=True,
use_flash_attn="auto",
intermediate_size=22016,
no_bias=True,
tie_word_embeddings=False,
**kwargs,
):
self.vocab_size = vocab_size
self.hidden_size = hidden_size
self.intermediate_size = intermediate_size
self.num_hidden_layers = num_hidden_layers
self.num_attention_heads = num_attention_heads
self.emb_dropout_prob = emb_dropout_prob
self.attn_dropout_prob = attn_dropout_prob
self.layer_norm_epsilon = layer_norm_epsilon
self.initializer_range = initializer_range
self.scale_attn_weights = scale_attn_weights
self.use_cache = use_cache
self.max_position_embeddings = max_position_embeddings
self.bf16 = bf16
self.fp16 = fp16
self.fp32 = fp32
self.kv_channels = kv_channels
self.rotary_pct = rotary_pct
self.rotary_emb_base = rotary_emb_base
self.use_dynamic_ntk = use_dynamic_ntk
self.use_logn_attn = use_logn_attn
self.use_flash_attn = use_flash_attn
self.no_bias = no_bias
super().__init__(
tie_word_embeddings=tie_word_embeddings,
**kwargs
)
```
## /eval/utils/utils_odyssey/generation_config.json
```json path="/eval/utils/utils_odyssey/generation_config.json"
{
"_from_model_config": true,
"transformers_version": "4.50.0"
}
```
## /eval/utils/utils_odyssey/qwen_generation_utils.py
```py path="/eval/utils/utils_odyssey/qwen_generation_utils.py"
# Copyright (c) Alibaba Cloud.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
"""Generation support."""
from typing import Tuple, List, Union, Iterable
import numpy as np
import torch
import torch.nn.functional as F
from transformers import PreTrainedTokenizer
from transformers import logging
from transformers.generation import LogitsProcessor
logger = logging.get_logger(__name__)
# Types.
HistoryType = List[Tuple[str, str]]
TokensType = List[int]
BatchTokensType = List[List[int]]
def pad_batch(batch: BatchTokensType, pad_id: int, seq_length: int) -> BatchTokensType:
for tokens in batch:
context_length = len(tokens)
if context_length < seq_length:
tokens.extend([pad_id] * (seq_length - context_length))
return batch
def get_ltor_masks_and_position_ids(
data,
eod_token,
reset_position_ids,
reset_attention_mask,
eod_mask_loss,
):
"""Build masks and position id for left to right model."""
# Extract batch size and sequence length.
micro_batch_size, seq_length = data.size()
# Attention mask (lower triangular).
if reset_attention_mask:
att_mask_batch = micro_batch_size
else:
att_mask_batch = 1
attention_mask = torch.tril(
torch.ones((att_mask_batch, seq_length, seq_length), device=data.device)
).view(att_mask_batch, 1, seq_length, seq_length)
# Loss mask.
loss_mask = torch.ones(data.size(), dtype=torch.float, device=data.device)
if eod_mask_loss:
loss_mask[data == eod_token] = 0.0
# Position ids.
position_ids = torch.arange(seq_length, dtype=torch.long, device=data.device)
position_ids = position_ids.unsqueeze(0).expand_as(data)
# We need to clone as the ids will be modifed based on batch index.
if reset_position_ids:
position_ids = position_ids.clone()
if reset_position_ids or reset_attention_mask:
# Loop through the batches:
for b in range(micro_batch_size):
# Find indecies where EOD token is.
eod_index = position_ids[b, data[b] == eod_token]
# Detach indecies from positions if going to modify positions.
if reset_position_ids:
eod_index = eod_index.clone()
# Loop through EOD indecies:
prev_index = 0
for j in range(eod_index.size()[0]):
i = eod_index[j]
# Mask attention loss.
if reset_attention_mask:
attention_mask[b, 0, (i + 1) :, : (i + 1)] = 0
# Reset positions.
if reset_position_ids:
position_ids[b, (i + 1) :] -= i + 1 - prev_index
prev_index = i + 1
# Convert attention mask to binary:
attention_mask = attention_mask < 0.5
return attention_mask, loss_mask, position_ids
def get_batch(context_tokens: torch.LongTensor, eod_id: int):
"""Generate batch from context tokens."""
# Move to GPU.
tokens = context_tokens.contiguous().to(context_tokens.device)
# Get the attention mask and postition ids.
attention_mask, _, position_ids = get_ltor_masks_and_position_ids(
tokens,
eod_id,
reset_position_ids=False,
reset_attention_mask=False,
eod_mask_loss=False,
)
return tokens, attention_mask, position_ids
def get_stop_words_ids(chat_format, tokenizer):
if chat_format == "raw":
stop_words_ids = [tokenizer.encode("Human:"), [tokenizer.eod_id]]
elif chat_format == "chatml":
stop_words_ids = [[tokenizer.im_end_id], [tokenizer.im_start_id]]
else:
raise NotImplementedError(f"Unknown chat format {chat_format!r}")
return stop_words_ids
def make_context(
tokenizer: PreTrainedTokenizer,
query: str,
history: List[Tuple[str, str]] = None,
system: str = "",
max_window_size: int = 6144,
chat_format: str = "chatml",
):
if history is None:
history = []
if chat_format == "chatml":
im_start, im_end = "<|im_start|>", "<|im_end|>"
im_start_tokens = [tokenizer.im_start_id]
im_end_tokens = [tokenizer.im_end_id]
nl_tokens = tokenizer.encode("\n")
def _tokenize_str(role, content):
return f"{role}\n{content}", tokenizer.encode(
role, allowed_special=set(tokenizer.IMAGE_ST)
) + nl_tokens + tokenizer.encode(content, allowed_special=set(tokenizer.IMAGE_ST))
system_text, system_tokens_part = _tokenize_str("system", system)
system_tokens = im_start_tokens + system_tokens_part + im_end_tokens
raw_text = ""
context_tokens = []
for turn_query, turn_response in reversed(history):
query_text, query_tokens_part = _tokenize_str("user", turn_query)
query_tokens = im_start_tokens + query_tokens_part + im_end_tokens
if turn_response is not None:
response_text, response_tokens_part = _tokenize_str(
"assistant", turn_response
)
response_tokens = im_start_tokens + response_tokens_part + im_end_tokens
next_context_tokens = nl_tokens + query_tokens + nl_tokens + response_tokens
prev_chat = (
f"\n{im_start}{query_text}{im_end}\n{im_start}{response_text}{im_end}"
)
else:
next_context_tokens = nl_tokens + query_tokens + nl_tokens
prev_chat = f"\n{im_start}{query_text}{im_end}\n"
current_context_size = (
len(system_tokens) + len(next_context_tokens) + len(context_tokens)
)
if current_context_size < max_window_size:
context_tokens = next_context_tokens + context_tokens
raw_text = prev_chat + raw_text
else:
break
context_tokens = system_tokens + context_tokens
raw_text = f"{im_start}{system_text}{im_end}" + raw_text
context_tokens += (
nl_tokens
+ im_start_tokens
+ _tokenize_str("user", query)[1]
+ im_end_tokens
+ nl_tokens
+ im_start_tokens
+ tokenizer.encode("assistant")
+ nl_tokens
)
raw_text += f"\n{im_start}user\n{query}{im_end}\n{im_start}assistant\n"
elif chat_format == "raw":
raw_text = query
context_tokens = tokenizer.encode(raw_text)
else:
raise NotImplementedError(f"Unknown chat format {chat_format!r}")
return raw_text, context_tokens
def _decode_default(
tokens: List[int],
*,
stop_words: List[str],
eod_words: List[str],
tokenizer: PreTrainedTokenizer,
raw_text_len: int,
verbose: bool = False,
return_end_reason: bool = False,
errors: str='replace',
):
trim_decode_tokens = tokenizer.decode(tokens, errors=errors)[raw_text_len:]
if verbose:
print("\nRaw Generate: ", trim_decode_tokens)
end_reason = f"Gen length {len(tokens)}"
for stop_word in stop_words:
trim_decode_tokens = trim_decode_tokens.replace(stop_word, "").strip()
for eod_word in eod_words:
if eod_word in trim_decode_tokens:
end_reason = f"Gen {eod_word!r}"
trim_decode_tokens = trim_decode_tokens.split(eod_word)[0]
trim_decode_tokens = trim_decode_tokens.strip()
if verbose:
print("\nEnd Reason:", end_reason)
print("\nGenerate: ", trim_decode_tokens)
if return_end_reason:
return trim_decode_tokens, end_reason
else:
return trim_decode_tokens
def _decode_chatml(
tokens: List[int],
*,
stop_words: List[str],
eod_token_ids: List[int],
tokenizer: PreTrainedTokenizer,
raw_text_len: int,
context_length: int,
verbose: bool = False,
return_end_reason: bool = False,
errors: str='replace'
):
end_reason = f"Gen length {len(tokens)}"
eod_token_idx = context_length
for eod_token_idx in range(context_length, len(tokens)):
if tokens[eod_token_idx] in eod_token_ids:
end_reason = f"Gen {tokenizer.decode([tokens[eod_token_idx]])!r}"
break
trim_decode_tokens = tokenizer.decode(tokens[:eod_token_idx], errors=errors)[raw_text_len:]
if verbose:
print("\nRaw Generate w/o EOD:", tokenizer.decode(tokens, errors=errors)[raw_text_len:])
print("\nRaw Generate:", trim_decode_tokens)
print("\nEnd Reason:", end_reason)
for stop_word in stop_words:
trim_decode_tokens = trim_decode_tokens.replace(stop_word, "").strip()
trim_decode_tokens = trim_decode_tokens.strip()
if verbose:
print("\nGenerate:", trim_decode_tokens)
if return_end_reason:
return trim_decode_tokens, end_reason
else:
return trim_decode_tokens
def decode_tokens(
tokens: Union[torch.LongTensor, TokensType],
tokenizer: PreTrainedTokenizer,
raw_text_len: int,
context_length: int,
chat_format: str,
verbose: bool = False,
return_end_reason: bool = False,
errors: str="replace",
) -> str:
if torch.is_tensor(tokens):
tokens = tokens.cpu().numpy().tolist()
if chat_format == "chatml":
return _decode_chatml(
tokens,
stop_words=[],
eod_token_ids=[tokenizer.im_start_id, tokenizer.im_end_id],
tokenizer=tokenizer,
raw_text_len=raw_text_len,
context_length=context_length,
verbose=verbose,
return_end_reason=return_end_reason,
errors=errors,
)
elif chat_format == "raw":
return _decode_default(
tokens,
stop_words=["<|endoftext|>"],
eod_words=["<|endoftext|>"],
tokenizer=tokenizer,
raw_text_len=raw_text_len,
verbose=verbose,
return_end_reason=return_end_reason,
errors=errors,
)
else:
raise NotImplementedError(f"Unknown chat format {chat_format!r}")
class StopWordsLogitsProcessor(LogitsProcessor):
"""
:class:`transformers.LogitsProcessor` that enforces that when specified sequences appear, stop geration.
Args:
stop_words_ids (:obj:`List[List[int]]`):
List of list of token ids of stop ids. In order to get the tokens of the words
that should not appear in the generated text, use :obj:`tokenizer(bad_word,
add_prefix_space=True).input_ids`.
eos_token_id (:obj:`int`):
The id of the `end-of-sequence` token.
"""
def __init__(self, stop_words_ids: Iterable[Iterable[int]], eos_token_id: int):
if not isinstance(stop_words_ids, List) or len(stop_words_ids) == 0:
raise ValueError(
f"`stop_words_ids` has to be a non-emtpy list, but is {stop_words_ids}."
)
if any(not isinstance(bad_word_ids, list) for bad_word_ids in stop_words_ids):
raise ValueError(
f"`stop_words_ids` has to be a list of lists, but is {stop_words_ids}."
)
if any(
any(
(not isinstance(token_id, (int, np.integer)) or token_id < 0)
for token_id in stop_word_ids
)
for stop_word_ids in stop_words_ids
):
raise ValueError(
f"Each list in `stop_words_ids` has to be a list of positive integers, but is {stop_words_ids}."
)
self.stop_words_ids = list(
filter(
lambda bad_token_seq: bad_token_seq != [eos_token_id], stop_words_ids
)
)
self.eos_token_id = eos_token_id
for stop_token_seq in self.stop_words_ids:
assert (
len(stop_token_seq) > 0
), "Stop words token sequences {} cannot have an empty list".format(
stop_words_ids
)
def __call__(
self, input_ids: torch.LongTensor, scores: torch.FloatTensor
) -> torch.FloatTensor:
stopped_samples = self._calc_stopped_samples(input_ids)
for i, should_stop in enumerate(stopped_samples):
if should_stop:
scores[i, self.eos_token_id] = float(2**15)
return scores
def _tokens_match(self, prev_tokens: torch.LongTensor, tokens: List[int]) -> bool:
if len(tokens) == 0:
# if bad word tokens is just one token always ban it
return True
elif len(tokens) > len(prev_tokens):
# if bad word tokens are longer then prev input_ids they can't be equal
return False
elif prev_tokens[-len(tokens) :].tolist() == tokens:
# if tokens match
return True
else:
return False
def _calc_stopped_samples(self, prev_input_ids: Iterable[int]) -> Iterable[int]:
stopped_samples = []
for prev_input_ids_slice in prev_input_ids:
match = False
for stop_token_seq in self.stop_words_ids:
if self._tokens_match(prev_input_ids_slice, stop_token_seq):
# if tokens do not match continue
match = True
break
stopped_samples.append(match)
return stopped_samples
def top_k_logits(logits, top_k=0, top_p=0.0, filter_value=-float("Inf")):
"""This function has been mostly taken from huggingface conversational
ai code at
https://medium.com/huggingface/how-to-build-a-state-of-the-art-
conversational-ai-with-transfer-learning-2d818ac26313"""
if top_k > 0:
# Remove all tokens with a probability less than the
# last token of the top-k
indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None]
logits[indices_to_remove] = filter_value
if top_p > 0.0:
# Cconvert to 1D
sorted_logits, sorted_indices = torch.sort(logits, descending=True, dim=-1)
cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
# Remove tokens with cumulative probability above the threshold
sorted_indices_to_remove = cumulative_probs > top_p
# Shift the indices to the right to keep also the first token
# above the threshold
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
sorted_indices_to_remove[..., 0] = 0
for i in range(sorted_indices.size(0)):
indices_to_remove = sorted_indices[i][sorted_indices_to_remove[i]]
logits[i][indices_to_remove] = filter_value
return logits
def switch(val1, val2, boolean):
boolean = boolean.type_as(val1)
return (1 - boolean) * val1 + boolean * val2
```
## /eval/utils/utils_odyssey/special_tokens_map.json
```json path="/eval/utils/utils_odyssey/special_tokens_map.json"
{}
```
## /eval/utils/utils_odyssey/tokenizer_config.json
```json path="/eval/utils/utils_odyssey/tokenizer_config.json"
{
"added_tokens_decoder": {},
"auto_map": {
"AutoTokenizer": [
"Qwen/Qwen-VL-Chat--tokenization_qwen.QWenTokenizer",
null
]
},
"clean_up_tokenization_spaces": false,
"extra_special_tokens": {},
"model_max_length": 8192,
"tokenizer_class": "QWenTokenizer"
}
```
## /eval/utils/utils_odyssey/visual.py
```py path="/eval/utils/utils_odyssey/visual.py"
# Copyright (c) Alibaba Cloud.
#
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
from collections import OrderedDict
import math
import requests
from io import BytesIO
from functools import partial
from PIL import Image
from typing import Callable, Optional, Sequence, Tuple, List
import numpy as np
import torch
from torch import nn
from torch.nn import functional as F
from torch.nn.init import trunc_normal_
from torchvision import transforms
from torchvision.transforms import InterpolationMode
def get_abs_pos(abs_pos, tgt_size):
# abs_pos: L, C
# tgt_size: M
# return: M, C
src_size = int(math.sqrt(abs_pos.size(0)))
tgt_size = int(math.sqrt(tgt_size))
dtype = abs_pos.dtype
if src_size != tgt_size:
return F.interpolate(
abs_pos.float().reshape(1, src_size, src_size, -1).permute(0, 3, 1, 2),
size=(tgt_size, tgt_size),
mode="bicubic",
align_corners=False,
).permute(0, 2, 3, 1).flatten(0, 2).to(dtype=dtype)
else:
return abs_pos
# https://github.com/facebookresearch/mae/blob/efb2a8062c206524e35e47d04501ed4f544c0ae8/util/pos_embed.py#L20
def get_2d_sincos_pos_embed(embed_dim, grid_size, cls_token=False):
"""
grid_size: int of the grid height and width
return:
pos_embed: [grid_size*grid_size, embed_dim] or [1+grid_size*grid_size, embed_dim] (w/ or w/o cls_token)
"""
grid_h = np.arange(grid_size, dtype=np.float32)
grid_w = np.arange(grid_size, dtype=np.float32)
grid = np.meshgrid(grid_w, grid_h) # here w goes first
grid = np.stack(grid, axis=0)
grid = grid.reshape([2, 1, grid_size, grid_size])
pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
if cls_token:
pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], axis=0)
return pos_embed
def get_2d_sincos_pos_embed_from_grid(embed_dim, grid):
assert embed_dim % 2 == 0
# use half of dimensions to encode grid_h
emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
emb = np.concatenate([emb_h, emb_w], axis=1) # (H*W, D)
return emb
def get_1d_sincos_pos_embed_from_grid(embed_dim, pos):
"""
embed_dim: output dimension for each position
pos: a list of positions to be encoded: size (M,)
out: (M, D)
"""
assert embed_dim % 2 == 0
omega = np.arange(embed_dim // 2, dtype=np.float32)
omega /= embed_dim / 2.
omega = 1. / 10000**omega # (D/2,)
pos = pos.reshape(-1) # (M,)
out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product
emb_sin = np.sin(out) # (M, D/2)
emb_cos = np.cos(out) # (M, D/2)
emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D)
return emb
class Resampler(nn.Module):
"""
A 2D perceiver-resampler network with one cross attention layers by
(grid_size**2) learnable queries and 2d sincos pos_emb
Outputs:
A tensor with the shape of (grid_size**2, embed_dim)
"""
def __init__(
self,
grid_size,
embed_dim,
num_heads,
kv_dim=None,
norm_layer=nn.LayerNorm
):
super().__init__()
self.num_queries = grid_size ** 2
self.embed_dim = embed_dim
self.num_heads = num_heads
self.pos_embed = nn.Parameter(
torch.from_numpy(get_2d_sincos_pos_embed(embed_dim, grid_size)).float()
).requires_grad_(False)
self.query = nn.Parameter(torch.zeros(self.num_queries, embed_dim))
trunc_normal_(self.query, std=.02)
if kv_dim is not None and kv_dim != embed_dim:
self.kv_proj = nn.Linear(kv_dim, embed_dim, bias=False)
else:
self.kv_proj = nn.Identity()
self.attn = nn.MultiheadAttention(embed_dim, num_heads)
self.ln_q = norm_layer(embed_dim)
self.ln_kv = norm_layer(embed_dim)
# self.apply(self._init_weights)
def _init_weights(self, m):
if isinstance(m, nn.Linear):
trunc_normal_(m.weight, std=.02)
if isinstance(m, nn.Linear) and m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.LayerNorm):
nn.init.constant_(m.bias, 0)
nn.init.constant_(m.weight, 1.0)
def forward(self, x, attn_mask=None):
pos_embed = get_abs_pos(self.pos_embed, x.size(1))
x = self.kv_proj(x)
x = self.ln_kv(x).permute(1, 0, 2)
N = x.shape[1]
q = self.ln_q(self.query)
out = self.attn(
self._repeat(q, N) + self.pos_embed.unsqueeze(1),
x + pos_embed.unsqueeze(1),
x,
attn_mask=attn_mask)[0]
return out.permute(1, 0, 2)
def _repeat(self, query, N: int):
return query.unsqueeze(1).repeat(1, N, 1)
class VisualAttention(nn.Module):
"""self-attention layer class.
Self-attention layer takes input with size [s, b, h]
and returns output of the same size.
"""
def __init__(self, embed_dim, num_heads,
bias=True, kdim=None, vdim=None):
super(VisualAttention, self).__init__()
self.embed_dim = embed_dim
self.kdim = kdim if kdim is not None else embed_dim
self.vdim = vdim if vdim is not None else embed_dim
self._qkv_same_embed_dim = self.kdim == embed_dim and self.vdim == embed_dim
self.num_heads = num_heads
# Per attention head and per partition values.
assert embed_dim % num_heads == 0
self.hidden_size_per_attention_head = embed_dim // num_heads
self.num_attention_heads_per_partition = num_heads
self.hidden_size_per_partition = embed_dim
# Strided linear layer.
assert self._qkv_same_embed_dim, 'Only Support SelfAttention Currently'
self.in_proj = nn.Linear(embed_dim, 3 * embed_dim)
self.out_proj = nn.Linear(embed_dim, embed_dim)
self.norm_factor = math.sqrt(self.hidden_size_per_attention_head)
def forward(self, query, key, value, attn_mask = None):
# query/key/value: [sq, b, h]
sq, b, _ = query.size()
assert torch.allclose(query, key), 'Only Support Self-Attention Currently'
sk = sq
mixed_x_layer = self.in_proj(query)
# [sq, b, (np * 3 * hn)] --> [sq, b, np, 3 * hn]
new_tensor_shape = mixed_x_layer.size()[:-1] + \
(self.num_attention_heads_per_partition,
3 * self.hidden_size_per_attention_head)
mixed_x_layer = mixed_x_layer.view(*new_tensor_shape)
# [sq, b, np, 3 * hn] --> 3 [sq, b, np, hn]
query_layer, key_layer, value_layer = mixed_x_layer.split(
self.hidden_size_per_attention_head, dim=-1)
# [sq, b, np, hn] -> [sq, b * np, hn]
query_layer = query_layer.view(sq,
b * self.num_attention_heads_per_partition,
self.hidden_size_per_attention_head).transpose(0, 1)
# [sk, b, np, hn] -> [sk, b * np, hn]
key_layer = key_layer.view(sk,
b * self.num_attention_heads_per_partition,
self.hidden_size_per_attention_head).transpose(0, 1)
q_scaled = query_layer / self.norm_factor
if attn_mask is not None:
attention_probs = torch.baddbmm(attn_mask, q_scaled, key_layer.transpose(-2, -1))
else:
attention_probs = torch.bmm(q_scaled, key_layer.transpose(-2, -1))
attention_probs = attention_probs.softmax(dim=-1)
value_layer = value_layer.view(sk,
b * self.num_attention_heads_per_partition,
self.hidden_size_per_attention_head).transpose(0, 1)
# matmul: [b * np, sq, hn]
context_layer = torch.bmm(attention_probs, value_layer)
# change view [b, np, sq, hn]
context_layer = context_layer.view(b,
self.num_attention_heads_per_partition,
sq, self.hidden_size_per_attention_head)
# [b, np, sq, hn] --> [sq, b, np, hn]
context_layer = context_layer.permute(2, 0, 1, 3).contiguous()
# [sq, b, np, hn] --> [sq, b, hp]
new_context_layer_shape = context_layer.size()[:-2] + \
(self.hidden_size_per_partition,)
context_layer = context_layer.view(*new_context_layer_shape)
output = self.out_proj(context_layer)
return output
class VisualAttentionBlock(nn.Module):
def __init__(
self,
d_model: int,
n_head: int,
mlp_ratio: float = 4.0,
act_layer: Callable = nn.GELU,
norm_layer: Callable = nn.LayerNorm,
is_cross_attention: bool = False,
):
super().__init__()
self.ln_1 = norm_layer(d_model)
if is_cross_attention:
self.ln_1_kv = norm_layer(d_model)
self.ln_2 = norm_layer(d_model)
mlp_width = int(d_model * mlp_ratio)
self.attn = VisualAttention(d_model, n_head)
self.mlp = nn.Sequential(OrderedDict([
("c_fc", nn.Linear(d_model, mlp_width)),
("gelu", act_layer()),
("c_proj", nn.Linear(mlp_width, d_model))
]))
def attention(
self,
q_x: torch.Tensor,
k_x: Optional[torch.Tensor] = None,
v_x: Optional[torch.Tensor] = None,
attn_mask: Optional[torch.Tensor] = None,
):
k_x = k_x if k_x is not None else q_x
v_x = v_x if v_x is not None else q_x
attn_mask = attn_mask.to(q_x.dtype) if attn_mask is not None else None
return self.attn(q_x, k_x, v_x, attn_mask=attn_mask)
def forward(
self,
q_x: torch.Tensor,
k_x: Optional[torch.Tensor] = None,
v_x: Optional[torch.Tensor] = None,
attn_mask: Optional[torch.Tensor] = None,
):
k_x = self.ln_1_kv(k_x) if hasattr(self, "ln_1_kv") and k_x is not None else None
v_x = self.ln_1_kv(v_x) if hasattr(self, "ln_1_kv") and v_x is not None else None
x = q_x + self.attention(q_x=self.ln_1(q_x), k_x=k_x, v_x=v_x, attn_mask=attn_mask)
x = x + self.mlp(self.ln_2(x))
return x
class TransformerBlock(nn.Module):
def __init__(
self,
width: int,
layers: int,
heads: int,
mlp_ratio: float = 4.0,
act_layer: Callable = nn.GELU,
norm_layer: Callable = nn.LayerNorm,
):
super().__init__()
self.width = width
self.layers = layers
self.resblocks = nn.ModuleList([
VisualAttentionBlock(
width, heads, mlp_ratio, act_layer=act_layer, norm_layer=norm_layer)
for _ in range(layers)
])
def get_cast_dtype(self) -> torch.dtype:
return self.resblocks[0].mlp.c_fc.weight.dtype
def get_cast_device(self) -> torch.device:
return self.resblocks[0].mlp.c_fc.weight.device
def forward(self, x: torch.Tensor, attn_mask: Optional[torch.Tensor] = None):
for r in self.resblocks:
x = r(x, attn_mask=attn_mask)
return x
class VisionTransformer(nn.Module):
def __init__(
self,
image_size: int,
patch_size: int,
width: int,
layers: int,
heads: int,
mlp_ratio: float,
n_queries: int = 256,
output_dim: int = 512,
**kwargs
):
super().__init__()
image_height, image_width = self.image_size = (image_size, image_size)
patch_height, patch_width = self.patch_size = (patch_size, patch_size)
self.grid_size = (image_height // patch_height, image_width // patch_width)
self.output_dim = output_dim
mean = (0.48145466, 0.4578275, 0.40821073)
std = (0.26862954, 0.26130258, 0.27577711)
self.image_transform = transforms.Compose([
transforms.Resize(
(image_size, image_size),
interpolation=InterpolationMode.BICUBIC
),
transforms.ToTensor(),
transforms.Normalize(mean=mean, std=std),
])
self.conv1 = nn.Conv2d(in_channels=3, out_channels=width, kernel_size=patch_size, stride=patch_size, bias=False)
# class embeddings and positional embeddings
scale = width ** -0.5
self.positional_embedding = nn.Parameter(scale * torch.randn(256, width))
norm_layer = partial(nn.LayerNorm, eps=1e-6)
act_layer = nn.GELU
self.ln_pre = norm_layer(width)
self.transformer = TransformerBlock(
width,
layers,
heads,
mlp_ratio,
act_layer=act_layer,
norm_layer=norm_layer,
)
self.attn_pool = Resampler(
grid_size=int(math.sqrt(n_queries)),
embed_dim=output_dim,
num_heads=output_dim // 128,
kv_dim=width,
norm_layer=norm_layer,
)
self.ln_post = norm_layer(output_dim)
self.proj = nn.Parameter((output_dim** -0.5) * torch.randn(output_dim, output_dim))
def forward(self, x: torch.Tensor):
x = x.to(
dtype=self.transformer.get_cast_dtype(),
device=self.transformer.get_cast_device(),
)
# to patches
x = self.conv1(x) # shape = [*, width, grid, grid]
x = x.reshape(x.shape[0], x.shape[1], -1) # shape = [*, width, grid ** 2]
x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width]
x = x + get_abs_pos(self.positional_embedding, x.size(1))
x = self.ln_pre(x)
x = x.permute(1, 0, 2) # NLD -> LND
x = self.transformer(x)
x = x.permute(1, 0, 2) # LND -> NLD
x = self.attn_pool(x)
x = self.ln_post(x)
x = x @ self.proj
return x
def encode(self, image_paths: List[str]):
images = []
for image_path in image_paths:
if image_path.startswith("http://") or image_path.startswith("https://"):
image = Image.open(requests.get(image_path, stream=True).raw)
else:
image = Image.open(image_path)
image = image.convert("RGB")
images.append(self.image_transform(image))
images = torch.stack(images, dim=0)
return self(images)
```
## /image_hash/commands.txt
app1.jpg openapp
app2.jpg openapp
app3.jpg openapp
app4.jpg openapp
app5.png openapp
search1.jpg entersearch
search2.jpg entersearch
search3.jpg entersearch
search4.jpg entersearch
search5.jpg entersearch
search6.jpg entersearch
search7.jpg entersearch
search8.jpg entersearch
search9.jpg entersearch
search10.jpg entersearch
page1.jpg enterpage
page2.jpg enterpage
page3.jpg enterpage
page4.jpg enterpage
page5.jpg enterpage
page6.jpg enterpage
page7.jpg enterpage
page8.jpg enterpage
## /image_hash/image/app1.jpg
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/image_hash/image/app1.jpg
## /image_hash/image/app2.jpg
Binary file available at https://raw.githubusercontent.com/OpenBMB/AppCopilot/refs/heads/main/image_hash/image/app2.jpg
## /omni_parser/__init__.py
```py path="/omni_parser/__init__.py"
```
The content has been capped at 50000 tokens. The user could consider applying other filters to refine the result. The better and more specific the context, the better the LLM can follow instructions. If the context seems verbose, the user can refine the filter using uithub. Thank you for using https://uithub.com - Perfect LLM context for any GitHub repo.