// BogieNN_v5_DirectLibTorchGNN.cpp
// ============================================================================
// Zorro S 3.11+ / Zorro64 / GPU LibTorch C++
//
// DIRECT LIBTORCH GRAPH NEURAL NETWORK VERSION
//
// NO ZORRO MACHINE-LEARNING FUNCTIONS ARE USED.
//
// Specifically, this file contains NO:
// any Zorro machine-learning training or prediction interface
//
// Zorro is used only for:
// - market/history simulation,
// - indicators and price series,
// - WFO scheduling,
// - Train/Test/Trade mode,
// - trade execution and account management.
//
// LibTorch directly owns:
// - training sample collection,
// - class/direction balancing,
// - GPU training,
// - model validation / early stopping,
// - WFO model persistence,
// - model loading,
// - graph inference,
// - directional output,
// - learned market-state classification.
//
// WFO lifecycle used by this script:
// [Train]
// Every WFO training cycle is a separate Zorro simulation run.
// Samples are collected bar-by-bar.
// At EXITRUN:
// LibTorch trains the GNN.
// Model is saved to Data\BogieGNN_Direct_WFOxx.pt.
//
// [Test]
// WFOCycle tells us the current OOS test segment.
// The corresponding .pt model is loaded directly by LibTorch.
// Predictions are made directly with model->forward().
//
// [Trade]
// The last trained WFO model is loaded.
//
// Five graph nodes:
// 0 Trend specialist
// 1 Mean-reversion specialist
// 2 Hybrid-trend specialist
// 3 Hybrid-mean specialist
// 4 Market-state specialist
//
// 5 nodes x 8 features = 40 graph signals.
//
// ============================================================================
#ifndef NOMINMAX
#define NOMINMAX
#endif
#ifndef BOGIE_ENABLE_CUDA
#define BOGIE_ENABLE_CUDA 1
#endif
#include <torch/torch.h>
#include <torch/serialize.h>
#if BOGIE_ENABLE_CUDA
#include <torch/cuda.h>
#endif
#include <algorithm>
#include <cmath>
#include <cstdint>
#include <cstdlib>
#include <fstream>
#include <limits>
#include <memory>
#include <sstream>
#include <string>
#include <tuple>
#include <vector>
// LibTorch owns namespace `at`; Zorro declares a global function named `at`.
// Preserve the working include arrangement from the confirmed Zorro64 build.
#define at zorro_at
#include <zorro.h>
#undef at
// ============================================================================
// USER SETTINGS
// ============================================================================
cstr BogieAsset = "EUR/USD";
int BacktestStart = 2010;
int BacktestEnd = 2026;
int StrategyBarPeriod = 60;
// WFO scheduling is still provided by Zorro.
// It does NOT perform any neural training.
int WFOCycles = 8;
int WFOTrainPercent = 85;
// One GPU should normally be controlled by one training Zorro process.
// Keep WFO GPU training sequential unless you deliberately distribute cycles
// over separate GPUs yourself.
int DirectGPUTrainingCores = 1;
// State-dependent future horizons.
int TrendPredictionHorizonBars = 12;
int MeanPredictionHorizonBars = 4;
int NeutralPredictionHorizonBars = 6;
int MaxPredictionHorizonBars = 12;
// Direction score thresholds; GNN output is scaled to approximately -100..100.
var TrendConfidenceThreshold = 20.0;
var MeanConfidenceThreshold = 20.0;
var NeutralConfidenceThreshold = 35.0;
// Learned state probability required for specialist authority.
var StateProbabilityMin = 0.50;
int AllowNeutralTrades = 0;
int RequireDIConfirmation = 1;
// ------------------------ Trend specialist parameters -------------------------
int NocPeriod = 12;
int EmaPeriod = 5;
int ATRPeriod = 14;
int ADXPeriod = 14;
int AroonPeriod = 25;
int MMIPeriod = 100;
var TrendStandaloneADXMin = 20.0;
var TrendStandaloneMMIMax = 65.0;
var TrendStopATRMult = 2.5;
var TrendTrailATRMult = 3.0;
// --------------------- Mean-reversion specialist parameters -------------------
int RSIPeriod = 14;
int BBPeriod = 20;
int MeanPeriod = 20;
var MeanStandaloneMMIMin = 58.0;
var MeanStandaloneADXMax = 28.0;
var MeanLongExtremeStandalone = -0.45;
var MeanShortExtremeStandalone = 0.45;
var MeanLongRSIMaxStandalone = 40.0;
var MeanShortRSIMinStandalone = 60.0;
var MeanStopATRMult = 1.6;
// -------------------------- Hybrid state parameters ---------------------------
var HybridTrendADXMin = 25.0;
var HybridTrendMMIMax = 64.0;
var HybridMeanADXMax = 20.0;
var HybridMeanMMIMin = 60.0;
var HybridMeanLongExtreme = -0.45;
var HybridMeanShortExtreme = 0.45;
var HybridMeanLongRSIMax = 40.0;
var HybridMeanShortRSIMin = 60.0;
var RegimeScaleADX = 15.0;
var RegimeScaleMMI = 15.0;
// ------------------------------ Trade settings --------------------------------
var NeutralStopATRMult = 2.0;
var TrailGuardPips = 0.5;
int UseMM = 1;
int MiniAcct = 0;
var RiskPercent = 5.0;
var FixedAmount = 0.10;
// MQL4 weekday numbering.
int NoTradeDay_1 = 0;
int NoTradeDay_2 = 0;
int CloseOnNoTradeDay = 0;
int UseDiagnostics = 1;
// ----------------------------- Target settings --------------------------------
// Continuous target:
// (future price - current price) / (TargetATRScale * ATR)
// clipped to -1..+1.
var TargetATRScale = 2.0;
// --------------------------- LibTorch hyperparameters --------------------------
int GNNHidden = 32;
int GNNEpochs = 80;
int GNNBatchSize = 128;
int GNNPatience = 12;
int GNNValidationPercent = 15;
var GNNLearningRate = 0.001;
var GNNWeightDecay = 0.0001;
// Shared graph learns direction + current market state.
var GNNStateLossWeight = 0.20;
// Explicit replacement for Zorro's old +BALANCED behavior.
// 1 = inverse-frequency weighting of positive/negative direction samples.
int UseBalancedDirectionLoss = 1;
// CPU workers used by LibTorch around GPU operations.
int LibTorchThreads = 2;
// Minimum number of samples required to train a cycle.
int MinTrainingSamples = 200;
// ------------------------------- Model files ----------------------------------
// Relative to the Zorro root folder.
cstr DirectModelPrefix = "Data\\BogieGNN_Direct_WFO";
// ============================================================================
// GRAPH DIMENSIONS
// ============================================================================
static const int GRAPH_NODE_COUNT = 5;
static const int GRAPH_NODE_FEATURES = 8;
static const int GRAPH_SIGNAL_COUNT = 40;
// ============================================================================
// PREDICTION DIAGNOSTICS
// ============================================================================
var GNNStateTrendProbability = 0.0;
var GNNStateMeanProbability = 0.0;
var GNNStateNeutralProbability = 1.0;
var GNNNodeTrendAuthority = 0.0;
var GNNNodeMeanAuthority = 0.0;
var GNNNodeHybridTrendAuthority = 0.0;
var GNNNodeHybridMeanAuthority = 0.0;
var GNNNodeStateAuthority = 0.0;
// ============================================================================
// DIRECT LIBTORCH GLOBAL STATE
// ============================================================================
static torch::Device GNNDevice(torch::kCPU);
#if BOGIE_ENABLE_CUDA
static HMODULE GTorchCudaModule = 0;
#endif
class GraphAttentionBlockImpl;
class BogieGraphNetImpl;
static std::shared_ptr<BogieGraphNetImpl> GNNModel;
// Direct training dataset for the current Zorro WFO training run.
static std::vector<float> GNNTrainX;
static std::vector<float> GNNTrainY;
static std::vector<int64_t> GNNTrainState;
// Current model lifecycle.
static int GNNLoadedCycle = 0;
static int GNNLibTorchReady = 0;
// ============================================================================
// GRAPH ATTENTION LAYER
// ============================================================================
class GraphAttentionBlockImpl : public torch::nn::Module
{
public:
torch::nn::Linear Q{nullptr};
torch::nn::Linear K{nullptr};
torch::nn::Linear V{nullptr};
torch::nn::Linear Out{nullptr};
torch::nn::LayerNorm Norm{nullptr};
torch::Tensor EdgeBias;
int Hidden = 0;
GraphAttentionBlockImpl(int HiddenSize)
{
Hidden = HiddenSize;
Q = register_module(
"q",
torch::nn::Linear(
torch::nn::LinearOptions(Hidden,Hidden).bias(false)));
K = register_module(
"k",
torch::nn::Linear(
torch::nn::LinearOptions(Hidden,Hidden).bias(false)));
V = register_module(
"v",
torch::nn::Linear(
torch::nn::LinearOptions(Hidden,Hidden).bias(false)));
Out = register_module(
"out",
torch::nn::Linear(Hidden,Hidden));
Norm = register_module(
"norm",
torch::nn::LayerNorm(
torch::nn::LayerNormOptions(
std::vector<int64_t>{Hidden})));
// Initial graph topology:
// Trend <-> HybridTrend <-> State
// Mean <-> HybridMean <-> State
//
// All entries remain trainable.
std::vector<float> InitialBias = {
0.50f,-0.75f, 0.75f,-0.75f, 1.00f,
-0.75f, 0.50f,-0.75f, 0.75f, 1.00f,
0.75f,-0.50f, 0.50f,-0.25f, 1.00f,
-0.50f, 0.75f,-0.25f, 0.50f, 1.00f,
1.00f, 1.00f, 1.00f, 1.00f, 0.50f
};
EdgeBias =
torch::tensor(
InitialBias,
torch::TensorOptions().dtype(torch::kFloat32))
.reshape({GRAPH_NODE_COUNT,GRAPH_NODE_COUNT});
EdgeBias =
register_parameter(
"edge_bias",
EdgeBias);
}
std::pair<torch::Tensor,torch::Tensor> forward(torch::Tensor H)
{
torch::Tensor Query;
torch::Tensor Key;
torch::Tensor Value;
torch::Tensor Logits;
torch::Tensor Attention;
torch::Tensor Messages;
torch::Tensor Updated;
Query = Q->forward(H);
Key = K->forward(H);
Value = V->forward(H);
Logits =
torch::matmul(
Query,
Key.transpose(1,2));
Logits =
Logits/
std::sqrt((double)Hidden);
Logits =
Logits+
EdgeBias.unsqueeze(0);
Attention =
torch::softmax(
Logits,
-1);
Messages =
torch::matmul(
Attention,
Value);
Updated =
torch::relu(
Out->forward(Messages));
Updated =
Norm->forward(
H+Updated);
return std::make_pair(
Updated,
Attention);
}
};
// ============================================================================
// FIVE-NODE GRAPH NETWORK
// ============================================================================
class BogieGraphNetImpl : public torch::nn::Module
{
public:
torch::nn::Linear TrendEncoder{nullptr};
torch::nn::Linear MeanEncoder{nullptr};
torch::nn::Linear HybridTrendEncoder{nullptr};
torch::nn::Linear HybridMeanEncoder{nullptr};
torch::nn::Linear StateEncoder{nullptr};
torch::Tensor NodeEmbedding;
std::shared_ptr<GraphAttentionBlockImpl> Graph1;
std::shared_ptr<GraphAttentionBlockImpl> Graph2;
torch::nn::Linear PoolGate{nullptr};
torch::nn::Sequential DirectionHead;
torch::nn::Sequential StateHead;
int Hidden = 0;
BogieGraphNetImpl(int HiddenSize = 32)
{
Hidden = HiddenSize;
TrendEncoder =
register_module(
"trend_encoder",
torch::nn::Linear(
GRAPH_NODE_FEATURES,
Hidden));
MeanEncoder =
register_module(
"mean_encoder",
torch::nn::Linear(
GRAPH_NODE_FEATURES,
Hidden));
HybridTrendEncoder =
register_module(
"hybrid_trend_encoder",
torch::nn::Linear(
GRAPH_NODE_FEATURES,
Hidden));
HybridMeanEncoder =
register_module(
"hybrid_mean_encoder",
torch::nn::Linear(
GRAPH_NODE_FEATURES,
Hidden));
StateEncoder =
register_module(
"state_encoder",
torch::nn::Linear(
GRAPH_NODE_FEATURES,
Hidden));
NodeEmbedding =
register_parameter(
"node_embedding",
0.05*
torch::randn(
{GRAPH_NODE_COUNT,Hidden},
torch::TensorOptions().dtype(torch::kFloat32)));
Graph1 =
register_module(
"graph1",
std::make_shared<GraphAttentionBlockImpl>(Hidden));
Graph2 =
register_module(
"graph2",
std::make_shared<GraphAttentionBlockImpl>(Hidden));
PoolGate =
register_module(
"pool_gate",
torch::nn::Linear(
Hidden,
1));
DirectionHead =
register_module(
"direction_head",
torch::nn::Sequential(
torch::nn::Linear(Hidden,32),
torch::nn::ReLU(),
torch::nn::Dropout(0.10),
torch::nn::Linear(32,16),
torch::nn::ReLU(),
torch::nn::Linear(16,1),
torch::nn::Tanh()));
StateHead =
register_module(
"state_head",
torch::nn::Sequential(
torch::nn::Linear(Hidden,16),
torch::nn::ReLU(),
torch::nn::Linear(16,3)));
}
std::tuple<
torch::Tensor,
torch::Tensor,
torch::Tensor> forward(torch::Tensor X)
{
torch::Tensor Nodes;
torch::Tensor TrendNode;
torch::Tensor MeanNode;
torch::Tensor HybridTrendNode;
torch::Tensor HybridMeanNode;
torch::Tensor StateNode;
torch::Tensor H;
std::pair<torch::Tensor,torch::Tensor> G1;
std::pair<torch::Tensor,torch::Tensor> G2;
torch::Tensor PoolLogits;
torch::Tensor PoolWeights;
torch::Tensor GraphState;
torch::Tensor Direction;
torch::Tensor StateLogits;
if(X.dim() == 1)
X = X.unsqueeze(0);
Nodes =
X.reshape(
{
X.size(0),
GRAPH_NODE_COUNT,
GRAPH_NODE_FEATURES
});
TrendNode =
torch::relu(
TrendEncoder->forward(
Nodes.select(1,0)));
MeanNode =
torch::relu(
MeanEncoder->forward(
Nodes.select(1,1)));
HybridTrendNode =
torch::relu(
HybridTrendEncoder->forward(
Nodes.select(1,2)));
HybridMeanNode =
torch::relu(
HybridMeanEncoder->forward(
Nodes.select(1,3)));
StateNode =
torch::relu(
StateEncoder->forward(
Nodes.select(1,4)));
H =
torch::stack(
{
TrendNode,
MeanNode,
HybridTrendNode,
HybridMeanNode,
StateNode
},
1);
H =
H+
NodeEmbedding.unsqueeze(0);
G1 = Graph1->forward(H);
H = G1.first;
G2 = Graph2->forward(H);
H = G2.first;
PoolLogits =
PoolGate->forward(H)
.squeeze(-1);
PoolWeights =
torch::softmax(
PoolLogits,
-1);
GraphState =
torch::sum(
PoolWeights.unsqueeze(-1)*H,
1);
Direction =
DirectionHead->forward(GraphState)
.squeeze(-1);
StateLogits =
StateHead->forward(
GraphState);
return std::make_tuple(
Direction,
StateLogits,
PoolWeights);
}
};
// ============================================================================
// LIBTORCH INITIALIZATION
// ============================================================================
static void ClearDirectTrainingData()
{
GNNTrainX.clear();
GNNTrainY.clear();
GNNTrainState.clear();
}
static int InitializeDirectLibTorch()
{
#if BOGIE_ENABLE_CUDA
// Register CUDA hooks before the first LibTorch API initializes its context.
if(!GTorchCudaModule)
GTorchCudaModule =
LoadLibraryA(
"torch_cuda.dll");
if(!GTorchCudaModule)
{
printf(
"\nDirect GNN: torch_cuda.dll load failed (Windows error %lu)",
GetLastError());
}
#endif
if(LibTorchThreads < 1)
LibTorchThreads = 1;
torch::manual_seed(365);
torch::set_num_threads(
LibTorchThreads);
#if BOGIE_ENABLE_CUDA
if(torch::cuda::is_available())
{
GNNDevice =
torch::Device(
torch::kCUDA);
printf(
"\nDirect LibTorch GNN initialized on CUDA");
}
else
{
GNNDevice =
torch::Device(
torch::kCPU);
printf(
"\nDirect LibTorch GNN: CUDA unavailable; using CPU");
}
#else
GNNDevice =
torch::Device(
torch::kCPU);
printf(
"\nDirect LibTorch GNN initialized on CPU");
#endif
GNNModel.reset();
GNNLoadedCycle = 0;
ClearDirectTrainingData();
GNNLibTorchReady = 1;
return 1;
}
static std::shared_ptr<BogieGraphNetImpl> CreateDirectModel(int SeedOffset)
{
std::shared_ptr<BogieGraphNetImpl> Net;
torch::manual_seed(
365+
SeedOffset);
Net =
std::make_shared<BogieGraphNetImpl>(
GNNHidden);
Net->to(
GNNDevice);
return Net;
}
// ============================================================================
// DIRECT TRAINING SAMPLE COLLECTION
// ============================================================================
static void AddDirectTrainingSample(
const var* Signals,
var Target,
int TrainingState)
{
int i;
int64_t StateLabel;
if(!Signals)
return;
for(i=0; i<GRAPH_SIGNAL_COUNT; i++)
{
var V = Signals[i];
if(V > 1.0)
V = 1.0;
if(V < -1.0)
V = -1.0;
GNNTrainX.push_back(
(float)V);
}
if(Target > 1.0)
Target = 1.0;
if(Target < -1.0)
Target = -1.0;
GNNTrainY.push_back(
(float)Target);
// 0 = trend
// 1 = mean reversion
// 2 = neutral
StateLabel = 2;
if(TrainingState == 1)
StateLabel = 0;
else if(TrainingState == -1)
StateLabel = 1;
GNNTrainState.push_back(
StateLabel);
}
// ============================================================================
// MODEL PATHS
// ============================================================================
static std::string DirectModelPath(int Cycle)
{
std::ostringstream Path;
Path <<
DirectModelPrefix;
if(Cycle < 10)
Path << "0";
Path <<
Cycle <<
".pt";
return Path.str();
}
static int DirectFileExists(const std::string& Path)
{
std::ifstream File(
Path.c_str(),
std::ios::binary);
if(File.good())
return 1;
return 0;
}
// ============================================================================
// BEST-MODEL MEMORY SNAPSHOT
// ============================================================================
static std::string DirectModelToMemory(
std::shared_ptr<BogieGraphNetImpl> Net)
{
torch::serialize::OutputArchive Archive;
std::ostringstream Stream(
std::ios::out |
std::ios::binary);
Net->save(
Archive);
Archive.save_to(
Stream);
return Stream.str();
}
static void DirectModelFromMemory(
std::shared_ptr<BogieGraphNetImpl> Net,
const std::string& Blob)
{
torch::serialize::InputArchive Archive;
std::istringstream Stream(
Blob,
std::ios::in |
std::ios::binary);
Archive.load_from(
Stream,
GNNDevice);
Net->load(
Archive);
}
// ============================================================================
// DIRECT MODEL SAVE / LOAD
// ============================================================================
static int SaveDirectModel(
std::shared_ptr<BogieGraphNetImpl> Net,
int Cycle)
{
torch::serialize::OutputArchive Archive;
std::string Path;
if(!Net)
return 0;
Path =
DirectModelPath(
Cycle);
try
{
Net->save(
Archive);
Archive.save_to(
Path);
}
catch(const c10::Error& E)
{
printf(
"\nDirect GNN SAVE ERROR cycle %i: %s",
Cycle,
E.what());
return 0;
}
catch(const std::exception& E)
{
printf(
"\nDirect GNN SAVE ERROR cycle %i: %s",
Cycle,
E.what());
return 0;
}
printf(
"\nDirect GNN saved WFO cycle %i -> %s",
Cycle,
Path.c_str());
return 1;
}
static int LoadDirectModel(int Cycle)
{
torch::serialize::InputArchive Archive;
std::shared_ptr<BogieGraphNetImpl> Net;
std::string Path;
Path =
DirectModelPath(
Cycle);
if(!DirectFileExists(Path))
{
printf(
"\nDirect GNN LOAD ERROR: model file not found: %s",
Path.c_str());
return 0;
}
try
{
Archive.load_from(
Path,
GNNDevice);
Net =
CreateDirectModel(
Cycle);
Net->load(
Archive);
Net->to(
GNNDevice);
Net->eval();
}
catch(const c10::Error& E)
{
printf(
"\nDirect GNN LOAD ERROR cycle %i: %s",
Cycle,
E.what());
return 0;
}
catch(const std::exception& E)
{
printf(
"\nDirect GNN LOAD ERROR cycle %i: %s",
Cycle,
E.what());
return 0;
}
GNNModel = Net;
GNNLoadedCycle = Cycle;
printf(
"\nDirect GNN loaded WFO cycle %i from %s",
Cycle,
Path.c_str());
return 1;
}
static int DesiredDirectModelCycle()
{
int Cycle;
Cycle = WFOCycle;
// In live mode WFOCycle can be 0; use the last trained cycle.
if(Cycle <= 0)
{
Cycle = WFOCycles;
if(Cycle < 0)
Cycle = -Cycle;
}
if(Cycle <= 0)
Cycle = 1;
return Cycle;
}
static int EnsureDirectModelLoaded()
{
int Cycle;
Cycle =
DesiredDirectModelCycle();
if(
GNNModel &&
GNNLoadedCycle == Cycle)
{
return 1;
}
return LoadDirectModel(
Cycle);
}
// ============================================================================
// DIRECT TRAINING
// ============================================================================
static var TrainDirectGNNCycle(int Cycle)
{
int Rows;
int ValidRows;
int TrainRows;
int Epoch;
int StaleEpochs;
int PositiveCount;
int NegativeCount;
int i;
double BestValidation;
std::string BestBlob;
std::vector<float> DirectionWeights;
torch::Tensor XAll;
torch::Tensor YAll;
torch::Tensor SAll;
torch::Tensor WAll;
torch::Tensor XTrain;
torch::Tensor YTrain;
torch::Tensor STrain;
torch::Tensor WTrain;
torch::Tensor XValid;
torch::Tensor YValid;
torch::Tensor SValid;
std::shared_ptr<BogieGraphNetImpl> Net;
Rows =
(int)GNNTrainY.size();
if(
Rows < MinTrainingSamples ||
(int)GNNTrainState.size() != Rows ||
(int)GNNTrainX.size() !=
Rows*GRAPH_SIGNAL_COUNT)
{
printf(
"\nDirect GNN TRAIN ERROR: cycle %i has %i usable rows",
Cycle,
Rows);
return 0;
}
// ------------------------------------------------------------------------
// Explicit replacement for +BALANCED:
// calculate inverse-frequency weights by target sign.
// ------------------------------------------------------------------------
PositiveCount = 0;
NegativeCount = 0;
for(i=0; i<Rows; i++)
{
if(GNNTrainY[i] >= 0.0f)
PositiveCount++;
else
NegativeCount++;
}
DirectionWeights.resize(
Rows,
1.0f);
if(
UseBalancedDirectionLoss &&
PositiveCount > 0 &&
NegativeCount > 0)
{
float PositiveWeight;
float NegativeWeight;
PositiveWeight =
(float)Rows/
(2.0f*
(float)PositiveCount);
NegativeWeight =
(float)Rows/
(2.0f*
(float)NegativeCount);
for(i=0; i<Rows; i++)
{
if(GNNTrainY[i] >= 0.0f)
{
DirectionWeights[i] =
PositiveWeight;
}
else
{
DirectionWeights[i] =
NegativeWeight;
}
}
}
// ------------------------------------------------------------------------
// Copy direct C++ vectors into LibTorch tensors.
// ------------------------------------------------------------------------
XAll =
torch::from_blob(
GNNTrainX.data(),
{
Rows,
GRAPH_SIGNAL_COUNT
},
torch::TensorOptions()
.dtype(torch::kFloat32))
.clone()
.to(GNNDevice);
YAll =
torch::from_blob(
GNNTrainY.data(),
{Rows},
torch::TensorOptions()
.dtype(torch::kFloat32))
.clone()
.to(GNNDevice);
SAll =
torch::from_blob(
GNNTrainState.data(),
{Rows},
torch::TensorOptions()
.dtype(torch::kInt64))
.clone()
.to(GNNDevice);
WAll =
torch::from_blob(
DirectionWeights.data(),
{Rows},
torch::TensorOptions()
.dtype(torch::kFloat32))
.clone()
.to(GNNDevice);
// Chronological validation block.
ValidRows =
Rows*
GNNValidationPercent/
100;
if(ValidRows < 1)
ValidRows = 1;
if(ValidRows > Rows/3)
ValidRows = Rows/3;
TrainRows =
Rows-
ValidRows;
XTrain =
XAll.narrow(
0,
0,
TrainRows);
YTrain =
YAll.narrow(
0,
0,
TrainRows);
STrain =
SAll.narrow(
0,
0,
TrainRows);
WTrain =
WAll.narrow(
0,
0,
TrainRows);
XValid =
XAll.narrow(
0,
TrainRows,
ValidRows);
YValid =
YAll.narrow(
0,
TrainRows,
ValidRows);
SValid =
SAll.narrow(
0,
TrainRows,
ValidRows);
Net =
CreateDirectModel(
Cycle);
torch::optim::AdamW Optimizer(
Net->parameters(),
torch::optim::AdamWOptions(
GNNLearningRate)
.weight_decay(
GNNWeightDecay));
torch::nn::CrossEntropyLoss StateLoss;
BestValidation =
std::numeric_limits<double>::infinity();
StaleEpochs = 0;
// ------------------------------------------------------------------------
// GPU training.
// ------------------------------------------------------------------------
for(Epoch=0; Epoch<GNNEpochs; Epoch++)
{
torch::Tensor Permutation;
int Start;
double EpochLoss;
int Batches;
if(!wait(0))
return 0;
EpochLoss = 0.0;
Batches = 0;
Net->train();
Permutation =
torch::randperm(
TrainRows,
torch::TensorOptions()
.dtype(torch::kInt64)
.device(GNNDevice));
for(
Start=0;
Start<TrainRows;
Start += GNNBatchSize)
{
int Count;
torch::Tensor Index;
torch::Tensor XB;
torch::Tensor YB;
torch::Tensor SB;
torch::Tensor WB;
std::tuple<
torch::Tensor,
torch::Tensor,
torch::Tensor> Output;
torch::Tensor DirectionPred;
torch::Tensor StateLogits;
torch::Tensor DirectionError;
torch::Tensor LossDirection;
torch::Tensor LossState;
torch::Tensor Loss;
Count =
GNNBatchSize;
if(
Start+
Count >
TrainRows)
{
Count =
TrainRows-
Start;
}
Index =
Permutation.narrow(
0,
Start,
Count);
XB =
XTrain.index_select(
0,
Index);
YB =
YTrain.index_select(
0,
Index);
SB =
STrain.index_select(
0,
Index);
WB =
WTrain.index_select(
0,
Index);
Output =
Net->forward(
XB);
DirectionPred =
std::get<0>(
Output);
StateLogits =
std::get<1>(
Output);
DirectionError =
DirectionPred-
YB;
LossDirection =
torch::mean(
DirectionError*
DirectionError*
WB);
LossState =
StateLoss(
StateLogits,
SB);
Loss =
LossDirection+
GNNStateLossWeight*
LossState;
Optimizer.zero_grad();
Loss.backward();
torch::nn::utils::clip_grad_norm_(
Net->parameters(),
2.0);
Optimizer.step();
EpochLoss +=
Loss.item<double>();
Batches++;
}
// --------------------------------------------------------------------
// Validation.
// --------------------------------------------------------------------
Net->eval();
{
torch::NoGradGuard NoGrad;
std::tuple<
torch::Tensor,
torch::Tensor,
torch::Tensor> ValidOutput;
torch::Tensor ValidDirection;
torch::Tensor ValidStateLogits;
torch::Tensor ValidError;
torch::Tensor ValidDirectionLoss;
torch::Tensor ValidStateLoss;
torch::Tensor ValidLoss;
double Validation;
ValidOutput =
Net->forward(
XValid);
ValidDirection =
std::get<0>(
ValidOutput);
ValidStateLogits =
std::get<1>(
ValidOutput);
ValidError =
ValidDirection-
YValid;
ValidDirectionLoss =
torch::mean(
ValidError*
ValidError);
ValidStateLoss =
StateLoss(
ValidStateLogits,
SValid);
ValidLoss =
ValidDirectionLoss+
GNNStateLossWeight*
ValidStateLoss;
Validation =
ValidLoss.item<double>();
if(
Validation <
BestValidation-
0.000001)
{
BestValidation =
Validation;
BestBlob =
DirectModelToMemory(
Net);
StaleEpochs = 0;
}
else
{
StaleEpochs++;
}
if(
Epoch % 10 == 0 ||
Epoch ==
GNNEpochs-1)
{
double AverageTrainLoss;
AverageTrainLoss = 0.0;
if(Batches > 0)
{
AverageTrainLoss =
EpochLoss/
Batches;
}
printf(
"\nDirect GNN WFO %i epoch %i train %.6f valid %.6f",
Cycle,
Epoch,
AverageTrainLoss,
Validation);
}
}
if(
StaleEpochs >=
GNNPatience)
{
printf(
"\nDirect GNN WFO %i early stop at epoch %i",
Cycle,
Epoch);
break;
}
}
if(!BestBlob.empty())
{
DirectModelFromMemory(
Net,
BestBlob);
}
Net->to(
GNNDevice);
Net->eval();
GNNModel = Net;
GNNLoadedCycle = Cycle;
if(!SaveDirectModel(
Net,
Cycle))
{
return 0;
}
printf(
"\nDirect GNN WFO %i trained with %i rows; best validation %.6f; positive %i negative %i",
Cycle,
Rows,
BestValidation,
PositiveCount,
NegativeCount);
if(BestValidation <= 0.0)
return 0.0001;
return
BestValidation*
100.0;
}
// ============================================================================
// DIRECT INFERENCE
// ============================================================================
static var PredictDirectGNN(
const var* Signals)
{
std::vector<float> Input;
torch::Tensor X;
std::tuple<
torch::Tensor,
torch::Tensor,
torch::Tensor> Output;
torch::Tensor Direction;
torch::Tensor StateLogits;
torch::Tensor PoolWeights;
torch::Tensor StateProb;
int i;
double Score;
if(
!GNNModel ||
!Signals)
{
return 0.0;
}
Input.resize(
GRAPH_SIGNAL_COUNT);
for(i=0; i<GRAPH_SIGNAL_COUNT; i++)
{
var V = Signals[i];
if(V > 1.0)
V = 1.0;
if(V < -1.0)
V = -1.0;
Input[i] =
(float)V;
}
X =
torch::from_blob(
Input.data(),
{
1,
GRAPH_SIGNAL_COUNT
},
torch::TensorOptions()
.dtype(torch::kFloat32))
.clone()
.to(GNNDevice);
GNNModel->eval();
{
torch::NoGradGuard NoGrad;
Output =
GNNModel->forward(
X);
Direction =
std::get<0>(
Output);
StateLogits =
std::get<1>(
Output);
PoolWeights =
std::get<2>(
Output);
StateProb =
torch::softmax(
StateLogits,
1);
Score =
Direction
.to(torch::kCPU)
.item<float>()*
100.0;
GNNStateTrendProbability =
StateProb[0][0]
.to(torch::kCPU)
.item<float>();
GNNStateMeanProbability =
StateProb[0][1]
.to(torch::kCPU)
.item<float>();
GNNStateNeutralProbability =
StateProb[0][2]
.to(torch::kCPU)
.item<float>();
GNNNodeTrendAuthority =
PoolWeights[0][0]
.to(torch::kCPU)
.item<float>();
GNNNodeMeanAuthority =
PoolWeights[0][1]
.to(torch::kCPU)
.item<float>();
GNNNodeHybridTrendAuthority =
PoolWeights[0][2]
.to(torch::kCPU)
.item<float>();
GNNNodeHybridMeanAuthority =
PoolWeights[0][3]
.to(torch::kCPU)
.item<float>();
GNNNodeStateAuthority =
PoolWeights[0][4]
.to(torch::kCPU)
.item<float>();
}
return Score;
}
// ============================================================================
// BOGIE / MARKET FEATURES
// ============================================================================
var BogieRangePosition()
{
var Highest;
var Lowest;
var Range;
var Position;
Highest =
HH(
NocPeriod,
0);
Lowest =
LL(
NocPeriod,
0);
Range =
Highest-
Lowest;
if(Range <= 0.0)
return 0.0;
Position =
2.0*
(priceC(0)-Lowest)/
Range-
1.0;
return
clamp(
Position,
-1.0,
1.0);
}
// Return 0..1 causal trend authority.
var TrendAuthority(
var ADXNow,
var MMINow,
var ADXThreshold,
var MMIThreshold)
{
var A;
var B;
A =
clamp(
(ADXNow-ADXThreshold)/
RegimeScaleADX,
0.0,
1.0);
B =
clamp(
(MMIThreshold-MMINow)/
RegimeScaleMMI,
0.0,
1.0);
if(A < B)
return A;
return B;
}
// Return 0..1 causal mean-reversion authority.
var MeanAuthority(
var ADXNow,
var MMINow,
var ADXThreshold,
var MMIThreshold)
{
var A;
var B;
A =
clamp(
(ADXThreshold-ADXNow)/
RegimeScaleADX,
0.0,
1.0);
B =
clamp(
(MMINow-MMIThreshold)/
RegimeScaleMMI,
0.0,
1.0);
if(A < B)
return A;
return B;
}
// ============================================================================
// BUILD 5 x 8 GRAPH
// ============================================================================
void BuildFiveNodeGraph(
var RawBogie,
var* SmoothSeries,
var ATRNow,
var PlusNow,
var MinusNow,
var ADXNow,
var AroonNow,
var MMINow,
var RSINow,
var BBOscNow,
var MeanNow,
var* Signals)
{
var ATRSafe;
var Momentum6;
var Momentum3;
var DISpread;
var MeanDeviation;
var TrendStandaloneAuth;
var MeanStandaloneAuth;
var HybridTrendAuth;
var HybridMeanAuth;
var StateSigned;
var StateStrength;
ATRSafe = ATRNow;
if(ATRSafe < PIP)
ATRSafe = PIP;
Momentum6 =
(priceC(0)-priceC(6))/
(3.0*ATRSafe);
Momentum3 =
(priceC(0)-priceC(3))/
(2.0*ATRSafe);
DISpread =
(PlusNow-MinusNow)/
100.0;
MeanDeviation =
(priceC(0)-MeanNow)/
(2.0*ATRSafe);
TrendStandaloneAuth =
TrendAuthority(
ADXNow,
MMINow,
TrendStandaloneADXMin,
TrendStandaloneMMIMax);
MeanStandaloneAuth =
MeanAuthority(
ADXNow,
MMINow,
MeanStandaloneADXMax,
MeanStandaloneMMIMin);
HybridTrendAuth =
TrendAuthority(
ADXNow,
MMINow,
HybridTrendADXMin,
HybridTrendMMIMax);
HybridMeanAuth =
MeanAuthority(
ADXNow,
MMINow,
HybridMeanADXMax,
HybridMeanMMIMin);
StateSigned =
HybridTrendAuth-
HybridMeanAuth;
StateStrength =
HybridTrendAuth;
if(
HybridMeanAuth >
StateStrength)
{
StateStrength =
HybridMeanAuth;
}
// Node 0: standalone TrendML.
Signals[0] =
clamp(
SmoothSeries[0],
-1.0,
1.0);
Signals[1] =
clamp(
2.0*
(
SmoothSeries[0]-
SmoothSeries[1]
),
-1.0,
1.0);
Signals[2] =
clamp(
SmoothSeries[0]-
SmoothSeries[3],
-1.0,
1.0);
Signals[3] =
clamp(
Momentum6,
-1.0,
1.0);
Signals[4] =
clamp(
DISpread,
-1.0,
1.0);
Signals[5] =
clamp(
(ADXNow-25.0)/
25.0,
-1.0,
1.0);
Signals[6] =
clamp(
AroonNow/
100.0,
-1.0,
1.0);
Signals[7] =
clamp(
(75.0-MMINow)/
25.0,
-1.0,
1.0);
// Node 1: standalone MeanReversionML.
Signals[8] =
clamp(
SmoothSeries[0],
-1.0,
1.0);
Signals[9] =
clamp(
RawBogie,
-1.0,
1.0);
Signals[10] =
clamp(
(RSINow-50.0)/
50.0,
-1.0,
1.0);
Signals[11] =
clamp(
(BBOscNow-50.0)/
50.0,
-1.0,
1.0);
Signals[12] =
clamp(
MeanDeviation,
-1.0,
1.0);
Signals[13] =
clamp(
Momentum3,
-1.0,
1.0);
Signals[14] =
clamp(
(MMINow-50.0)/
25.0,
-1.0,
1.0);
Signals[15] =
clamp(
(25.0-ADXNow)/
25.0,
-1.0,
1.0);
// Node 2: Hybrid trend specialist.
Signals[16] =
clamp(
SmoothSeries[0],
-1.0,
1.0);
Signals[17] =
clamp(
2.0*
(
SmoothSeries[0]-
SmoothSeries[1]
),
-1.0,
1.0);
Signals[18] =
clamp(
Momentum6,
-1.0,
1.0);
Signals[19] =
clamp(
DISpread,
-1.0,
1.0);
Signals[20] =
clamp(
AroonNow/
100.0,
-1.0,
1.0);
Signals[21] =
clamp(
2.0*
HybridTrendAuth-
1.0,
-1.0,
1.0);
Signals[22] =
clamp(
(ADXNow-
HybridTrendADXMin)/
25.0,
-1.0,
1.0);
Signals[23] =
clamp(
(HybridTrendMMIMax-
MMINow)/
25.0,
-1.0,
1.0);
// Node 3: Hybrid mean specialist.
Signals[24] =
clamp(
SmoothSeries[0],
-1.0,
1.0);
Signals[25] =
clamp(
RawBogie,
-1.0,
1.0);
Signals[26] =
clamp(
(RSINow-50.0)/
50.0,
-1.0,
1.0);
Signals[27] =
clamp(
(BBOscNow-50.0)/
50.0,
-1.0,
1.0);
Signals[28] =
clamp(
MeanDeviation,
-1.0,
1.0);
Signals[29] =
clamp(
2.0*
HybridMeanAuth-
1.0,
-1.0,
1.0);
Signals[30] =
clamp(
(HybridMeanADXMax-
ADXNow)/
25.0,
-1.0,
1.0);
Signals[31] =
clamp(
(MMINow-
HybridMeanMMIMin)/
25.0,
-1.0,
1.0);
// Node 4: Market-state node.
Signals[32] =
clamp(
HybridTrendAuth,
-1.0,
1.0);
Signals[33] =
clamp(
HybridMeanAuth,
-1.0,
1.0);
Signals[34] =
clamp(
StateSigned,
-1.0,
1.0);
Signals[35] =
clamp(
StateStrength,
-1.0,
1.0);
Signals[36] =
clamp(
(ADXNow-25.0)/
25.0,
-1.0,
1.0);
Signals[37] =
clamp(
(MMINow-50.0)/
25.0,
-1.0,
1.0);
Signals[38] =
clamp(
DISpread,
-1.0,
1.0);
Signals[39] =
clamp(
SmoothSeries[0],
-1.0,
1.0);
}
// ============================================================================
// CAUSAL TRAINING STATE
// ============================================================================
int CausalTrainingState(
var ADXNow,
var MMINow)
{
var TrendState;
var MeanState;
TrendState =
TrendAuthority(
ADXNow,
MMINow,
HybridTrendADXMin,
HybridTrendMMIMax);
MeanState =
MeanAuthority(
ADXNow,
MMINow,
HybridMeanADXMax,
HybridMeanMMIMin);
if(
TrendState >= 0.35 &&
TrendState >=
MeanState+
0.10)
{
return 1;
}
if(
MeanState >= 0.35 &&
MeanState >=
TrendState+
0.10)
{
return -1;
}
return 0;
}
// ============================================================================
// LEARNED GNN STATE
// ============================================================================
int LearnedGNNState()
{
if(
GNNStateTrendProbability >=
StateProbabilityMin &&
GNNStateTrendProbability >
GNNStateMeanProbability &&
GNNStateTrendProbability >
GNNStateNeutralProbability)
{
return 1;
}
if(
GNNStateMeanProbability >=
StateProbabilityMin &&
GNNStateMeanProbability >
GNNStateTrendProbability &&
GNNStateMeanProbability >
GNNStateNeutralProbability)
{
return -1;
}
return 0;
}
// ============================================================================
// CALENDAR
// ============================================================================
int MQLDayToZorro(int MqlDay)
{
if(MqlDay == 0)
return 7;
return MqlDay;
}
int TradeAllowedToday()
{
int DayNow;
DayNow =
dow(0);
if(
DayNow ==
MQLDayToZorro(
NoTradeDay_1))
{
return 0;
}
if(
DayNow ==
MQLDayToZorro(
NoTradeDay_2))
{
return 0;
}
return 1;
}
// ============================================================================
// POSITION SIZING
// ============================================================================
var LegacyBogieAmount()
{
var AmountValue;
var FreeMargin;
var Step;
if(!UseMM)
return FixedAmount;
FreeMargin =
Equity-
MarginVal;
if(FreeMargin < 0.0)
FreeMargin = 0.0;
AmountValue =
FreeMargin*
RiskPercent/
100.0/
1000.0;
if(MiniAcct)
{
Step = 0.01;
AmountValue =
roundto(
AmountValue,
Step);
if(AmountValue < 0.01)
AmountValue = 0.01;
}
else
{
Step = 0.10;
AmountValue =
roundto(
AmountValue,
Step);
if(AmountValue < 0.10)
AmountValue = 0.10;
}
if(AmountValue > 50.0)
AmountValue = 50.0;
return AmountValue;
}
void ConfigureTrendTrade(var ATRNow)
{
Amount =
LegacyBogieAmount();
Risk = 0;
Stop =
TrendStopATRMult*
ATRNow;
TakeProfit = 0;
Trail = 0;
}
void ConfigureMeanTrade(var ATRNow)
{
Amount =
LegacyBogieAmount();
Risk = 0;
Stop =
MeanStopATRMult*
ATRNow;
TakeProfit = 0;
Trail = 0;
}
void ConfigureNeutralTrade(var ATRNow)
{
Amount =
LegacyBogieAmount();
Risk = 0;
Stop =
NeutralStopATRMult*
ATRNow;
TakeProfit = 0;
Trail = 0;
}
// ============================================================================
// TREND TRAILING TMF
// ============================================================================
int GraphTrendTrailTMF(
var TrailDistance,
var GuardDistance)
{
var Candidate;
if(!TradeIsOpen)
return 0;
if(TrailDistance <= 0.0)
return 0;
if(TradeIsShort)
{
Candidate =
priceC(0)+
TrailDistance;
if(
TradeStopLimit >
Candidate+
GuardDistance)
{
TradeStopLimit =
Candidate;
}
}
else
{
Candidate =
priceC(0)-
TrailDistance;
if(
TradeStopLimit <
Candidate-
GuardDistance)
{
TradeStopLimit =
Candidate;
}
}
return 0;
}
// ============================================================================
// DIAGNOSTICS
// ============================================================================
void LogGraphEntry(
cstr Side,
int Regime,
var GNNScore,
var ADXNow,
var MMINow)
{
if(!UseDiagnostics)
return;
printf(
"\n%s Bar %i %s | DirectGNN %.2f | State %i "
"| P(T) %.3f P(M) %.3f P(N) %.3f "
"| Nodes T %.3f M %.3f HT %.3f HM %.3f S %.3f "
"| ADX %.2f MMI %.2f | Amount %.3f",
Asset,
Bar,
Side,
GNNScore,
Regime,
GNNStateTrendProbability,
GNNStateMeanProbability,
GNNStateNeutralProbability,
GNNNodeTrendAuthority,
GNNNodeMeanAuthority,
GNNNodeHybridTrendAuthority,
GNNNodeHybridMeanAuthority,
GNNNodeStateAuthority,
ADXNow,
MMINow,
Amount);
}
// ============================================================================
// ZORRO STRATEGY
// ============================================================================
DLLFUNC void run()
{
var RawBogie;
var SmoothNow;
var ATRNow;
var PlusNow;
var MinusNow;
var ADXNow;
var AroonNow;
var MMINow;
var RSINow;
var BBOscNow;
var MeanNow;
var GNNTarget;
var GNNScore;
var FutureMove;
var TargetScale;
var TrailDistance;
var GuardDistance;
var GraphSignals[GRAPH_SIGNAL_COUNT] = {
0,0,0,0,0,0,0,0,
0,0,0,0,0,0,0,0,
0,0,0,0,0,0,0,0,
0,0,0,0,0,0,0,0,
0,0,0,0,0,0,0,0
};
var* PriceSeries;
var* RawSeries;
var* SmoothSeries;
int TrainingState;
int LearnedState;
int TargetHorizon;
int Allowed;
int LongSignal;
int ShortSignal;
int Cycle;
if(is(FIRSTINITRUN))
require(-3.11);
// ------------------------------------------------------------------------
// Zorro simulation / WFO configuration only.
// No Zorro machine-learning subsystem is enabled.
// ------------------------------------------------------------------------
if(is(INITRUN))
{
set(RECALCULATE);
set(TICKS);
set(LOGFILE);
set(OPENEND);
if(Train)
set(PEEK);
// GPU LibTorch training should normally be sequential.
if(Train)
NumCores =
DirectGPUTrainingCores;
if(!InitializeDirectLibTorch())
{
printf(
"\nDirect LibTorch initialization failed");
return;
}
}
BarPeriod =
StrategyBarPeriod;
LookBack = 300;
Capital = 10000;
StartDate =
BacktestStart;
EndDate =
BacktestEnd;
NumWFOCycles =
WFOCycles;
DataSplit =
WFOTrainPercent;
// This is a general Zorro WFO control, not a neural function.
// It prevents trading at the start of the OOS segment after future-peeking
// targets were used in the preceding training period.
DataHorizon =
MaxPredictionHorizonBars;
asset(
(string)BogieAsset);
algo(
(string)"BogieDirectGNN");
Hedge = 0;
MaxLong = 1;
MaxShort = 1;
// ------------------------------------------------------------------------
// Direct training finalization.
//
// Zorro documents that each WFO training cycle has its own EXITRUN.
// Train LibTorch directly here at the end of the WFO training run.
// ------------------------------------------------------------------------
if(is(EXITRUN))
{
if(Train)
{
Cycle =
WFOCycle;
if(Cycle <= 0)
Cycle = 1;
TrainDirectGNNCycle(
Cycle);
}
GNNModel.reset();
GNNLoadedCycle = 0;
return;
}
// ------------------------------------------------------------------------
// Feature pipeline.
// ------------------------------------------------------------------------
PriceSeries =
series(
priceC(0),
256);
RawBogie =
BogieRangePosition();
RawSeries =
series(
RawBogie,
64);
SmoothNow =
EMA(
RawSeries,
EmaPeriod);
SmoothSeries =
series(
SmoothNow,
64);
ATRNow =
ATR(
ATRPeriod);
PlusNow =
PlusDI(
ADXPeriod);
MinusNow =
MinusDI(
ADXPeriod);
ADXNow =
ADX(
ADXPeriod);
AroonNow =
AroonOsc(
AroonPeriod);
MMINow =
MMI(
PriceSeries,
MMIPeriod);
RSINow =
RSI(
PriceSeries,
RSIPeriod);
BBOscNow =
BBOsc(
PriceSeries,
BBPeriod,
2.0,
MAType_SMA);
MeanNow =
EMA(
PriceSeries,
MeanPeriod);
BuildFiveNodeGraph(
RawBogie,
SmoothSeries,
ATRNow,
PlusNow,
MinusNow,
ADXNow,
AroonNow,
MMINow,
RSINow,
BBOscNow,
MeanNow,
GraphSignals);
if(is(LOOKBACK))
return;
// ------------------------------------------------------------------------
// Causal state determines training target horizon.
// ------------------------------------------------------------------------
TrainingState =
CausalTrainingState(
ADXNow,
MMINow);
TargetHorizon =
NeutralPredictionHorizonBars;
if(TrainingState == 1)
{
TargetHorizon =
TrendPredictionHorizonBars;
}
else if(TrainingState == -1)
{
TargetHorizon =
MeanPredictionHorizonBars;
}
// ------------------------------------------------------------------------
// TRAIN MODE:
// collect samples directly into C++ vectors.
// Samples go directly into the LibTorch training dataset.
// ------------------------------------------------------------------------
if(Train)
{
FutureMove =
priceC(
-TargetHorizon)-
priceC(0);
TargetScale =
TargetATRScale*
ATRNow;
if(TargetScale < PIP)
TargetScale = PIP;
GNNTarget =
FutureMove/
TargetScale;
GNNTarget =
clamp(
GNNTarget,
-1.0,
1.0);
AddDirectTrainingSample(
GraphSignals,
GNNTarget,
TrainingState);
return;
}
// ------------------------------------------------------------------------
// TEST / TRADE MODE:
// load the active WFO model and run direct LibTorch inference.
// ------------------------------------------------------------------------
if(!GNNLibTorchReady)
{
if(!InitializeDirectLibTorch())
return;
}
if(!EnsureDirectModelLoaded())
return;
GNNScore =
PredictDirectGNN(
GraphSignals);
LearnedState =
LearnedGNNState();
// ------------------------------------------------------------------------
// Diagnostic plots.
// ------------------------------------------------------------------------
plot(
"Direct GNN",
GNNScore,
NEW,
BLUE);
plot(
"P Trend",
100.0*
GNNStateTrendProbability,
0,
GREEN);
plot(
"P Mean",
-100.0*
GNNStateMeanProbability,
0,
RED);
plot(
"State Node",
100.0*
GNNNodeStateAuthority,
0,
BLACK);
// ------------------------------------------------------------------------
// Regime-aware mean-reversion exits.
// ------------------------------------------------------------------------
if(LearnedState == -1)
{
if(
NumOpenLong > 0 &&
priceC(0) >= MeanNow)
{
exitLong();
return;
}
if(
NumOpenShort > 0 &&
priceC(0) <= MeanNow)
{
exitShort();
return;
}
}
// ------------------------------------------------------------------------
// Direction + learned state -> trade signal.
// ------------------------------------------------------------------------
LongSignal = 0;
ShortSignal = 0;
if(LearnedState == 1)
{
if(
GNNScore >
TrendConfidenceThreshold)
{
if(
!RequireDIConfirmation ||
PlusNow >
MinusNow)
{
LongSignal = 1;
}
}
else if(
GNNScore <
-TrendConfidenceThreshold)
{
if(
!RequireDIConfirmation ||
MinusNow >
PlusNow)
{
ShortSignal = 1;
}
}
}
else if(LearnedState == -1)
{
if(
GNNScore >
MeanConfidenceThreshold &&
SmoothNow <=
HybridMeanLongExtreme &&
RSINow <=
HybridMeanLongRSIMax)
{
LongSignal = 1;
}
else if(
GNNScore <
-MeanConfidenceThreshold &&
SmoothNow >=
HybridMeanShortExtreme &&
RSINow >=
HybridMeanShortRSIMin)
{
ShortSignal = 1;
}
}
else if(AllowNeutralTrades)
{
if(
GNNScore >
NeutralConfidenceThreshold)
{
LongSignal = 1;
}
else if(
GNNScore <
-NeutralConfidenceThreshold)
{
ShortSignal = 1;
}
}
// ------------------------------------------------------------------------
// Calendar.
// ------------------------------------------------------------------------
Allowed =
TradeAllowedToday();
if(!Allowed)
{
if(CloseOnNoTradeDay)
{
if(NumOpenLong > 0)
exitLong();
if(NumOpenShort > 0)
exitShort();
}
return;
}
// ------------------------------------------------------------------------
// LONG.
// ------------------------------------------------------------------------
if(LongSignal)
{
if(NumOpenShort > 0)
{
exitShort();
return;
}
if(NumOpenLong == 0)
{
if(LearnedState == 1)
{
ConfigureTrendTrade(
ATRNow);
TrailDistance =
TrendTrailATRMult*
ATRNow;
GuardDistance =
TrailGuardPips*
PIP;
LogGraphEntry(
"LONG",
LearnedState,
GNNScore,
ADXNow,
MMINow);
enterLong(
GraphTrendTrailTMF,
TrailDistance,
GuardDistance);
}
else if(LearnedState == -1)
{
ConfigureMeanTrade(
ATRNow);
LogGraphEntry(
"LONG",
LearnedState,
GNNScore,
ADXNow,
MMINow);
enterLong();
}
else if(AllowNeutralTrades)
{
ConfigureNeutralTrade(
ATRNow);
LogGraphEntry(
"LONG",
LearnedState,
GNNScore,
ADXNow,
MMINow);
enterLong();
}
}
return;
}
// ------------------------------------------------------------------------
// SHORT.
// ------------------------------------------------------------------------
if(ShortSignal)
{
if(NumOpenLong > 0)
{
exitLong();
return;
}
if(NumOpenShort == 0)
{
if(LearnedState == 1)
{
ConfigureTrendTrade(
ATRNow);
TrailDistance =
TrendTrailATRMult*
ATRNow;
GuardDistance =
TrailGuardPips*
PIP;
LogGraphEntry(
"SHORT",
LearnedState,
GNNScore,
ADXNow,
MMINow);
enterShort(
GraphTrendTrailTMF,
TrailDistance,
GuardDistance);
}
else if(LearnedState == -1)
{
ConfigureMeanTrade(
ATRNow);
LogGraphEntry(
"SHORT",
LearnedState,
GNNScore,
ADXNow,
MMINow);
enterShort();
}
else if(AllowNeutralTrades)
{
ConfigureNeutralTrade(
ATRNow);
LogGraphEntry(
"SHORT",
LearnedState,
GNNScore,
ADXNow,
MMINow);
enterShort();
}
}
return;
}
}