Training a Decision Transformer to Play
Pong Game in C/C++

Materials available at: http://forejune.co/cuda/

Recap

  • Training a GPT Transformer to Play Tic-Tac-Toe

  • Pong Game
  • Decision Transformer

  • GPT (Decoder-only) Transformer models a sequence of tokens.
  • Decision Transformer models a sequence of (return, state, action) triples

    GPT (Decoder-only) Predicts Next Token
     

    Decision Transformer Predicts the Next Action
  • 	The Pong environment generates the next state, while the Decision Transformer generates the action.
    
    	The key assumption is that history always ends with an ACTION query position. 
    	In other words, the last timestep in history has:
    
    	RTG   STATE   ACTION_TOKEN
    
    	so P.back() is the model's prediction for the action.
    	
    Return-To-Go

  • Return-To-Go (RTG): The future reward remaining from the current timestep onward.

  • RTG is the sum of all future rewards from time step t until the episode terminates.

  • Let T = terminal time step of an episode
    (the total trajectory length is T + 1)
    rk = reward at time step k

  • Return-To-Go at time step t:

  • Return-To-Go at time step t+1:

  • Recursive Formulas:

  • Inference Update Rule
    • Environment yields a reward rt
    • The update RTG for the next time step by subtracting the reward

  • Boundary Condition
    When the episode ends, no more reward to sum:

      Theoretically only!
  • Decision Transformer for Pong

    
    	struct PongStruct
    	{
        		double returnToGo;  // RTG
    
        		double ballX;
        		double ballY;
    
      		double ballDX;
        		double ballDY;
    
        		double leftY;       // left paddle position
        		double rightY;
    
        		int action;        // UP, DOWN, STAY
    	};
    		

    C/C++ Implementation

    1. Create header pong.h

    #ifndef _PONG_H_
    #define _PONG_H_
    
    // Totally seven different tokens (we only need action tokens)
    const int UP = 0;
    const int DOWN = 1;
    const int STAY = 2;
    const int EOS = 3;	    // End of stream
    // Add special tokens for Decision Transformer
    const int RTG_TOKEN = 4;    // Not used for embedding lookup
    const int STATE_TOKEN = 5;  // Not used for embedding lookup
    const int ACTION_TOKEN = 6; // Not used for embedding lookup
    
    const int ACTION_OFFSET = 1;  // Actions are stored at indices 1, 2, 3
    			      // index 0 can store -1
    struct PongStruct
    {
        double returnToGo;
    
        double ballX;
        double ballY;
    
        double ballDX;
        double ballDY;
    
        double leftY;	// left paddle position
        double rightY;
    
        int action;        // UP, DOWN, STAY
    };
    #endif
    	

    2. Modify Embedding Class

    class Embedding
    {
    protected:
        int nTokens;        // number of words
        int dim;            // embedding dimension
        matrixd em_matrix;  // embedding matrix;
        // New members for Decision Transformer
        matrixd W_rtg;      // 1 x dim - projects scalar RTG to embedding dim
        matrixd W_state;    // 6 x dim - projects 6 state dims to embedding dim
    
    public:
        Embedding(int sequence_length, int embeddingDimension);
        // create embedding matrix with tokenIDs
        matrixd embed(const vector<int>& tokenIDs);
        int get_embedDim() const;
        int get_nTokens() const;
        void backProp(const vector<int>& tokenIDs, const matrixd& dL_dX,
                             double eta);
        // Added for DecisionTransformer
        matrixd embed(const vector<PongStruct>& seq);
        void backProp(const vector<PongStruct>& tokens, const matrixd& dL_dX,
                             double eta);
        vector project_rtg(const double returnToGo);
        vector project_state(const PongStruct& token);
        void saveWeights(FILE *fp);
        void readWeights(FILE *fp);
    };
    
    Embedding::Embedding(int sequence_length, int embeddingDimension)
    {
      nTokens = sequence_length;
      dim = embeddingDimension;
      initMatrix(em_matrix, nTokens, dim);
    
       // Initialize continuous value projections
       initMatrix(W_rtg, 1, dim);      // 1 x dim
       initMatrix(W_state, 6, dim);    // 6 x dim
    }
    
    // create embedding matrix with a sequence of PongStruct (trajectory)
    matrixd Embedding::embed(const vector& seq) {
        matrixd embeddings;
        embeddings.reserve(seq.size()); //set capacity of of embedding vector
        for (int i = 0; i < seq.size(); i++) {
            embeddings.push_back(project_rtg(seq[i].returnToGo));
            embeddings.push_back(project_state(seq[i]));
            embeddings.push_back(em_matrix[seq[i].action + ACTION_OFFSET]);
        }
    
        return embeddings;
    }
    
    // Project a scalar (return-to-go) to embedding dimension
    vector<double> Embedding::project_rtg(double returnToGo)
    {
       vector<double> result(dim, 0.0);
    
       // result = W_rtg^T x returnToGo
       // W_rtg : 1xdim, input: 1x1, result: dim x 1
       for (int j = 0; j < dim; j++)
         result[j] = returnToGo * W_rtg[0][j];
    
       return result;
    }
    
    // Project state vector to embedding dimension
    vector<double> Embedding::project_state(const PongStruct& pong)
    {
       // Build state vector (6 dimensions)
       vector<double> pong_state = {
          pong.ballX,  pong.ballY,
          pong.ballDX, pong.ballDY,
          pong.leftY,  pong.rightY
        };
    
        int n = pong_state.size();
        // result = W_state^T * pong_state 
        // W_state: 6 x dim, pong_state: 6 x 1, result: dim x 1
        vector<double> result(dim, 0.0);
    
        for (int j = 0; j < dim; j++) {
            result[j] = 0.0;
            for (int i = 0; i < n; i++) {
                result[j] += pong_state[i] * W_state[i][j];
            }
        }
    
        return result;
    }
    
    void Embedding::backProp(const vector<PongStruct>& pongs,
                             const matrixd& dL_dX, double eta)
    {
        // Backward pass for PongStruct
        int nt = pongs.size();
    
        for (int t = 0; t < nt; t++) {
          int rtg_pos = 3*t;        // RTG embedding position
          int state_pos = 3*t + 1;  // State embedding position
          int action_pos = 3*t + 2; // Action embedding position
    
           // Gradient for RTG projection weights
           vector<double> dL_drtg = dL_dX[rtg_pos];  // size: dim
           // dL/dW_rtg = returnToGo^T * dL/d(rtg_embedding)
           // Since input is scalar: dL/dW_rtg[0][j] = returnToGo * dL_drtg[j]
           for (int j = 0; j < dim; j++)
             W_rtg[0][j] -= eta * pongs[t].returnToGo * dL_drtg[j];
    
           // Gradient for state projection weights
            vector<double> dL_dstate = dL_dX[state_pos];  // dim
    
            vector<double> pong_state = {pongs[t].ballX, pongs[t].ballY,
                    pongs[t].ballDX, pongs[t].ballDY,
                    pongs[t].leftY,  pongs[t].rightY
                };
    
            // dL/dW_state = pong_state^T * dL/d(state_embedding)
            for (int i = 0; i < 6; i++)
              for (int j = 0; j < dim; j++)
                        W_state[i][j] -= eta * pong_state[i] * dL_dstate[j];
    
            // Gradient for action embedding (discrete lookup)
             int action_token = pongs[t].action + ACTION_OFFSET;
             vector<double> dL_daction = dL_dX[action_pos];  // dim
    
             // Update the specific row in em_matrix
             for (int j = 0; j < dim; j++)
                em_matrix[action_token][j] -= eta * dL_daction[j];
    
        }
    }

    3. Extend GPT Transformer

    class DecisionTransformer : public MiniGPT 
    {
    public:
        DecisionTransformer(int vocab_size, int seq_len, int dModel,
                           int nHeads, int d_ff, int nLayers);
    
        // Forward pass for Decision Transformer
        matrixd forward(const vector& trajectory);
    
        // Training step for Decision Transformer
        double trainStep(const vector& trajectory, double eta);
        
        // Predict single action
        int predictAction(const vector& history);
    };
    
    // seq_len = 3 x number_of_timesteps
    DecisionTransformer::DecisionTransformer(int vocab_size, int seq_len, int dModel,
                           int nHeads, int d_ff, int nLayers)
            : MiniGPT(vocab_size, seq_len, dModel, nHeads, d_ff, nLayers) 
    {
    }
    
    // Forward pass for Decision Transformer
    matrixd DecisionTransformer::forward(const vector<PongStruct>& trajectory) 
    {					// nt = trajectory.size()
       matrixd X = embed.embed(trajectory); // 3*nt x dModel 
    					// (rtg,state,action) per time step
       X = pe.addPE(X);
    
       for (int i = 0; i < layers.size(); i++) 
           X = layers[i].forward(X);
    
       matrixd P = output.forward(X);  // 3*nt x vocab_size
    
       return P;
    }
    
    double DecisionTransformer::trainStep(const vector<PongStruct>& trajectory, double eta)
    {
        int nt = trajectory.size();
        int total_tokens = nt * 3;  // 3 tokens per timestep (RTG, State, Action)
       
        /*
          Build target sequence:
            1. The model predicts the next token given previous tokens.
            2. We only care about action predictions.
            3. target[i] = token[i+1] for positions where token[i+1] is an action.
            4. target[i] = -1 for positions where we ignore the prediction
        */
        vector<int> target(total_tokens, -1);  // -1 => ignore
    					   //
        for (int t = 0; t < nt; t++) {
            int action_pos = t * 3 + 2;  // action position in triplet (RTG, state, action) 
    
            // Target for position (action_pos - 1) is this action
            // because we predict action given RTG and state,
    	// so target action is the current action 
            if (action_pos > 0  && action_pos - 1 < total_tokens) 
                target[action_pos - 1] = trajectory[t].action;  // t*3 + 1
            
        }
    
        for (int i = 0; i < target.size(); i++) {
           if (target[i] >= vocab_size) {
                cout << "ERROR: target[" << i << "] = " << target[i]
                     << " exceeds vocab_size " << vocab_size << endl;
                exit(1);
           }
        }
    
        // Embedding forward
        matrixd X = embed.embed(trajectory);
    
        // Positional encoding
        X = pe.addPE(X);
    
        // Transformer forward
        vector<matrixd> cache;
    
        for (int i = 0; i < layers.size(); i++) {
            cache.push_back(X);   // Cache input for backward
            X = layers[i].forward(X);
        }
    
        // Output forward
        matrixd P = output.forward(X);
    
        // Loss computation
        double loss = crossEntropy(P, target);
    
        // Output backward
        matrixd dL_dX = output.backward(X, P, target, eta);
    
        // Transformer backward
        for (int i = layers.size() - 1; i >= 0; i--) {
            dL_dX = layers[i].backward(cache[i], dL_dX, eta);
        }
    
        // Embedding backward
        embed.backProp(trajectory, dL_dX, eta);
    
        return loss;
    }
    
    int DecisionTransformer::predictAction(const vector<PongStruct>& history)
    {
        matrixd P = forward(history);
        int last = P.size() - 1;
    
        // Define the mapping from our action enum to the model's logit index
        const int actionTokens[3] = {UP, DOWN, STAY};
        int bestAction = UP;
        double bestValue = -1e9;
    
        for (int i = 0; i < 3; i++ ) // action is 0, 1, 2
        {
    	int action = actionTokens[i];
            int logitIndex = action;
            if (P[last][logitIndex] > bestValue)
            {
                bestValue = P[last][logitIndex];
                bestAction = action; // Store the actual enum (0, 1, or 2)
            }
        }
    
        return bestAction; // Returns UP(0), DOWN(1), or STAY(2)
    }
    
    	

    Collecting Data

     /*
     * pongt.cpp : Gathering data, saving in trajectories.txt
     * http://forejune.co/cuda
     */
    
    #include <GL/gl.h>
    #include <GL/glut.h>
    #include <string>
    #include <vector>
    #include <iostream>
    #include "pong.h"
    
    using namespace std;
    
    // Game variables
    double ballX = 0.0;	// horizontal position of ball
    double ballY = 0.0;	// vertical position of ball
    
    double ballDX = 0.15;	// change in X of ball
    double ballDY = 0.10;	// change in Y of ball
    
    double leftY = 0.0;	// left paddle position
    double rightY = 0.0;	// rifht paddle position
    
    int   action = STAY;
    
    const double pW  = 0.5;	// paddle width
    const double pH = 3.0;   // paddle height
    
    double scoreLeft  = 0;
    double scoreRight = 0;
    double returnToGo = 11;
    double reward = 0;
    int maxTimeSteps = 100;	   // maximum number of time steips in a trajectory
    int maxTrajectories = 400; // maximum number of trajectories to be collected
    int nT = 0;	// number of trajectories recorded
    FILE *fp = NULL;
    
    bool pause = false;
    
    vector<PongStruct> trajectory; 
    
    void drawPaddle(double x, double y)
    {
      // (x, y) is the lower left vertex coordinates
      glRectf(x, y, x+pW, y+pH);
    }
    
    // draw String
    void drawString(double x, double y, const string &s)
    {
        glRasterPos2f(x, y);
    
        for (int i = 0; i < s.length(); i++)
            glutBitmapCharacter(GLUT_BITMAP_HELVETICA_18, s[i]);
    }
    
    void drawBall(double x, double y)
    {
      glPushMatrix();
      glTranslatef(x, y, 0);
      glutSolidSphere(0.5, 16, 16);
      glPopMatrix();
    }
    
    // a trajectory at one time step: {returnToGo, ballX, ballY, ballDX, ballDY, leftY, rightY, action}
    void saveTrajectory()
    {
      for (const auto &pong : trajectory){
        fprintf(fp, "\n%6.3f ", pong.returnToGo);
        fprintf(fp, "%6.3f %6.3f %6.3f %6.3f ", pong.ballX, pong.ballY, pong.ballDX, pong.ballDY);
        fprintf(fp, "%6.3f %6.3f %4d ", pong.leftY, pong.rightY, pong.action);
      }
      fprintf(fp, "\n%6.3f ", -101.0);   // signifies end of trajectory
      ++nT;
      if ( nT >= maxTrajectories ) {
        fclose( fp );
        cout << "Saved " << nT << " trajectories" << endl;
        exit(0);
      }
    }
    
    // Initialization
    void init(void)
    {
        glClearColor(1, 1, 1, 0);
    
        glMatrixMode(GL_PROJECTION);
        glLoadIdentity();
        gluOrtho2D(-10, 10, -10, 10);
    
        glMatrixMode(GL_MODELVIEW);
        glLoadIdentity();
        if ( (fp = fopen("trajectories.txt", "wt")) == NULL ) {
          cout << "Error opening file trajectories.txt" << endl;
          exit ( 1 );
        }
    }
    
    // Draw scene
    void display(void)
    {
        glClear(GL_COLOR_BUFFER_BIT);
    
        double x, y;		//paddle lower left corner
        glColor3f(1, 0, 0); //red color
    
        // Left paddle
        x = -9.0;
        y = leftY - pH/2;
        drawPaddle(x, y);
    
        // Right paddle
        x = 8.5;
        y = rightY - pH / 2;
        drawPaddle(x, y);
    
        // Ball
        glColor3f(0, 1, 0);
        drawBall(ballX, ballY);
    
        // Score
        char str[40];
        sprintf(str,"%d   :   %d", (int)scoreLeft, (int)scoreRight);
        glColor3f(0, 0, 0);  // black color
        drawString(-1.0, 9.0, str);
    
        glutSwapBuffers();
    }
    
    void animate()
    {
        if (pause) return;
        if ( nT >= maxTrajectories ) return;
    
        // Record the current state s(t) and compute r(t)
        // action will be determined below based on this current state.
        // We need to store the state before moving.
    
        // Expert decides action (based on current state)
        if (ballY > leftY) 
            action = UP;
        else 
            action = DOWN; // Expert always moves, no STAY needed for data collection
    
        // Execute the action (move left paddle)
        if (action == UP) leftY += 0.08;
        else leftY -= 0.08;
    
        // Keep left paddle inside the screen
        if (leftY > 10.0 - pH / 2) leftY = 10.0 - pH / 2;
        if (leftY < -10.0 + pH / 2) leftY = -10.0 + pH / 2;
    
        //  Push the current triplet (rtg(t), state(t), action(t)
        PongStruct current_triplet = {
            returnToGo,         // rtg at time t
            ballX,              // state at time t (before moving ball)
            ballY,              // state at time t
            ballDX,             // state at time t
            ballDY,             // state at time t
            leftY,              // state at time t (after paddle move)
    			    //  paddle move is part of the environment transition
            rightY,             // state at time t
            action              // action at time t
        };
        trajectory.push_back(current_triplet);
        if (trajectory.size() > maxTimeSteps)
            trajectory.erase(trajectory.begin());
    
        // Move the right paddle and the ball
        if (ballY > rightY) rightY += 0.08;
        else rightY -= 0.08;
        if (rightY > 10.0 - pH / 2) rightY = 10.0 - pH / 2;
        if (rightY < -10.0 + pH / 2) rightY = -10.0 + pH / 2;
    
        // Move the ball
        ballX += ballDX;
        ballY += ballDY;
    
        // Collisions
        if (ballY > 9.5) { ballY = 9.5; ballDY = -ballDY; }
        else if (ballY < -9.5) { ballY = -9.5; ballDY = -ballDY; }
    
        if (ballX <= -8.25 && ballY >= leftY - pH/2 && ballY <= leftY + pH/2 && ballDX < 0)
            ballDX = -ballDX;
        else if (ballX >= 8.25 && ballY >= rightY - pH/2 && ballY <= rightY + pH/2 && ballDX > 0)
            ballDX = -ballDX;
    
        //  Determine reward, r(t) and handle goals
        reward = 0.0;
        bool scored = false;
    
        if (ballX < -10) {
            scoreRight++;
            reward = -1.0;
            scored = true;
        } else if (ballX > 10) {
            scoreLeft++;
            reward = +1.0;
            scored = true;
        }
    
        // Update rtg for the next timestep: rtg(t+1)
        returnToGo -= reward; // rtg{t+1} = rtg(t) - r(t)
    
        // Episode reset (if goal scored)
        if (scored) {
            PongStruct terminal_step = {
                returnToGo,  // This is now rtg(t+1) = rtg(t) - r(t),
                ballX,              
                ballY,
                ballDX,
                ballDY,
                leftY,
                rightY,
                action      // Action that led to this terminal state
            };
            trajectory.push_back(terminal_step);
            if (trajectory.size() > maxTimeSteps)
                trajectory.erase(trajectory.begin());
    
            // Save the complete trajectory to file
            saveTrajectory();
    
            // Reset ball and clear trajectory for the new episode
            ballX = 0; ballY = 0;
            ballDX = -ballDX;
            
            // Reset rtg for the new episode
            returnToGo = 11.0; 
            
            // Clear trajectory and push the new initial state (with STAY action)
            trajectory.clear();
            PongStruct initial_step = {
                returnToGo,
                ballX, ballY, ballDX, ballDY,
                leftY, rightY,
                STAY  // No action taken yet for the new episode start
            };
            trajectory.push_back(initial_step);
        }
    
        glutPostRedisplay();
    }
    
    // Keyboard control
    void keyboard(unsigned char key,int x, int y)
    {
        switch(key)
        {
            case 27:	// ESC
    	    if ( fp != NULL ) {
    	      fclose( fp );
    	      cout << "File closed! Saved " << nT << " trajectories" << endl;
    	    }
                exit(0);
    	    break;
    	case 'p':	// toggles pause
    	    pause = pause ? false : true;
    	    break;
        }
    }
    
    void specialKey(int key, int x, int y)
    {
        switch(key)
        {
    	case GLUT_KEY_UP:
                leftY += 0.5;
                break;
    
            case GLUT_KEY_DOWN:
                leftY -= 0.5;
                break;
        }
    }
    
    // Visibility callback
    void timerHandle ( int value )
    {
       animate();
       glutPostRedisplay();
       // call timerHandle 25 ms later, 0 is passed to timerHandle, not used here 
       glutTimerFunc (25, timerHandle, 0);
    }
    
    void visHandle( int visible )
    {
       if (visible == GLUT_VISIBLE)
          timerHandle ( 0 );
       else
          ;
    }
    
    int main(int argc, char *argv[])
    {
        glutInit(&argc, argv);
    
        glutInitDisplayMode(GLUT_DOUBLE | GLUT_RGB);
    
        glutInitWindowSize(500,500);
        glutCreateWindow("Pong Game");
    
        glutDisplayFunc(display);
        glutVisibilityFunc(visHandle);
        glutKeyboardFunc(keyboard);
        glutSpecialFunc(specialKey);
    
        init();
    
        glutMainLoop();
    
        return 0;
    }
    
    	

    Training the Transformer

    // testMain.cpp : Training DecisionTransformer
    // http://forejune.co/cuda/
    
    #include "util.h"
    #include "transGpt.h"
    #include "decision.h"
    
    using namespace std;
    
    struct Database {
        vector<vector<PongStruct>> sequences;
    };
    
    void addSample(Database &db, const vector<PongStruct> &seq)
    {
        if (seq.size() < 2)
            return;
    
        db.sequences.push_back(seq);
    }
    
    int buildDatabase(Database &db, char fname[], const int maxT)
    {
      FILE *fp;
    
      if ( (fp = fopen(fname, "rt")) == NULL ) {
        printf("\nError opening file %s\n", fname);
        return -1;
      }
      double returnToGo = 0, ballX, ballY, ballDX, ballDY, leftY, rightY;
      int action;
      int n = 0;
      while ( n < maxT ){
        vector<PongStruct> seq;
        while ( true ) {
          fscanf(fp, "%lf", &returnToGo);
          if ( returnToGo < -100 )
             break;
          fscanf(fp, "%lf", &ballX);
          fscanf(fp, "%lf", &ballY);
          fscanf(fp, "%lf", &ballDX);
          fscanf(fp, "%lf", &ballDY);
          fscanf(fp, "%lf", &leftY);
          fscanf(fp, "%lf", &rightY);
          fscanf(fp, "%d", &action);
          seq.push_back(PongStruct (returnToGo, ballX, ballY, ballDX, ballDY,
                                  leftY, rightY, action));
        }
        addSample(db, seq);
        n++;
      }
    
      if ( fp != NULL )
        fclose( fp );
    
      return 1;
    }
    
    void printSeq (vector<PongStruct> seq)
    {
       for (const auto &p : seq)
          printf("\n%6.3f %6.3f %6.3f %6.3f %6.3f %6.3f %6.3f %4d",
    	     p.returnToGo, p.ballX, p.ballY, p.ballDX, p.ballDY, 
    	     p.leftY, p.rightY, p.action);
    }
    
    int main(int argc, char *argv[])
    {
       if ( argc < 2) {
          cout << "Usage: " << argv[0] << " # of trajectories " << endl;
          return 1; 
       }
       int maxT = atoi (argv[1]);  // maximum number of trajectories to read
       if ( maxT < 0 ) 
          maxT = 10;
       else if ( maxT > 150 )
          maxT = 150;
    
       // Build  database 
       Database db;
       char fname[] = "trajectories.txt";
       if ( buildDatabase(db, fname, maxT) < 0 ){
         cout << "build database failed!" << endl;
         return 1;
       }
    
       cout << "Number of training sequences: " << db.sequences.size() << endl;
    
       // Model
       int vocab_size = 7;
       int max_steps = 32;
       int seq_len = 3 * max_steps;
       int dModel = 32;
       int nHeads = 4;
       int d_ff = 64;
       int nLayers = 2;
    
       DecisionTransformer model(vocab_size, seq_len, dModel, nHeads, d_ff, nLayers);
       // Extract dataset
       vector<vector<PongStruct>> dataset = db.sequences;
    
       // shuffle support
       random_device rd;
       mt19937 gen(rd());
    
       // Training
       int nEpoches = 150;
       cout << "\n==== Training Decision Transformer Pong Game ====\n";
    
       for (int epoch = 0; epoch < nEpoches; epoch++)
       {
          shuffle(dataset.begin(), dataset.end(), gen);
          double loss = 0.0;
          for (int i = 0; i < dataset.size(); i++) {
    	 vector<PongStruct>seq = dataset[i];
             loss += model.trainStep(seq, 0.01);
          }
    
          if (epoch % 2 == 0)
            cout << "Epoch " << epoch << " Loss: " << loss / dataset.size() << endl;
       }
       model.saveWeights( (char *) "weights.txt" );
    
       cout << "Weights saved in weights.txt!" << endl;
    
       return 0;
    }	

    Makefile:

    PROG    = testMain
    #source codes
    SRCS =  $(PROG).cpp
    #substitute .cpp by .o to obtain object filenames
    OBJS = $(SRCS:.cpp=.o) util.o  transGpt.o  decision.o
    
    #$< evaluates to the target's dependencies, 
    #$@ evaluates to the target
    
    $(PROG): $(OBJS)
            g++ -o $@ $(OBJS)  
    
    $(OBJS): 
            g++ -c  -std=c++20 $*.cpp 
    
    clean:
            rm $(OBJS) $(PROG)
    
    

    Sample Output:

    Number of training sequences: 150
    
    ==== Training Decision Transformer Pong Game ====
    Epoch 0 Loss: 0.424162
    Epoch 4 Loss: 0.0209024
    Epoch 8 Loss: 0.00973791
    Epoch 12 Loss: 0.00544495
    Epoch 16 Loss: 0.00372065
    Epoch 20 Loss: 0.00274737
    Epoch 24 Loss: 0.0021282
    Epoch 28 Loss: 0.00171588
    Epoch 32 Loss: 0.00142454
    Epoch 36 Loss: 0.00121189
    Epoch 40 Loss: 0.00105091
    Epoch 44 Loss: 0.000922877
    Epoch 48 Loss: 0.000822407
    Epoch 52 Loss: 0.000742089
    Epoch 56 Loss: 0.000675655
    Epoch 60 Loss: 0.000619225
    Epoch 64 Loss: 0.000571053
    Epoch 68 Loss: 0.000529348
    Epoch 72 Loss: 0.000493154
    Epoch 76 Loss: 0.000461549
    Epoch 80 Loss: 0.000433255
    Epoch 84 Loss: 0.000408239
    Epoch 88 Loss: 0.000385804
    Epoch 92 Loss: 0.00036562
    Epoch 96 Loss: 0.000347291
    Epoch 100 Loss: 0.00033068
    Epoch 104 Loss: 0.000315483
    Epoch 108 Loss: 0.000301638
    Epoch 112 Loss: 0.000288843
    Epoch 116 Loss: 0.000277057
    Epoch 120 Loss: 0.000266185
    Epoch 124 Loss: 0.000256119
    Epoch 128 Loss: 0.000246676
    Epoch 132 Loss: 0.000237902
    Epoch 136 Loss: 0.000229716
    Epoch 140 Loss: 0.000222069
    Epoch 144 Loss: 0.000214866
    Epoch 148 Loss: 0.000208154
    

    Playing Pong with DT

    /*
     * pongDT.cpp -- Playing Pong with a Decision Transformer
     * http://forejune.co/cuda
     */
    
    #include <GL/gl.h>
    #include <GL/glut.h>
    #include <string>
    #include "util.h"
    #include "transGpt.h"
    #include "decision.h"
    
    using namespace std;
    
    // Game variables
    double ballX = 0.0;	// horizontal position of ball
    double ballY = 0.0;	// vertical position of ball
    
    double ballDX = 0.15;	// change in X of ball
    double ballDY = 0.10;	// change in Y of ball
    
    double leftY = 0.0;	// left paddle position
    double rightY = 0.0;	// rifht paddle position
    
    int   action = STAY;
    
    const double pW  = 0.5;	// paddle width
    const double pH = 3.0;   // paddle height
    
    double scoreLeft  = 0;
    double scoreRight = 0;
    double returnToGo = 11;
    int maxTimeSteps = 150;
    
    bool pause = false;
    
    // model(vocab_size, seq_len, dModel, nHeads, d_ff, nLayers);
    DecisionTransformer model(7, 3*maxTimeSteps, 32, 4, 64, 2);
    vector<PongStruct> history; 
    
    void drawPaddle(double x, double y)
    {
      // (x, y) is the lower left vertex coordinates
      glRectf(x, y, x+pW, y+pH);
    }
    
    // draw String
    void drawString(double x, double y, const string &s)
    {
        glRasterPos2f(x, y);
    
        for (int i = 0; i < s.length(); i++)
            glutBitmapCharacter(GLUT_BITMAP_HELVETICA_18, s[i]);
    }
    
    void drawBall(double x, double y)
    {
      glPushMatrix();
      glTranslatef(x, y, 0);
      glutSolidSphere(0.5, 16, 16);
      glPopMatrix();
    }
    
    void resetBall() {
        ballX = 0; ballY = 0;
        ballDX = -ballDX;
        returnToGo = 11.0;      // Reset target return
        history.clear();
        PongStruct pong = {returnToGo, ballX, ballY, ballDX, ballDY, leftY, rightY, STAY};
        history.push_back(pong);
    }
    
    // Initialization
    void init(void)
    {
        glClearColor(1, 1, 1, 0);
    
        glMatrixMode(GL_PROJECTION);
        glLoadIdentity();
        gluOrtho2D(-10, 10, -10, 10);
    
        glMatrixMode(GL_MODELVIEW);
        glLoadIdentity();
        // get weights of DecisionTransformer 
        model.readWeights( (char *) "weights.txt" );
        resetBall();
    }
    
    // Draw scene
    void display(void)
    {
        glClear(GL_COLOR_BUFFER_BIT);
    
        double x, y;		//paddle lower left corner
        glColor3f(1, 0, 0); //red color
    
        // Left paddle
        x = -9.0;
        y = leftY - pH/2;
        drawPaddle(x, y);
    
        // Right paddle
        x = 8.5;
        y = rightY - pH / 2;
        drawPaddle(x, y);
    
        // Ball
        glColor3f(0, 1, 0);
        drawBall(ballX, ballY);
    
        // Score
        char str[40];
        sprintf(str,"%d   :   %d", (int)scoreLeft, (int)scoreRight);
        glColor3f(0, 0, 0);  // black color
        drawString(-1.0, 9.0, str);
    
        glutSwapBuffers();
    }
    
    void animate()
    {
        if (pause)
            return;
    
        // Construct the current state
        // The action is not known yet, so use -1.
        PongStruct pong = {
            returnToGo,
            ballX,
            ballY,
            ballDX,
            ballDY,
            leftY,
            rightY,
            -1
        };
    
        history.push_back(pong);
        if (history.size() > maxTimeSteps)
          history.erase(history.begin());
      
        // Decision Transformer predicts the action
        int action = model.predictAction(history);
    
        // Store the predicted action in the history
        history.back().action = action;
    
        // Execute the action on the left paddle
        switch (action)
        {
            case UP:
                leftY += 0.08;
                break;
    
            case DOWN:
                leftY -= 0.08;
                break;
    
            case STAY:
                break;
        }
    
        // Keep left paddle inside the screen
        if (leftY > 10.0 - pH / 2)
            leftY = 10.0 - pH / 2;
    
        if (leftY < -10.0 + pH / 2)
            leftY = -10.0 + pH / 2;
    
        // Simple AI controls the right paddle
        if (ballY > rightY)
            rightY += 0.08;
        else
            rightY -= 0.08;
    
        // Keep right paddle inside the screen
        if (rightY > 10.0 - pH / 2)
            rightY = 10.0 - pH / 2;
    
        if (rightY < -10.0 + pH / 2)
            rightY = -10.0 + pH / 2;
    
        // Move the ball
        ballX += ballDX;
        ballY += ballDY;
    
        // Top/bottom collision
        if (ballY > 9.5)
        {
            ballY = 9.5;
            ballDY = -ballDY;
        } else if (ballY < -9.5)
        {
            ballY = -9.5;
            ballDY = -ballDY;
        }
    
        // Left paddle collision
        if (ballX <= -8.25 && ballY >= leftY - pH / 2 &&
          		ballY <= leftY + pH / 2 && ballDX < 0)
            ballDX = -ballDX;
        else if (ballX >= 8.25 && ballY >= rightY - pH / 2 &&
                 	ballY <= rightY + pH / 2 && ballDX > 0)
          	ballDX = -ballDX;
        
        // Determine reward
        double reward = 0.0;
        if (ballX < -10)
        {
            // DT-controlled player missed
            scoreRight++;
            reward = -1.0;   // not used
            resetBall();
        }
        else if (ballX > 10)
        {
            // DT-controlled player scored
            scoreLeft++;
            reward = +1.0;   // not used
            resetBall();
        }
    
        // Update Return-To-Go
        returnToGo -= reward;   
    
        glutPostRedisplay();
    }
    
    // Keyboard control
    void keyboard(unsigned char key,int x, int y)
    {
        switch(key)
        {
            case 27:
                exit(0);
                break;
    
            case 'r':	// reset
                scoreLeft = scoreRight = 0;
                resetBall();
                break;
    
    	case 'p':	// toggles pause
    	    pause = pause ? false : true;
    	    break;
        }
    }
    
    // Visibility callback
    void timerHandle ( int value )
    {
       animate();
       glutPostRedisplay();
       // call timerHandle 25 ms later, 0 is passed to timerHandle, not used here 
       glutTimerFunc (25, timerHandle, 0);
    }
    
    void visHandle( int visible )
    {
       if (visible == GLUT_VISIBLE)
          timerHandle ( 0 );
       else
          ;
    }
    
    int main(int argc, char *argv[])
    {
        glutInit(&argc, argv);
    
        glutInitDisplayMode(GLUT_DOUBLE | GLUT_RGB);
    
        glutInitWindowSize(500,500);
        glutCreateWindow("Pong Game");
    
        glutDisplayFunc(display);
        glutVisibilityFunc(visHandle);
        glutKeyboardFunc(keyboard);
    
        init();
        glutMainLoop();
    
        return 0;
    }