Code
// 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;
	}
}

Last edited by TipmyPip; 09/09/26 14:09.