Warrior_EA/DirectML/WarriorDML.h
AnimateDread 6a687cda41 feat: add SGD+momentum optimizer and input-driven hyperparameters
Replace hardcoded lr and momentum with new input variables for Adam and
SGD+momentum. Add OpenCL kernel LSTM_UpdateWeightsMomentum alongside the
existing Adam kernel. Update comments and revert beta1 to book default 0.9.
2026-07-18 14:56:41 -04:00

127 lines
7.6 KiB
C

//+------------------------------------------------------------------+
//| Warrior_EA |
//| D3D12 compute fallback for GPUs|
//| without OpenCL (DirectX12/D3D)|
//+------------------------------------------------------------------+
// Flat C ABI so MQL5 can #import this DLL directly. Every buffer is a
// plain array of 64-bit doubles on the MQL5 side; internally the GPU
// works in 32-bit floats (see WarriorDML.cpp for why) and this layer
// does the conversion, so callers never see the difference.
#pragma once
#define WARRIORDML_API extern "C" __declspec(dllexport)
// Every exported function (besides DML_Init) takes a DmlHandle as its first
// argument - an opaque pointer to a heap-allocated, self-contained context
// (its own D3D12 device/queue/command list/fence, compiled kernels and
// buffer table) that DML_Init() allocates and the caller (CDirectMLMy on the
// MQL5 side) is responsible for remembering and passing back on every
// subsequent call, then releasing via DML_Shutdown(). No state is shared
// between contexts and this DLL keeps no global/static mutable state of its
// own, so it can be loaded any number of times and driven by any number of
// instances/threads in parallel - a driver fault, a device-removal, or a
// wedged call against one context can never poison another context's calls,
// unlike a process-wide singleton device would.
// A plain `long long` (not C++ `long`, which is only 32 bits on Windows) so
// it round-trips exactly through MQL5's 64-bit `long` on the #import side.
typedef long long DmlHandle;
// Lifecycle. DML_Init() returns a non-zero handle on success (a usable D3D12
// device was created), 0 otherwise - callers should fall back to CPU on 0. A
// failed attempt allocates no context to hand back (there is nothing to pass
// to DML_Shutdown() or DML_GetLastError()), so the failure stage is instead
// written directly to *outError (may be NULL if the caller doesn't care) so
// the MQL5 log can show *why* (e.g. "no DX12 adapter" vs "kernel compile
// failed") without this DLL needing to keep any state alive past the failure.
// See DML_GetLastError() below for the error codes.
WARRIORDML_API DmlHandle __stdcall DML_Init(int *outError);
WARRIORDML_API void __stdcall DML_Shutdown(DmlHandle ctx);
// 0 = no error / not yet run
// 1 = CreateDXGIFactory1 failed
// 2 = no hardware D3D_FEATURE_LEVEL_11_0 adapter found (only software/none enumerated)
// 3 = command queue creation failed
// 4 = command allocator creation failed
// 5 = command list creation failed
// 6 = fence creation failed
// 7 = fence event creation failed
// 8 = one or more compute kernels failed to compile/create (HLSL/PSO)
// 9 = device lost (TDR/removal/wedge) after a successful Init, or a fault
// during a kernel dispatch
// Reads the error off a live context (e.g. after a DML_* call returns 0) -
// see DML_Init()'s outError parameter for reading an Init() failure itself.
WARRIORDML_API int __stdcall DML_GetLastError(DmlHandle ctx);
// Buffer management. Handles are small non-negative integers scoped to ctx;
// -1 means failure. Buffers are GPU-resident; Write uploads, Read reads back.
WARRIORDML_API int __stdcall DML_BufferCreate(DmlHandle ctx, int elementCount);
WARRIORDML_API int __stdcall DML_BufferWrite(DmlHandle ctx, int handle, const double *data, int count);
WARRIORDML_API int __stdcall DML_BufferRead(DmlHandle ctx, int handle, double *data, int count);
WARRIORDML_API void __stdcall DML_BufferFree(DmlHandle ctx, int handle);
// Compute kernels - argument order/semantics mirror AI\Network.cl 1:1.
// activation: 0 = TANH, 1 = SIGMOID.
WARRIORDML_API int __stdcall DML_FeedForward(DmlHandle ctx, int wHandle, int iHandle, int oHandle,
int inputs, int activation);
WARRIORDML_API int __stdcall DML_CalcOutputGradient(DmlHandle ctx, int tHandle, int oHandle, int igHandle,
int activation, int count);
WARRIORDML_API int __stdcall DML_CalcHiddenGradient(DmlHandle ctx, int wHandle, int gHandle, int oHandle, int igHandle,
int outputs, int activation, int count);
WARRIORDML_API int __stdcall DML_UpdateWeightsMomentum(DmlHandle ctx, int wHandle, int gHandle, int iHandle, int dwHandle,
int inputs, double learningRate, double momentum, int neurons);
WARRIORDML_API int __stdcall DML_UpdateWeightsAdam(DmlHandle ctx, int wHandle, int gHandle, int iHandle,
int mHandle, int vHandle,
int inputs, double lt, double b1, double b2, int neurons);
// Convolution layer kernels (CNeuronConvOCL) - mirror AI\Network.cl's
// FeedForwardConv/CalcHiddenGradientConv/UpdateWeightsConv* 1:1. A single
// (windowIn+1)*windowOut weight block is shared across every sliding position.
// activation: 0 = TANH, 1 = SIGMOID, 2 = PReLU (fixed param 0.01).
WARRIORDML_API int __stdcall DML_FeedForwardConv(DmlHandle ctx, int wHandle, int iHandle, int oHandle,
int inputs, int step, int windowIn, int windowOut, int activation, int positions);
WARRIORDML_API int __stdcall DML_CalcHiddenGradientConv(DmlHandle ctx, int wHandle, int gHandle, int oHandle, int igHandle,
int outputs, int step, int windowIn, int windowOut, int activation, int inputCount);
WARRIORDML_API int __stdcall DML_UpdateWeightsConvMomentum(DmlHandle ctx, int wHandle, int gHandle, int iHandle, int dwHandle,
int inputs, double learningRate, double momentum, int windowIn, int windowOut, int step);
WARRIORDML_API int __stdcall DML_UpdateWeightsConvAdam(DmlHandle ctx, int wHandle, int gHandle, int iHandle,
int mHandle, int vHandle,
int inputs, double lt, double b1, double b2, int windowIn, int windowOut, int step);
// LSTM layer kernels (CNeuronLSTMOCL) - single-timestep-truncated BPTT,
// Adam-only. Derived from the standard LSTM equations, NOT ported from any
// reference (see AI\Network.cl for why). Gate order: forget, input, output,
// candidate; weight layout per gate is a H*(H+I+1) block (H weights on
// h_prev, I weights on the current input, 1 bias), 4 such blocks total.
WARRIORDML_API int __stdcall DML_LSTMGates(DmlHandle ctx, int wHandle, int hiddenPrevHandle, int inputsHandle,
int concatenatedHandle, int hiddenSize, int inputSize);
WARRIORDML_API int __stdcall DML_LSTMState(DmlHandle ctx, int concatenatedHandle, int memoryHandle, int hiddenPrevHandle,
int hiddenCacheHandle, int outputHandle, int hiddenSize);
WARRIORDML_API int __stdcall DML_LSTMGateGradient(DmlHandle ctx, int gradientHandle, int memoryHandle, int concatenatedHandle,
int concatenatedGradientHandle, int hiddenSize);
WARRIORDML_API int __stdcall DML_LSTMWeightsGradient(DmlHandle ctx, int concatenatedGradientHandle, int hiddenCacheHandle,
int inputsHandle, int weightsGradientHandle, int hiddenSize, int inputSize);
WARRIORDML_API int __stdcall DML_LSTMInputsGradient(DmlHandle ctx, int concatenatedGradientHandle, int wHandle,
int inputsGradientHandle, int hiddenSize, int inputSize);
WARRIORDML_API int __stdcall DML_LSTMUpdateWeightsAdam(DmlHandle ctx, int wHandle, int weightsGradientHandle,
int mHandle, int vHandle, double l, double b1, double b2, int total);
WARRIORDML_API int __stdcall DML_LSTMUpdateWeightsMomentum(DmlHandle ctx, int wHandle, int weightsGradientHandle,
int dwHandle, double learningRate, double momentum, int total);
// Max-pooling layer kernels (CNeuronPoolOCL) - no weights.
WARRIORDML_API int __stdcall DML_FeedForwardProof(DmlHandle ctx, int iHandle, int oHandle, int inputs, int window, int step, int outputs);
WARRIORDML_API int __stdcall DML_CalcInputGradientProof(DmlHandle ctx, int iHandle, int gHandle, int oHandle, int igHandle,
int outputs, int window, int step, int inputs);