I trained an AI Model to Detect Trading Candlesticks (from scratch using ViTs)

Nicholas RenotteAbout 8 min readJul 27, 2025Watch original
THE SUMMARYAI-generated

Key Concepts

Vision Transformer (ViT), Candlestick Pattern Recognition, Deep Learning, PyTorch, Data Augmentation (CutMix, MixUp, Color Jitter), Image Patching, Positional Encoding, Multi-Head Attention, Transfer Learning, Real-time Prediction, Plotly Dashboards, Git Large File Storage (LFS).

1. Getting Started: Running the Pre-trained Model (5 Minutes)

  • Goal: To quickly get the candlestick detection running on a local machine (MacBook, no GPU required).
  • Steps:
    1. Clone the GitHub repository: git clone github.com/nicknack/vitcandlestick.
    2. Verify Checkpoint Size: Ensure the 25_model checkpoint file is approximately 100MB (102.3MB). If smaller, use Git LFS.
    3. Download Checkpoints (if needed): Navigate to the cloned directory (cd vitcandlesticks) and run git lfs pull to download the full checkpoint files.
    4. Configure Candlestick Data Source (Optional): In source/candlesticks.py, modify the stock_code variable (e.g., from "AMZN" for Amazon to "AAPL" for Apple) to change the stock data source. The interval variable can also be adjusted to control the candlestick update frequency.
    5. Run the Candlestick Dashboard: Execute uv venv source/utils/candlestick.py. This creates a virtual environment, installs dependencies (specified in tommoal file), and starts a Plotly dashboard on http://127.0.0.1:8050.
    6. Run the Candlestick Prediction: Open a new terminal, navigate to the cloned directory, and execute uv venv source/candlestick_prediction.py. This loads the pre-trained model (25_model.pt) and starts making predictions based on screen captures.
    7. Adjust Screen Capture Region: The prediction script captures a region of the screen. Ensure the Plotly dashboard with the candlesticks is within this region. The script is optimized for a 4K screen, so adjustments may be needed for lower resolutions.
  • Output: The script displays live candlestick pattern predictions (Dogee, Bullish Engulfing, Bearish Engulfing, Morning Star, Evening Star) overlaid on the candlestick chart.

2. Data Preparation: Transforming Raw Images into Image Patches

  • Rationale for Training from Scratch:
    • Customizing model parameters (size, architecture).
    • Training on different data sources (MetaTrader, TradingView).
    • Detecting different candlestick patterns.
  • Data Collection and Labeling:
    • Dataset consists of over 2,000 labeled images, including hand-labeled screenshots and synthetically generated data.
    • Data is split into train_data and test_data folders to prevent data leakage.
    • Synthetic data is generated using utils/data_generation_screen.py.
    • Images are annotated with labels indicating the candlestick pattern.
  • Label Encoding:
    • Originally, labels included a "nothing" class (0), but it was later removed.
    • Final label encoding:
      • Dogee: 0
      • Bullish Engulfing: 1
      • Bearish Engulfing: 2
      • Morning Star: 3
      • Evening Star: 4
  • Data Loading and Transformation (source/data.py):
    • CLFDataset Class: Core data class for loading and processing images and labels.
    • Label Handling: The class reads labels from labels.csv, drops the "nothing" class (0), and adjusts the remaining labels to the 0-4 range.
    • Data Augmentation (Albumentations):
      • Color Jitter: Adjusts brightness, contrast, saturation, and hue to improve model generalization. Hue adjustment was initially problematic due to color inversions (red to green), so it was disabled.
      • Applied only to the training and validation partitions.
    • Image Processing:
      • Resizes images to 224x224 pixels.
      • Crops the image to focus on the last eight candlesticks.
      • Normalizes pixel values using ImageNet statistics (mean and standard deviation).
    • Image Patching:
      • Divides the cropped image into patches of 8x8 pixels.
      • Flattens each patch into a 192-dimensional vector (8x8x3).
  • Data Exploration:
    • The output_samples parameter in data.py can be set to True to output transformed images to the images/transformed_images folder.
    • The output_patches parameter can be set to True to visualize the image patches using Matplotlib.
  • Key Functions:
    • __init__: Initializes the dataset, loads labels, and sets up transformations.
    • __len__: Returns the length of the dataset.
    • __getitem__: Loads an image, applies transformations, and returns the image tensor and label.

3. Building the Vision Transformer (ViT) Model (source/model.py)

  • Model Architecture: Implements a Vision Transformer (ViT) based on the "Attention is All You Need" paper, but adapted for image classification.
  • Core Classes:
    • ViT: The main class that encapsulates the entire ViT model.
    • MLPBlock: Implements a multi-layer perceptron (feed-forward neural network) with linear layers, GELU activation, and dropout.
    • EncoderBlock: Implements a single transformer encoder block with layer normalization, multi-head self-attention, and skip connections.
    • Encoder: Stacks multiple EncoderBlock layers sequentially.
    • PositionalEncoding: Adds positional information to the image patch embeddings using a sinusoidal function.
  • Model Breakdown:
    1. Patch Conversion: The input image is divided into 8x8 pixel patches, resulting in 117 patches (for a 104x72 image).
    2. Patch Flattening: Each patch is flattened into a 192-dimensional vector.
    3. Linear Embedding: The flattened patch vectors are passed through a linear layer to create an embedding with 1024 features.
    4. Positional Encoding: Positional encodings are added to the patch embeddings to provide information about the location of each patch. The sinusoidal encoding function from the original transformer paper is used.
    5. Class Token: A class token is appended to the sequence of patch embeddings, increasing the sequence length from 117 to 118. This token is used for classification.
    6. Encoder Blocks: The sequence of patch embeddings is passed through multiple encoder blocks (3 in this implementation). Each encoder block consists of:
      • Layer normalization.
      • Multi-head self-attention.
      • Skip connection (adding the input to the output of the self-attention layer).
      • MLP block (feed-forward neural network).
      • Skip connection (adding the input to the output of the MLP block).
    7. Classification Head: The class token is extracted from the output of the encoder, and passed through a linear layer to produce the final classification output (5 classes: Dogee, Bullish Engulfing, Bearish Engulfing, Morning Star, Evening Star).
  • Customization:
    • The number of encoder layers can be adjusted.
    • The number of candlestick patterns (output classes) can be changed.
    • The patch size can be modified, but the image height and width must be divisible by the patch size.
  • Visualization:
    • The stack_component variable in model.py can be used to visualize different layers of the model (encoder, encoder block, MLP block, positional embedding, ViT).
  • Key Functions:
    • forward: Defines the forward pass of the ViT model.
    • PositionalEncoding: Implements the sinusoidal positional encoding function.

4. Training the Model (source/train.py)

  • Data Partitioning:
    • The training data is split into training (70%) and validation (30%) sets.
    • A separate test dataset is used for final evaluation.
  • Data Augmentation (Batch-Level):
    • CutMix and MixUp: Applied at the batch level to improve model generalization.
      • MixUp: Blends two images and their corresponding labels.
      • CutMix: Cuts and pastes patches from different images.
    • A random choice is used to apply CutMix (25%), MixUp (25%), or no augmentation (50%) to each batch.
  • Training Setup:
    • Manual seed is set to 42 for reproducibility.
    • Cross-entropy loss is used as the loss function.
    • Adam optimizer is used with a learning rate of 1e-4 (0.0001).
    • Cosine annealing with warm restarts learning rate scheduler is used to adjust the learning rate during training.
  • Training Loop:
    • The model is trained for a specified number of epochs (default: 100).
    • For each epoch:
      • The model is set to training mode (model.train()).
      • The training data is iterated over in batches.
      • CutMix or MixUp augmentation is applied to the batch (if selected).
      • The model makes predictions on the batch.
      • The loss is calculated.
      • Gradients are zeroed out.
      • Gradients are computed using backpropagation.
      • The optimizer updates the model parameters.
      • The learning rate scheduler is updated.
      • The same process is repeated for the validation dataset (with the model set to evaluation mode model.eval()).
      • The same process is repeated for the test dataset (with the model set to evaluation mode model.eval()).
      • The model is saved every five epochs to the checkpoints folder.
  • Monitoring:
    • Training, validation, and test losses are tracked and displayed during training.
  • Key Functions:
    • collate_fn: Applies CutMix or MixUp augmentation to a batch of data.

5. Testing the Model (source/test.py)

  • Purpose: To evaluate the performance of the trained model on the test dataset.
  • Process:
    1. The script loads a trained model from a specified checkpoint file.
    2. The script loads the test dataset.
    3. The script makes predictions on a batch of images from the test dataset.
    4. The script displays the actual and predicted labels for each image in a grid.
    5. The script saves the grid to a file named results.png.
  • Customization:
    • The model_path variable can be used to specify the path to the checkpoint file.
    • The data_path variable can be used to specify the path to the test dataset.

6. Real-time Prediction

  • Process:
    1. Start the candlestick dashboard (source/utils/candlesticks.py).
    2. Run the candlestick prediction script (source/candlestick_prediction.py).
    3. The prediction script captures a region of the screen where the candlestick dashboard is running.
    4. The script loads a trained model from a specified checkpoint file.
    5. The script makes predictions on the captured screen region in real-time.
    6. The script displays the predicted candlestick patterns overlaid on the candlestick chart.
  • Customization:
    • The checkpoint variable in candlestick_prediction.py can be used to specify the path to the checkpoint file.
    • The stock_code variable in candlesticks.py can be used to change the stock data source.
    • The interval variable in candlesticks.py can be adjusted to control the candlestick update frequency.
    • The screen capture region in candlestick_prediction.py may need to be adjusted depending on the screen resolution.

7. Conclusion

The video demonstrates how to build a real-time candlestick pattern detection system using a Vision Transformer (ViT) trained from scratch. It covers data preparation, model architecture, training, testing, and real-time prediction. The presenter emphasizes the importance of data augmentation techniques like CutMix and MixUp for improving model generalization. The code is available on GitHub, allowing viewers to replicate the results and customize the system for their own needs. The presenter also highlights the challenges of training ViTs and the importance of careful hyperparameter tuning.

AI summaries can miss context or contain errors. Check important details against the original video.

Go a little deeper.

Have a question about this video? Load its transcript to open the video chat.