Detalles técnicos de GRPO
[!TIP] Esta sección profundiza en los detalles técnicos y matemáticos de GRPO. Fue escrita por Shirin Yamani.
Profundicemos en nuestra comprensión de GRPO para que podamos mejorar el proceso de entrenamiento de nuestro modelo.
GRPO evalúa directamente las respuestas generadas por el modelo comparándolas dentro de grupos de generación para optimizar el modelo de política, en lugar de entrenar un modelo de valor (Crítico) separado. ¡Este enfoque conduce a una reducción significativa en el costo computacional!
GRPO se puede aplicar a cualquier tarea verificable donde la corrección de la respuesta se pueda determinar. Por ejemplo, en el razonamiento matemático, la corrección de la respuesta se puede verificar fácilmente comparándola con la verdad fundamental.
Antes de sumergirnos en los detalles técnicos, visualicemos cómo funciona GRPO a un alto nivel:

Ahora que tenemos una visión general visual, desglosemos cómo funciona GRPO paso a paso.
El algoritmo GRPO
La innovación central de GRPO es su enfoque para evaluar y aprender de múltiples respuestas generadas simultáneamente. En lugar de depender de un modelo de recompensa separado, compara las salidas dentro del mismo grupo para determinar cuáles deben reforzarse.
Repasemos cada paso del algoritmo en detalle:
Paso 1: Muestreo de grupo
El primer paso es generar múltiples respuestas posibles para cada pregunta. Esto crea un conjunto diverso de salidas que se pueden comparar entre sí.
Para cada pregunta \( q \), el modelo generará \( G \) salidas (tamaño del grupo) de la política entrenada: { \( {o_1, o_2, o_3, \dots, o_G}\pi_{\theta_{\text{old}}} \) }, \( G=8 \) donde cada \( o_i \) representa una finalización del modelo.
Ejemplo
Para hacerlo más concreto, veamos un problema aritmético simple:
Pregunta
\( q \) : \( \text{Calcula}\space2 + 2 \times 6 \)
Salidas
\( (G = 8) \): \( {o_1:14 \text{ (correcta)}, o_2:16 \text{ (incorrecta)}, o_3:10 \text{ (incorrecta)}, \ldots, o_8:14 \text{ (correcta)}} \)
Observa cómo algunas de las respuestas generadas son correctas (14) mientras que otras son incorrectas (16 o 10). Esta diversidad es crucial para el siguiente paso.
Paso 2: Cálculo de la ventaja
Una vez que tenemos múltiples respuestas, necesitamos una forma de determinar cuáles son mejores que otras. Aquí es donde entra en juego el cálculo de la ventaja.
Distribución de recompensas
Primero, asignamos una puntuación de recompensa a cada respuesta generada. En este ejemplo, usaremos un modelo de recompensa, pero como aprendimos en la sección anterior, podemos usar cualquier función que devuelva una recompensa.
Asigna una puntuación RM a cada una de las respuestas generadas basándose en la corrección \( r_i \) (por ejemplo, 1 para una respuesta correcta, 0 para una respuesta incorrecta) y luego para cada \( r_i \) calcula el siguiente valor de Ventaja.
Fórmula del valor de ventaja
La clave de GRPO es que no necesitamos medidas absolutas de calidad, podemos comparar las salidas dentro del mismo grupo. Esto se hace usando la estandarización:
$$A_i = \frac{r_i - \text{media}({r_1, r_2, \ldots, r_G})}{\text{desviación_estándar}({r_1, r_2, \ldots, r_G})}$$
Ejemplo
Continuando con nuestro ejemplo aritmético para el mismo ejemplo anterior, imagina que tenemos 8 respuestas, 4 de las cuales son correctas y el resto incorrectas, por lo tanto;
| Métrica | Valor |
|---|---|
| Promedio del grupo | \( media(r_i) = 0.5 \) |
| Desviación estándar | \( desviación_estándar(r_i) = 0.53 \) |
| Valor de ventaja (respuesta correcta) | \( A_i = \frac{1 - 0.5}{0.53}= 0.94 \) |
| Valor de ventaja (respuesta incorrecta) | \( A_i = \frac{0 - 0.5}{0.53}= -0.94 \) |
Interpretación
Ahora que hemos calculado los valores de ventaja, entendamos lo que significan:
Esta estandarización (es decir, la ponderación \( A_i \)) permite al modelo evaluar el rendimiento relativo de cada respuesta, guiando el proceso de optimización hacia respuestas favorables que son mejores que el promedio (recompensa alta) y desalentando aquellas que son peores. Por ejemplo, si \( A_i > 0 \), entonces \( o_i \) es una respuesta mejor que el nivel promedio dentro de su grupo; y si \( A_i < 0 \), entonces la calidad de la respuesta es menor que el promedio (es decir, baja calidad/rendimiento).
Para el ejemplo anterior, si \( A_i = 0.94 \text{(salida correcta)} \) entonces durante los pasos de optimización su probabilidad de generación aumentará.
Con nuestros valores de ventaja calculados, ahora estamos listos para actualizar la política.
Paso 3: Actualización de la política
El paso final es usar estos valores de ventaja para actualizar nuestro modelo para que sea más probable que genere buenas respuestas en el futuro.
La función objetivo para la actualización de la política es:
$$J_{GRPO}(\theta) = \left[\frac{1}{G} \sum_{i=1}^{G} \min \left( \frac{\pi_{\theta}(o_i|q)}{\pi_{\theta_{old}}(o_i|q)} A_i \text{clip}\left( \frac{\pi_{\theta}(o_i|q)}{\pi_{\theta_{old}}(o_i|q)}, 1 - \epsilon, 1 + \epsilon \right) A_i \right)\right]- \beta D_{KL}(\pi_{\theta} || \pi_{ref})$$
Esta fórmula puede parecer intimidante al principio, pero está construida a partir de varios componentes que cumplen un propósito importante cada uno. Desglosémoslos uno por uno.
Componentes clave de la función objetivo
La función de actualización de GRPO combina varias técnicas para garantizar un aprendizaje estable y efectivo. Examinemos cada componente:
1. Relación de probabilidad
La relación de probabilidad se define como:
\( \left(\frac{\pi_{\theta}(o_i|q)}{\pi_{\theta_{old}}(o_i|q)}\right) \)
Intuitivamente, la fórmula compara cuánto difiere la probabilidad de respuesta del nuevo modelo de la probabilidad de respuesta del modelo antiguo, incorporando una preferencia por las respuestas que mejoran el resultado esperado.
Interpretación
- Si \( \text{ratio} > 1 \), el nuevo modelo asigna una probabilidad más alta a la respuesta \( o_i \) que el modelo antiguo.
- Si \( \text{ratio} < 1 \), el nuevo modelo asigna una probabilidad más baja a \( o_i \)
Esta relación nos permite controlar cuánto cambia el modelo en cada paso, lo que nos lleva al siguiente componente.
2. Función de recorte (Clip Function)
La función de recorte se define como:
\( \text{clip}\left( \frac{\pi_{\theta}(o_i|q)}{\pi_{\theta_{old}}(o_i|q)}, 1 - \epsilon, 1 + \epsilon\right) \)
Limita la relación discutida anteriormente para que esté dentro de \( [1 - \epsilon, 1 + \epsilon] \) para evitar/controlar cambios drásticos o actualizaciones descabelladas y alejarse demasiado de la política antigua. En otras palabras, limita cuánto puede aumentar la relación de probabilidad para ayudar a mantener la estabilidad, evitando actualizaciones que alejen demasiado al nuevo modelo del antiguo.
Ejemplo (ε = 0.2)
Veamos dos escenarios diferentes para entender mejor esta función de recorte:
- Caso 1: si la nueva política tiene una probabilidad de 0.9 para una respuesta específica y la política antigua tiene una probabilidad de 0.5, significa que esta respuesta está siendo reforzada por la nueva política para tener una probabilidad más alta, pero dentro de un límite controlado que es el recorte para evitar cambios drásticos.
- \( \text{Relación}: \frac{\pi_{\theta}(o_i|q)}{\pi_{\theta_{old}}(o_i|q)} = \frac{0.9}{0.5} = 1.8 → \text{Recorte}\space1.2 \) (límite superior 1.2)
- Caso 2: Si la nueva política no está a favor de una respuesta (probabilidad más baja, por ejemplo, 0.2), lo que significa que si la respuesta no es beneficiosa, el aumento podría ser incorrecto y el modelo sería penalizado.
- \( \text{Relación}: \frac{\pi_{\theta}(o_i|q)}{\pi_{\theta_{old}}(o_i|q)} = \frac{0.2}{0.5} = 0.4 →\text{Recorte}\space0.8 \) (límite inferior 0.8)
Interpretación
- La fórmula anima al nuevo modelo a favorecer las respuestas que el modelo antiguo infravaloró si mejoran el resultado.
- Si el modelo antiguo ya favorecía una respuesta con alta probabilidad, el nuevo modelo aún puede reforzarla pero solo dentro de un límite controlado \( [1 - \epsilon, 1 + \epsilon] \), \( \text{(por ejemplo, }\epsilon = 0.2, \space \text{entonces} \space [0.8-1.2]) \).
- Si el modelo antiguo sobreestimó una respuesta que funciona mal, el nuevo modelo es desalentado de mantener esa alta probabilidad.
- Por lo tanto, intuitivamente, al incorporar la relación de probabilidad, la función objetivo asegura que las actualizaciones de la política sean proporcionales a la ventaja \( A_i \) mientras se moderan para evitar cambios drásticos.
Aunque la función de recorte ayuda a prevenir cambios drásticos, necesitamos una salvaguarda más para asegurar que nuestro modelo no se desvíe demasiado de su comportamiento original.
3. Divergencia KL
El término de divergencia KL es:
\( \beta D_{KL}(\pi_{\theta} || \pi_{ref}) \)
En el término de divergencia KL, \( \pi_{ref} \) es básicamente la salida del modelo pre-actualización, per_token_logps y \( \pi_{\theta} \) es la salida del nuevo modelo, new_per_token_logps. Teóricamente, la divergencia KL se minimiza para evitar que el modelo se desvíe demasiado de su comportamiento original durante la optimización. Esto ayuda a lograr un equilibrio entre mejorar el rendimiento basándose en la señal de recompensa y mantener la coherencia. En este contexto, minimizar la divergencia KL reduce el riesgo de que el modelo genere texto sin sentido o, en el caso del razonamiento matemático, produzca respuestas extremadamente incorrectas.
Interpretación
- Una penalización por divergencia KL mantiene las salidas del modelo cerca de su distribución original, evitando cambios extremos.
- En lugar de desviarse hacia salidas completamente irracionales, el modelo refinaría su comprensión al mismo tiempo que permitiría cierta exploración.
Definición matemática
Para aquellos interesados en los detalles matemáticos, veamos la definición formal:
Recuerda que la distancia KL se define de la siguiente manera: $$D_{KL}(P || Q) = \sum_{x \in X} P(x) \log \frac{P(x)}{Q(x)}$$ En RLHF, las dos distribuciones de interés suelen ser la distribución de la nueva versión del modelo, P(x), y una distribución de la política de referencia, Q(x).
El papel del parámetro β
El coeficiente \( \beta \) controla la fuerza con la que aplicamos la restricción de divergencia KL:
- Beta más alto (penalización KL más fuerte)
- Mayor restricción en las actualizaciones de la política. El modelo permanece cerca de su distribución de referencia.
- Puede ralentizar la adaptación: El modelo puede tener dificultades para explorar mejores respuestas.
- Beta más bajo (penalización KL más débil)
- Más libertad para actualizar la política: El modelo puede desviarse más de la referencia.
- Adaptación más rápida pero riesgo de inestabilidad: El modelo podría aprender comportamientos de "reward-hacking".
- Riesgo de sobreoptimización: Si el modelo de recompensa es defectuoso, la política podría generar salidas sin sentido.
- El documento original DeepSeekMath estableció este \( \beta= 0.04 \)
Ahora que entendemos los componentes de GRPO, veamos cómo funcionan juntos en un ejemplo completo.
Ejemplo práctico con GRPO
Para consolidar nuestra comprensión de GRPO, repasemos un ejemplo completo de principio a fin.
Problema de ejemplo
$$\text{P: Calcula}\space2 + 2 \times 6$$
Paso 1: Muestreo de grupo
Primero, generamos múltiples respuestas de nuestro modelo.
Genera \( (G = 8) \) respuestas, \( 4 \) de las cuales son la respuesta correcta (\( 14, \text{recompensa=} 1 \)) y \( 4 \) incorrectas \( \text{(recompensa= 0)} \), por lo tanto:
$${o_1:14(correcta), o_2:10 (incorrecta), o_3:16 (incorrecta), ... o_G:14(correcta)}$$
Paso 2: Cálculo de la ventaja
A continuación, calculamos los valores de ventaja para determinar qué respuestas son mejores que el promedio:
| Estadística | Valor |
|---|---|
| Promedio del grupo | \( media(r_i) = 0.5 \) |
| Desviación estándar | \( desviación_estándar(r_i) = 0.53 \) |
| Valor de ventaja (respuesta correcta) | \( A_i = \frac{1 - 0.5}{0.53}= 0.94 \) |
| Valor de ventaja (respuesta incorrecta) | \( A_i = \frac{0 - 0.5}{0.53}= -0.94 \) |
Paso 3: Actualización de la política
Finalmente, actualizamos nuestro modelo para reforzar las respuestas correctas:
- Suponiendo que la probabilidad de la política antigua (\( \pi_{\theta_{old}} \)) para una salida correcta \( o_1 \) es \( 0.5 \) y la nueva política la aumenta a \( 0.7 \), entonces: $$\text{Relación}: \frac{0.7}{0.5} = 1.4 →\text{después del recorte}\space1.2 \space (\epsilon = 0.2)$$
- Luego, cuando la función objetivo se vuelve a ponderar, el modelo tiende a reforzar la generación de la salida correcta, y la \( \text{Divergencia KL} \) limita la desviación de la política de referencia.
Con la comprensión teórica establecida, veamos cómo se puede implementar GRPO en código.
Ejemplo de implementación
Unamos todo en un ejemplo práctico. El siguiente código demuestra cómo implementar GRPO en PyTorch.
1. Cargando el modelo y generando respuestas
Primero, necesitamos cargar un modelo y generar múltiples respuestas para una pregunta dada:
from transformers import AutoModelForCausalLM, AutoTokenizer
# Load the model and tokenizer
model_name = "Qwen/Qwen2-Math-1.5B"
model = AutoModelForCausalLM.from_pretrained(model_name)
tokenizer = AutoTokenizer.from_pretrained(model_name)
model.eval()
# Move model to GPU if available
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device)
# Input prompt
prompt = "Solve y = 2x + 1 for x = 2, y = " # Correct answer: 5
inputs = tokenizer(prompt, return_tensors="pt", padding=True)
input_ids = inputs["input_ids"].to(device) # Shape: (1, prompt_len)
attention_mask = inputs["attention_mask"].to(device)
# Step 1: Generate 8 responses (B = 2 groups, G = 4 responses per group)
batch_size, num_generations = 2, 4
outputs = model.generate(
input_ids=input_ids, # Shape: (1, prompt_len)
attention_mask=attention_mask,
max_new_tokens=1, # seq_len = 1 (single token per response)
num_return_sequences=batch_size * num_generations, # 8 responses total
do_sample=True,
top_k=10,
temperature=0.7,
pad_token_id=tokenizer.eos_token_id,
return_dict_in_generate=True,
output_scores=True,
)
esta generación inicial (antes de cualquier paso) producirá algo como esto:
Output 1: 5.0
Output 2: 6.0
Output 3: 7.0
Output 4: 5.0
Output 5: 10.0
Output 6: 2.0
Output 7: 5.0
Output 8: 5.0
2. Calculando recompensas
Ahora, necesitamos determinar qué respuestas son correctas y asignar recompensas en consecuencia:
Con GRPO, con el mismo prompt de muestra, generamos múltiples finalizaciones. Así, por ejemplo, para nuestros prompts de "Solve y = 2x + 1 for x = 2, y = " y Solve y = 2x + 1 for x = 4, y = " tenemos dos grupos de salidas generadas para el prompt dado, uno es, digamos
[5, 6, 7, 5]y el otro es[10, 2, 9, 9]mientras que la respuesta correcta es 5 y 9.
Ten en cuenta que en la práctica estas puntuaciones de recompensa se logran mediante una función de recompensa basada en reglas que asigna recompensas basándose en la corrección de la respuesta o un modelo de red neuronal más complejo que puede entrenarse para asignar recompensas basándose en la corrección de la respuesta o una mezcla de ambos. Pero por simplicidad, digamos que nuestra recompensa por respuesta es 1 si la respuesta es correcta y 0 si es incorrecta, por lo tanto;
reward_1 = [1, 0, 0, 1]
reward_2 = [0, 0, 1, 1]
a continuación, obtenemos la media y la desviación estándar de las recompensas por grupo;
# Shape: (B * G,) = (8,) bc we have 2 groups of 4 generations that we flatten
rewards = torch.tensor([1, 0, 0, 1, 0, 0, 1, 1], dtype=torch.float32)
num_generations = 4
# Group rewards: Shape (B, G) = 2, 4)
rewards_grouped = rewards.view(-1, num_generations)
# Mean per group: Shape (B,) = (2,)
mean_grouped_rewards = rewards_grouped.mean(dim=1)
# Std per group: Shape (B,) = (2,)
std_grouped_rewards = rewards_grouped.std(dim=1)
# Broadcast to match rewards and normalize: Shape (B * G,) = (8,)
# why we need to broadcast? because we need to calculate the advantage values for each response within the group
mean_grouped_rewards = mean_grouped_rewards.repeat_interleave(num_generations, dim=0)
std_grouped_rewards = std_grouped_rewards.repeat_interleave(num_generations, dim=0)
esto producirá:
Grouped Rewards: tensor([[1., 0., 0., 1.],
[0., 0., 1., 1.]])
Mean per group: tensor([0.5000, 0.5000])
Std per group: tensor([0.5774, 0.5774])
Broadcasted Mean: tensor([0.5000, 0.5000, 0.5000, 0.5000, 0.5000, 0.5000, 0.5000, 0.5000])
Broadcasted Std: tensor([0.5774, 0.5774, 0.5774, 0.5774, 0.5774, 0.5774, 0.5774, 0.5774])
Ahora podemos calcular los valores de ventaja para cada respuesta:
# Advantages: Shape (B * G,) = (8,)
advantages = (rewards - mean_grouped_rewards) / (std_grouped_rewards + 1e-8)
esto producirá:
Advantages: tensor([ 0.8659, -0.8660, -0.8660, 0.8659, -0.8660, -0.8660, 0.8659, 0.8659])
que proviene de la fórmula de Ventaja anterior, entonces:
For reward_1 = [1, 0, 0, 1]:
1 - 0.5 / 0.5774 ≈ 0.8659
0 - 0.5 / 0.5774 ≈ -0.8660
For reward_2 = [0, 0, 1, 1]: Same pattern.
sin embargo, la forma aquí es (B*G,) = (8,) pero en la práctica, necesitamos tener la forma de (B, G) = (2, 4) para que coincida con la forma de los logits, ¿verdad? Por lo tanto, necesitamos expandir el tensor de ventajas para que tenga la forma de (B*G, 1) = (8, 1) para que coincida con la forma de los logits.
# Shape (B * G, 1) = (8, 1) to match the logits shape
advantages = advantages.unsqueeze(1)
lo que producirá:
Advantages: tensor([[ 0.8659],
[-0.8660],
[-0.8660],
[ 0.8659],
[-0.8660],
[-0.8660],
[ 0.8659],
[ 0.8659]])
Ahora estamos listos, pasemos al siguiente paso de actualizar el modelo de política basándonos en los valores de ventaja.
3. Actualizando la política
Finalmente, usamos los valores de ventaja para actualizar nuestro modelo:
# Compute probability ratio between new and old policies
ratio = torch.exp(
new_per_token_logps - per_token_logps
) # Shape: (B*G, seq_len) seq_len is the length of the output i.e. the num of generated tokens so here for simplicity let's assume it is 1 # (8, 1)
Ten en cuenta que el per_token_logps se puede lograr pasando las salidas generadas al modelo y obteniendo los logits para luego aplicar la función softmax y obtener las probabilidades F.softmax(logits, dim=-1).
# Clipping Function
eps = self.cliprange # e.g. 0.2
pg_losses1 = -advantages * ratio # Shape: (B*G, seq_len) #(8, 1)
pg_losses2 = -advantages * torch.clamp(
ratio, 1.0 - eps, 1.0 + eps
) # Shape: (B*G, seq_len) #(8, 1)
pg_loss_max = torch.max(pg_losses1, pg_losses2) # Shape: (B*G, seq_len) #(8, 1)
# Now Combine with KL penalty # Shape: (B*G, seq_len) #(8, 1)
per_token_loss = pg_loss_max + self.beta * per_token_kl
per_token_kl también se puede calcular de la siguiente manera:
# Shape: (B*G, seq_len) #(8, 1)
per_token_kl = F.kl_div(
F.log_softmax(new_per_token_logps, dim=-1),
F.softmax(per_token_logps, dim=-1),
reduction="none",
).sum(dim=-1, keepdim=True)
El ejemplo completo se puede encontrar aquí. GRPO también está implementado por el excelente equipo de TRL, puedes consultar la implementación TRL/GRPO_trainer para más detalles.
Resumen y próximos pasos
¡Felicidades! Ahora has aprendido sobre la Optimización de Políticas Relativas de Grupo (GRPO). Para recapitular lo que hemos cubierto:
- GRPO compara múltiples salidas dentro de un grupo para determinar cuáles son mejores que otras, sin requerir un modelo de valor separado.
- El cálculo de la ventaja estandariza las recompensas para identificar qué respuestas están por encima o por debajo del promedio.
- La actualización de la política utiliza una función objetivo recortada con una penalización de divergencia KL para garantizar un aprendizaje estable.
Este enfoque es particularmente potente para tareas de razonamiento matemático, donde la corrección puede verificarse objetivamente. El método GRPO permite un entrenamiento más eficiente en comparación con los enfoques tradicionales de RLHF que requieren un modelo crítico separado.
A medida que continúes explorando GRPO, considera experimentar con diferentes tamaños de grupo, funciones de recompensa y coeficientes de penalización KL para ver cómo afectan el rendimiento de tu modelo.
¡Feliz entrenamiento! 🚀