Image-Text-to-Text
PEFT
Safetensors
android
android-world
ui-automation
accessibility
vision-language
lora
reinforcement-learning
Instructions to use dmitchelljackson/cerebellum-qwen35-history-actions-lora with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use dmitchelljackson/cerebellum-qwen35-history-actions-lora with PEFT:
from peft import PeftModel from transformers import AutoModelForCausalLM base_model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen3.5-0.8B") model = PeftModel.from_pretrained(base_model, "dmitchelljackson/cerebellum-qwen35-history-actions-lora") - Notebooks
- Google Colab
- Kaggle
Download harness_code/rl_harness/apk_state_collector.py from dmitchelljackson/cerebellum-qwen35-history-actions-lora: direct link, hf CLI and curl.
- Browser
- Download file 9.84 kB
-
https://huggingface.co/dmitchelljackson/cerebellum-qwen35-history-actions-lora/resolve/main/harness_code/rl_harness/apk_state_collector.py
- Command line
-
hf download hf://dmitchelljackson/cerebellum-qwen35-history-actions-lora/harness_code/rl_harness/apk_state_collector.py
-
curl -L -o apk_state_collector.py https://huggingface.co/dmitchelljackson/cerebellum-qwen35-history-actions-lora/resolve/main/harness_code/rl_harness/apk_state_collector.py
9.84 kB
| from __future__ import annotations | |
| import base64 | |
| import json | |
| import subprocess | |
| import time | |
| import urllib.request | |
| from io import BytesIO | |
| from typing import Any | |
| from PIL import Image | |
| from data.datasets.android_control_a11y import A11yNode, A11yWindow | |
| from data.datasets.android_control_a11y import compress_windowed_tree, expand_range_targets, get_action_targets | |
| from rl_harness.types import HarnessFrame | |
| class ApkStateCollector: | |
| """Model-facing state source from the Cerebellum Android collector APK.""" | |
| def __init__( | |
| self, | |
| adb_path: str = "/opt/android/platform-tools/adb", | |
| serial: str | None = None, | |
| host: str = "127.0.0.1", | |
| port: int = 8765, | |
| device_port: int = 8765, | |
| max_tree_chars: int = 5000, | |
| retries: int = 6, | |
| retry_sleep: float = 0.5, | |
| ): | |
| self.adb_path = adb_path | |
| self.serial = serial | |
| self.host = host | |
| self.port = port | |
| self.device_port = device_port | |
| self.max_tree_chars = max_tree_chars | |
| self.retries = retries | |
| self.retry_sleep = retry_sleep | |
| def _adb_cmd(self, *args: str) -> list[str]: | |
| cmd = [self.adb_path] | |
| if self.serial: | |
| cmd.extend(["-s", self.serial]) | |
| else: | |
| cmd.append("-e") | |
| cmd.extend(args) | |
| return cmd | |
| def ensure_forward(self) -> None: | |
| subprocess.run( | |
| self._adb_cmd("forward", f"tcp:{self.port}", f"tcp:{self.device_port}"), | |
| stdout=subprocess.PIPE, | |
| stderr=subprocess.PIPE, | |
| check=False, | |
| timeout=10.0, | |
| ) | |
| def restart_service(self) -> None: | |
| service = "com.cerebellum.collector/.CollectorAccessibilityService" | |
| current = subprocess.run( | |
| self._adb_cmd("shell", "settings", "get", "secure", "enabled_accessibility_services"), | |
| stdout=subprocess.PIPE, | |
| stderr=subprocess.PIPE, | |
| text=True, | |
| check=False, | |
| timeout=10.0, | |
| ).stdout.strip() | |
| services = [s for s in current.split(":") if s and s != "null"] | |
| if service not in services: | |
| services.append(service) | |
| subprocess.run( | |
| self._adb_cmd( | |
| "shell", | |
| "settings", | |
| "put", | |
| "secure", | |
| "enabled_accessibility_services", | |
| ":".join(services), | |
| ), | |
| stdout=subprocess.PIPE, | |
| stderr=subprocess.PIPE, | |
| check=False, | |
| timeout=10.0, | |
| ) | |
| subprocess.run( | |
| self._adb_cmd("shell", "settings", "put", "secure", "accessibility_enabled", "1"), | |
| stdout=subprocess.PIPE, | |
| stderr=subprocess.PIPE, | |
| check=False, | |
| timeout=10.0, | |
| ) | |
| self.ensure_forward() | |
| def _fetch_frame(self) -> dict[str, Any]: | |
| self.ensure_forward() | |
| with urllib.request.urlopen(f"http://{self.host}:{self.port}/frame", timeout=5.0) as resp: | |
| return json.loads(resp.read().decode("utf-8")) | |
| def _decode_image(self, row: dict[str, Any]) -> Image.Image: | |
| b64 = row.get("screenshot_png_b64") | |
| if not b64: | |
| raise RuntimeError("collector APK returned no screenshot") | |
| return Image.open(BytesIO(base64.b64decode(b64))).convert("RGB") | |
| def _parse_node(self, row: dict[str, Any], nodes: list[A11yNode]) -> A11yNode: | |
| bounds = row.get("bounds") | |
| node = A11yNode() | |
| node.node_id = len(nodes) | |
| node.bounds = tuple(int(x) for x in bounds) if bounds and len(bounds) == 4 else None | |
| node.text = row.get("text") or "" | |
| node.content_desc = row.get("content_desc") or "" | |
| node.resource_id = row.get("view_id") or "" | |
| node.class_name = row.get("class_name") or "" | |
| node.visible = bool(row.get("visible", True)) | |
| node.is_enabled = bool(row.get("enabled", True)) | |
| node.is_focused = bool(row.get("focused", False)) | |
| node.is_editable = bool(row.get("editable", False)) | |
| node.is_scrollable = bool(row.get("scrollable", False)) | |
| node.is_clickable = bool(row.get("clickable", False)) | |
| node.is_long_clickable = bool(row.get("long_clickable", False)) | |
| node.is_checkable = bool(row.get("checkable", False)) | |
| node.is_checked = bool(row.get("checked", False)) | |
| node.is_selected = bool(row.get("selected", False)) | |
| node.is_password = bool(row.get("password", False)) | |
| range_info = row.get("range") or {} | |
| if range_info: | |
| node.range_current = range_info.get("current") | |
| node.range_min = range_info.get("min") | |
| node.range_max = range_info.get("max") | |
| node.range_type = range_info.get("type") | |
| node.action_ids = [int(x) for x in row.get("actions") or [] if isinstance(x, int)] | |
| # SFT parser sets these from the accessibility action list. The APK | |
| # service exposes the equivalent high-level flags; use them as the | |
| # action-space source so get_action_targets() sees the same fields. | |
| node.clickable = node.is_clickable | |
| node.long_clickable = node.is_long_clickable | |
| node.scrollable = node.is_scrollable | |
| nodes.append(node) | |
| for child in row.get("children") or []: | |
| child_node = self._parse_node(child, nodes) | |
| node.child_ids.append(child_node.node_id) | |
| node.children.append(child_node) | |
| return node | |
| def _package_candidates(self, row: dict[str, Any]) -> set[str]: | |
| packages = set() | |
| def walk(node: dict[str, Any]) -> None: | |
| package_name = node.get("package") or node.get("package_name") | |
| if package_name: | |
| packages.add(str(package_name)) | |
| resource_id = node.get("view_id") or "" | |
| if ":id/" in resource_id: | |
| packages.add(resource_id.split(":id/", 1)[0]) | |
| for child in node.get("children") or []: | |
| walk(child) | |
| for window in row.get("windows") or []: | |
| root = window.get("root") | |
| if root: | |
| walk(root) | |
| return packages | |
| def _windows(self, row: dict[str, Any]) -> list[A11yWindow]: | |
| windows: list[A11yWindow] = [] | |
| for window in row.get("windows") or []: | |
| nodes: list[A11yNode] = [] | |
| root = window.get("root") | |
| if root: | |
| self._parse_node(root, nodes) | |
| windows.append( | |
| A11yWindow( | |
| idx=int(window.get("id", len(windows))), | |
| title=window.get("title") or "", | |
| type_=window.get("type"), | |
| layer=window.get("layer"), | |
| active=window.get("active"), | |
| focused=window.get("focused"), | |
| nodes=nodes, | |
| ) | |
| ) | |
| return windows | |
| def capture_frame(self) -> HarnessFrame: | |
| last_error: Exception | None = None | |
| for attempt in range(self.retries): | |
| try: | |
| row = self._fetch_frame() | |
| image = self._decode_image(row) | |
| windows = self._windows(row) | |
| expand_range_targets(windows, zones=5) | |
| tappable = get_action_targets(windows) | |
| tree = compress_windowed_tree(windows, tappable) | |
| if self.max_tree_chars and len(tree) > self.max_tree_chars: | |
| tree = tree[: self.max_tree_chars].rstrip() + "\n... tree truncated ..." | |
| return HarnessFrame( | |
| image=image, | |
| tappable=tappable, | |
| compressed_tree=tree, | |
| packages=sorted(self._package_candidates(row)), | |
| ) | |
| except Exception as exc: | |
| last_error = exc | |
| if attempt >= 1: | |
| self.restart_service() | |
| time.sleep(self.retry_sleep) | |
| raise RuntimeError(f"APK state capture failed: {last_error}") from last_error | |
| def _frame_signature(frame: HarnessFrame) -> tuple: | |
| targets = [] | |
| for node in frame.tappable: | |
| targets.append(( | |
| getattr(node, "bounds", None), | |
| getattr(node, "text", "") or "", | |
| getattr(node, "content_desc", "") or "", | |
| getattr(node, "resource_id", "") or "", | |
| getattr(node, "class_name", "") or "", | |
| )) | |
| return (frame.image.size, frame.compressed_tree, tuple(targets)) | |
| def capture_stable_frame( | |
| self, | |
| max_wait_s: float = 1.0, | |
| interval_s: float = 0.2, | |
| min_captures: int = 2, | |
| min_wait_s: float = 0.5, | |
| ) -> HarnessFrame: | |
| """Capture after the APK tree has stopped changing. | |
| Accessibility nodes often arrive a few hundred milliseconds after the | |
| screenshot, especially immediately after opening an app. Returning that | |
| partial tree gives the model a different label space from SFT. This | |
| bounded poll waits for the model-facing frame signature to repeat. | |
| """ | |
| start = time.time() | |
| deadline = start + max(0.0, max_wait_s) | |
| previous_sig = None | |
| last_frame = None | |
| captures = 0 | |
| while True: | |
| frame = self.capture_frame() | |
| captures += 1 | |
| sig = self._frame_signature(frame) | |
| if ( | |
| captures >= min_captures | |
| and previous_sig == sig | |
| and time.time() - start >= max(0.0, min_wait_s) | |
| ): | |
| return frame | |
| previous_sig = sig | |
| last_frame = frame | |
| if time.time() >= deadline: | |
| return last_frame | |
| time.sleep(max(0.01, interval_s)) | |