//+------------------------------------------------------------------+ //| Warrior_EA | //| D3D12 compute fallback kernels, semantics mirror | //| AI\Network.cl exactly (see that file for reference). | //+------------------------------------------------------------------+ // Why float instead of double: DXGI_FORMAT/typed UAV double support is // spotty across consumer GPUs (D3D12_FEATURE_DATA_D3D12_OPTIONS:: // DoublePrecisionFloatShaderOps is frequently false on gaming cards), // while 32-bit float compute is universally supported on any DX12/D3D12 // capable device. This tier only exists for hardware that has neither // OpenCL nor a "real" GPU compute story, so we trade a little precision // for something that actually runs; the OpenCL and CPU paths remain // full double precision. Conversion happens at the buffer boundary so // MQL5 never has to know. #include "WarriorDML.h" #include #include #include #include #include #include #include #include #pragma comment(lib, "d3d12.lib") #pragma comment(lib, "dxgi.lib") #pragma comment(lib, "d3dcompiler.lib") using Microsoft::WRL::ComPtr; namespace { struct GpuBuffer { ComPtr gpu; // default heap, UAV state, RWStructuredBuffer ComPtr upload; // upload heap staging buffer, sized to match ComPtr readback; // readback heap staging buffer, sized to match UINT count = 0; bool inUse = false; }; struct Kernel { ComPtr rootSig; ComPtr pso; }; //+------------------------------------------------------------------+ //| Everything one CDirectMLMy/CNet instance needs, heap-allocated by | //| DML_Init() and handed back as an opaque DmlHandle. No part of | //| this DLL keeps any global or static mutable state - every export | //| below (besides DML_Init) operates only on the DmlContext its | //| caller passes in, so a device-removal, a driver fault, or a | //| wedged wait can only ever strand its own context (and only the | //| one GPU device it owns), never another instance's, and the DLL | //| can be loaded/used by any number of instances or threads in | //| parallel with zero cross-talk - including each getting its own | //| D3D12 device, which multiple instances in one process can hold | //| concurrently without contention. | //+------------------------------------------------------------------+ struct DmlContext { // Guards buffer-table bookkeeping and the command-list submit+wait sequence for THIS context's // device only. Unlike WarriorCPU.cpp's pool (which can dispatch many chunks across many worker // threads at once), this context has exactly one shared command allocator/list, so the whole // submit+wait sequence must stay serialized end to end - narrowing the lock scope here would let // a second call reset the command allocator while the GPU is still using it, which is undefined // behavior at the driver level. A driver TDR normally recovers within ~2s and CheckDeviceLost() // below would catch it anyway; WaitForGPU()'s bounded timeout only fires if the wait itself never // returns (a genuinely wedged driver) - without it, MQL5's DLL-call watchdog could kill this // thread mid-wait while `mutex` is held, leaving it locked forever. Because this mutex belongs to // one context alone, that can only strand this one instance, never any other chart's calls. std::mutex mutex; ComPtr device; ComPtr queue; ComPtr allocator; ComPtr cmdList; ComPtr fence; HANDLE fenceEvent = nullptr; UINT64 fenceValue = 0; // Set once GetDeviceRemovedReason() confirms the adapter is gone (driver crash/reset/TDR, GPU // physically unplugged on a hot-swap rig, etc.). Every dispatch keeps trying a dead device // otherwise - each call pays a full driver-call timeout before FAILED() finally trips, so a lost // GPU silently turns training into a multi-second-per-call stall instead of a clean, fast, // detectable failure. Latches until this context is torn down (DML_Shutdown) and a fresh one is // created by another DML_Init() call - there is no in-place recovery. bool deviceLost = false; int lastError = 0; std::vector buffers; Kernel kFeedForward, kOutputGradient, kHiddenGradient, kMomentum, kAdam; Kernel kFeedForwardConv, kHiddenGradientConv, kMomentumConv, kAdamConv; Kernel kLstmGates, kLstmState, kLstmGateGradient, kLstmWeightsGradient, kLstmInputsGradient, kLstmUpdateWeightsAdam, kLstmUpdateWeightsMomentum; Kernel kFeedForwardProof, kCalcInputGradientProof; ~DmlContext() { if(fenceEvent) CloseHandle(fenceEvent); } }; void CheckDeviceLost(DmlContext *ctx) { if(ctx->device && ctx->device->GetDeviceRemovedReason() != S_OK) { ctx->deviceLost = true; ctx->lastError = 9; } } //+------------------------------------------------------------------+ //| GPU sync helper - waits for all work submitted so far | //+------------------------------------------------------------------+ bool WaitForGPU(DmlContext *ctx) { if(!ctx->queue || !ctx->fence) return false; ctx->fenceValue++; if(FAILED(ctx->queue->Signal(ctx->fence.Get(), ctx->fenceValue))) { CheckDeviceLost(ctx); return false; } if(ctx->fence->GetCompletedValue() < ctx->fenceValue) { if(FAILED(ctx->fence->SetEventOnCompletion(ctx->fenceValue, ctx->fenceEvent))) { CheckDeviceLost(ctx); return false; } const DWORD FENCE_WAIT_TIMEOUT_MS = 15000; if(WaitForSingleObject(ctx->fenceEvent, FENCE_WAIT_TIMEOUT_MS) != WAIT_OBJECT_0) { ctx->deviceLost = true; // wedged/unresponsive device - stop trusting it ctx->lastError = 9; return false; } } CheckDeviceLost(ctx); return !ctx->deviceLost; } bool ResetCommandList(DmlContext *ctx) { if(FAILED(ctx->allocator->Reset())) return false; if(FAILED(ctx->cmdList->Reset(ctx->allocator.Get(), nullptr))) return false; return true; } //+------------------------------------------------------------------+ //| Build a root signature: N root UAVs (u0..u(N-1)) + a root | //| constants block (b0) holding constDwords 32-bit values. | //+------------------------------------------------------------------+ bool CreateRootSignature(DmlContext *ctx, int numUAV, int constDwords, ComPtr &outSig) { std::vector params; params.reserve(numUAV + 1); for(int i = 0; i < numUAV; i++) { D3D12_ROOT_PARAMETER p = {}; p.ParameterType = D3D12_ROOT_PARAMETER_TYPE_UAV; p.ShaderVisibility = D3D12_SHADER_VISIBILITY_ALL; p.Descriptor.ShaderRegister = i; p.Descriptor.RegisterSpace = 0; params.push_back(p); } if(constDwords > 0) { D3D12_ROOT_PARAMETER p = {}; p.ParameterType = D3D12_ROOT_PARAMETER_TYPE_32BIT_CONSTANTS; p.ShaderVisibility = D3D12_SHADER_VISIBILITY_ALL; p.Constants.ShaderRegister = 0; p.Constants.RegisterSpace = 0; p.Constants.Num32BitValues = constDwords; params.push_back(p); } D3D12_ROOT_SIGNATURE_DESC desc = {}; desc.NumParameters = (UINT)params.size(); desc.pParameters = params.data(); desc.Flags = D3D12_ROOT_SIGNATURE_FLAG_NONE; ComPtr blob, error; if(FAILED(D3D12SerializeRootSignature(&desc, D3D_ROOT_SIGNATURE_VERSION_1, &blob, &error))) return false; return SUCCEEDED(ctx->device->CreateRootSignature(0, blob->GetBufferPointer(), blob->GetBufferSize(), IID_PPV_ARGS(&outSig))); } bool CompileKernel(DmlContext *ctx, const char *hlsl, int numUAV, int constDwords, Kernel &out) { ComPtr code, error; HRESULT hr = D3DCompile(hlsl, strlen(hlsl), nullptr, nullptr, nullptr, "main", "cs_5_0", 0, 0, &code, &error); if(FAILED(hr)) return false; // error->GetBufferPointer() has the message if debugging is needed if(!CreateRootSignature(ctx, numUAV, constDwords, out.rootSig)) return false; D3D12_COMPUTE_PIPELINE_STATE_DESC psoDesc = {}; psoDesc.pRootSignature = out.rootSig.Get(); psoDesc.CS = {code->GetBufferPointer(), code->GetBufferSize()}; return SUCCEEDED(ctx->device->CreateComputePipelineState(&psoDesc, IID_PPV_ARGS(&out.pso))); } //+------------------------------------------------------------------+ //| HLSL sources - one-to-one with the OpenCL kernels in Network.cl. | //| Scalar-per-thread (no double4 batching - GPU is wide enough that | //| the manual SIMD packing in the OpenCL version buys little here, | //| and it keeps this file auditable against the .cl reference). | //+------------------------------------------------------------------+ const char *kHlslFeedForward = R"( RWStructuredBuffer matrix_w : register(u0); RWStructuredBuffer matrix_i : register(u1); RWStructuredBuffer matrix_o : register(u2); cbuffer Consts : register(b0) { int inputs; int activation; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int i = (int)id.x; int shift = (inputs + 1) * i; float sum = 0.0; for(int k = 0; k < inputs; k++) sum += matrix_i[k] * matrix_w[shift + k]; sum += matrix_w[shift + inputs]; // bias if(activation == 0) sum = tanh(sum); else if(activation == 1) sum = 1.0 / (1.0 + exp(-clamp(sum, -50.0, 50.0))); matrix_o[i] = sum; } )"; const char *kHlslOutputGradient = R"( RWStructuredBuffer matrix_t : register(u0); RWStructuredBuffer matrix_o : register(u1); RWStructuredBuffer matrix_ig : register(u2); cbuffer Consts : register(b0) { int activation; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int i = (int)id.x; float out_v = matrix_o[i]; float temp = 0.0; if(activation == 0) { // Not multiplied by the tanh derivative (1-out^2) - see WarriorCPU.cpp's // CPU_CalcOutputGradient for why (it vanishes exactly at the +-1 // targets a bounded output neuron needs to reach, e.g. a regression // head's buy/sell extremes). temp = clamp(matrix_t[i], -1.0, 1.0) - out_v; } else if(activation == 1) { // Also deliberately NOT multiplied by the sigmoid derivative out_v*(1-out_v) - see // WarriorCPU.cpp's CPU_CalcOutputGradient for the full reasoning. This is the // classification output layer (3 neurons, one-hot 0/1 targets); (target-out) unscaled // is exactly the sigmoid+binary-cross-entropy gradient, whereas re-multiplying by // out_v*(1-out_v) damps toward zero as out_v approaches 0 or 1 - exactly where a // 0/1-target neuron needs to converge, and was observed in practice as all 3 output // neurons converging to (and never escaping) an identical value. temp = clamp(matrix_t[i], 0.0, 1.0) - out_v; } else { // NONE (raw logits - the 3-class softmax classification output layer) and PRELU (2, never // actually used as an OUTPUT activation in this codebase) both fall here. Identity/NONE has // derivative 1 everywhere, so the plain (target-out) error passes through unscaled - // matching CNeuron::calcOutputGradients()'s unconditional formula on the MQL5-side // un-accelerated fallback. Leaving this case unhandled left temp at its 0.0 initializer, // i.e. every softmax-classification output neuron got a permanently zero gradient on this // DLL - and since kHlslHiddenGradient propagates backward FROM the output layer's gradient, // that silently froze the whole network (not just the output layer) at its random initial // weights whenever the DirectML/GPU tier was active. temp = matrix_t[i] - out_v; } matrix_ig[i] = temp; } )"; const char *kHlslHiddenGradient = R"( RWStructuredBuffer matrix_w : register(u0); RWStructuredBuffer matrix_g : register(u1); RWStructuredBuffer matrix_o : register(u2); RWStructuredBuffer matrix_ig : register(u3); cbuffer Consts : register(b0) { int outputs; int activation; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int i = (int)id.x; int shift = (outputs + 1) * i; float out_v = matrix_o[i]; float sum = 0.0; for(int k = 0; k < outputs; k++) sum += matrix_g[k] * matrix_w[shift + k]; if(activation == 0) { sum = clamp(sum + out_v, -1.0, 1.0) - out_v; sum = sum * max(1e-4, 1.0 - out_v * out_v); } else if(activation == 1) { sum = clamp(sum + out_v, 0.0, 1.0) - out_v; sum = sum * max(1e-4, out_v * (1.0 - out_v)); } else if(activation == 2 && out_v < 0.0) { // PReLU, param=0.01 - see WarriorCPU.cpp's matching CPU_CalcHiddenGradient comment for the // full rationale. This branch was missing, so every hidden PReLU layer's backprop gradient // was left completely unscaled (1.0 instead of 0.01) whenever out_v was negative, on this // (DirectML GPU) tier. sum = sum * 0.01; } matrix_ig[i] = sum; } )"; // Bias weight (j == inputs) is intentionally never touched here - this // matches the existing OpenCL UpdateWeightsMomentum kernel's dispatch // range (0..inputs-1) exactly, so SGD behaves identically on both backends. // The +-100 clamp on every matrix_w write below (here and in kHlslAdam/kHlslMomentumConv/ // kHlslAdamConv/kHlslLstmUpdateWeightsAdam) mirrors AI\Network.cl's MAX_WEIGHT: without it a gradient // spike can drive a weight to +-Infinity, and the next Adam step turns that into NaN (Inf/Inf) that // propagates through every FeedForward sum touching it and never recovers. Tightened from 1e6 - see // WarriorCPU.cpp's matching MAX_WEIGHT comment for why that ceiling was too loose to ever engage. // Every Adam kernel below (kHlslAdam/kHlslAdamConv/kHlslLstmUpdateWeightsAdam) also applies a // decoupled AdamW-style weight decay (the inline "- l * 0.001 * matrix_w[...]" term) - see // WarriorCPU.cpp's WEIGHT_DECAY comment for the full rationale. HLSL string kernels can't share a // C++ #define, so the 0.001 is inlined directly; keep it in sync with WarriorCPU.cpp/Network.cl's // WEIGHT_DECAY if that ever changes. Same reasoning for the inline 1e-4 MIN_ACTIVATION_DERIVATIVE // floor used in the tanh/sigmoid hidden-gradient and LSTM gate kernels below - keep in sync with // WarriorCPU.cpp/Network.cl/Network.mqh's MIN_ACTIVATION_DERIVATIVE if that ever changes. The Adam // kernels below apply their step unconditionally - the sign-agreement gate that used to sit on every // matrix_w write was removed (see WarriorCPU.cpp's CPU_UpdateWeightsAdam comment for why: it // rectified the one-hot softmax-CCE gradient stream into an all-Neutral collapse). const char *kHlslMomentum = R"( RWStructuredBuffer matrix_w : register(u0); RWStructuredBuffer matrix_g : register(u1); RWStructuredBuffer matrix_i : register(u2); RWStructuredBuffer matrix_dw : register(u3); cbuffer Consts : register(b0) { int inputs; float learning_rate; float momentum; } [numthreads(8,8,1)] void main(uint3 id : SV_DispatchThreadID) { int i = (int)id.x; int j = (int)id.y; if(j >= inputs) return; int wi = i * (inputs + 1) + j; float delta = learning_rate * matrix_g[i] * matrix_i[j] + momentum * matrix_dw[wi]; matrix_dw[wi] = delta; matrix_w[wi] = clamp(matrix_w[wi] + delta, -100.0, 100.0); } )"; const char *kHlslAdam = R"( RWStructuredBuffer matrix_w : register(u0); RWStructuredBuffer matrix_g : register(u1); RWStructuredBuffer matrix_i : register(u2); RWStructuredBuffer matrix_m : register(u3); RWStructuredBuffer matrix_v : register(u4); cbuffer Consts : register(b0) { int inputs; float l; float b1; float b2; } [numthreads(8,8,1)] void main(uint3 id : SV_DispatchThreadID) { int i = (int)id.x; int j = (int)id.y; if(j > inputs) return; int wi = i * (inputs + 1) + j; float inp = (j < inputs) ? matrix_i[j] : 1.0; float g = matrix_g[i] * inp; float mt = b1 * matrix_m[wi] + (1.0 - b1) * g; float vt = sqrt(b2 * matrix_v[wi] + (1.0 - b2) * g * g); float delta = clamp(l * mt / (vt > 0.0 ? vt : l * 10.0) - l * 0.001 * matrix_w[wi], -0.1, 0.1); matrix_w[wi] = clamp(matrix_w[wi] + delta, -100.0, 100.0); matrix_m[wi] = mt; matrix_v[wi] = vt; } )"; //+------------------------------------------------------------------+ //| Max-pooling layer - no weights, mirrors AI\Network.cl's | //| FeedForwardProof/CalcInputGradientProof. | //+------------------------------------------------------------------+ const char *kHlslFeedForwardProof = R"( RWStructuredBuffer matrix_i : register(u0); RWStructuredBuffer matrix_o : register(u1); cbuffer Consts : register(b0) { int inputs; int window; int step; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int i = (int)id.x; int pos = i * step; float result = matrix_i[pos]; for(int k = 1; k < window; k++) { int shift = k + pos; if(shift >= inputs) break; result = max(result, matrix_i[shift]); } matrix_o[i] = result; } )"; const char *kHlslCalcInputGradientProof = R"( RWStructuredBuffer matrix_i : register(u0); RWStructuredBuffer matrix_g : register(u1); RWStructuredBuffer matrix_o : register(u2); RWStructuredBuffer matrix_ig : register(u3); cbuffer Consts : register(b0) { int outputs; int window; int step; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int i = (int)id.x; float value = matrix_i[i]; int start = i - window + step; start = (start - start % step) / step; int stop = (i - i % step) / step + 1; float prevGrad = 0.0; for(int out = max(0, start); out < min(outputs, stop); out++) if(value == matrix_o[out]) prevGrad += matrix_g[out]; matrix_ig[i] = prevGrad; } )"; //+------------------------------------------------------------------+ //| Convolution kernels - ported alongside AI\Network.cl's Conv | //| kernels. A single (windowIn+1)*windowOut weight block is shared | //| across every sliding position. | //+------------------------------------------------------------------+ const char *kHlslFeedForwardConv = R"( RWStructuredBuffer matrix_w : register(u0); RWStructuredBuffer matrix_i : register(u1); RWStructuredBuffer matrix_o : register(u2); cbuffer Consts : register(b0) { int inputs; int step; int windowIn; int windowOut; int activation; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int i = (int)id.x; int shiftOut = windowOut * i; int shiftIn = step * i; for(int out = 0; out < windowOut; out++) { int shift = (windowIn + 1) * out; int stop = (windowIn <= (inputs - shiftIn)) ? windowIn : (inputs - shiftIn); float sum = 0.0; for(int k = 0; k < stop; k++) sum += matrix_i[shiftIn + k] * matrix_w[shift + k]; sum += matrix_w[shift + windowIn]; if(activation == 0) sum = tanh(sum); else if(activation == 1) sum = 1.0 / (1.0 + exp(-clamp(sum, -50.0, 50.0))); else if(activation == 2 && sum < 0) sum *= 0.01; matrix_o[out + shiftOut] = sum; } } )"; const char *kHlslHiddenGradientConv = R"( RWStructuredBuffer matrix_w : register(u0); RWStructuredBuffer matrix_g : register(u1); RWStructuredBuffer matrix_o : register(u2); RWStructuredBuffer matrix_ig : register(u3); cbuffer Consts : register(b0) { int outputs; int step; int windowIn; int windowOut; int activation; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int i = (int)id.x; float out_v = matrix_o[i]; int start = i - windowIn + step; start = max((start - start % step) / step, 0); int stop = (i - i % step) / step + 1; if(stop > (outputs / windowOut)) stop = outputs / windowOut; float sum = 0.0; for(int h = 0; h < windowOut; h++) { for(int k = start; k < stop; k++) { int shiftW = (stop - k - 1) * step + i % step + h * (windowIn + 1); int shiftG = k * windowOut + h; if(shiftG >= outputs || shiftW >= (windowIn + 1) * windowOut) break; sum += matrix_g[shiftG] * matrix_w[shiftW]; } } if(activation == 0) { sum = clamp(sum + out_v, -1.0, 1.0) - out_v; sum = sum * max(1e-4, 1.0 - out_v * out_v); } else if(activation == 1) sum = (clamp(sum + out_v, 0.0, 1.0) - out_v) * max(1e-4, out_v * (1.0 - out_v)); else if(activation == 2 && out_v < 0) sum *= 0.01; matrix_ig[i] = sum; } )"; const char *kHlslMomentumConv = R"( RWStructuredBuffer matrix_w : register(u0); RWStructuredBuffer matrix_g : register(u1); RWStructuredBuffer matrix_i : register(u2); RWStructuredBuffer matrix_dw : register(u3); cbuffer Consts : register(b0) { int inputs; float learningRate; float momentum; int windowIn; int windowOut; int step; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int i = (int)id.x; if(i >= (windowIn + 1) * windowOut) return; int shift = i % (windowIn + 1); int shiftOut = (i - shift) / (windowIn + 1); int total = (inputs - windowIn) % step; total = (inputs - windowIn - total) / step + (total > 0 ? 1 : 0); float grad = 0.0; for(int t = 0; t < total; t++) { if(shift != windowIn && (shift + t * step) >= inputs) break; grad += matrix_g[t * windowOut + shiftOut] * (shift == windowIn ? 1.0 : matrix_i[shift + t * step]); } float delta = learningRate * grad + momentum * matrix_dw[i]; matrix_dw[i] = delta; matrix_w[i] = clamp(matrix_w[i] + delta, -100.0, 100.0); } )"; const char *kHlslAdamConv = R"( RWStructuredBuffer matrix_w : register(u0); RWStructuredBuffer matrix_g : register(u1); RWStructuredBuffer matrix_i : register(u2); RWStructuredBuffer matrix_m : register(u3); RWStructuredBuffer matrix_v : register(u4); cbuffer Consts : register(b0) { int inputs; float l; float b1; float b2; int windowIn; int windowOut; int step; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int i = (int)id.x; if(i > windowIn) return; int total = (inputs - (windowIn - step)) % step; total = (inputs - (windowIn - step) - total) / step + (total > 0 ? 1 : 0); for(int out = 0; out < windowOut; out++) { int shiftW = i + out * (windowIn + 1); float grad = 0.0; for(int t = 0; t < total; t++) { if(i != windowIn && (i + t * step) >= inputs) break; grad += matrix_g[t * windowOut + out] * (i == windowIn ? 1.0 : matrix_i[i + t * step]); } float mt = b1 * matrix_m[shiftW] + (1.0 - b1) * grad; float vt = sqrt(b2 * matrix_v[shiftW] + (1.0 - b2) * grad * grad); float delta = clamp(l * mt / (vt > 0.0 ? vt : l * 10.0) - l * 0.001 * matrix_w[shiftW], -0.1, 0.1); matrix_w[shiftW] = clamp(matrix_w[shiftW] + delta, -100.0, 100.0); matrix_m[shiftW] = mt; matrix_v[shiftW] = vt; } } )"; //+------------------------------------------------------------------+ //| LSTM kernels - mirror AI\Network.cl's LSTM_* kernels 1:1. See the | //| comment above those kernels for the derivation notes (single- | //| timestep-truncated BPTT, Adam-only, NOT ported from any reference).| //+------------------------------------------------------------------+ const char *kHlslLstmGates = R"( RWStructuredBuffer matrix_w : register(u0); RWStructuredBuffer hidden_prev : register(u1); RWStructuredBuffer inputs : register(u2); RWStructuredBuffer concatenated : register(u3); cbuffer Consts : register(b0) { int hiddenSize; int inputSize; } [numthreads(8,4,1)] void main(uint3 id : SV_DispatchThreadID) { int hid = (int)id.x; int gate = (int)id.y; if(hid >= hiddenSize || gate >= 4) return; int perGate = hiddenSize * (hiddenSize + inputSize + 1); int shift = gate * perGate + hid * (hiddenSize + inputSize + 1); float sum = 0.0; for(int k = 0; k < hiddenSize; k++) sum += hidden_prev[k] * matrix_w[shift + k]; for(int k = 0; k < inputSize; k++) sum += inputs[k] * matrix_w[shift + hiddenSize + k]; sum += matrix_w[shift + hiddenSize + inputSize]; float val = (gate < 3) ? (1.0 / (1.0 + exp(-sum))) : tanh(sum); concatenated[gate * hiddenSize + hid] = val; } )"; const char *kHlslLstmState = R"( RWStructuredBuffer concatenated : register(u0); RWStructuredBuffer memory : register(u1); RWStructuredBuffer hidden_prev : register(u2); RWStructuredBuffer hidden_cache : register(u3); RWStructuredBuffer output : register(u4); cbuffer Consts : register(b0) { int hiddenSize; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int i = (int)id.x; if(i >= hiddenSize) return; float f = concatenated[i]; float ii = concatenated[hiddenSize + i]; float o = concatenated[2 * hiddenSize + i]; float g = concatenated[3 * hiddenSize + i]; float c_prev = memory[i]; memory[hiddenSize + i] = c_prev; float c_t = f * c_prev + ii * g; memory[i] = c_t; hidden_cache[i] = hidden_prev[i]; output[i] = o * tanh(c_t); } )"; const char *kHlslLstmGateGradient = R"( RWStructuredBuffer gradient : register(u0); RWStructuredBuffer memory : register(u1); RWStructuredBuffer concatenated : register(u2); RWStructuredBuffer concatenated_gradient : register(u3); cbuffer Consts : register(b0) { int hiddenSize; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int i = (int)id.x; if(i >= hiddenSize) return; float c_t = memory[i]; float c_prev = memory[hiddenSize + i]; float f = concatenated[i]; float ii = concatenated[hiddenSize + i]; float o = concatenated[2 * hiddenSize + i]; float g = concatenated[3 * hiddenSize + i]; float t = tanh(c_t); float dh = gradient[i]; float dc = dh * o * max(1e-4, 1.0 - t * t); concatenated_gradient[2 * hiddenSize + i] = dh * t * max(1e-4, o * (1.0 - o)); concatenated_gradient[i] = dc * c_prev * max(1e-4, f * (1.0 - f)); concatenated_gradient[hiddenSize + i] = dc * g * max(1e-4, ii * (1.0 - ii)); concatenated_gradient[3 * hiddenSize + i] = dc * ii * max(1e-4, 1.0 - g * g); } )"; const char *kHlslLstmWeightsGradient = R"( RWStructuredBuffer concatenated_gradient : register(u0); RWStructuredBuffer hidden_cache : register(u1); RWStructuredBuffer inputs : register(u2); RWStructuredBuffer weights_gradient : register(u3); cbuffer Consts : register(b0) { int hiddenSize; int inputSize; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int wi = (int)id.x; int perGate = hiddenSize * (hiddenSize + inputSize + 1); if(wi >= 4 * perGate) return; int gate = wi / perGate; int rem = wi % perGate; int hid = rem / (hiddenSize + inputSize + 1); int k = rem % (hiddenSize + inputSize + 1); float inp = (k < hiddenSize) ? hidden_cache[k] : ((k < hiddenSize + inputSize) ? inputs[k - hiddenSize] : 1.0); weights_gradient[wi] = concatenated_gradient[gate * hiddenSize + hid] * inp; } )"; const char *kHlslLstmInputsGradient = R"( RWStructuredBuffer concatenated_gradient : register(u0); RWStructuredBuffer matrix_w : register(u1); RWStructuredBuffer inputs_gradient : register(u2); cbuffer Consts : register(b0) { int hiddenSize; int inputSize; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int j = (int)id.x; if(j >= inputSize) return; int perGate = hiddenSize * (hiddenSize + inputSize + 1); float sum = 0.0; for(int gate = 0; gate < 4; gate++) for(int hid = 0; hid < hiddenSize; hid++) sum += concatenated_gradient[gate * hiddenSize + hid] * matrix_w[gate * perGate + hid * (hiddenSize + inputSize + 1) + hiddenSize + j]; inputs_gradient[j] = sum; } )"; const char *kHlslLstmUpdateWeightsAdam = R"( RWStructuredBuffer matrix_w : register(u0); RWStructuredBuffer weights_gradient : register(u1); RWStructuredBuffer matrix_m : register(u2); RWStructuredBuffer matrix_v : register(u3); cbuffer Consts : register(b0) { float l; float b1; float b2; int total; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int wi = (int)id.x; if(wi >= total) return; float g = weights_gradient[wi]; float mt = b1 * matrix_m[wi] + (1.0 - b1) * g; float vt = sqrt(b2 * matrix_v[wi] + (1.0 - b2) * g * g); float delta = clamp(l * mt / (vt > 0.0 ? vt : l * 10.0) - l * 0.001 * matrix_w[wi], -0.1, 0.1); matrix_w[wi] = clamp(matrix_w[wi] + delta, -100.0, 100.0); matrix_m[wi] = mt; matrix_v[wi] = vt; } )"; // SGD+momentum counterpart to kHlslLstmUpdateWeightsAdam above - same flat // (already-elementwise) weights_gradient/matrix_dw layout as // CPU_LSTMUpdateWeightsMomentum, just the classic heavy-ball update. const char *kHlslLstmUpdateWeightsMomentum = R"( RWStructuredBuffer matrix_w : register(u0); RWStructuredBuffer weights_gradient : register(u1); RWStructuredBuffer matrix_dw : register(u2); cbuffer Consts : register(b0) { float learning_rate; float momentum; int total; } [numthreads(64,1,1)] void main(uint3 id : SV_DispatchThreadID) { int wi = (int)id.x; if(wi >= total) return; float delta = learning_rate * weights_gradient[wi] + momentum * matrix_dw[wi]; matrix_dw[wi] = delta; matrix_w[wi] = clamp(matrix_w[wi] + delta, -100.0, 100.0); } )"; //+------------------------------------------------------------------+ //| Buffer helpers | //+------------------------------------------------------------------+ int AllocBufferSlot(DmlContext *ctx) { for(size_t i = 0; i < ctx->buffers.size(); i++) if(!ctx->buffers[i].inUse) return (int)i; ctx->buffers.push_back(GpuBuffer()); return (int)ctx->buffers.size() - 1; } bool CreateUploadBuffer(DmlContext *ctx, UINT64 size, ComPtr &out) { D3D12_HEAP_PROPERTIES heap = {}; heap.Type = D3D12_HEAP_TYPE_UPLOAD; D3D12_RESOURCE_DESC desc = {}; desc.Dimension = D3D12_RESOURCE_DIMENSION_BUFFER; desc.Width = size; desc.Height = 1; desc.DepthOrArraySize = 1; desc.MipLevels = 1; desc.SampleDesc.Count = 1; desc.Layout = D3D12_TEXTURE_LAYOUT_ROW_MAJOR; return SUCCEEDED(ctx->device->CreateCommittedResource(&heap, D3D12_HEAP_FLAG_NONE, &desc, D3D12_RESOURCE_STATE_GENERIC_READ, nullptr, IID_PPV_ARGS(&out))); } bool CreateReadbackBuffer(DmlContext *ctx, UINT64 size, ComPtr &out) { D3D12_HEAP_PROPERTIES heap = {}; heap.Type = D3D12_HEAP_TYPE_READBACK; D3D12_RESOURCE_DESC desc = {}; desc.Dimension = D3D12_RESOURCE_DIMENSION_BUFFER; desc.Width = size; desc.Height = 1; desc.DepthOrArraySize = 1; desc.MipLevels = 1; desc.SampleDesc.Count = 1; desc.Layout = D3D12_TEXTURE_LAYOUT_ROW_MAJOR; return SUCCEEDED(ctx->device->CreateCommittedResource(&heap, D3D12_HEAP_FLAG_NONE, &desc, D3D12_RESOURCE_STATE_COPY_DEST, nullptr, IID_PPV_ARGS(&out))); } bool CreateDefaultUAVBuffer(DmlContext *ctx, UINT64 size, ComPtr &out) { D3D12_HEAP_PROPERTIES heap = {}; heap.Type = D3D12_HEAP_TYPE_DEFAULT; D3D12_RESOURCE_DESC desc = {}; desc.Dimension = D3D12_RESOURCE_DIMENSION_BUFFER; desc.Width = size; desc.Height = 1; desc.DepthOrArraySize = 1; desc.MipLevels = 1; desc.SampleDesc.Count = 1; desc.Layout = D3D12_TEXTURE_LAYOUT_ROW_MAJOR; desc.Flags = D3D12_RESOURCE_FLAG_ALLOW_UNORDERED_ACCESS; return SUCCEEDED(ctx->device->CreateCommittedResource(&heap, D3D12_HEAP_FLAG_NONE, &desc, D3D12_RESOURCE_STATE_UNORDERED_ACCESS, nullptr, IID_PPV_ARGS(&out))); } } // namespace //+------------------------------------------------------------------+ //| Exported API | //+------------------------------------------------------------------+ WARRIORDML_API int __stdcall DML_GetLastError(DmlHandle handle) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); return ctx->lastError; } namespace { // Device creation and shader compilation run entirely inside the GPU vendor's // driver, which is third-party, closed-source code this project has no // control over - a buggy driver crashing here must not be allowed to take // the whole terminal down with it (extracted to its own function, no locals // with destructors, so it can hold a __try - see WarriorCPU.cpp's SehStartPool // for the same MSVC C2712 constraint this works around). bool DML_InitUnguarded(DmlContext *ctx); bool DML_InitSeh(DmlContext *ctx) { __try { return DML_InitUnguarded(ctx); } __except(EXCEPTION_EXECUTE_HANDLER) { ctx->lastError = 9; return false; } } } WARRIORDML_API DmlHandle __stdcall DML_Init(int *outError) { DmlContext *ctx = new DmlContext(); if(!DML_InitSeh(ctx)) { if(outError) *outError = ctx->lastError; delete ctx; return 0; } if(outError) *outError = 0; return reinterpret_cast(ctx); } namespace { bool DML_InitUnguarded(DmlContext *ctx) { ComPtr factory; if(FAILED(CreateDXGIFactory1(IID_PPV_ARGS(&factory)))) { ctx->lastError = 1; return false; } ComPtr adapter; ComPtr device; for(UINT i = 0; factory->EnumAdapters1(i, &adapter) != DXGI_ERROR_NOT_FOUND; i++) { DXGI_ADAPTER_DESC1 desc; adapter->GetDesc1(&desc); if(desc.Flags & DXGI_ADAPTER_FLAG_SOFTWARE) continue; if(SUCCEEDED(D3D12CreateDevice(adapter.Get(), D3D_FEATURE_LEVEL_11_0, IID_PPV_ARGS(&device)))) break; device.Reset(); } if(!device) { ctx->lastError = 2; // no usable DX12 hardware adapter - caller should fall back to CPU return false; } ctx->device = device; D3D12_COMMAND_QUEUE_DESC qDesc = {}; qDesc.Type = D3D12_COMMAND_LIST_TYPE_COMPUTE; if(FAILED(ctx->device->CreateCommandQueue(&qDesc, IID_PPV_ARGS(&ctx->queue)))) { ctx->lastError = 3; return false; } if(FAILED(ctx->device->CreateCommandAllocator(D3D12_COMMAND_LIST_TYPE_COMPUTE, IID_PPV_ARGS(&ctx->allocator)))) { ctx->lastError = 4; return false; } if(FAILED(ctx->device->CreateCommandList(0, D3D12_COMMAND_LIST_TYPE_COMPUTE, ctx->allocator.Get(), nullptr, IID_PPV_ARGS(&ctx->cmdList)))) { ctx->lastError = 5; return false; } ctx->cmdList->Close(); if(FAILED(ctx->device->CreateFence(0, D3D12_FENCE_FLAG_NONE, IID_PPV_ARGS(&ctx->fence)))) { ctx->lastError = 6; return false; } ctx->fenceValue = 0; ctx->fenceEvent = CreateEvent(nullptr, FALSE, FALSE, nullptr); if(!ctx->fenceEvent) { ctx->lastError = 7; return false; } if(!CompileKernel(ctx, kHlslFeedForward, 3, 2, ctx->kFeedForward) || !CompileKernel(ctx, kHlslOutputGradient, 3, 1, ctx->kOutputGradient) || !CompileKernel(ctx, kHlslHiddenGradient, 4, 2, ctx->kHiddenGradient) || !CompileKernel(ctx, kHlslMomentum, 4, 3, ctx->kMomentum) || !CompileKernel(ctx, kHlslAdam, 5, 4, ctx->kAdam) || !CompileKernel(ctx, kHlslFeedForwardConv, 3, 5, ctx->kFeedForwardConv) || !CompileKernel(ctx, kHlslHiddenGradientConv, 4, 5, ctx->kHiddenGradientConv) || !CompileKernel(ctx, kHlslMomentumConv, 4, 6, ctx->kMomentumConv) || !CompileKernel(ctx, kHlslAdamConv, 5, 7, ctx->kAdamConv) || !CompileKernel(ctx, kHlslLstmGates, 4, 2, ctx->kLstmGates) || !CompileKernel(ctx, kHlslLstmState, 5, 1, ctx->kLstmState) || !CompileKernel(ctx, kHlslLstmGateGradient, 4, 1, ctx->kLstmGateGradient) || !CompileKernel(ctx, kHlslLstmWeightsGradient, 4, 2, ctx->kLstmWeightsGradient) || !CompileKernel(ctx, kHlslLstmInputsGradient, 3, 2, ctx->kLstmInputsGradient) || !CompileKernel(ctx, kHlslLstmUpdateWeightsAdam, 4, 4, ctx->kLstmUpdateWeightsAdam) || !CompileKernel(ctx, kHlslLstmUpdateWeightsMomentum, 3, 3, ctx->kLstmUpdateWeightsMomentum) || !CompileKernel(ctx, kHlslFeedForwardProof, 2, 3, ctx->kFeedForwardProof) || !CompileKernel(ctx, kHlslCalcInputGradientProof, 4, 3, ctx->kCalcInputGradientProof)) { ctx->lastError = 8; return false; } return true; } } // namespace WARRIORDML_API void __stdcall DML_Shutdown(DmlHandle handle) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return; WaitForGPU(ctx); delete ctx; // ComPtr members release the device/queue/kernels/buffers themselves } WARRIORDML_API int __stdcall DML_BufferCreate(DmlHandle handle, int elementCount) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return -1; std::lock_guard lock(ctx->mutex); if(ctx->deviceLost || elementCount <= 0) return -1; int slot = AllocBufferSlot(ctx); GpuBuffer &b = ctx->buffers[slot]; UINT64 size = (UINT64)elementCount * sizeof(float); if(!CreateDefaultUAVBuffer(ctx, size, b.gpu) || !CreateUploadBuffer(ctx, size, b.upload) || !CreateReadbackBuffer(ctx, size, b.readback)) { b = GpuBuffer(); return -1; } b.count = (UINT)elementCount; b.inUse = true; return slot; } WARRIORDML_API void __stdcall DML_BufferFree(DmlHandle handle, int bufHandle) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return; std::lock_guard lock(ctx->mutex); if(bufHandle < 0 || (size_t)bufHandle >= ctx->buffers.size()) return; ctx->buffers[bufHandle] = GpuBuffer(); } WARRIORDML_API int __stdcall DML_BufferWrite(DmlHandle handle, int bufHandle, const double *data, int count) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); if(ctx->deviceLost || bufHandle < 0 || (size_t)bufHandle >= ctx->buffers.size()) return 0; GpuBuffer &b = ctx->buffers[bufHandle]; if(!b.inUse || (UINT)count > b.count) return 0; // Same rationale as WarriorCPU.cpp's CPU_BufferWrite: a NaN/Inf weight or // gradient doesn't crash the GPU path, it silently poisons every value it // touches from then on with no error ever surfaced. Reject at the boundary. for(int i = 0; i < count; i++) if(!std::isfinite(data[i])) return 0; std::vector tmp((size_t)count); for(int i = 0; i < count; i++) tmp[i] = (float)data[i]; void *mapped = nullptr; D3D12_RANGE noRead = {0, 0}; if(FAILED(b.upload->Map(0, &noRead, &mapped))) return 0; memcpy(mapped, tmp.data(), tmp.size() * sizeof(float)); b.upload->Unmap(0, nullptr); if(!ResetCommandList(ctx)) return 0; ctx->cmdList->CopyBufferRegion(b.gpu.Get(), 0, b.upload.Get(), 0, tmp.size() * sizeof(float)); ctx->cmdList->Close(); ID3D12CommandList *lists[] = {ctx->cmdList.Get()}; ctx->queue->ExecuteCommandLists(1, lists); return WaitForGPU(ctx) ? 1 : 0; } WARRIORDML_API int __stdcall DML_BufferRead(DmlHandle handle, int bufHandle, double *data, int count) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); if(ctx->deviceLost || bufHandle < 0 || (size_t)bufHandle >= ctx->buffers.size()) return 0; GpuBuffer &b = ctx->buffers[bufHandle]; if(!b.inUse || (UINT)count > b.count) return 0; if(!ResetCommandList(ctx)) return 0; ctx->cmdList->CopyBufferRegion(b.readback.Get(), 0, b.gpu.Get(), 0, (UINT64)count * sizeof(float)); ctx->cmdList->Close(); ID3D12CommandList *lists[] = {ctx->cmdList.Get()}; ctx->queue->ExecuteCommandLists(1, lists); if(!WaitForGPU(ctx)) return 0; void *mapped = nullptr; D3D12_RANGE readRange = {0, (SIZE_T)count * sizeof(float)}; if(FAILED(b.readback->Map(0, &readRange, &mapped))) return 0; const float *src = (const float *)mapped; for(int i = 0; i < count; i++) data[i] = (double)src[i]; D3D12_RANGE noWrite = {0, 0}; b.readback->Unmap(0, &noWrite); return 1; } namespace { bool DispatchKernelUnguarded(DmlContext *ctx, Kernel &k, const std::vector &bufferHandles, const void *constants, UINT constDwords, UINT gx, UINT gy) { if(ctx->deviceLost) return false; for(int h : bufferHandles) if(h < 0 || (size_t)h >= ctx->buffers.size() || !ctx->buffers[h].inUse) return false; if(!ResetCommandList(ctx)) return false; ctx->cmdList->SetPipelineState(k.pso.Get()); ctx->cmdList->SetComputeRootSignature(k.rootSig.Get()); for(size_t i = 0; i < bufferHandles.size(); i++) ctx->cmdList->SetComputeRootUnorderedAccessView((UINT)i, ctx->buffers[bufferHandles[i]].gpu->GetGPUVirtualAddress()); if(constDwords > 0) ctx->cmdList->SetComputeRoot32BitConstants((UINT)bufferHandles.size(), constDwords, constants, 0); ctx->cmdList->Dispatch(gx, gy, 1); ctx->cmdList->Close(); ID3D12CommandList *lists[] = {ctx->cmdList.Get()}; ctx->queue->ExecuteCommandLists(1, lists); return WaitForGPU(ctx); } // Fault isolation boundary, same rationale as WarriorCPU.cpp's SehCallFn: a // driver bug or a genuinely bad COM call here must degrade to "this dispatch // failed" rather than an access violation that takes the whole terminal down // with it. Every one of the 16 kernel dispatch functions below routes through // this single choke point. bool DispatchKernel(DmlContext *ctx, Kernel &k, const std::vector &bufferHandles, const void *constants, UINT constDwords, UINT gx, UINT gy) { __try { return DispatchKernelUnguarded(ctx, k, bufferHandles, constants, constDwords, gx, gy); } __except(EXCEPTION_EXECUTE_HANDLER) { ctx->deviceLost = true; // unknown-state GPU context - don't trust it further ctx->lastError = 9; return false; } } UINT CeilDiv(int a, int b) { return (UINT)((a + b - 1) / b); } } WARRIORDML_API int __stdcall DML_FeedForward(DmlHandle handle, int wHandle, int iHandle, int oHandle, int inputs, int activation) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int inputs; int activation; } consts = {inputs, activation}; // oHandle is caller-supplied and unchecked at this point - indexing ctx->buffers[oHandle] // directly (as this used to) is an out-of-bounds vector access for a bad handle, // crashing before DispatchKernel ever gets to run its own bounds check below. bool valid = oHandle >= 0 && (size_t)oHandle < ctx->buffers.size() && ctx->buffers[oHandle].inUse; int neurons = valid ? (int)ctx->buffers[oHandle].count : 0; return DispatchKernel(ctx, ctx->kFeedForward, {wHandle, iHandle, oHandle}, &consts, 2, CeilDiv(neurons, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_CalcOutputGradient(DmlHandle handle, int tHandle, int oHandle, int igHandle, int activation, int count) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int activation; } consts = {activation}; return DispatchKernel(ctx, ctx->kOutputGradient, {tHandle, oHandle, igHandle}, &consts, 1, CeilDiv(count, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_CalcHiddenGradient(DmlHandle handle, int wHandle, int gHandle, int oHandle, int igHandle, int outputs, int activation, int count) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int outputs; int activation; } consts = {outputs, activation}; return DispatchKernel(ctx, ctx->kHiddenGradient, {wHandle, gHandle, oHandle, igHandle}, &consts, 2, CeilDiv(count, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_UpdateWeightsMomentum(DmlHandle handle, int wHandle, int gHandle, int iHandle, int dwHandle, int inputs, double learningRate, double momentum, int neurons) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int inputs; float learningRate; float momentum; } consts = {inputs, (float)learningRate, (float)momentum}; return DispatchKernel(ctx, ctx->kMomentum, {wHandle, gHandle, iHandle, dwHandle}, &consts, 3, CeilDiv(neurons, 8), CeilDiv(inputs, 8)) ? 1 : 0; } WARRIORDML_API int __stdcall DML_UpdateWeightsAdam(DmlHandle handle, int wHandle, int gHandle, int iHandle, int mHandle, int vHandle, int inputs, double lt, double b1, double b2, int neurons) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int inputs; float l; float b1; float b2; } consts = {inputs, (float)lt, (float)b1, (float)b2}; return DispatchKernel(ctx, ctx->kAdam, {wHandle, gHandle, iHandle, mHandle, vHandle}, &consts, 4, CeilDiv(neurons, 8), CeilDiv(inputs + 1, 8)) ? 1 : 0; } WARRIORDML_API int __stdcall DML_FeedForwardConv(DmlHandle handle, int wHandle, int iHandle, int oHandle, int inputs, int step, int windowIn, int windowOut, int activation, int positions) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int inputs; int step; int windowIn; int windowOut; int activation; } consts = {inputs, step, windowIn, windowOut, activation}; return DispatchKernel(ctx, ctx->kFeedForwardConv, {wHandle, iHandle, oHandle}, &consts, 5, CeilDiv(positions, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_CalcHiddenGradientConv(DmlHandle handle, int wHandle, int gHandle, int oHandle, int igHandle, int outputs, int step, int windowIn, int windowOut, int activation, int inputCount) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int outputs; int step; int windowIn; int windowOut; int activation; } consts = {outputs, step, windowIn, windowOut, activation}; return DispatchKernel(ctx, ctx->kHiddenGradientConv, {wHandle, gHandle, oHandle, igHandle}, &consts, 5, CeilDiv(inputCount, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_UpdateWeightsConvMomentum(DmlHandle handle, int wHandle, int gHandle, int iHandle, int dwHandle, int inputs, double learningRate, double momentum, int windowIn, int windowOut, int step) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int inputs; float learningRate; float momentum; int windowIn; int windowOut; int step; } consts = {inputs, (float)learningRate, (float)momentum, windowIn, windowOut, step}; int total = (windowIn + 1) * windowOut; return DispatchKernel(ctx, ctx->kMomentumConv, {wHandle, gHandle, iHandle, dwHandle}, &consts, 6, CeilDiv(total, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_UpdateWeightsConvAdam(DmlHandle handle, int wHandle, int gHandle, int iHandle, int mHandle, int vHandle, int inputs, double lt, double b1, double b2, int windowIn, int windowOut, int step) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int inputs; float l; float b1; float b2; int windowIn; int windowOut; int step; } consts = {inputs, (float)lt, (float)b1, (float)b2, windowIn, windowOut, step}; return DispatchKernel(ctx, ctx->kAdamConv, {wHandle, gHandle, iHandle, mHandle, vHandle}, &consts, 7, CeilDiv(windowIn + 1, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_LSTMGates(DmlHandle handle, int wHandle, int hiddenPrevHandle, int inputsHandle, int concatenatedHandle, int hiddenSize, int inputSize) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int hiddenSize; int inputSize; } consts = {hiddenSize, inputSize}; return DispatchKernel(ctx, ctx->kLstmGates, {wHandle, hiddenPrevHandle, inputsHandle, concatenatedHandle}, &consts, 2, CeilDiv(hiddenSize, 8), CeilDiv(4, 4)) ? 1 : 0; } WARRIORDML_API int __stdcall DML_LSTMState(DmlHandle handle, int concatenatedHandle, int memoryHandle, int hiddenPrevHandle, int hiddenCacheHandle, int outputHandle, int hiddenSize) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int hiddenSize; } consts = {hiddenSize}; return DispatchKernel(ctx, ctx->kLstmState, {concatenatedHandle, memoryHandle, hiddenPrevHandle, hiddenCacheHandle, outputHandle}, &consts, 1, CeilDiv(hiddenSize, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_LSTMGateGradient(DmlHandle handle, int gradientHandle, int memoryHandle, int concatenatedHandle, int concatenatedGradientHandle, int hiddenSize) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int hiddenSize; } consts = {hiddenSize}; return DispatchKernel(ctx, ctx->kLstmGateGradient, {gradientHandle, memoryHandle, concatenatedHandle, concatenatedGradientHandle}, &consts, 1, CeilDiv(hiddenSize, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_LSTMWeightsGradient(DmlHandle handle, int concatenatedGradientHandle, int hiddenCacheHandle, int inputsHandle, int weightsGradientHandle, int hiddenSize, int inputSize) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int hiddenSize; int inputSize; } consts = {hiddenSize, inputSize}; int total = 4 * hiddenSize * (hiddenSize + inputSize + 1); return DispatchKernel(ctx, ctx->kLstmWeightsGradient, {concatenatedGradientHandle, hiddenCacheHandle, inputsHandle, weightsGradientHandle}, &consts, 2, CeilDiv(total, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_LSTMInputsGradient(DmlHandle handle, int concatenatedGradientHandle, int wHandle, int inputsGradientHandle, int hiddenSize, int inputSize) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int hiddenSize; int inputSize; } consts = {hiddenSize, inputSize}; return DispatchKernel(ctx, ctx->kLstmInputsGradient, {concatenatedGradientHandle, wHandle, inputsGradientHandle}, &consts, 2, CeilDiv(inputSize, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_LSTMUpdateWeightsAdam(DmlHandle handle, int wHandle, int weightsGradientHandle, int mHandle, int vHandle, double l, double b1, double b2, int total) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { float l; float b1; float b2; int total; } consts = {(float)l, (float)b1, (float)b2, total}; return DispatchKernel(ctx, ctx->kLstmUpdateWeightsAdam, {wHandle, weightsGradientHandle, mHandle, vHandle}, &consts, 4, CeilDiv(total, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_LSTMUpdateWeightsMomentum(DmlHandle handle, int wHandle, int weightsGradientHandle, int dwHandle, double learningRate, double momentum, int total) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { float learningRate; float momentum; int total; } consts = {(float)learningRate, (float)momentum, total}; return DispatchKernel(ctx, ctx->kLstmUpdateWeightsMomentum, {wHandle, weightsGradientHandle, dwHandle}, &consts, 3, CeilDiv(total, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_FeedForwardProof(DmlHandle handle, int iHandle, int oHandle, int inputs, int window, int step, int outputs) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int inputs; int window; int step; } consts = {inputs, window, step}; return DispatchKernel(ctx, ctx->kFeedForwardProof, {iHandle, oHandle}, &consts, 3, CeilDiv(outputs, 64), 1) ? 1 : 0; } WARRIORDML_API int __stdcall DML_CalcInputGradientProof(DmlHandle handle, int iHandle, int gHandle, int oHandle, int igHandle, int outputs, int window, int step, int inputs) { DmlContext *ctx = reinterpret_cast(handle); if(!ctx) return 0; std::lock_guard lock(ctx->mutex); struct { int outputs; int window; int step; } consts = {outputs, window, step}; return DispatchKernel(ctx, ctx->kCalcInputGradientProof, {iHandle, gHandle, oHandle, igHandle}, &consts, 3, CeilDiv(inputs, 64), 1) ? 1 : 0; }