Chuyển đến nội dung chính

レッスン 15: RL の運用 — RL エージェントの展開と監視

RL ポリシーを実稼働環境に展開します。 ONNX、TorchScript で提供されるモデル。安全上の制約。オンラインとオフラインの比較報酬のドリフトを監視します。 A/B テスト ポリシー。

🧠 AI と ML — レッスン 14 レッスン 15: RL の運用 — 導入と監視 RL エージェント

強化学習: 基礎から高度まで

パート 4: RLHF、LLM の調整と生産

xdev.asia

はじめに

RL エージェントを運用環境にデプロイすることは、教師あり ML とは大きく異なります。安全性の制約、オンライン学習、報酬の監視、ポリシーのバージョン管理を処理する必要があります。


1. モデルのエクスポートと提供

ONNX エクスポート

import torch
from stable_baselines3 import PPO

model = PPO.load("best_model")

# Export policy to ONNX
dummy_input = torch.randn(1, model.observation_space.shape[0])
torch.onnx.export(
    model.policy, dummy_input, "policy.onnx",
    input_names=["observation"],
    output_names=["action"],
    dynamic_axes={"observation": {0: "batch"}, "action": {0: "batch"}}
)

FastAPI サービス

from fastapi import FastAPI
import onnxruntime as ort
import numpy as np

app = FastAPI()
session = ort.InferenceSession("policy.onnx")

@app.post("/predict")
def predict(observation: list[float]):
    obs = np.array([observation], dtype=np.float32)
    result = session.run(None, {"observation": obs})
    action = int(np.argmax(result[0]))
    return {"action": action}

2. 安全上の制約

class SafeRLPolicy:
    def __init__(self, model, constraints):
        self.model = model
        self.constraints = constraints
    
    def predict(self, observation):
        action, _ = self.model.predict(observation, deterministic=True)
        
        # Check safety constraints
        if self.constraints.is_unsafe(observation, action):
            action = self.constraints.safe_fallback(observation)
            self.log_safety_override(observation, action)
        
        return action
    
    def log_safety_override(self, obs, action):
        # Track safety overrides for monitoring
        pass

3. オフライン R

ログに記録されたデータからトレーニング — 環境との対話は不要:

# Conservative Q-Learning (CQL)
from d3rlpy.algos import CQLConfig

cql = CQLConfig().create(device="cuda")
cql.fit(
    offline_dataset,
    n_steps=100_000,
    evaluators={"environment": gym_evaluator}
)

4. モニタリングと A/B テスト

class RLMonitor:
    def __init__(self):
        self.rewards = []
        self.actions = []
    
    def log_step(self, obs, action, reward):
        self.rewards.append(reward)
        self.actions.append(action)
    
    def check_drift(self, window=1000):
        recent = self.rewards[-window:]
        historical = self.rewards[-2*window:-window]
        # Statistical test for reward drift
        from scipy.stats import ks_2samp
        statistic, p_value = ks_2samp(recent, historical)
        if p_value < 0.05:
            alert("Reward distribution drift detected!")

5. 製造チェックリスト

側面アプローチ
モデル形式ONNX または TorchScript
サービングFastAPI + 非同期推論
安全性ハード制約 + フォールバック ポリシー
モニタリング報酬追跡、ドリフト検出
バージョン管理ポリシー バージョン レジストリ
ロールバック以前のポリシーに瞬時に切り替える
A/B テストカナリア展開
ロギング完全な状態、アクション、報酬のトレース

概要

トピック重要なポイント
エクスポートクロスプラットフォーム導入のための ONNX
安全性常にフォールバック ポリシーを使用する
モニタリング報酬の分布を経時的に追跡する
オフライン RLオンラインが不可能な場合にログからトレーニングする