Reinforcement Learning: обучение с подкреплением для управления LLM

opensourceairlrlhfllmmachine-learningit
← Back to Blog

Введение: Reinforcement Learning

Что такое RL?

Reinforcement Learning (RL) — это раздел машинного обучения, где агент учится принимать решения, взаимодействуя с окружением.

Цикл RL:

  1. Agent наблюдает состояние s_t
  2. Agent выбирает действие a_t
  3. Environment даёт награду r_t
  4. Environment переходит в s_{t+1}

Цель: максимизировать сумму наград: max E[Σ γ^t * r_t]

Факт: AlphaGo обыграл чемпиона мира Ли Седоля 4-1, используя RL + MCTS.


Основы RL

Markov Decision Process (MDP)

MDP = (S, A, P, R, γ)

S = множество состояний
A = множество действий
P(s'|s,a) = вероятность перехода
R(s,a) = функция награды
γ = дисконт (0 ≤ γ ≤ 1)

Markov Property:
  P(s_{t+1} | s_t, a_t, ..., s_0, a_0) = P(s_{t+1} | s_t, a_t)
  
  Будущее зависит только от текущего состояния!

Value Functions

State Value Function (V):
  Vπ(s) = Eπ[Σ γ^t * r_t | s_0 = s]
  
  "Насколько хороша эта состояние при политике π"

Action Value Function (Q):
  Qπ(s, a) = Eπ[Σ γ^t * r_t | s_0 = s, a_0 = a]
  
  "Насколько хороша это действие в этом состоянии"

Bellman Equation:
  Q(s,a) = R(s,a) + γ · Σ_{s'} P(s'|s,a) · max_a' Q(s',a')

RL Algorithms

Q-Learning

Q-Learning: Off-policy temporal difference

Update rule:
  Q(s,a) ← Q(s,a) + α · [r + γ · max_a' Q(s',a') - Q(s,a)]
  
  α = learning rate
  r + γ · max_a' Q(s',a') = target

Epsilon-Greedy Exploration:
  π(a|s) = ε           с случайным действием
  π(a|s) = 1 - ε + ε/K с лучшим действием
  π(a|s) = ε/K         с другими действиями
class QLearning:
    def __init__(self, state_dim, action_dim, 
                 lr=0.1, gamma=0.99, epsilon=0.1):
        self.q_table = {}
        self.lr = lr
        self.gamma = gamma
        self.epsilon = epsilon
    
    def get_q(self, state, action):
        return self.q_table.get((state, action), 0.0)
    
    def select_action(self, state):
        if random.random() < self.epsilon:
            return random.randint(0, self.action_dim - 1)
        q_values = [self.get_q(state, a) for a in range(self.action_dim)]
        return q_values.index(max(q_values))
    
    def update(self, state, action, reward, next_state):
        current_q = self.get_q(state, action)
        next_q = max(self.get_q(next_state, a) for a in range(self.action_dim))
        target = reward + self.gamma * next_q
        new_q = current_q + self.lr * (target - current_q)
        self.q_table[(state, action)] = new_q
        return new_q - current_q

Deep Q-Network (DQN)

DQN: Q-Learning с нейронной сетью

Ключевые инновации:
  1. Experience Replay: хранение (s,a,r,s')
  2. Target Network: отдельная сеть для targets
  3. Gradient clipping: ограничение размера обновления

Архитектура:
  Input: state (image или features)
  Output: Q(s, a₁), Q(s, a₂), ..., Q(s, aₙ)
class DQN:
    def __init__(self, state_dim, action_dim, hidden_dim=256):
        self.q_network = torch.nn.Sequential(
            torch.nn.Linear(state_dim, hidden_dim),
            torch.nn.ReLU(),
            torch.nn.Linear(hidden_dim, action_dim)
        )
        self.target_network = copy.deepcopy(self.q_network)
        self.buffer = deque(maxlen=100000)
        self.optimizer = torch.optim.Adam(self.q_network.parameters(), lr=0.001)
        self.gamma = 0.99
    
    def select_action(self, state, epsilon=0.1):
        if random.random() < epsilon:
            return random.randint(0, self.action_dim - 1)
        with torch.no_grad():
            q_values = self.q_network(state)
            return q_values.argmax(dim=-1).item()
    
    def train_step(self, batch_size=64):
        batch = random.sample(self.buffer, batch_size)
        states, actions, rewards, next_states, dones = zip(*batch)
        
        current_q = self.q_network(states).gather(1, actions)
        with torch.no_grad():
            next_q = self.target_network(next_states).max(dim=1)[0]
            target_q = rewards + self.gamma * next_q * (1 - dones)
        
        loss = torch.nn.functional.mse_loss(current_q, target_q)
        self.optimizer.zero_grad()
        loss.backward()
        torch.nn.utils.clip_grad_norm_(self.q_network.parameters(), max_norm=1.0)
        self.optimizer.step()
        return loss.item()

Policy Gradient (REINFORCE)

Policy Gradient: Прямая оптимизация политики

Policy πθ(a|s): параметризована θ

Objective J(θ) = E[Σ γ^t · r_t]

Gradient:
  ∇θ J(θ) = E[(Σ γ^t · r_t) · ∇θ log πθ(a|s)]
  
  Просто масштабируем log-probability по возврату
class PolicyGradient:
    def __init__(self, state_dim, action_dim, lr=0.001):
        self.policy = torch.nn.Sequential(
            torch.nn.Linear(state_dim, 256),
            torch.nn.ReLU(),
            torch.nn.Linear(256, action_dim),
            torch.nn.Softmax(dim=-1)
        )
        self.optimizer = torch.optim.Adam(self.policy.parameters(), lr=lr)
        self.gamma = 0.99
    
    def select_action(self, state):
        probs = self.policy(state)
        dist = torch.distributions.Categorical(probs)
        action = dist.sample()
        return action, dist.log_prob(action)
    
    def train(self, log_probs, returns):
        discounted_returns = self._compute_discounted_returns(returns)
        policy_loss = sum(-log_prob * G for log_prob, G in zip(log_probs, discounted_returns))
        policy_loss = policy_loss.mean()
        
        self.optimizer.zero_grad()
        policy_loss.backward()
        self.optimizer.step()
        return policy_loss.item()

Actor-Critic

Actor-Critic: Policy gradient + Value function

Actor: policy πθ(a|s) — выбирает действие
Critic: value Vφ(s) — оценивает состояние

Update:
  δ = r + γ·V(s') - V(s)  (TD error)
  
  Actor: ∇θ log πθ(a|s) · δ
  Critic: (δ)²
  
Critic снижает дисперсию policy gradient
class ActorCritic:
    def __init__(self, state_dim, action_dim, lr=0.001):
        self.actor = torch.nn.Sequential(
            torch.nn.Linear(state_dim, 256),
            torch.nn.ReLU(),
            torch.nn.Linear(256, action_dim),
            torch.nn.Softmax(dim=-1)
        )
        self.critic = torch.nn.Sequential(
            torch.nn.Linear(state_dim, 256),
            torch.nn.ReLU(),
            torch.nn.Linear(256, 1)
        )
        self.gamma = 0.99
    
    def select_action(self, state):
        probs = self.actor(state)
        dist = torch.distributions.Categorical(probs)
        action = dist.sample()
        return action, dist.log_prob(action)
    
    def update(self, state, action, reward, next_state, done, log_prob, advantage):
        current_v = self.critic(state).squeeze(-1)
        target_v = reward + self.gamma * self.critic(next_state) * (1 - done)
        
        critic_loss = torch.nn.functional.mse_loss(current_v, target_v.detach())
        actor_loss = -log_prob * advantage.detach()
        
        self.critic.zero_grad()
        critic_loss.backward()
        self.actor.zero_grad()
        actor_loss.backward()

PPO (Proximal Policy Optimization)

PPO: State-of-the-art policy optimization

r(θ) = πθ(a|s) / πθ_old(a|s)

L^CLIP(θ) = E[min(r(θ)·A, clip(r(θ), 1-ε, 1+ε)·A)]

ε = 0.1 или 0.2 (clip range)

Преимущества:
  - Стабильное обучение
  - Не нужны сложные second-order методы
  - Работает в большинстве окружений
class PPO:
    def __init__(self, state_dim, action_dim, lr=3e-4, 
                 epsilon=0.2, gamma=0.99, k_epochs=10):
        self.epsilon = epsilon
        self.gamma = gamma
        self.k_epochs = k_epochs
        
        self.actor = torch.nn.Sequential(
            torch.nn.Linear(state_dim, 256),
            torch.nn.Tanh(),
            torch.nn.Linear(256, action_dim),
            torch.nn.Softmax(dim=-1)
        )
        self.critic = torch.nn.Sequential(
            torch.nn.Linear(state_dim, 256),
            torch.nn.Tanh(),
            torch.nn.Linear(256, 1)
        )
        self.optimizer = torch.optim.Adam(
            list(self.actor.parameters()) + list(self.critic.parameters()), lr=lr
        )
    
    def train_step(self, states, actions, old_log_probs, advantages, returns):
        advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
        
        for _ in range(self.k_epochs):
            probs = self.actor(states)
            dist = torch.distributions.Categorical(probs)
            log_probs = dist.log_prob(actions).sum(-1, keepdim=True)
            ratios = torch.exp(log_probs - old_log_probs)
            
            surr1 = ratios * advantages
            surr2 = torch.clamp(ratios, 1 - self.epsilon, 1 + self.epsilon) * advantages
            policy_loss = -torch.min(surr1, surr2).mean()
            
            value_loss = torch.nn.functional.mse_loss(
                self.critic(states).squeeze(-1), returns
            )
            
            loss = policy_loss + 0.5 * value_loss
            self.optimizer.zero_grad()
            loss.backward()
            self.optimizer.step()
        
        return loss.item()

RLHF (Reinforcement Learning from Human Feedback)

RLHF Pipeline

RLHF = Обучение LLM на основе человеческих предпочтений

3 шага:

1. Supervised Fine-Tuning (SFT)
   Обучение модели на примерах от людей
   
2. Reward Model Training
   Сбор предпочтений: "Response A > Response B"
   Обучение reward model предсказывать предпочтения
   
3. PPO Fine-Tuning
   Использование reward model как сигнала награды
   Дообучение LLM с PPO

Step 1: SFT

Supervised Fine-Tuning:

Данные: пары (prompt, response) от людей

Loss:
  L = -Σ log πθ(response_i | prompt)
  
  Standard cross-entropy on human demonstrations

Step 2: Reward Model

Reward Model:

Данные предпочтений: (x, y_w, y_l)
  x = prompt
  y_w = preferred response ("win")
  y_l = losing response ("lose")

Reward model R_φ(y|x) predicts preference:

Loss:
  L = -log σ(R_φ(x, y_w) - R_φ(x, y_l))
  
  Reward model учит: R(y_w) > R(y_l)
class RewardModel:
    def __init__(self, base_model, hidden_dim=512):
        self.base = base_model  # Fine-tuned LLM
        self.reward_head = torch.nn.Linear(
            self.base.config.hidden_size, 1
        )
    
    def forward(self, input_ids, attention_mask):
        outputs = self.base(
            input_ids=input_ids, 
            attention_mask=attention_mask,
            output_hidden_states=True
        )
        last_hidden = outputs.hidden_states[-1]
        # Mean pooling
        mask = attention_mask.unsqueeze(-1).float()
        mean_hidden = (last_hidden * mask).sum(1) / mask.sum(1)
        reward = self.reward_head(mean_hidden).squeeze(-1)
        return reward
    
    def train_step(self, prompt_ids, win_ids, lose_ids, masks_win, masks_lose):
        r_win = self.forward(win_ids, masks_win)
        r_lose = self.forward(lose_ids, masks_lose)
        
        loss = -torch.nn.functional.logsigmoid(r_win - r_lose).mean()
        return loss.item()

Step 3: PPO with Reward

PPO Fine-Tuning LLM:

State: prompt context
Action: next token
Reward: from Reward Model

KL Penalty:
  L = L_PPO - β · KL[πθ || π_ref]
  
  β = KL penalty coefficient (обычно 0.01-0.1)
  π_ref = SFT model (reference)
  
  Prevents model from drifting too far from SFT
class RLHFTrainer:
    def __init__(self, policy_model, reference_model, reward_model, 
                 beta=0.05, epsilon=0.2, gamma=0.99):
        self.policy = policy_model      # LLM being trained
        self.reference = reference_model  # SFT model
        self.reward = reward_model        # Reward model
        self.beta = beta
        self.epsilon = epsilon
        self.gamma = gamma
    
    def compute_reward(self, generated_ids, prompt_ids, prompt_mask):
        """Compute PPO reward + KL penalty"""
        # Reward from reward model
        r = self.reward(generated_ids, prompt_mask)
        
        # KL penalty
        with torch.no_grad():
            ref_logits = self.reference(generated_ids).log_softmax(dim=-1)
            pol_logits = self.policy(generated_ids).log_softmax(dim=-1)
            kl = (pol_logits.exp() * (pol_logits - ref_logits)).sum(-1).mean(-1)
        
        # Combined reward
        combined_reward = r - self.beta * kl
        return combined_reward
    
    def ppo_update(self, prompts, responses, rewards, advantages):
        """PPO update step for LLM"""
        # Standard PPO clip loss
        ratios = torch.exp(log_probs - old_log_probs)
        surr1 = ratios * advantages
        surr2 = torch.clamp(ratios, 1-self.epsilon, 1+self.epsilon) * advantages
        policy_loss = -torch.min(surr1, surr2).mean()
        
        # KL penalty in loss
        kl_penalty = self.compute_kl_penalty(prompts, responses)
        total_loss = policy_loss + self.beta * kl_penalty
        
        total_loss.backward()
        self.optimizer.step()
        
        return total_loss.item()

RL в управлении LLM

Resource Management через RL

RL для управления LLM ресурсами:

State:
  - GPU memory usage
  - Queue length
  - Request priorities
  - Temperature / top_p settings

Action:
  - Batch size adjustment
  - KV cache management
  - Request scheduling
  - Quantization level selection

Reward:
  + throughput
  - latency penalty
  - OOM penalty
  - priority-weighted completion
class LLMResourceAgent:
    """RL agent for LLM resource management"""
    
    def __init__(self, gpu_memory, batch_size=32):
        self.gpu_memory = gpu_memory
        self.batch_size = batch_size
        self.ppo = PPO(state_dim=8, action_dim=4)
        
        # State features:
        # [mem_usage, queue_len, avg_priority, temp, 
        #  kv_cache_usage, requests_per_sec, error_rate, gpu_temp]
        
        # Action: [batch_size_delta, kv_cache_threshold, 
        #          quantization_level, scheduling_policy]
    
    def get_state(self):
        """Observe environment state"""
        return [
            self.gpu_memory.used / self.gpu_memory.total,
            len(self.request_queue),
            np.mean([r.priority for r in self.request_queue]),
            self.current_temperature,
            self.kv_cache.usage_ratio,
            self.metrics.requests_per_sec,
            self.metrics.error_rate,
            self.gpu_temperature / 100.0
        ]
    
    def step(self, action):
        """Execute action and return reward"""
        # Decode action
        batch_delta = action[0] * 8  # -8 to +8
        self.batch_size = max(1, min(128, self.batch_size + batch_delta))
        
        # Run inference batch
        results = self.process_batch(self.batch_size)
        
        # Compute reward
        throughput = len(results) / self.processing_time
        avg_latency = np.mean([r.latency for r in results])
        
        reward = throughput - 0.1 * avg_latency - 10 * (self.gpu_memory.used > 0.95)
        
        return reward

Популярные фреймворки

OpenAI Baselines / Stable Baselines3

import stable_baselines3 as sb3
from stable_baselines3 import PPO
from stable_baselines3.common.envs import SimpleDocEnv

# Create environment
env = SimpleDocEnv()

# Train PPO
model = PPO("MlpPolicy", env, verbose=1, 
            tensorboard_log="./logs/")

model.learn(total_timesteps=100000)

# Evaluate
obs = env.reset()
for _ in range(1000):
    action, _ = model.predict(obs, deterministic=True)
    obs, reward, done, _ = env.step(action)
    if done:
        obs = env.reset()

Hugging FaceTRL (TRL)

from trl import PPOConfig, PPOTrainer, GPT2LMHeadModel
from transformers import AutoTokenizer

# Load models
policy = GPT2LMHeadModel.from_pretrained("gpt2-medium")
reference = GPT2LMHeadModel.from_pretrained("gpt2-medium")
tokenizer = AutoTokenizer.from_pretrained("gpt2-medium")

# TRL PPO config
config = PPOConfig(
    model_name="gpt2-medium-rlhf",
    learning_rate=1.41e-5,
    batch_size=8,
    mini_batch_size=2,
    ppo_epochs=4,
    seed=42,
)

# Create PPO trainer
ppo_trainer = PPOTrainer(
    config=config,
    model=policy,
    ref_model=reference,
    tokenizer=tokenizer,
    dataset=dataset,  # (prompt, response) pairs
)

# Train
for epoch, batch in train_dataloader:
    response = ppo_trainer.generate(
        batch["prompt"], 
        length_sampler=torch.randint(20, 100, (1,))
    )
    stats = ppo_trainer.step(
        batch["prompt"], response, reward
    )

Заключение

Reinforcement Learning — мощный инструмент для:

  1. RLHF — выравнивание LLM с человеческими предпочтениями
  2. Resource Management — автоматическая оптимизация GPU памяти и батчей
  3. Scheduling — умное планирование запросов
  4. Quantization — адаптивный выбор уровня квантования

Ключевые алгоритмы:

  • Q-Learning — базовый, для дискретных пространств
  • DQN — с neural network, для сложных состояний
  • PPO — state-of-the-art для continuous control
  • RLHF — 3-шаговый пайплайн для alignment

Open Source инструменты:

  • Hugging Face TRL — RLHF на LLM
  • Stable Baselines3 — PPO, A2C, SAC
  • Ray RLlib — distributed RL
  • OpenAI Spinning Up — обучение RL