Skip to content

Trust Guided Decision Transformer

Conference: NeurIPS2026
arXiv: 2609.31586
Area: Reinforcement Learning
Keywords: offline reinforcement learning, Decision Transformer, context reliability, conformal calibration, value guidance

TL;DR

TGDT first filters trustworthy history suffixes using rolling next-state prediction errors on realized transitions, then ranks their candidate actions with a frozen IQL critic, improving return and persistent-error behavior in long-horizon navigation without guaranteeing closed-loop coverage or safety.

Background & Motivation

Decision Transformer (DT) formulates offline reinforcement learning as sequence prediction conditioned on returns-to-go, states, and past actions. During training, input windows come from recorded trajectories; during deployment, the history gradually becomes a trajectory induced by the policy's own actions. Even after the training loss converges, this self-generated history can drift away from the offline window distribution over a long episode. The problem is then not merely a poor current action: the model is making decisions from a context insufficiently supported by its training data, which the authors call rollout context mismatch.

Value-guided methods add a critic to improve actions given a history, behavior regularization limits action drift, and methods such as Elastic Decision Transformer (EDT) allow online history-length adaptation. However, a high Q value does not establish that the context generating an action is reliable. Once a long history drifts, the critic can still prefer its candidate action, allowing an unsupported conditioning input to influence subsequent control. Conversely, clearing the history whenever error increases can discard information needed for trajectory stitching.

The authors first diagnose vanilla DT with a separate, frozen transition probe, observing sustained high-error segments rather than only isolated spikes, and then integrate an online reliability signal into DT. Core Idea: decide which contexts are trustworthy before deciding which actions have the highest value, filtering contexts with offline-calibrated historical errors and applying value ranking only within the trusted set.

Method

Overall Architecture

TGDT retains DT's return-conditioned policy but changes its auxiliary training components and the order of inference-time context selection. Offline, an IQL critic is trained and frozen, followed by a DT with a bounded residual action head and an internal reliability head; held-out trajectories provide the prediction-error threshold. At deployment, each step reads several recent suffixes of the same realized trajectory, rejects unreliable suffixes using error records through the previous step, and lets the critic select the highest-value action among the survivors.

The default maximum context length is \(K=20\), with candidate set \(\mathcal{L}=\{1,5,10,20\}\). These suffixes are not four simulated trajectories, but four truncated views of one realized history. The next-state head is only a reliability sensor: it does not generate model rollouts or synthetic training transitions, and no policy is optimized inside the learned dynamics model.

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    D["Offline trajectories"] -.->|Training supervision: recorded actions and next states| A["Near-data actions and<br/>internal reliability head"]
    D -.->|Held-out trajectories: teacher-forced errors| B["Offline error calibration"]
    H["Realized history and return-to-go"] --> A
    A -.->|Offline rolling errors| B
    A -->|Candidate action from each suffix| C["Trust before value"]
    B -.->|Fixed threshold| C
    Q["Pretrained and frozen<br/>IQL critic"] -.->|Training value term| A
    Q -->|Q values of trusted candidates| C
    C -->|Execute one action| E["Environment reveals the real next state"]
    E --> F["Same-action error feedback"]
    F -->|Update queues: filter next step| C

Dashed edges denote offline supervision, calibration, or frozen training signals; solid edges denote inference-time data flow. The external probe used to diagnose vanilla DT is absent from this deployment path, which uses the internal reliability head. The four contribution nodes use the same names as the key designs below.

Key Designs

1. Near-data actions and internal reliability head: make error more indicative of context problems

If the action head drifts far from offline actions in pursuit of high critic scores, state-prediction error can increase because of action-distribution shift rather than corrupted history. TGDT therefore starts from a base action prediction, bounds its residual with \(\delta_{\max}\tanh(\cdot)\), and applies tanh to the sum of the base prediction and residual. Behavior cloning pulls the final action toward the recorded action, a residual penalty limits local corrections, and a small value term favors higher-value behavior. These constraints reduce confounding in the error signal; they do not guarantee that actions remain within the dataset's support.

The internal reliability head predicts the next state from the action token's hidden representation and receives a state-dimension-normalized mean squared error loss. Sharing the context representation with action learning lets the trained model propose actions and predict transitions from the current history. The paper's Eq. 4 writes next-state prediction from the action-token representation, whereas runtime Eq. 12 explicitly conditions prediction on the executed action. The paper does not fully explain the token injection, recomputation, or caching that connects these descriptions, so this implementation detail remains unresolved rather than being filled in with an invented network connection.

The IQL critic is pretrained on offline data and remains frozen thereafter. During DT training, its value signal can influence the action head through the proposed action, but the critic parameters do not update; during inference, it compares a finite set of DT proposals rather than searching unconstrained continuous action space. The external probe establishes that vanilla DT also exhibits context mismatch: it is neither a teacher for the internal head nor an additional model required at deployment.

2. Offline error calibration: establish a data-grounded scale for sustained anomalies, not a closed-loop guarantee

State dimensionality, dynamics difficulty, and predictor fit differ across tasks, making a shared arbitrary absolute-error threshold inappropriate. TGDT uses teacher forcing on held-out offline trajectories, computes normalized next-state prediction error per step, and averages the most recent \(K_e=10\) errors. This smooths isolated spikes and emphasizes persistent deviations. It also calibrates the same rolling statistic used during deployment rather than thresholding individual errors directly.

Given \(n\) held-out rolling scores, the paper uses the split-conformal order statistic:

\[ \tau_\alpha=S^{\mathrm{val}}_{(\lceil(n+1)(1-\alpha)\rceil)},\qquad \alpha=0.05. \]

This finite-sample calibration rule is approximately the 95th percentile of held-out errors, but it does not imply that only 5% of deployment steps will be unreliable. Closed-loop transitions depend on previous actions, overlapping windows introduce temporal correlation, and evaluation actions differ from those of the behavior policy. The authors explicitly make no online coverage claim: the threshold is an empirical reliability reference grounded in offline data, to be judged by return and sustained above-threshold errors.

3. Trust before value: restrict the context set before allowing critic ranking

At time \(t\), the model independently reads each candidate suffix and proposes an action. Each suffix also maintains its own historical error queue. Filtering can use only errors observed through \(t-1\), not a next state that has yet to occur. Let \(e_j^{(L)}\) denote suffix \(L\)'s normalized prediction error on the realized transition. The selection rule is:

\[ S_{t-1}^{(L)}=\frac{1}{K_e}\sum_{j=t-K_e}^{t-1}e_j^{(L)},\qquad \mathcal{A}_t=\{L\in\mathcal{L}:S_{t-1}^{(L)}\leq\tau_\alpha\},\qquad L_t^*=\arg\max_{L\in\mathcal{A}_t}Q_\phi(s_t,\hat a_t^{(L)}). \]

This is a hard gate, not an error penalty subtracted from Q values. An action generated by an unreliable context cannot enter normal ranking even if it receives the highest Q value. The controlled critic-only comparison uses the same trained model, frozen critic, and candidate lengths, changing only whether filtering happens first. It therefore isolates the ordering contribution relatively directly. Critic-only captures value-driven elastic context selection, but is not a fully matched reproduction of EDT itself.

If every suffix is rejected, TGDT falls back to \(L=1\) so that an action is always defined. This shortest context may itself be unreliable, so the fallback is not a safety controller or proof that a reliable action always exists. The main text does not sufficiently specify error-queue initialization or boundary handling before a full window is available; no implementation convention is invented here.

4. Same-action error feedback: every suffix evaluates the same transition that actually occurred

After the selected action is executed, the environment reveals the true next state. Following Eq. 12, each suffix's reliability prediction should condition on that executed action and be compared against the same observed state. Unselected suffixes do not execute their own candidate actions, so their candidate-conditioned predictions cannot simply be compared with the observed state and treated as errors of actions never taken.

\[ e_t^{(L)}=\frac{1}{d_s}\|\tilde s_{t+1}^{(L)}-s_{t+1}\|^2. \]

Appending these errors produces \(S_t^{(L)}\) for the next decision. They are different context interpretations of the same realized trajectory, not counterfactual rollouts. This separates choosing the current action from past evidence and updating evidence after execution, avoiding future leakage. How the internal head receives the executed action remains subject to the implementation ambiguity between Eq. 4 and Eq. 12 noted above.

A Worked Example

Consider a Maze2D navigation decision after the rollout has drifted from offline trajectories. Suppose the historical rolling errors for lengths 20 and 10 exceed the threshold, while those for lengths 5 and 1 do not. Even if length 20 proposes the action with the highest Q value overall, it is rejected first. The critic compares only the length-5 and length-1 candidates; if the former has higher value, it is executed rather than unconditionally resetting to the shortest history.

When the environment reveals the next state, all four suffixes update their errors around that one executed action. If the length-10 or length-20 scores later fall below the threshold, they can re-enter ranking; long histories are not permanently disabled. This is an illustrative flow, not an additional numerical experiment, and explains why TGDT is not equivalent to hard reset.

Loss & Training

The training objective combines behavior imitation, frozen-critic value guidance, residual regularization, and state prediction:

\[ \mathcal{L}=w_t\|\hat a_t-a_t\|^2-\lambda_Q Q_\phi(s_t,\hat a_t)+\beta\|\Delta_t\|^2+\lambda_s\frac{1}{d_s}\|\hat s_{t+1}-s_{t+1}\|^2. \]

The weight \(w_t\) can use normalized advantages; this is an optional training recipe, not evidence that every reported result necessarily uses one identical advantage-weighting variant. Defaults are \(\lambda_Q=0.01\), \(\lambda_s=1.0\), \(\delta_{\max}=0.05\), and \(\beta=0.05\), keeping value corrections local while training a useful reliability head.

Appendix A.1 specifies a 3-layer Transformer with 1 attention head, 128-dimensional embeddings, batch size 64, AdamW, learning rate and weight decay both \(10^{-4}\), and gradient clipping at 0.25. IQL uses expectile 0.7 and discount 0.99; the critic is frozen after 20,000 training steps. These components improve candidate quality, but the runtime rule determines which history conditions each evaluation action.

Key Experimental Results

Main Results

The evaluation covers state-based D4RL Maze2D, AntMaze, MuJoCo, Adroit, and Kitchen. The authors' methods use 3 independent training seeds and 100 evaluation episodes per seed, reporting means and standard errors across seeds, not standard deviations; internal execution modes use paired evaluation seeds. External baselines primarily use their papers' best published results, with missing entries reproduced by the authors, so the cross-method table is not a strictly uniform retraining comparison.

The following representative results come from Table 1. D4RL normalized scores are not a universal success-rate measure; the MuJoCo row averages 9 tasks.

Task DT VDT TGDT Comparison boundary
maze2d-umaze-v1 31.0 \(88.0\pm4.6\) \(92.0\pm0.4\) TGDT exceeds VDT here
maze2d-medium-v1 8.2 \(60.3\pm0.5\) \(66.1\pm0.9\) Substantial long-horizon navigation improvement
antmaze-umaze-v0 59.2 \(100.0\pm5.5\) \(98.8\pm2.3\) TGDT has a lower mean than VDT
antmaze-umaze-diverse-v0 66.2 \(100.0\pm4.7\) \(95.2\pm3.1\) TGDT has a lower mean than VDT
antmaze-medium-diverse-v0 7.5 \(30.0\pm2.8\) \(60.0\pm5.2\) Still below IQL's 70.0
MuJoCo average 76.2 84.1 86.3 Higher average does not imply winning every task
hopper-medium-replay-v2 82.7 \(96.0\pm1.9\) \(95.5\pm2.2\) TGDT has a lower mean here
pen-human-v1 79.5 \(126.7\pm4.3\) \(123.2\pm5.2\) TGDT also does not exceed VDT here

The Maze2D average row in Table 1 contains arithmetic inconsistencies: COMBO's two entries are 76.4 and 38.5, but its reported average is 72.5; DC's entries are 20.1 and 38.2, but its reported average is 57.6. Their simple two-row averages would be 57.45 and 29.15, respectively. The source values and inconsistencies are retained rather than replacing author-reported results with recalculated values or using these averages to claim superiority.

Ablation Study

This separate table uses the internal maze2d-medium-v1 ablation in Appendix Table 5, not the 66.1 result in Table 1. The parenthetical source metric is the fraction of steps where the selected context's rolling error exceeds the threshold; it is not the longest consecutive above-threshold run.

Training / execution configuration Normalized score, mean \(\pm\) standard error Above-threshold step fraction
DT / none \(8.7\pm0.8\) 45.2%
DT + SP / none \(10.2\pm0.9\) 40.9%
DT + Critic / none \(35.9\pm1.7\) 37.6%
DT + Critic + SP / none \(56.7\pm1.5\) 29.8%
Same full model / full context \(56.3\pm1.4\) 30.1%
Same full model / hard reset \(58.2\pm1.6\) 11.5%
Same full model / critic-only \(60.1\pm1.3\) 28.7%
Same full model / TGDT \(68.6\pm1.5\) 8.6%

In Table 5, TGDT improves over critic-only by 8.5 points on medium, while umaze is 91.3 versus 87.1, a 4.2-point gain. The appendix prose summarizes these gains as “5–7 points,” which does not fully match the table. TGDT scores also differ between Table 1 and Table 5 without a sufficiently explicit explanation, so their sources remain separate. The full model's training-side none and execution-side full-context entries likewise differ slightly and are not forced into agreement.

Key Findings

  • Adding state prediction alone raises the medium score only from 8.7 to 10.2. The full trained model reaches 56.7 but still has 29.8% above-threshold steps. Measuring error does not by itself control context.
  • The same-model critic-only versus TGDT comparison best isolates filtering: the score rises from 60.1 to 68.6 and the violation fraction falls from 28.7% to 8.6%. This does not imply return improves on every task: Table 5 reports 111.3 for TGDT versus 112.0 for critic-only on hopper-medium-expert.
  • Figure 2 describes a substantial reduction in medium's longest consecutive violation run; the introduction gives an approximate factor of 8. The cache lacks the figure's complete numerical values, so this approximate factor is not converted into a precise table entry.
  • Appendix A.2 reports Pearson correlation 0.82, threshold agreement 0.89, and probe-positive recall 0.78 between the external probe and internal head on medium. Each uses its own offline threshold; they should not share one absolute-error threshold.
  • Appendix A.6 measures \(0.62\pm0.03\) ms/step for full context, \(2.35\pm0.08\) for critic-only, and \(2.48\pm0.09\) for four-suffix TGDT on an RTX A4500 20GB. Default TGDT takes approximately 4 times the full-context runtime; dense 20-suffix selection takes \(11.4\pm0.4\) ms/step.

Highlights & Insights

  • The method makes the critic's boundary explicit: value estimation ranks acceptable candidates rather than certifying their conditioning contexts. The improvement comes from decision ordering, not a larger model or new critic.
  • Same-executed-action feedback across suffixes can transfer to other history-conditioned policies. It compares memory lengths without pretending that the environment executed multiple actions simultaneously.
  • The external probe and internal head have distinct roles: the former provides architecture-independent diagnostic evidence, while the latter supplies the runtime signal. This supports the diagnosis more strongly than reporting only the proposed auxiliary head's error, but does not prove that context mismatch is the sole cause of error.

Limitations & Future Work

  • The offline conformal-style threshold has no closed-loop coverage guarantee. Correlated windows, policy-induced state distributions, and changed action distributions violate exchangeability; unimplemented non-exchangeable calibration guarantees should not be attributed to TGDT.
  • Error can also reflect stochastic dynamics, predictor bias, or action-distribution changes. Near-data action constraints reduce confounding but do not make error a necessary and sufficient test for context mismatch.
  • Evaluation is limited to state-based D4RL, without image observations, real robots, nonstationary dynamics, or strict latency budgets. A small absolute multi-suffix latency does not make its relative cost negligible.
  • The shortest-suffix fallback has no safety certificate. Queue startup and the mechanism for explicitly conditioning the internal head on the executed action are underspecified, affecting independent reproduction.
  • Average-row arithmetic, main-table versus internal-ablation differences, and the “5–7 points” summary require clarification. Uniform baseline reruns and more training seeds would strengthen robustness.
  • vs DT: DT predicts return-conditioned actions from history; TGDT adds an internal reliability sensor and online suffix filtering. The central change is managing conditioning inputs at deployment, not merely improving imitation loss.
  • vs EDT / critic-only: Both allow changing history length, but TGDT filters by error before ranking by Q. The strictly controlled baseline is the authors' own critic-only variant; its shortcomings should not automatically be attributed to the complete EDT method.
  • vs TD3+BC / ReBRAC / VDT: Behavior regularization and value guidance are established training ideas. TGDT uses them to maintain near-data candidates and then adds reliability gating; the training recipe itself is not its main claimed novelty.
  • vs MOPO / MOReL / COMBO: These methods use dynamics models for model-based policy learning or rollouts. TGDT reads next-state predictions only as reliability signals and does not replace the environment with predicted trajectories. It belongs to offline reinforcement learning with sequence policies and runtime context control, not model-based planning.

Rating

  • Novelty: 4/5 — The trust-before-value ordering is clear; most auxiliary ingredients are established.
  • Experimental Thoroughness: 3/5 — Multiple domains, paired execution ablations, and runtime analysis are useful, but evaluation uses only 3 seeds and state inputs, with numerical inconsistencies.
  • Writing Quality: 3/5 — The diagnosis-to-intervention narrative is clear; action-conditioning implementation and several table–text discrepancies need clarification.
  • Value: 4/5 — A reusable context-reliability approach for long-horizon sequence policies, not a safety guarantee.