2025-02-v2

unslothai/unsloth2025-02-v2Feb 20, 2025by danielhanchen

AI Summary

Introduces GRPO (Group Relative Policy Optimization) with significant memory efficiency improvements and long context support.

Key Highlights

  • 90% less memory usage for GRPO compared to TRL + FA2
  • Long context support (20K context)
  • vLLM fast inference integration
  • Custom reward functions support

New Features

  • GRPO implementation
  • Long context support
  • vLLM integration
  • Memory efficient RL

Full Release Notes

# 90% less memory usage GRPO
## Update Unsloth via `pip install --upgrade --no-cache-dir unsloth unsloth_zoo`

More details in blog post: https://unsloth.ai/blog/grpo

Llama 3.1 8B GRPO Colab: https://colab.research.google.com/github/unslothai/notebooks/blob/main/nb/Llama3.1_(8B)-GRPO.ipynb

Metric | Unsloth | TRL + FA2
-- | -- | --
Training Memory Cost (GB) | 42GB | 414GB
GRPO Memory Cost (GB) | 9.8GB | 78.3GB
Inference Cost (GB) | 0GB | 16GB
Inference KV Cache for 20K context (GB) | 2.5GB | 2.5GB
Total Memory Usage | 54.3GB (90% less) | 510.8GB

You automatically get 90% less memory usage! Also all reward logs for individual reward functions will show up.
![Screenshot_2025-02-20_at_04-52-52_Copy_of_Yet_another_copy_of_Llama3 1_(8B)-GRPO ipynb_-_Colab_5lpAL05rCEjw67tij45ua](https://github.com/user-attachments/assets/3d16d5a7-5f15-41a4-9d14-461b1112f29a)

Script to run GRPO:
```python
!pip install unsloth vllm
from unsloth import FastLanguageModel
import torch
max_seq_length = 1024 # Can increase for longer reasoning traces
lora_rank = 32 # Larger rank = smarter, but slower

model, tokenizer = FastLanguageModel.from_pretrained(
    model_name = "meta-llama/meta-Llama-3.1-8B-Instruct",
    max_seq_length = max_seq_length,
    load_in_4bit = True, # False for LoRA 16bit
    fast_inference = True, # Enable vLLM fast inference
    max_lora_rank = lora_rank,
    gpu_memory_utilization = 0.6, # Reduce if out of memory
)

model = FastLanguageModel.get_peft_model(
    model,
    r = lora_rank, # Choose any number > 0 ! Suggested 8, 16, 32, 64, 128
    target_modules = [
        "q_proj", "k_proj", "v_proj", "o_proj",
        "gate_proj", "up_proj", "down_proj",
    ], # Remove QKVO if out of memory
    lora_alpha = lora_rank,
    use_gradient_checkpointing = "unsloth", # Enable long context finetuning
    random_state = 3407,
)
import re
from datasets import load_dataset, Dataset
global COUNTER
COUNTER = 0
global PRINT_EVERY
PRINT_EVERY = 20

# Load and prep dataset
SYSTEM_PROMPT = """
Respond in the following format:
<reasoning>
...
</reasoning>
<answer>
...
</answer>
"""

XML_COT_FORMAT = """\
<reasoning>
{reasoning}
</reasoning>
<answer>
{answer}
</answer>
"""

def extract_xml_answer(text: str) -> str:
    answer = text.split("<answer>")[-1]
    answer = answer.split("</answer>")[0]
    return answer.strip()

def extract_hash_answer(text: str) -> str | None:
    if "####" not in text:
        return None
    return text.split("####")[1].strip()

# uncomment middle messages for 1-shot prompting
def get_gsm8k_questions(split = "train") -> Dataset:
    data = load_dataset('openai/gsm8k', 'main')[split] # type: ignore
    data = data.map(lambda x: { # type: ignore
        'prompt': [
            {'role': 'system', 'content': SYSTEM_PROMPT},
            {'role': 'user', 'content': x['question']}
        ],
        'answer': extract_hash_answer(x['answer'])
    }) # type: ignore
    return data # type: ignore

dataset = get_gsm8k_questions()

# Reward functions
def correctness_reward_func(prompts, completions, answer, **kwargs) -> list[float]:
    responses = [completion[0]['content'] for completion in completions]
    q = prompts[0][-1]['content']
    extracted_responses = [extract_xml_answer(r) for r in responses]
    global COUNTER
    if COUNTER % PRINT_EVERY == 0:
        print('-'*20, f"Question:\n{q}", f"\nAnswer:\n{answer[0]}", f"\nResponse:\n{responses[0]}", f"\nExtracted:\n{extracted_responses[0]}")
    COUNTER += 1
    return [2.0 if r == a else 0.0 for r, a in zip(extracted_responses, answer)]

def int_reward_func(completions, **kwargs) -> list[float]:
    responses = [completion[0]['content'] for completion in completions]
    extracted_responses = [extract_xml_answer(r) for r in responses]
    return [0.5 if r.isdigit() else 0.0 for r in extracted_responses]

def strict_format_reward_func(completions, **kwargs) -> list[float]:
    """Reward function that checks if the completion has a specific format."""
    pattern = r"^<reasoning>\n.*?\n</reasoning>\n<answer>\n.*?\n</answer>\n$"
    responses = [completion[0]["content"] for completion in completions]
    matches = [re.match(pattern, r) for r in responses]
    return [0.5 if match else 0.0 for match in matches]

def soft_format_reward_func(completions, **kwargs) -> list[float]:
    """Reward function that checks if the completion has a specific format."""
    pattern = r"<reasoning>.*?</reasoning>\s*<answer>.*?</answer>"
    responses = [completion[0]["content"] for completion in completions]
    matches = [re.match(pattern, r) for r in responses]
    return [0.5 if match else 0.0 for match in matches]

def count_xml(text) -> float:
    count = 0.0
    if text.count("<reasoning>\n") == 1:
        count += 0.125
    if text.count("\n</reasoning>\n") == 1:
        count += 0.125
    if text.count("\n<answer>\n") == 1:
        count += 0.125
        count -= len(text.split("\n</answer>\n")[-1])*0.001
    if text.count("\n</answer>") == 1:
        count += 0.125
        count -= (len(text.split("\n</answer>")[-1]) - 1)*0.001
    return count

def xmlcount_reward_func(completions, **kwargs) -> list[float]:
    contents = [completion[0]["content"] for completion in completions]
    return [count_xml(c) for c in contents]

max_prompt_length = 256
from trl import GRPOConfig, GRPOTrainer

# Optional extra params for vLLM
from unsloth import vLLMSamplingParams
vllm_sampling_params = vLLMSamplingParams(
    min_p = 0.01,
    seed = 3407,
)
training_args = GRPOConfig(
    learning_rate = 5e-6,
    warmup_ratio = 0.1,
    lr_scheduler_type = "cosine",
    optim = "adamw_8bit",
    per_device_train_batch_size = 1,
    gradient_accumulation_steps = 1, # Increase to 4 for smoother training
    num_generations = 6, # Decrease if out of memory
    max_prompt_length = max_prompt_length,
    max_completion_length = max_seq_length - max_prompt_length,
    # num_train_epochs = 1, # Set to 1 for a full training run
    max_steps = 250,
    report_to = "none", # Can use Weights & Biases
    vllm_sampling_params = vllm_sampling_params, # Optional
    temperature = 1.0,
)
trainer = GRPOTrainer(
    model = model,
    processing_class = tokenizer,
    reward_funcs = [
        xmlcount_reward_func,
        soft_format_reward_func,
        strict_format_reward_func,
        int_reward_func,
        correctness_reward_func,
    ],
    args = training_args,
    train_dataset = dataset,
)
trainer.train()
```


## What's Changed
* GRPO Bug fixes by @danielhanchen in https://github.com/unslothai/unsloth/pull/1623
* Fixes Triton url in README.md by @DiogoNeves in https://github.com/unslothai/unsloth/pull/1607
* Update README.md by @shimmyshimmer in https://github.com/unslothai/unsloth/pull/1654
* Update README.md by @shimmyshimmer in https://github.com/unslothai/unsloth/pull/1688
* Fix bugs by @danielhanchen in https://github.com/unslothai/unsloth/pull/1701
* Fix bugs by @danielhanchen in https://github.com/unslothai/unsloth/pull/1706
* Memory efficient GRPO, DPO etc by @danielhanchen in https://github.com/unslothai/unsloth/pull/1716
* Add GRPO metrics by @danielhanchen in https://github.com/unslothai/unsloth/pull/1718
* llama-quantize on WINDOWS WSL error fix - edit save.py (gguf saving breaks) by @everythingisc00l in https://github.com/unslothai/unsloth/pull/1649
* Update rl_replacements.py by @SethHWeidman in https://github.com/unslothai/unsloth/pull/1754
* Update README.md by @danielhanchen in https://github.com/unslothai/unsloth/pull/1768
* fix an import error by @NinoRisteski in https://github.com/unslothai/unsloth/pull/1767
* Gemma Mask convert to float by @Erland366 in https://github.com/unslothai/unsloth/pull/1762
* [Windows Support] Add latest `xformers` wheels to pyproject.toml by @versipellis in https://github.com/unslothai/unsloth/pull/1753
* Memory Efficient GRPO by @danielhanchen in https://github.com/unslothai/unsloth/pull/1773

## New Contributors
* @DiogoNeves made their first contribution in https://github.com/unslothai/unsloth/pull/1607
* @everythingisc00l made their first contribution in https://github.com/unslothai/unsloth/pull/1649
* @SethHWeidman made their first contribution in https://github.com/unslothai/unsloth/pull/1754
* @versipellis made their first contribution in https://github.com/unslothai/unsloth/pull/1753

**Full Changelog**: https://github.com/unslothai/unsloth/compare/2025-02...2025-02-v2