Bridging Risk Approximation Gaps in Model Predictive Task Sampling via In-Context Modeling
Abstract
Reinforcement learning (RL) is essential for adaptive decision-making, e.g., robust Meta-RL and domain randomization (DR), and efficient post-training of large language models (LLMs), where tasks are commonly structured as Markov decision processes or prompts. A central bottleneck in all these settings is the cost of agent-environment interactions required for task selection and policy optimization. Model predictive task sampling (MPTS) has emerged as a promising approach to improve sample efficiency by querying informative tasks via a risk predictive model (RPM) trained on optimization history. However, existing RPMs suffer from two failure modes: (i) sparse optimization histories, which limit RPMs' generalization across the task space, and (ii) a one-step lag in risk approximation, which causes biased difficulty estimation as the policy evolves. This work formalizes these failure modes, derives lagged-risk decompositions that expose the resulting temporal bias, and recasts risk prediction as an in-context sequence modeling problem. The developed RPM integrates a one-step look-ahead mechanism with a temporal-weighted sliding window to mitigate data non-stationarity and scarcity, without requiring additional rollouts. Empirically, our approach yields more reliable difficulty estimates and consistent performance gains across robust Meta-RL, DR, and prompt curriculum RL post-training of LLMs.