// Numerical gradient check for the fused sequence-LSTM kernels // (CPU_LSTMSeqForward / CPU_LSTMSeqBackward in WarriorCPU.cpp). // // Why this exists: those kernels implement backpropagation-through-time from // scratch. A wrong BPTT does not crash and does not look wrong in a training // log - it just learns worse than it should, which is indistinguishable from // "this architecture does not suit the data". The only way to know the gradient // is right is to check it against finite differences. // // Loss is L = sum(k * h_k) over the final hidden state, so dL/dh_{T-1} = [0,1,2,...], // an asymmetric vector - a symmetric one can hide index-swap bugs (e.g. reading // gate f where i was meant). // // Build (from a VS x64 developer prompt, in this folder): // cl /nologo /EHsc /O2 /std:c++17 lstm_seq_gradcheck.cpp WarriorCPU.lib // Run: lstm_seq_gradcheck.exe -> prints max relative error, exits 0 on PASS. #include #include #include #include #include "WarriorCPU.h" static const int H = 5; // hidden units static const int IW = 3; // per-step input width static const int T = 4; // timesteps - must be > 1 or the recurrence is untested struct Ctx { CpuHandle h; int w, in, cg, cc, ch, out, og, wg, ig; }; // Runs the forward pass and returns L = sum(k * h_k) over the final hidden state. static double Forward(Ctx &c, const std::vector &weights, const std::vector &inputs) { CPU_BufferWrite(c.h, c.w, weights.data(), (int)weights.size()); CPU_BufferWrite(c.h, c.in, inputs.data(), (int)inputs.size()); if(!CPU_LSTMSeqForward(c.h, c.w, c.in, c.cg, c.cc, c.ch, c.out, H, IW, T)) { std::printf("FAIL: SeqForward returned 0 (err=%d)\n", CPU_GetLastError(c.h)); std::exit(2); } std::vector out((size_t)H, 0.0); CPU_BufferRead(c.h, c.out, out.data(), H); double loss = 0.0; for(int k = 0; k < H; k++) loss += k * out[k]; return loss; } int main() { Ctx c; c.h = CPU_Init(2); const int perGate = H * (H + IW + 1); const int wTotal = 4 * perGate; c.w = CPU_BufferCreate(c.h, wTotal); c.in = CPU_BufferCreate(c.h, T * IW); c.cg = CPU_BufferCreate(c.h, T * 4 * H); c.cc = CPU_BufferCreate(c.h, T * H); c.ch = CPU_BufferCreate(c.h, T * H); c.out = CPU_BufferCreate(c.h, H); c.og = CPU_BufferCreate(c.h, H); c.wg = CPU_BufferCreate(c.h, wTotal); c.ig = CPU_BufferCreate(c.h, T * IW); std::mt19937 rng(12345); std::uniform_real_distribution dist(-0.6, 0.6); std::vector weights((size_t)wTotal), inputs((size_t)(T * IW)); for(auto &v : weights) v = dist(rng); for(auto &v : inputs) v = dist(rng); // dL/dh_{T-1} for L = sum(k * h_k) std::vector dOut((size_t)H); for(int k = 0; k < H; k++) dOut[k] = (double)k; CPU_BufferWrite(c.h, c.og, dOut.data(), H); Forward(c, weights, inputs); // populate the caches the backward pass reads if(!CPU_LSTMSeqBackward(c.h, c.w, c.in, c.cg, c.cc, c.ch, c.og, c.wg, c.ig, H, IW, T)) { std::printf("FAIL: SeqBackward returned 0 (err=%d)\n", CPU_GetLastError(c.h)); return 2; } std::vector dW((size_t)wTotal, 0.0), dX((size_t)(T * IW), 0.0); CPU_BufferRead(c.h, c.wg, dW.data(), wTotal); CPU_BufferRead(c.h, c.ig, dX.data(), T * IW); // Central differences. The kernels floor gate derivatives at MIN_ACTIVATION_DERIVATIVE (1e-4), // which is a deliberate deviation from the true derivative near saturation, so a weight whose // gate is saturated will legitimately disagree. Skip those rather than report a false failure: // what is being verified here is the BPTT wiring - which term flows where - not the floor. const double EPS = 1e-6; double worstW = 0.0, worstX = 0.0; int checkedW = 0, skipped = 0; for(int i = 0; i < wTotal; i += 7) // stride: 4*perGate is 720 here, no need to check every one { double save = weights[i]; weights[i] = save + EPS; double lp = Forward(c, weights, inputs); weights[i] = save - EPS; double lm = Forward(c, weights, inputs); weights[i] = save; double num = (lp - lm) / (2 * EPS); double den = std::max(1.0, std::max(std::fabs(num), std::fabs(dW[i]))); double rel = std::fabs(num - dW[i]) / den; if(rel > 1e-4 && std::fabs(num) < 1e-3 && std::fabs(dW[i]) < 1e-3) { skipped++; // both ~zero: dominated by the derivative floor continue; } checkedW++; if(rel > worstW) worstW = rel; } for(int i = 0; i < T * IW; i++) { double save = inputs[i]; inputs[i] = save + EPS; double lp = Forward(c, weights, inputs); inputs[i] = save - EPS; double lm = Forward(c, weights, inputs); inputs[i] = save; double num = (lp - lm) / (2 * EPS); double den = std::max(1.0, std::max(std::fabs(num), std::fabs(dX[i]))); double rel = std::fabs(num - dX[i]) / den; if(rel > worstX) worstX = rel; } std::printf("checked %d weights (%d skipped near derivative floor), %d inputs\n", checkedW, skipped, T * IW); std::printf("max relative error dW=%.3e dX=%.3e\n", worstW, worstX); // Guard against the check trivially passing on an all-zero gradient. double mag = 0.0; for(double v : dX) mag += std::fabs(v); std::printf("sum|dX| = %.6f (must be non-zero, else nothing was tested)\n", mag); bool pass = (worstW < 1e-5 && worstX < 1e-5 && mag > 1e-6); std::printf("%s\n", pass ? "PASS" : "FAIL"); CPU_Shutdown(c.h); return pass ? 0 : 1; }