[論文レビュー] Multi-task Batch Reinforcement Learning with Metric Learning
本稿では、オフラインデータセットにおける分布シフトにさらされても頑健なタスク推論を向上させるために、三重項損失を用いたメトリクス学習を行うMBMLという、マルチタスクバッチ強化学習手法を提案する。報酬の再ラベル付けによりハードネガティブ例を生成することで、方策はスパuriousな状態-行動相関に依存するのではなく報酬に注目するよう強制され、未観測タスクの微調整において最大80%の高速化を達成する。
We tackle the Multi-task Batch Reinforcement Learning problem. Given multiple datasets collected from different tasks, we train a multi-task policy to perform well in unseen tasks sampled from the same distribution. The task identities of the unseen tasks are not provided. To perform well, the policy must infer the task identity from collected transitions by modelling its dependency on states, actions and rewards. Because the different datasets may have state-action distributions with large divergence, the task inference module can learn to ignore the rewards and spuriously correlate $ extit{only}$ state-action pairs to the task identity, leading to poor test time performance. To robustify task inference, we propose a novel application of the triplet loss. To mine hard negative examples, we relabel the transitions from the training tasks by approximating their reward functions. When we allow further training on the unseen tasks, using the trained policy as an initialization leads to significantly faster convergence compared to randomly initialized policies (up to $80\%$ improvement and across 5 different Mujoco task distributions). We name our method $ extbf{MBML}$ ($ extbf{M} ext{ulti-task}$ $ extbf{B} ext{atch}$ RL with $ extbf{M} ext{etric}$ $ extbf{L} ext{earning}$).
研究の動機と目的
- オフラインデータセットに大きな状態-行動分布シフトが存在する場合のマルチタスクバッチ強化学習における一般化性能の低さという課題に対処すること。
- データ分布の重複が少ない状況で、タスク推論モジュールが状態-行動ペアとタスクIDの間でスパurious相関を学習するリスクを軽減すること。
- 真のタスクラベルにアクセスできない状況でも、コンテキスト遷移からタスクIDを推定できるように方策を設計し、未学習タスクにおけるゼロショット性能を向上させること。
- 訓練済みマルチタスク方策を有効な初期化として用いることで、新しいタスクにおける微調整を高速化すること。
- 報酬再ラベル化によるハードネガティブマイニングを伴うメトリクス学習が、標準的手法よりも優れた一般化性能を達成することを示すこと。
提案手法
- 正例と負例のコンテキスト遷移を区別できるようにするタスク埋め込みネットワークを訓練するための三重項損失目的関数を導入する。
- 各タスクの報酬関数を近似し、他のタスクの遷移をこれらの近似報酬で再ラベルすることで、ハードネガティブ例を生成する。
- 三重項損失が埋め込み空間内で異なるタスクを分離するよう促進するように、方策ネットワークは状態-行動ペアと報酬を用いてタスクIDを推定するように訓練される。
- 共有表現とタスク固有のヘッドを持つマルチタスクQネットワークアーキテクチャを用い、多様なタスクからのオフラインデータで訓練する。
- SACなどのオフポリシー学習アルゴリズムと統合し、オンポリシー微調整の初期化としてマルチタスク方策を用いる。
- 微調整性能の評価にSACの変種を用い、アブレーションスタディにより報酬再ラベル化と三重項損失の両方が必要不可欠であることを示す。
実験結果
リサーチクエスチョン
- RQ1分布シフト下でハードネガティブマイニングを伴うメトリクス学習が、マルチタスクバッチRLにおけるタスク推論の頑健性を向上させ得るか?
- RQ2他のタスクの報酬を再ラベルすることで、方策が状態-行動ペアとタスクIDの間でスパurious相関を学習するのを防げるか?
- RQ3本手法は、AntDir, AntGoal, HDir, UGoal, HalfCheetahVel, WalkerParam などのすべての評価タスク分布において、標準的なオフラインRLベースライン(PEARL やコンテキスト付きBCQ)を上回るゼロショット一般化性能を示すか?
- RQ4マルチタスク方策が、新しいタスクにおける微調整の初期化としてどれほど有効であるか?
- RQ5三重項損失、報酬再ラベル化、マルチタスク学習の各コンポonentが全体の性能に果たす寄与度はどの程度か?
主な発見
- 提案手法MBMLは、未学習タスクにおける微調整において、ランダム初期化された方策と比較して最大80%のサンプル効率の向上を達成する。
- アブレーションスタディの結果、三重項損失または報酬再ラベル化を削除すると顕著な性能低下が生じ、両者の必要性が裏付けられる。
- AntDir, AntGoal, HDir, UGoal, HalfCheetahVel, WalkerParam などのすべての評価タスク分布において、PEARL やコンテキスト付きBCQといった強力なベースラインを上回る性能を示す。
- 推論時に真のタスクID や報酬関数にアクセスできない状況でも、本手法は未学習タスクへの一般化が良好に機能する。
- MBMLモデルの訓練時間はタスク分布によって3.9~26.6時間の範囲にのぼるが、MBMLで初期化されたSACの微調整は著しく高速化される。
より良い研究を、今すぐ始めましょう
論文の読解から最終レビューまで、研究時間を劇的に削減しましょう。
クレジットカード登録不要
このレビューはAIが作成し、人間の編集者が確認しました。