""" GUI-Shift GRPO Training Script Based on: GUI-Shift (arxiv 2505.12493) + VLM-R1 framework Trains Qwen2.5-VL-7B on K-step GUI Transition task with rule-based rewards. """ import os import re import json import pathlib from dataclasses import dataclass, field from typing import Optional from datasets import Dataset from transformers import Qwen2_5_VLForConditionalGeneration, AutoProcessor from trl import GRPOTrainer, GRPOConfig from trl import ModelConfig, ScriptArguments, TrlParser, get_peft_config # ----------------------- Monkey-patch for Qwen2.5-VL flash attention ----------------------- from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import Qwen2_5_VLVisionFlashAttention2, apply_rotary_pos_emb_flashatt, flash_attn_varlen_func import torch from transformers.utils import logging logger = logging.get_logger(__name__) def custom_forward( self, hidden_states: torch.Tensor, cu_seqlens: torch.Tensor, rotary_pos_emb: Optional[torch.Tensor] = None, position_embeddings: Optional[tuple] = None, ) -> torch.Tensor: seq_length = hidden_states.shape[0] q, k, v = self.qkv(hidden_states).reshape(seq_length, 3, self.num_heads, -1).permute(1, 0, 2, 3).unbind(0) if position_embeddings is None: logger.warning_once( "The attention layers in this model are transitioning from computing the RoPE embeddings internally " "through `rotary_pos_emb` (2D tensor of RoPE theta values), to using externally computed " "`position_embeddings` (Tuple of tensors, containing cos and sin). In v4.54 `rotary_pos_emb` will be " "removed and `position_embeddings` will be mandatory." ) emb = torch.cat((rotary_pos_emb, rotary_pos_emb), dim=-1) cos = emb.cos().float() sin = emb.sin().float() else: cos, sin = position_embeddings cos = cos.to(torch.float) sin = sin.to(torch.float) q, k = apply_rotary_pos_emb_flashatt(q.unsqueeze(0), k.unsqueeze(0), cos, sin) q = q.squeeze(0) k = k.squeeze(0) max_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).max().item() attn_output = flash_attn_varlen_func(q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen).reshape(seq_length, -1) attn_output = self.proj(attn_output) return attn_output Qwen2_5_VLVisionFlashAttention2.forward = custom_forward # ----------------------- Config ----------------------- @dataclass class GRPOScriptArguments(ScriptArguments): data_file_paths: str = field(default=None, metadata={"help": "Paths to data files, separated by ':'"}) image_folders: str = field(default=None, metadata={"help": "Paths to image folders, separated by ':'"}) val_split_ratio: float = field(default=0.1, metadata={"help": "Ratio of validation split"}) reward_funcs: list[str] = field( default_factory=lambda: ["action_reward", "format_reward"], metadata={"help": "List of reward functions"}, ) max_pixels: Optional[int] = field(default=12845056, metadata={"help": "Maximum number of pixels for the image"}) min_pixels: Optional[int] = field(default=3136, metadata={"help": "Minimum number of pixels for the image"}) @dataclass class GRPOModelConfig(ModelConfig): freeze_vision_modules: bool = False # ----------------------- Reward Functions ----------------------- QUESTION_PROMPT = ( "{Question} Output the final answer in tags. " "Do not output any extra text." ) def format_reward(completions, **kwargs): pattern = r".*?" completion_contents = [completion[0]["content"] for completion in completions] matches = [re.search(pattern, content, re.DOTALL) is not None for content in completion_contents] return [1.0 if match else 0.0 for match in matches] def is_point_in_bbox(point_x, point_y, bbox): x1, y1, x2, y2 = bbox return x1 <= point_x <= x2 and y1 <= point_y <= y2 def action_reward(completions, solution, **kwargs): contents = [completion[0]["content"] for completion in completions] rewards = [] for content, sol in zip(contents, solution): reward = 0.0 try: content_match = re.search(r'(.*?)', content, re.DOTALL) answer = content_match.group(1).strip() if content_match else content.strip() try: gt_action = json.loads(sol) except Exception: rewards.append(0.0) continue try: json_match = re.search(r'```(?:json)?\s*(.*?)\s*```', answer, re.DOTALL) if json_match: answer = json_match.group(1) answer = answer.replace('\\n', '').strip() pred_action = json.loads(answer) if not isinstance(pred_action, dict): rewards.append(0.0) continue except Exception: rewards.append(0.0) continue if "action_type" not in pred_action: rewards.append(0.0) continue if pred_action["action_type"].lower() != gt_action["action_type"].lower(): rewards.append(0.0) continue action_type = pred_action["action_type"] if action_type == "click": if "x" not in pred_action or "y" not in pred_action: rewards.append(0.0) continue pred_x = float(pred_action["x"]) pred_y = float(pred_action["y"]) if "bbox" in gt_action: bbox = [float(x) for x in gt_action["bbox"]] if is_point_in_bbox(pred_x, pred_y, bbox): reward = 1.0 else: gt_x = float(gt_action["x"]) gt_y = float(gt_action["y"]) if pred_x == gt_x and pred_y == gt_y: reward = 1.0 elif action_type == "scroll": if "direction" in pred_action and pred_action["direction"].strip().lower() == gt_action["direction"].strip().lower(): reward = 1.0 elif action_type == "open_app": if "app_name" in pred_action and pred_action["app_name"].strip().lower() == gt_action["app_name"].strip().lower(): reward = 1.0 elif action_type == "navigate_back": reward = 1.0 elif action_type == "navigate_home": reward = 1.0 elif action_type == "wait": reward = 1.0 elif action_type == "input_text": if "text" in pred_action and pred_action["text"].strip().lower() == gt_action["text"].strip().lower(): reward = 1.0 elif action_type == "long_press": if "x" not in pred_action or "y" not in pred_action: rewards.append(0.0) continue pred_x = float(pred_action["x"]) pred_y = float(pred_action["y"]) if "bbox" in gt_action: bbox = [float(x) for x in gt_action["bbox"]] if is_point_in_bbox(pred_x, pred_y, bbox): reward = 1.0 else: gt_x = float(gt_action["x"]) gt_y = float(gt_action["y"]) if pred_x == gt_x and pred_y == gt_y: reward = 1.0 except Exception: reward = 0.0 rewards.append(reward) return rewards # ----------------------- Main ----------------------- def main(script_args, training_args, model_args): reward_funcs_registry = { "action_reward": action_reward, "format_reward": format_reward, } reward_funcs = [reward_funcs_registry[func] for func in script_args.reward_funcs] print("reward_funcs:", reward_funcs) data_files = script_args.data_file_paths.split(":") image_folders = script_args.image_folders.split(":") if len(data_files) != len(image_folders): raise ValueError("Number of data files must match number of image folders") all_data = [] for data_file, image_folder in zip(data_files, image_folders): with open(data_file, 'r') as f: for line in f: item = json.loads(line) if 'image' in item: if isinstance(item['image'], str): item['image_path'] = [os.path.join(image_folder, item['image'])] del item['image'] elif isinstance(item['image'], list): item['image_path'] = [os.path.join(image_folder, img) for img in item['image']] del item['image'] else: raise ValueError(f"Unsupported image type: {type(item['image'])}") item['problem'] = item['conversations'][0]['value'].replace('', '') action_data = item['conversations'][1]['value'] if isinstance(action_data, str): try: action_json = json.loads(action_data) item['solution'] = json.dumps(action_json) except: item['solution'] = action_data.replace('', '').replace('', '').strip() else: item['solution'] = json.dumps(action_data) del item['conversations'] all_data.append(item) dataset = Dataset.from_list(all_data) def make_conversation_from_jsonl(example): return { 'image_path': [p for p in example['image_path']], 'problem': example['problem'], 'solution': example['solution'], 'prompt': [{ 'role': 'user', 'content': [ *({'type': 'image', 'text': None} for _ in range(len(example['image_path']))), {'type': 'text', 'text': QUESTION_PROMPT.format(Question=example['problem'])} ] }] } dataset = dataset.map(make_conversation_from_jsonl, num_proc=8) splits = {'train': dataset} if script_args.val_split_ratio > 0: train_val_split = dataset.train_test_split(test_size=script_args.val_split_ratio) splits['train'] = train_val_split['train'] splits['validation'] = train_val_split['test'] processing_class = AutoProcessor.from_pretrained( model_args.model_name_or_path, max_pixels=script_args.max_pixels, min_pixels=script_args.min_pixels, ) model = Qwen2_5_VLForConditionalGeneration.from_pretrained( model_args.model_name_or_path, torch_dtype="bfloat16", attn_implementation="flash_attention_2", ) if model_args.freeze_vision_modules: print("Freezing vision modules...") for n, p in model.named_parameters(): if "visual" in n: p.requires_grad = False trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) print(f"Total trainable parameters: {trainable}") peft_config = get_peft_config(model_args) if peft_config is not None: from peft import get_peft_model model = get_peft_model(model, peft_config) trainer = GRPOTrainer( model=model, reward_funcs=reward_funcs, args=training_args, train_dataset=splits['train'], eval_dataset=splits.get('validation') if training_args.eval_strategy != "no" else None, processing_class=processing_class, ) if list(pathlib.Path(training_args.output_dir).glob("checkpoint-*")): trainer.train(resume_from_checkpoint=True) else: trainer.train() trainer.save_model(training_args.output_dir) if training_args.push_to_hub: trainer.push_to_hub() if __name__ == "__main__": parser = TrlParser((GRPOScriptArguments, GRPOConfig, GRPOModelConfig)) script_args, training_args, model_args = parser.parse_args_and_config() main(script_args, training_args, model_args)