Decision Transformer performance degrades on long rollouts because the conditioning context drifts out of the training distribution. We show that this drift is visible through the model's own next state prediction error, which rises during rollout and stays elevated, giving a direct signal of when context has become unreliable. We introduce Trust Guided Decision Transformer (TGDT), which selects context before applying value guidance. At each step, TGDT evaluates several recent context suffixes using rolling next state prediction error, calibrated against held out offline data via split conformal prediction. It keeps only suffixes whose error stays within the calibrated threshold, then uses a frozen critic to choose the highest value action among the trusted suffixes. This reverses the order used by value only elastic selection, where the critic may choose an action generated from a context the model itself has flagged as unreliable. Experiments on D4RL navigation and locomotion tasks show that state prediction, critic guidance, and hard context reset each solve only part of the problem. TGDT reduces persistent high error runs and improves return over vanilla Decision Transformer, reset based context control, and value only context selection.
Figures & tables
Dataset
Value-Based Methods
Conditional Sequence Modeling Methods
Gym Tasks
BEAR
BCQ
CQL
IQL
MoRel
BC
DT
StAR
GDT
CGDT
DC
VDT
TGDT
halfcheetah-medium-replay-v2
38.6
34.8
37.5
44.1
40.2
36.6
36.6
36.8
40.5
40.4
41.3
39.4 ± 2.0
45.1 ± 1.7
hopper-medium-replay-v2
33.7
31.1
95.0
92.1
93.6
18.1
82.7
29.2
85.3
93.4
94.2
96.0 ± 1.9
95.5 ± 2.2
walker2d-medium-replay-v2
19.2
13.7
77.2
73.7
49.8
32.3
79.4
39.8
77.5
78.1
76.6
82.3 ± 2.1
83.9 ± 2.7
halfcheetah-medium-v2
41.7
41.5
44.0
47.4
42.1
42.6
42.6
42.9
42.9
43.0
43.0
43.9 ± 0.7
45.9 ± 0.9
hopper-medium-v2
52.1
65.1
58.5
63.8
95.4
52.9
67.6
59.5
77.1
96.9
92.5
98.3 ± 0.1
99.2 ± 0.3
Table 1 : Offline D4RL performance. Scores are normalized D4RL scores. TGDT denotes our trust-filtered execution rule. See Section 4.1 for the evaluation protocol and baseline sourcing.
Figure 1 : Rollout context mismatch and TGDT’s execution-time response. Left, vanilla DT shows persistent prediction error measured by a frozen diagnostic transition probe. Right, TGDT uses the reliability signal to adapt context length and reduce persistent high-error stretches.
Figure 2 : Ablation study on Maze2D medium . Left, training-side components improve the model but do not eliminate persistent high-error stretches. Right, execution-time reliability filtering is necessary to reduce persistent mismatch. Both panels report normalized D4RL score and the longest consecutive violation run.
Appendix figures & tables7 assets
Supplementary material from the paper’s appendix.
Appendix
Hyperparameter
Value
Maximum context length K
20
Candidate suffix set L
{1,5,10,20}
Reliability window Ke
10
Risk level α
0.05
Held-out threshold percentile
95th percentile
Hard-reset cooldown
10 steps
Appendix
Table 2 : Default hyperparameters used for TGDT. The first block lists execution-time reliability parameters. The second block lists critic and behavior-regularization parameters. The third block lists transformer training parameters. Unless otherwise stated, all sensitivity experiments vary one hyperparameter while keeping the remaining values fixed at the defaults shown here.
Dataset
Pearson r
Spearman ρ
Threshold agreement
Probe-positive recall
maze2d-umaze-v1
0.88
0.81
0.86
0.85
maze2d-medium-v1
0.82
0.75
0.89
0.78
Appendix
Table 3 : Agreement between the external diagnostic probe and TGDT’s internal reliability head on Maze2D tasks. Both signals are evaluated on the same closed-loop rollout trajectories. Each signal is compared against its own held-out reliability threshold, computed with α=0.05 . Higher values indicate stronger agreement between the diagnostic signal used to motivate TGDT and the internal signal used by TGDT at runtime.
Figure 3 : Sensitivity to the reliability window Ke and risk level α on Maze2D-medium and Maze2D-umaze. Each setting is evaluated over three independently trained seeds with 100 evaluation episodes per seed. Blue curves show normalized D4RL score, with shaded 95% confidence intervals across seeds. Orange dashed curves show the mean intervention rate per 1000 environment steps, where an intervention is a timestep with selected context length Lt<K . For Ke , intermediate windows perform best, while very short windows overreact and very long windows react late. For α , larger values correspond to lower held-out thresholds and more aggressive filtering. The default setting Ke=10,α=0.05 lies in a stable region.
Figure 4 : Training-side hyperparameter sensitivity on Maze2D tasks. Each subfigure contains four panels, varying one hyperparameter at a time while keeping the other three fixed at their default values: λQ=0.01 , λs=1.0 , δmax=0.05 , and β=0.05 . Blue curves show normalized D4RL score, with shaded uncertainty across three independently trained seeds. Orange dashed curves show violation fraction, defined as the fraction of evaluation timesteps where the rolling reliability score exceeds the held-out threshold. Across both Maze2D tasks, moderate critic guidance, a trained reliability head, and bounded residual corrections improve return while keeping rollout reliability stable. Overly large critic weight or residual bound can increase action drift, raise violation fraction, and reduce performance.
Figure 5 : Context-length usage on Maze2D tasks. Bars show the fraction of evaluation timesteps assigned to each candidate suffix length in L={1,5,10,20} . TGDT uses multiple suffix lengths rather than collapsing to the shortest context, indicating that the method performs reliability-aware context selection rather than constant hard reset.
Mode
Forward passes / step
ms / step
Relative time
Est. time / 100 eps
Full context
1
0.62±0.03
1.0×
37.2 s
Value only, ∣L∣=4
4
2.35±0.08
3.8×
2.35 min
TGDT, ∣L∣=4
4
2.48±0.09
4.0×
2.48 min
TGDT dense, ∣L∣=20
20
11.4±0.4
18.4×
11.4 min
Appendix
Table 4 : Evaluation-time overhead on maze2d-medium-v1 . Estimated time is computed for 100 evaluation episodes with 600 steps per episode, giving 60k environment steps.
Training ablations under none
Execution ablations on DT + Critic + SP
Dataset
DT
DT + SP
DT + Critic
DT + Critic + SP
Full context
Hard reset
Critic only
TGDT
MuJoCo medium-replay
halfcheetah-medium-replay-v2
35.2 ± 0.9 (31.4)
36.7 ± 0.8 (27.6)
40.1 ± 1.0 (29.8)
40.3 ± 0.9 (24.9)
39.8 ± 0.8 (25.1)
37.6 ± 1.2 (10.7)
40.8 ± 0.9 (23.6)
45.9 ± 1.1 (8.9)
hopper-medium-replay-v2
82.7 ± 1.2 (24.8)
83.1 ± 1.0 (22.1)
93.2 ± 0.9 (21.6)
93.8 ± 0.8 (18.4)
93.6 ± 0.9 (18.7)
89.5 ± 1.5 (7.9)
93.9 ± 0.8 (17.3)
94.8 ± 0.8 (6.8)
walker2d-medium-replay-v2
80.2 ± 1.4 (21.5)
79.1 ± 1.5 (20.2)
81.7 ± 1.2 (19.1)
82.1 ± 1.1 (16.8)
81.9 ± 1.2 (17.0)
83.1 ± 1.4 (7.5)
79.9 ± 1.6 (18.6)
84.4 ± 1.3 (6.9)
MuJoCo medium
Appendix
Table 5: Internal ablation across all evaluated tasks. Each entry reports normalized D4RL score, with the selected-suffix violation rate in parentheses. The violation rate is the percentage of evaluation timesteps at which the rolling next-state prediction error of the executed context exceeds the held-out reliability threshold τα . The first four columns are training-side ablations evaluated with standard full-context execution, denoted none : vanilla DT, DT with a state-prediction head, DT with critic-guided residual action prediction, and DT with both critic guidance and state prediction. The last four columns are execution-time ablations using the same trained DT + Critic + State Prediction model. Full context always uses the maximum context length K and disables execution-time control. Hard reset uses the full context unless the previous rolling score exceeds τα , in which case it resets the next decision to the shortest context L=1 . Critic only evaluates the same candidate suffix lengths L={1,5,10,K} as TGDT, but ranks all suffixes directly by Qϕ(st,a^t(L)) without checking reliability. TGDT first filters candidate suffixes using St−1(L)≤τα , then applies critic ranking only among the trusted suffixes. This table separates gains from training-side components from gains due to trust-guided execution.
Decision Transformer (DT) formulates offline reinforcement learning as autoregressive sequence modeling, achieving promising results by predicting actions from a sequence of Return-to-Go (RTG), state, and action tokens. However, RTG is a scalar that summarizes future rewards, containing far less information than typical state or action vectors, yet it consumes the same computational budget per token. Worse, the self-attention cost of Transformers grows quadratically with sequence length, so including RTG as a separate token adds unnecessary overhead. We propose SlimDT, which removes RTG from the autoregressive sequence. Instead, we inject RTG information into the state representations before the sequential modeling step, allowing the Transformer to process only a compact (state, action) sequence. This reduces the sequence length by one-third, directly improving inference efficiency. On the D4RL benchmark, SlimDT surpasses standard DT across various tasks and achieves performance comparable to existing state-of-the-art methods. Decoupling a sparse conditioning signal from an information-rich sequence thus yields both computational gains and higher task performance.
Yongyi Wang, Hanyu Liu, Lingfeng Li +6
School of Computer Science Peking University Beijing, China 100871
We investigate the ability of transformers to perform in-context reinforcement learning (ICRL), where a model must infer and execute learning algorithms from trajectory data without parameter updates. We show that a linear self-attention transformer block can provably implement policy-improvement methods, including semi-gradient SARSA and actor-critic, via explicit parameter constructions. Beyond existence, we design a teacher-mimicking training procedure, analyze its gradient-flow dynamics, and establish the first convergence guarantee in the ICRL literature: under suitable richness conditions on the training MDP distribution, gradient flow converges locally and exponentially to an optimal parameter manifold corresponding to the desired RL update. Empirically, training transformers on randomly generated tabular MDPs confirms these predictions: the learned models recover the parameter structure of our explicit constructions and, when deployed on unseen MDPs, deliver strong in-context control performance. Together, these results illuminate how transformer architectures internalize and execute classical reinforcement learning algorithms in context, bridging mechanistic understanding and training dynamics in ICRL.
Haodong Liang, Lifeng Lai
Department of Electrical and Computer Engineering, University of California, Davis
Obtaining the optimal action-value function in Markov decision processes is computationally intensive in large state--action spaces. In this study, we present statistically rigorous convergence results for a robust reinforcement learning algorithm warm-started by a transformer-based action-value function prediction, where natural language prompts encode task specifications. Our framework adopts the R-contamination model to characterize uncertainty in the state transition kernel, and employs conformal prediction to certify convergence via trajectory-level nonconformity scores constructed from the contracting Bellman residual. The resulting conformal quantile bounds the gap between the running and optimal action-value functions simultaneously over all iterations, thereby yielding a pre-certified stopping rule that requires little knowledge of the true transition kernel. Numerical case studies on perturbed maze environments of varying size and contamination level confirm that the transformer-based warm start measurably reduces the initial error and accelerates convergence, while the proposed conformal bounds track the true error trajectory more tightly than existing guarantees.
Suman Banerjee, Hiroyasu Tsukamoto
Department of Aerospace Engineering, The Grainger College of Engineering, University of Illinois Urbana-Champaign, Urbana, IL, USA