はじめに
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 | オンラインが不可能な場合にログからトレーニングする |