dmitchelljackson's picture
Upload Qwen3.5 AndroidWorld RL harness milestone 2026-06-09
b7df156 verified
Raw History Blame Contribute Delete
4.1 kB
from __future__ import annotations
import random
from collections import Counter
from data.datasets.android_control_a11y import ACTION_TYPES, AndroidControlA11yDataset
from rl_harness.types import PolicyPrediction, Reward
from scripts.train_qwen35_history_actions import ACTION_TYPES_BY_GROUP, action_code, group_for_sample
class OfflineA11yBanditEnv:
"""Single-step reward harness over held-out Android Control samples.
This is the cheap preflight path for RL plumbing. It validates prompt assembly,
constrained decode, action parsing, label lookup, reward accounting, and logs.
It does not claim trajectory credit because the dataset cannot branch after a
model action.
"""
def __init__(
self,
shards: list[int],
history_len: int = 4,
min_history: int = 4,
max_raw_tree_chars: int = 5000,
quota: str = "tp=5,k=1,scroll=3,wait=2,system=1",
seed: int = 1,
):
self.rng = random.Random(seed)
random.seed(seed)
excluded = set(range(20)) - set(shards)
self.ds = AndroidControlA11yDataset(history_len=history_len, start_shard=min(shards), excluded_shards=excluded)
self.min_history = min_history
self.max_raw_tree_chars = max_raw_tree_chars
self.quota = self._parse_quota(quota)
self.targets = []
for group, count in self.quota.items():
self.targets.extend([group] * count)
self._target_i = 0
@staticmethod
def _parse_quota(text: str) -> Counter:
out = Counter()
for part in text.split(","):
key, value = part.split("=", 1)
out[key.strip()] = int(value)
return out
def sample(self) -> dict:
while True:
if self._target_i % max(len(self.targets), 1) == 0:
self.rng.shuffle(self.targets)
wanted = self.targets[self._target_i % len(self.targets)]
self._target_i += 1
sample = self.ds.sample_random(allowed_action_types=ACTION_TYPES_BY_GROUP[wanted])
if len(sample.get("history", [])) < self.min_history:
continue
if self.max_raw_tree_chars and len(sample.get("compressed_tree", "")) > self.max_raw_tree_chars:
continue
if group_for_sample(sample) != wanted:
continue
return sample
@staticmethod
def _score_strings(gt: str, pred: str, group: str) -> tuple[bool, bool, str]:
gt = gt.strip()
pred = pred.strip()
gt_action = gt.split(maxsplit=1)[0]
pred_action = pred.split(maxsplit=1)[0] if pred else ""
action_ok = pred_action == gt_action
if not action_ok:
return False, action_ok, "wrong_action"
if group == "k":
if "<eos>" not in pred:
return False, action_ok, "missing_eos"
pred_text = pred.split("<eos>", 1)[0].strip().lower()
gt_text = gt.lower()
return pred_text == gt_text, action_ok, "wrong_value"
return pred == gt, action_ok, "wrong_value"
def reward(self, sample: dict, prediction: PolicyPrediction, label_map: dict[int, str]) -> Reward:
gt = action_code(sample, label_map)
if gt is None:
return Reward(0.0, True, "unscoreable")
pred = prediction.raw.strip()
group = group_for_sample(sample)
exact, action_ok, miss_reason = self._score_strings(gt, pred, group)
if exact:
value = 1.0
reason = "exact"
elif action_ok:
value = 0.25
reason = f"right_action_{miss_reason}"
else:
value = 0.0
reason = miss_reason
return Reward(
value=value,
done=True,
reason=reason,
details={
"gt": gt,
"pred": pred,
"group": group,
"act_type": ACTION_TYPES[sample["act_type_idx"]],
"action_ok": action_ok,
"exact": exact,
},
)