MCPcopy Create free account
hub / github.com/Kwai-Klear/AR-GRPO

github.com/Kwai-Klear/AR-GRPO @main

Chat with this repo
repository ↗ · DeepWiki ↗ · + Follow
773 symbols 2,772 edges 84 files 106 documented · 14% updated 9mo ago★ 53

Browse by type

Functions 650 Types & classes 121 Endpoints 2
What it actually does AI analysis from the code graph — generated when you open this
loading…
README

AR-GRPO: Training Autoregressive Image Generation Models via Reinforcement Learning

Last Commit Maintenance Contributing

[arXiv] | [Codes] | [🤗 Huggingface]

Shihao Yuan, Yahui Liu♥, Yang Yue, Jingyuan Zhang, Wangmeng Zuo♥, Qi Wang, Fuzheng Zhang, Guorui Zhou

♥ Corresponding author

Abstract

Inspired by the success of reinforcement learning (RL) in refining large language models (LLMs), we propose AR-GRPO, an approach to integrate online RL training into autoregressive (AR) image generation models. We adapt the Group Relative Policy Optimization (GRPO) algorithm to refine the vanilla autoregressive models’ outputs by carefully designed reward functions that evaluate generated images across multiple quality dimensions, including perceptual quality, realism, and semantic fidelity. We conduct comprehensive experiments on both class-conditional (i.e., class-to-image) and text-conditional (i.e., text-to-image) image generation tasks, demonstrating that our RL-enhanced framework significantly improves both the image quality and human preference of generated images compared to the standard AR baselines. Our results show consistent improvements across various evaluation metrics, establishing the viability of RL-based optimization for AR image generation and opening new avenues for controllable and high-quality image synthesis.

Installation

Create a new Python environment with conda,

conda env create -f conda_environment.yaml

If you wish to install dependencies manually using pip, use following commands,

conda create -n ar-grpo python==3.10
conda activate ar-grpo
pip install -r requirements.txt

Evaluation

C2I evaluation

C2I evaluation with FID and other metrics is from OpenAI. For environment requirments, Please refer to requiremtes.txt. We recommand creating a new environment for evaluation to prevent conflicts of dependencies. Also, you should download ImageNet reference batch 256x256 before evaluation.

After environment installation, you can use test.sh for C2I inference on ImageNet.

bash test.sh MODEL_NAME GPU_RANK STEP CFG

For example,

bash test.sh my_c2i_model 0 100 2.0

You would find a .npz file in test_results/, then use fid_evaluation.py to collect metrics.

python fid_evaluation.py IMAGENET_REF_PATH/VIRTUAL_imagenet256_labeled.npz test_results/YOUR_RESULTS.npz

T2I evaluation

GenEval

Create a new environment and download pretrained models according to GenEval. You may need to change pretrained model path in benchmark/geneval/eval.sh to your downloaded model path. Then run test_t2i.sh to generate images from GenEval test prompts.

bash test_t2i.sh MODEL_NAME GPU_RANK STEP CFG

Then run eval.sh in GenEval to collect results.

cd ./benchmark/geneval
bash eval.sh test_results/GENERATED_IMAGES GPR_RANK

Other Metrics

For metrics like CLIP Score, HPSv2 and others, we implement them in a form of reward class in img_gen_grpo_rewards.py, so they can be used both in training and evaluation. For CLIP Score, HPSv2, PickScore, ImageReward and Aesthetic Score, their dependencies are included in requirements.txt already, so no other operations needed. For DeQA, please refer to reward-server to set a environment and deploy a backend server and modify address and port in reward_utils/deqa_client.py. For UniReward, please create a new sglang environment and start a server as below,

conda create -n sglang python=3.10
conda activate sglang
pip install "sglang[all]"

python -m sglang.launch_server --model-path CodeGoat24/UnifiedReward-7b-v1.5 --api-key unireward --port 17140 --chat-template chatml-llava --enable-p2p-check --mem-fraction-static 0.5

If you only want to test certain scores, you could modify test_t2i_rewards.py to remove unwanted metrics. Once you have all metrics you want, run following commands to generate images from DrawBench prompts.

bash test_t2i_rewards_drawbench.sh MODEL_NAME GPU_RANK STEP CFG

The results should be printed to the screen after evaluation is complete.

Training

Data Preparation

For T2I, the training prompts are already in dataset/coco_captions_30000.json. For C2I, please download ImageNet, and change the data_rt in training scripts.

GRPO Training

Download pretrained models from LlamaGen if you want to train your own model. Reward models requires openai/clip-vit-large-patch14 and Qwen/Qwen2.5-VL-3B-Instruct. Usually they should be downloaded automaticly from huggingface.co. If the auto download fails, you could download them manaully and change the path for corresponding rewards in reward_utils.

Change any parameters as you wish in training scripts. For C2I training, run following command,

bash run_c2i_train.sh

For T2I training, run following command,

bash run_t2i_train.sh

Citation

@article{yuan2025argrpo,
 title={AR-GRPO: Training Autoregressive Image Generation Models via Reinforcement Learning},
 author={Yuan, Shihao and Liu, Yahui and Yue, Yang and Zhang, Jingyuan and Zuo, Wangmeng and Wang, Qi and Zhang, Fuzheng and Zhou, Guorui},
 journal={arXiv preprint arXiv:2508.06924},
 year={2025}
}

Core symbols most depended-on inside this repo

browse all functions →

Shape

Method 393
Function 257
Class 121
Route 2

Languages

Python100%

Modules by API surface

img_gen_grpo_rewards.py77 symbols
fid_evaluation.py51 symbols
autoregressive/models/gpt.py47 symbols
tokenizer_lg/tokenizer_image/vq_model.py35 symbols
autoregressive/serve/gpt_model.py35 symbols
autoregressive/serve/model_runner.py33 symbols
modified_grpo_trainer.py32 symbols
autoregressive/serve/llm_engine.py24 symbols
autoregressive/modeling.py22 symbols
tokenizer_lg/vqgan/layer.py20 symbols
tokenizer_lg/tokenizer_image/discriminator.py20 symbols
autoregressive/serve/sampler.py20 symbols

For agents

$ claude mcp add AR-GRPO \
  -- python -m otcore.mcp_server <graph>

⬇ download graph artifact

Ask about this repo answers extend the page