Materials available at: http://forejune.co/cuda/
Recap
Decision Transformer
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
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
#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
|
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
|
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;
}
|