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 "" not in pred: return False, action_ok, "missing_eos" pred_text = pred.split("", 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, }, )