Skip to content
AI360Xpert
Beta

Decision Transformers and Trajectory Transformer

Instead of calculating values with recursive Bellman equations, Decision Transformers reframe reinforcement learning as language modeling—generating expert actions by conditioning on desired return-to-go.

Decision Transformers frame offline RL as autoregressive sequence modeling, predicting actions conditioned on return-to-go, states, and past context without Bellman updates.
Decision Transformers frame offline RL as autoregressive sequence modeling, predicting actions conditioned on return-to-go, states, and past context without Bellman updates.

Why Does This Exist?

For decades, reinforcement learning has been dominated by the dynamic programming paradigm. Algorithms like Q-learning, Actor-Critic, and policy gradients rely on temporal difference (TD) learning, updating value estimates recursively via Bellman backups:

Q(s,a)←r+γmax⁡a′Q(s′,a′)Q(s, a) \leftarrow r + \gamma \max_{a'} Q(s', a')

While mathematically elegant, TD learning suffers from notorious practical instabilities in the offline (batch) setting. Bootstrapping creates a fragile feedback loop where out-of-distribution function approximation errors compound uncontrollably, causing value estimates to diverge and policies to degrade. Practitioners have had to invent complex regularization techniques—such as conservative value penalties (CQL), batch constraints (BCQ), and expectile regressions (IQL)—simply to keep temporal difference updates stable.

Meanwhile, natural language processing underwent a revolution driven by Generative Pre-trained Transformers (GPT). Large causal Transformers demonstrated that predicting the next token in a sequence using multi-head self-attention scales predictably, trains stably via standard supervised cross-entropy or mean squared error (MSE), and captures intricate long-range dependencies without any recursive value bootstrapping.

In 2021, Chen et al. introduced the Decision Transformer (DT), alongside Janner et al.'s Trajectory Transformer (TT). They proposed a radical paradigm shift: reframe reinforcement learning not as dynamic programming, but as conditional autoregressive sequence modeling. By representing an agent's experience as an interleaved sequence of return targets, states, and actions, the agent bypasses Bellman equations entirely. To act expertly, the agent simply conditions the Transformer on a high desired return-to-go and lets autoregressive self-attention generate the exact sequence of actions that yields that result.

Think of It Like This

Prompting a Mystery Novelist

Imagine you want a story where a master detective solves an impossible heist:

Traditional Reinforcement Learning (TD Learning) behaves like an accountant trying to write a novel by evaluating every sentence with a calculator. At every comma, the accountant estimates the expected book sales: "If the detective turns left, expected sales are $40,000; if they turn right, expected sales are $42,000." If the accountant's formula makes a slight overestimation error in chapter 2, that error propagates to chapter 1, corrupting the entire plot until the story falls apart.

Decision Transformer behaves like prompting an elite author (a GPT language model) with the desired ending:

"Write a thriller where the detective catches the thief, recovers the stolen diamond, and achieves a flawless 100/100 score."

The language model does not calculate expected royalty values at every word. Instead, it draws upon thousands of previously read mystery novels (the offline trajectory dataset) and generates the natural, causal sequence of actions and dialogue that logically leads to that requested 100-point climax.

Where the analogy stops: Language models predict discrete word tokens from a vocabulary. Decision Transformers handle multi-modal continuous control trajectories composed of continuous state vectors, real-valued continuous action vectors, and scalar returns-to-go, requiring separate modality projection layers and continuous regression losses.

How It Actually Works

Trajectory Representation as Token Sequences

Instead of decomposing an offline dataset into independent transition tuples (st,at,rt,st+1)(s_t, a_t, r_t, s_{t+1}), the Decision Transformer treats episodes as coherent trajectory sequences.

A trajectory τ\tau of length TT is structured as an interleaved sequence of Return-to-Go, State, and Action tokens:

τ=(R^1,s1,a1, R^2,s2,a2, …, R^T,sT,aT)\tau = \left( \hat{R}_1, s_1, a_1,\, \hat{R}_2, s_2, a_2,\, \dots,\, \hat{R}_T, s_T, a_T \right)

Return-to-Go (R^t\hat{R}_t)

Rather than conditioning on the immediate reward rtr_t, the model conditions on the Return-to-Go (RTG), defined as the future cumulative return remaining from timestep tt until the end of the episode:

R^t=∑t′=tTrt′\hat{R}_t = \sum_{t'=t}^T r_{t'}

Conditioning on R^t\hat{R}_t provides a clear forward-looking objective: "What sequence of actions will accumulate R^t\hat{R}_t reward from this state onward?"

Architecture and Causal Self-Attention

The Decision Transformer adopts a standard GPT-style causal decoder architecture:

Inputs:     [ R̂₁ ]    [ s₁ ]    [ a₁ ]    [ R̂₂ ]    [ s₂ ]    ──> [ â₂ ] (Predict)              │         │         │         │         │Linear Proj: W_R       W_s       W_a       W_R       W_s              │         │         │         │         │Pos Embed:  + pos(1)  + pos(1)  + pos(1)  + pos(2)  + pos(2)              ▼         ▼         ▼         ▼         ▼┌───────────────────────────────────────────────────────────────┐│              Causal Multi-Head Self-Attention                 ││              (Lower-Triangular Mask: Past Only)               │└───────────────────────────────────────────────────────────────┘                                                      │                                                      ▼                                              Action Head W_act                                                      │                                                      ▼                                              Predicted Action â₂

1. Modality Embedding Projections

Because returns, states, and actions live in different vector spaces, each modality is projected into a shared embedding space of dimension dd:

eR,t=WRR^t+bR,es,t=Wsst+bs,ea,t=Waat+bae_{R, t} = W_R \hat{R}_t + b_R, \quad e_{s, t} = W_s s_t + b_s, \quad e_{a, t} = W_a a_t + b_a

2. Timestep Positional Embeddings

Unlike standard Transformers that use token index positions (1,2,3,…,3T1, 2, 3, \dots, 3T), the Decision Transformer uses timestep embeddings based on the environment step tt. All three modality tokens belonging to timestep tt receive the same learned positional vector Wpos(t)W_{\text{pos}}(t):

xR,t=eR,t+Wpos(t),xs,t=es,t+Wpos(t),xa,t=ea,t+Wpos(t)x_{R, t} = e_{R, t} + W_{\text{pos}}(t), \quad x_{s, t} = e_{s, t} + W_{\text{pos}}(t), \quad x_{a, t} = e_{a, t} + W_{\text{pos}}(t)

This allows the model to retain temporal alignment across the episode regardless of context truncation.

3. Causal Attention Mask

A lower-triangular causal attention mask ensures that token predictions at time tt only attend to historical context. Specifically, when predicting action ata_t, the model attends to:

{R^1,s1,a1,…,R^t−1,st−1,at−1,R^t,st}\{\hat{R}_1, s_1, a_1, \dots, \hat{R}_{t-1}, s_{t-1}, a_{t-1}, \hat{R}_t, s_t\}

The prediction of ata_t has access to the current state sts_t and the desired return R^t\hat{R}_t, but never sees future states or the true action ata_t.

Training Objective and Autoregressive Rollout

Training

The model is trained entirely through supervised learning on offline trajectories sampled from dataset D\mathcal{D}. For continuous action spaces, the objective is Mean Squared Error (MSE):

LDT(θ)=Eτ∼D[1T∑t=1T∥at−a^t(R^t,st,τ<t)∥2]\mathcal{L}_{\text{DT}}(\theta) = \mathbb{E}_{\tau \sim \mathcal{D}} \left[ \frac{1}{T} \sum_{t=1}^T \left\| a_t - \hat{a}_t\left(\hat{R}_t, s_t, \tau_{<t}\right) \right\|^2 \right]

For discrete actions, cross-entropy loss over action logits is used. There are no critics, no target networks, and no discount factors γ\gamma.

Autoregressive Inference (Evaluation)

During live deployment in an environment:

  1. Target Specification: The user sets an initial desired target return R^1\hat{R}_1 (e.g., the maximum return in the dataset or an expert performance threshold).
  2. State Observation: The agent observes initial state s1s_1.
  3. Action Generation: The prompt [R^1,s1][\hat{R}_1, s_1] is passed to the Decision Transformer to predict action a1a_1.
  4. Environment Step: The agent executes a1a_1, receiving reward r1r_1 and transition state s2s_2.
  5. Return-to-Go Decay: The agent decrements the return-to-go: R^2=R^1−r1\hat{R}_2 = \hat{R}_1 - r_1
  6. Context Append: The sequence is extended with [a1,R^2,s2][a_1, \hat{R}_2, s_2], and the process repeats autoregressively across a sliding context window of length KK (typically K=20K=20 timesteps).

Contrast with Trajectory Transformer (TT)

While both models frame RL as sequence modeling, the Trajectory Transformer (Janner et al., 2021) differs fundamentally in tokenization and planning:

CharacteristicDecision Transformer (DT)Trajectory Transformer (TT)
Prediction FocusDirect action generation conditioned on return: at∼P(at∣R^t,st,… )a_t \sim P(a_t \mid \hat{R}_t, s_t, \dots)Full joint distribution modeling: P(st,at,rt∣τ<t)P(s_t, a_t, r_t \mid \tau_{<t})
TokenizationContinuous linear projections per modality (R,s,a)(R, s, a)Discrete tokens via uniform binning or VQ-VAE
Planning MechanismFeedforward return-conditioning at inferenceBeam search planning across future predicted trajectory tokens
Role in RLPolicy generator (replaces actor)Generative world model (replaces simulator and dynamics)

Worked numerical example

Let us trace a concrete 3-step trajectory (T=3T=3) through return-to-go calculation, context windowing, action loss computation, and inference decay.

Step 1: Compute Returns-to-Go

Suppose a demonstrated trajectory yields rewards:

  • r1=2.0r_1 = 2.0
  • r2=3.0r_2 = 3.0
  • r3=5.0r_3 = 5.0

The total cumulative return is G=2.0+3.0+5.0=10.0G = 2.0 + 3.0 + 5.0 = 10.0.

The return-to-go for each step is computed backward: R^1=r1+r2+r3=2.0+3.0+5.0=10.0\hat{R}_1 = r_1 + r_2 + r_3 = 2.0 + 3.0 + 5.0 = 10.0 R^2=r2+r3=3.0+5.0=8.0(or R^1−r1=10.0−2.0=8.0)\hat{R}_2 = r_2 + r_3 = 3.0 + 5.0 = 8.0 \quad (\text{or } \hat{R}_1 - r_1 = 10.0 - 2.0 = 8.0) R^3=r3=5.0(or R^2−r2=8.0−3.0=5.0)\hat{R}_3 = r_3 = 5.0 \quad (\text{or } \hat{R}_2 - r_2 = 8.0 - 3.0 = 5.0)

Step 2: Context Input at Timestep t=2t=2

At timestep t=2t=2, the causal context sequence presented to the Transformer consists of 5 tokens:

Contextt=2=[(R^1=10.0), (s1), (a1), (R^2=8.0), (s2)]\text{Context}_{t=2} = \left[ (\hat{R}_1=10.0),\, (s_1),\, (a_1),\, (\hat{R}_2=8.0),\, (s_2) \right]

The Transformer processes these 5 tokens through causal self-attention. The query vector at the final position (s2s_2) attends to all 5 tokens to output the predicted action vector a^2\hat{a}_2.

Step 3: Compute Action Prediction Loss

Suppose the true 2-dimensional continuous action in the dataset is: a2=[0.50, 0.80]a_2 = [0.50,\, 0.80]

The Decision Transformer's linear action head outputs prediction: a^2=[0.45, 0.85]\hat{a}_2 = [0.45,\, 0.85]

The Mean Squared Error (MSE) loss is evaluated: L2=1D∑d=1D(a^2,d−a2,d)2=12[(0.45−0.50)2+(0.85−0.80)2]\mathcal{L}_2 = \frac{1}{D} \sum_{d=1}^D (\hat{a}_{2, d} - a_{2, d})^2 = \frac{1}{2} \left[ (0.45 - 0.50)^2 + (0.85 - 0.80)^2 \right] L2=12[(−0.05)2+(0.05)2]=12[0.0025+0.0025]=0.00502=0.0025\mathcal{L}_2 = \frac{1}{2} \left[ (-0.05)^2 + (0.05)^2 \right] = \frac{1}{2} [ 0.0025 + 0.0025 ] = \frac{0.0050}{2} = 0.0025

This supervised gradient updates Wa,Ws,WR,W_a, W_s, W_R, and the self-attention weights via standard Adam optimizer steps.

Step 4: Inference Return-to-Go Decay

During real-time evaluation:

  1. The user specifies target return R^1=10.0\hat{R}_1 = 10.0.
  2. The agent executes predicted action a1a_1 in the environment.
  3. The environment emits reward r1=2.0r_1 = 2.0.
  4. The agent updates its prompt return-to-go for step t=2t=2: R^2=R^1−r1=10.0−2.0=8.0\hat{R}_2 = \hat{R}_1 - r_1 = 10.0 - 2.0 = 8.0

The agent's internal goal is dynamically updated: it now knows that to satisfy the requested 10.0 overall performance, it must collect exactly 8.0 remaining points over the remainder of the episode.

Code

Below is a self-contained, type-hinted Python implementation of the core Decision Transformer trajectory evaluation engine, including return-to-go computation, causal attention masking, continuous action MSE loss, and autoregressive return decay:

from dataclasses import dataclassimport mathfrom typing import List, Tuple

@dataclassclass TrajectoryToken:    """Represents an individual token in the interleaved trajectory sequence."""
    modality: str  # 'RTG', 'STATE', or 'ACTION'    timestep: int    value: float

class DecisionTransformerTrajectoryEvaluator:    """Evaluates trajectory tokenization, causal masking, return-to-go decay, and action prediction loss."""
    def __init__(self, embedding_dim: int = 16) -> None:        self.embedding_dim = embedding_dim
    def compute_returns_to_go(self, rewards: List[float]) -> List[float]:        """Computes Return-to-go R_hat_t = sum_{t'=t}^T r_{t'} for each timestep."""        t_len = len(rewards)        rtg = [0.0] * t_len        running_sum = 0.0        for i in reversed(range(t_len)):            running_sum += rewards[i]            rtg[i] = running_sum        return rtg
    def construct_token_sequence(        self,        returns_to_go: List[float],        states: List[List[float]],        actions: List[List[float]],    ) -> List[TrajectoryToken]:        """Constructs interleaved token sequence: (R_hat_1, s_1, a_1, R_hat_2, s_2, a_2, ...)."""        tokens: List[TrajectoryToken] = []        for t in range(len(returns_to_go)):            step_idx = t + 1            tokens.append(TrajectoryToken("RTG", step_idx, returns_to_go[t]))            tokens.append(TrajectoryToken("STATE", step_idx, states[t][0]))            tokens.append(TrajectoryToken("ACTION", step_idx, actions[t][0]))        return tokens
    def build_causal_mask(self, num_tokens: int) -> List[List[int]]:        """Builds lower-triangular causal attention mask (1 = attend, 0 = mask)."""        mask = [            [1 if j <= i else 0 for j in range(num_tokens)]            for i in range(num_tokens)        ]        return mask
    def compute_action_loss_mse(        self, pred_action: List[float], target_action: List[float]    ) -> float:        """Computes Mean Squared Error loss between predicted and ground-truth action vectors."""        assert len(pred_action) == len(target_action)        dim = len(pred_action)        sq_diff_sum = sum(            (p - t) ** 2 for p, t in zip(pred_action, target_action)        )        return sq_diff_sum / dim
    def step_inference_rtg(        self, current_rtg: float, received_reward: float    ) -> float:        """Updates return-to-go for the next timestep: R_hat_{t+1} = R_hat_t - r_t."""        return current_rtg - received_reward

if __name__ == "__main__":    evaluator = DecisionTransformerTrajectoryEvaluator()
    # 1. Worked Numerical Example: T=3 trajectory    rewards = [2.0, 3.0, 5.0]    rtg = evaluator.compute_returns_to_go(rewards)
    print("=== Returns-to-Go Calculation ===")    print(f"Step Rewards:    {rewards}")    print(f"Returns-to-Go:   {rtg}")    assert rtg == [10.0, 8.0, 5.0]
    # 2. Context at t=2: Action Loss Evaluation    pred_a2 = [0.45, 0.85]    gt_a2 = [0.50, 0.80]    mse_loss = evaluator.compute_action_loss_mse(pred_a2, gt_a2)
    print("\n=== Action Prediction Loss (MSE) ===")    print(f"Predicted Action a_hat_2: {pred_a2}")    print(f"Ground Truth Action a_2:  {gt_a2}")    print(f"MSE Loss:                 {mse_loss:.6f}")    assert math.isclose(mse_loss, 0.0025)
    # 3. Inference RTG Decay    # Starting at R_hat_1 = 10.0, agent executes a_1 and receives r_1 = 2.0    rtg_t2 = evaluator.step_inference_rtg(        current_rtg=rtg[0], received_reward=rewards[0]    )    print("\n=== Inference Return-to-Go Decay ===")    print(        f"Prompt RTG at t=2: R_hat_2 = R_hat_1 - r_1 = {rtg[0]} - {rewards[0]} = {rtg_t2:.1f}"    )    assert math.isclose(rtg_t2, 8.0)
    # 4. Token Sequence and Causal Mask Verification    mock_states = [[1.0], [1.2], [1.5]]    mock_actions = [[0.5], [0.8], [0.9]]    tokens = evaluator.construct_token_sequence(rtg, mock_states, mock_actions)    causal_mask = evaluator.build_causal_mask(len(tokens))
    print("\n=== Causal Mask Verification ===")    print(f"Total Trajectory Tokens (3 steps x 3 modalities): {len(tokens)}")    print(f"Token 0 ({tokens[0].modality}_1) attend vector: {causal_mask[0]}")    print(f"Token 4 ({tokens[4].modality}_2) attend vector: {causal_mask[4]}")
    assert len(tokens) == 9    assert causal_mask[0] == [1, 0, 0, 0, 0, 0, 0, 0, 0]    assert causal_mask[4] == [1, 1, 1, 1, 1, 0, 0, 0, 0]    print("\nAll Decision Transformer assertions passed successfully!")

Expected output:

=== Returns-to-Go Calculation ===Step Rewards:    [2.0, 3.0, 5.0]Returns-to-Go:   [10.0, 8.0, 5.0]
=== Action Prediction Loss (MSE) ===Predicted Action a_hat_2: [0.45, 0.85]Ground Truth Action a_2:  [0.5, 0.8]MSE Loss:                 0.002500
=== Inference Return-to-Go Decay ===Prompt RTG at t=2: R_hat_2 = R_hat_1 - r_1 = 10.0 - 2.0 = 8.0
=== Causal Mask Verification ===Total Trajectory Tokens (3 steps x 3 modalities): 9Token 0 (RTG_1) attend vector: [1, 0, 0, 0, 0, 0, 0, 0, 0]Token 4 (STATE_2) attend vector: [1, 1, 1, 1, 1, 0, 0, 0, 0]
All Decision Transformer assertions passed successfully!

Watch Out For

The Trajectory Stitching Failure Mode

The Trap: In offline RL benchmarks (such as D4RL AntMaze), an agent often needs to "stitch" parts of multiple suboptimal trajectories. For example, trajectory 1 demonstrates navigating from start AA to checkpoint BB, while trajectory 2 demonstrates navigating from BB to goal CC. Temporal difference algorithms (like CQL or IQL) excel at this: Bellman updates propagate value backwards across states, stitching A→BA \to B and B→CB \to C together into an optimal A→CA \to C policy.

Because the vanilla Decision Transformer performs conditional sequence matching rather than dynamic programming, it cannot stitch fragments without having observed complete trajectories. If the dataset only contains separate, mediocre trajectories with low returns, prompting the model with an expert return R^>Rmax⁡\hat{R} > R_{\max} forces the Transformer out-of-distribution. Instead of discovering an optimal path, the model outputs erratic, uncoordinated actions.

The Symptom: Decision Transformer achieves state-of-the-art performance on continuous locomotion tasks (Gym MuJoCo Hopper/HalfCheetah) where datasets contain smooth, near-expert trajectories, but scores near zero on sparse-reward navigation tasks (AntMaze) that strictly require trajectory stitching.

The Fix:

  1. Hybridize with Value Learning: Use Q-guided Decision Transformers (Q-DT) or Dual-Value Transformers where token sampling is guided by learned temporal-difference critic evaluations.
  2. Data Augmentation: Augment the offline dataset using conservative model rollouts or sub-trajectory recombination before training.
  3. Trajectory Transformer with Beam Search: Use the Trajectory Transformer (TT) instead of DT; because TT explicitly models transition dynamics P(st+1∣st,at)P(s_{t+1} \mid s_t, a_t), beam search planning can stitch disjoint paths during inference.

The Quick Version

  • RL as Sequence Modeling: Decision Transformer replaces recursive Bellman backups with conditional autoregressive sequence modeling, using a GPT-style causal decoder to predict actions from historical context.
  • Interleaved Trajectory Tokens: Episodes are structured as token triplets (R^t,st,at)(\hat{R}_t, s_t, a_t) paired with timestep embeddings, where Return-to-Go R^t=∑t′=tTrt′\hat{R}_t = \sum_{t'=t}^T r_{t'} provides the forward conditioning target.
  • Pure Supervised Training: The model is trained using standard mean squared error (continuous actions) or cross-entropy (discrete actions), eliminating target networks, value critics, and offline Bellman instability.
  • Inference via Return Decay: The agent is prompted with a desired high return R^1\hat{R}_1 and decrements the target dynamically after each environment reward (R^t+1=R^t−rt\hat{R}_{t+1} = \hat{R}_t - r_t), guiding the Transformer to fulfill the overall episode objective.