Files
DoomRTX/d3d12/QD3D12NeuralPOM_CoopVectorTrainPersistent.hlsl
T
Justin Marshall 76f642286e Added modelviewer.
Added neural network code(turned off by default).
Lots of editor and rendering fixes.
2026-05-17 22:09:25 -07:00

446 lines
16 KiB
HLSL

// QD3D12 NeuralPOM persistent cooperative-vector training kernels.
// Compile with DXC SM 6.10, not D3DCompile:
// dxc -I <dxc-or-agility-sdk>/include/hlsl -T cs_6_10 -HV 2021 -enable-16bit-types -E CSAccumulateGradients QD3D12NeuralPOM_CoopVectorTrainPersistent.hlsl -Fo QD3D12NeuralPOM_CoopVectorTrainAccum.cso
// dxc -I <dxc-or-agility-sdk>/include/hlsl -T cs_6_10 -HV 2021 -enable-16bit-types -E CSAdamUpdate QD3D12NeuralPOM_CoopVectorTrainPersistent.hlsl -Fo QD3D12NeuralPOM_CoopVectorTrainAdam.cso
//
// The accumulate pass uses dx::linalg MatrixScope::Thread operations for the
// MLP forward pass, transposed backprop matvecs, and outer-product gradient
// accumulation. The Adam pass keeps the trained weights, latent grid, and Adam
// moments on the GPU between batches.
#ifndef QD3D12_NPOM_MAX_INPUT
#define QD3D12_NPOM_MAX_INPUT 32
#endif
#ifndef QD3D12_NPOM_MAX_HIDDEN
#define QD3D12_NPOM_MAX_HIDDEN 128
#endif
#ifndef QD3D12_NPOM_MAX_OUTPUT
#define QD3D12_NPOM_MAX_OUTPUT 16
#endif
#ifndef QD3D12_NPOM_MAX_LATENT_CHANNELS
#define QD3D12_NPOM_MAX_LATENT_CHANNELS 16
#endif
#ifndef QD3D12_NPOM_OUTPUT_COUNT
#define QD3D12_NPOM_OUTPUT_COUNT 13
#endif
#include "dx/linalg.h"
using namespace dx::linalg;
static const uint NPOM_MAX_I = QD3D12_NPOM_MAX_INPUT;
static const uint NPOM_MAX_H = QD3D12_NPOM_MAX_HIDDEN;
static const uint NPOM_MAX_O = QD3D12_NPOM_MAX_OUTPUT;
static const uint NPOM_MAX_C = QD3D12_NPOM_MAX_LATENT_CHANNELS;
static const uint NPOM_O = QD3D12_NPOM_OUTPUT_COUNT;
static const uint NPOM_STATUS_MAGIC = 0x52545643u; // 'CVTR'
StructuredBuffer<float4> gSamples : register(t0); // 6 float4 records per sample.
ByteAddressBuffer gWeights : register(t1); // padded row-major W1,b1,W2,b2,W3,b3.
ByteAddressBuffer gLatent : register(t2); // padded FP32 latent grid.
RWByteAddressBuffer gGradWeights : register(u0);
RWByteAddressBuffer gGradLatent : register(u1);
RWByteAddressBuffer gBatchStats : register(u2); // f32: lossSum, sampleCount, coopKernelSamples, reserved; u32 magic at byte 16.
RWByteAddressBuffer gWeightsRW : register(u3);
RWByteAddressBuffer gLatentRW : register(u4);
RWByteAddressBuffer gAdamMWeights : register(u5);
RWByteAddressBuffer gAdamVWeights : register(u6);
RWByteAddressBuffer gAdamMLatent : register(u7);
RWByteAddressBuffer gAdamVLatent : register(u8);
cbuffer CoopVecTrainCB : register(b0)
{
uint gSampleCount;
uint gInputCount;
uint gHiddenCount;
uint gOutputCount;
uint gLatentChannels;
uint gLatentResolution;
uint gOptimizerStep;
uint gWeightFloatCount;
uint gLatentFloatCount;
uint gW1OffsetBytes;
uint gB1OffsetBytes;
uint gW2OffsetBytes;
uint gB2OffsetBytes;
uint gW3OffsetBytes;
uint gB3OffsetBytes;
uint gLatentOffsetBytes;
float gNormalStrength01;
float gLearningRate;
float gGradientScale;
float gWeightL2;
float gLatentLearningRate;
float gLatentL2;
float gBeta1;
float gBeta2;
float gInvBiasCorrection1;
float gInvBiasCorrection2;
float gAdamEps;
uint gReserved0;
};
float LoadF32(ByteAddressBuffer b, uint byteOffset)
{
return asfloat(b.Load(byteOffset));
}
float LoadF32RW(RWByteAddressBuffer b, uint byteOffset)
{
return asfloat(b.Load(byteOffset));
}
void StoreF32(RWByteAddressBuffer b, uint byteOffset, float v)
{
b.Store(byteOffset, asuint(v));
}
void AtomicAddF32(RWByteAddressBuffer b, uint byteOffset, float v)
{
if (v == 0.0)
return;
uint oldBits = b.Load(byteOffset);
uint observed = 0u;
[allow_uav_condition]
for (;;)
{
uint newBits = asuint(asfloat(oldBits) + v);
b.InterlockedCompareExchange(byteOffset, oldBits, newBits, observed);
if (observed == oldBits)
break;
oldBits = observed;
}
}
void AtomicAccumulateVectorWeights(vector<float, QD3D12_NPOM_MAX_HIDDEN> v, uint byteOffset, uint activeCount)
{
[unroll]
for (uint i = 0; i < NPOM_MAX_H; ++i)
{
if (i < activeCount)
AtomicAddF32(gGradWeights, byteOffset + i * 4u, v[i]);
}
}
void AtomicAccumulateInputRow(vector<float, QD3D12_NPOM_MAX_INPUT> v, uint byteOffset, uint activeCount)
{
[unroll]
for (uint i = 0; i < NPOM_MAX_I; ++i)
{
if (i < activeCount)
AtomicAddF32(gGradWeights, byteOffset + i * 4u, v[i]);
}
}
void AtomicAccumulateLatent(vector<float, QD3D12_NPOM_MAX_LATENT_CHANNELS> v, uint byteOffset, uint activeCount)
{
[unroll]
for (uint i = 0; i < NPOM_MAX_C; ++i)
{
if (i < activeCount)
AtomicAddF32(gGradLatent, byteOffset + i * 4u, v[i]);
}
}
float SamplePacked(uint sampleIndex, uint scalarIndex)
{
float4 v = gSamples[sampleIndex * 6u + scalarIndex / 4u];
return v[scalarIndex & 3u];
}
uint LatentTexelBaseIndex(int x, int y)
{
int r = max((int)gLatentResolution, 1);
x %= r;
y %= r;
if (x < 0) x += r;
if (y < 0) y += r;
return ((uint)y * gLatentResolution + (uint)x) * NPOM_MAX_C;
}
vector<float, QD3D12_NPOM_MAX_LATENT_CHANNELS> LoadLatentVector(uint paddedFloatIndex)
{
vector<float, QD3D12_NPOM_MAX_LATENT_CHANNELS> v = (vector<float, QD3D12_NPOM_MAX_LATENT_CHANNELS>)0;
[unroll]
for (uint c = 0; c < NPOM_MAX_C; ++c)
{
if (c < gLatentChannels)
v[c] = LoadF32(gLatent, gLatentOffsetBytes + (paddedFloatIndex + c) * 4u);
}
return v;
}
vector<float, QD3D12_NPOM_MAX_HIDDEN> TanhHidden(vector<float, QD3D12_NPOM_MAX_HIDDEN> z)
{
vector<float, QD3D12_NPOM_MAX_HIDDEN> r = (vector<float, QD3D12_NPOM_MAX_HIDDEN>)0;
[unroll]
for (uint i = 0; i < NPOM_MAX_H; ++i)
r[i] = (i < gHiddenCount) ? tanh(z[i]) : 0.0;
return r;
}
float Sigmoid1(float x)
{
x = clamp(x, -40.0, 40.0);
return 1.0 / (1.0 + exp(-x));
}
vector<float, QD3D12_NPOM_MAX_OUTPUT> ActivateOutput(vector<float, QD3D12_NPOM_MAX_OUTPUT> z)
{
vector<float, QD3D12_NPOM_MAX_OUTPUT> y = (vector<float, QD3D12_NPOM_MAX_OUTPUT>)0;
[unroll]
for (uint i = 0; i < NPOM_MAX_O; ++i)
{
if (i >= NPOM_O)
y[i] = 0.0;
else if (i == 0u || i == 1u || i == 3u || i == 4u)
y[i] = tanh(z[i]);
else
y[i] = Sigmoid1(z[i]);
}
return y;
}
vector<float, QD3D12_NPOM_MAX_OUTPUT> OutputActivationDeriv(vector<float, QD3D12_NPOM_MAX_OUTPUT> y)
{
vector<float, QD3D12_NPOM_MAX_OUTPUT> d = (vector<float, QD3D12_NPOM_MAX_OUTPUT>)0;
[unroll]
for (uint i = 0; i < NPOM_MAX_O; ++i)
{
if (i >= NPOM_O)
d[i] = 0.0;
else if (i == 0u || i == 1u || i == 3u || i == 4u)
d[i] = 1.0 - y[i] * y[i];
else
d[i] = y[i] * (1.0 - y[i]);
}
return d;
}
vector<float, QD3D12_NPOM_MAX_OUTPUT> LossWeights(float viewZ)
{
vector<float, QD3D12_NPOM_MAX_OUTPUT> w = (vector<float, QD3D12_NPOM_MAX_OUTPUT>)0;
w[0] = 8.0; w[1] = 8.0; w[2] = 2.0;
w[3] = 2.5; w[4] = 2.5; w[5] = 0.8; w[6] = 1.5;
w[7] = 0.25; w[8] = 0.25; w[9] = 0.25;
w[10] = 0.20; w[11] = 0.20; w[12] = 0.20;
float grazingWeight = clamp(1.0 / max(viewZ, 0.16), 1.0, 4.0);
[unroll]
for (uint i = 0; i < 5u; ++i)
w[i] *= grazingWeight;
return w;
}
[numthreads(128, 1, 1)]
void CSAccumulateGradients(uint3 dispatchThreadId : SV_DispatchThreadID)
{
uint sampleIndex = dispatchThreadId.x;
if (sampleIndex >= gSampleCount)
return;
float uvx = SamplePacked(sampleIndex, 0u);
float uvy = SamplePacked(sampleIndex, 1u);
float viewX = SamplePacked(sampleIndex, 2u);
float viewY = SamplePacked(sampleIndex, 3u);
float viewZ = SamplePacked(sampleIndex, 4u);
float lightX = SamplePacked(sampleIndex, 5u);
float lightY = SamplePacked(sampleIndex, 6u);
float lightZ = SamplePacked(sampleIndex, 7u);
float distanceFade = SamplePacked(sampleIndex, 8u);
float fx = frac(uvx) * (float)gLatentResolution - 0.5;
float fy = frac(uvy) * (float)gLatentResolution - 0.5;
int x0 = (int)floor(fx);
int y0 = (int)floor(fy);
float tx = fx - (float)x0;
float ty = fy - (float)y0;
uint latentIdx0 = LatentTexelBaseIndex(x0, y0);
uint latentIdx1 = LatentTexelBaseIndex(x0 + 1, y0);
uint latentIdx2 = LatentTexelBaseIndex(x0, y0 + 1);
uint latentIdx3 = LatentTexelBaseIndex(x0 + 1, y0 + 1);
float lw0 = (1.0 - tx) * (1.0 - ty);
float lw1 = tx * (1.0 - ty);
float lw2 = (1.0 - tx) * ty;
float lw3 = tx * ty;
vector<float, QD3D12_NPOM_MAX_LATENT_CHANNELS> latent0 = LoadLatentVector(latentIdx0);
vector<float, QD3D12_NPOM_MAX_LATENT_CHANNELS> latent1 = LoadLatentVector(latentIdx1);
vector<float, QD3D12_NPOM_MAX_LATENT_CHANNELS> latent2 = LoadLatentVector(latentIdx2);
vector<float, QD3D12_NPOM_MAX_LATENT_CHANNELS> latent3 = LoadLatentVector(latentIdx3);
vector<float, QD3D12_NPOM_MAX_LATENT_CHANNELS> latent = latent0 * lw0 + latent1 * lw1 + latent2 * lw2 + latent3 * lw3;
vector<float, QD3D12_NPOM_MAX_INPUT> input = (vector<float, QD3D12_NPOM_MAX_INPUT>)0;
[unroll]
for (uint c = 0; c < NPOM_MAX_C; ++c)
{
if (c < gLatentChannels)
input[c] = latent[c];
}
input[gLatentChannels + 0u] = viewX;
input[gLatentChannels + 1u] = viewY;
input[gLatentChannels + 2u] = viewZ;
input[gLatentChannels + 3u] = lightX;
input[gLatentChannels + 4u] = lightY;
input[gLatentChannels + 5u] = lightZ;
input[gLatentChannels + 6u] = gNormalStrength01;
input[gLatentChannels + 7u] = distanceFade;
input[gLatentChannels + 8u] = 0.0;
using W1Ty = Matrix<ComponentType::F32, QD3D12_NPOM_MAX_HIDDEN, QD3D12_NPOM_MAX_INPUT, MatrixUse::A, MatrixScope::Thread>;
using W2Ty = Matrix<ComponentType::F32, QD3D12_NPOM_MAX_HIDDEN, QD3D12_NPOM_MAX_HIDDEN, MatrixUse::A, MatrixScope::Thread>;
using W3Ty = Matrix<ComponentType::F32, QD3D12_NPOM_MAX_OUTPUT, QD3D12_NPOM_MAX_HIDDEN, MatrixUse::A, MatrixScope::Thread>;
using W1TTy = Matrix<ComponentType::F32, QD3D12_NPOM_MAX_INPUT, QD3D12_NPOM_MAX_HIDDEN, MatrixUse::A, MatrixScope::Thread>;
using W2TTy = Matrix<ComponentType::F32, QD3D12_NPOM_MAX_HIDDEN, QD3D12_NPOM_MAX_HIDDEN, MatrixUse::A, MatrixScope::Thread>;
using W3TTy = Matrix<ComponentType::F32, QD3D12_NPOM_MAX_HIDDEN, QD3D12_NPOM_MAX_OUTPUT, MatrixUse::A, MatrixScope::Thread>;
W1Ty w1 = W1Ty::Load<MatrixLayout::RowMajor>(gWeights, gW1OffsetBytes, NPOM_MAX_I * 4u);
W2Ty w2 = W2Ty::Load<MatrixLayout::RowMajor>(gWeights, gW2OffsetBytes, NPOM_MAX_H * 4u);
W3Ty w3 = W3Ty::Load<MatrixLayout::RowMajor>(gWeights, gW3OffsetBytes, NPOM_MAX_H * 4u);
VectorRef<ComponentType::F32, QD3D12_NPOM_MAX_HIDDEN> b1 = { gWeights, gB1OffsetBytes };
VectorRef<ComponentType::F32, QD3D12_NPOM_MAX_HIDDEN> b2 = { gWeights, gB2OffsetBytes };
VectorRef<ComponentType::F32, QD3D12_NPOM_MAX_OUTPUT> b3 = { gWeights, gB3OffsetBytes };
vector<float, QD3D12_NPOM_MAX_HIDDEN> z1 = MultiplyAdd<float>(w1, input, b1);
vector<float, QD3D12_NPOM_MAX_HIDDEN> h1 = TanhHidden(z1);
vector<float, QD3D12_NPOM_MAX_HIDDEN> z2 = MultiplyAdd<float>(w2, h1, b2);
vector<float, QD3D12_NPOM_MAX_HIDDEN> h2 = TanhHidden(z2);
vector<float, QD3D12_NPOM_MAX_OUTPUT> z3 = MultiplyAdd<float>(w3, h2, b3);
vector<float, QD3D12_NPOM_MAX_OUTPUT> y = ActivateOutput(z3);
vector<float, QD3D12_NPOM_MAX_OUTPUT> target = (vector<float, QD3D12_NPOM_MAX_OUTPUT>)0;
[unroll]
for (uint o = 0; o < NPOM_O; ++o)
target[o] = SamplePacked(sampleIndex, 9u + o);
vector<float, QD3D12_NPOM_MAX_OUTPUT> lossW = LossWeights(viewZ);
vector<float, QD3D12_NPOM_MAX_OUTPUT> d = y - target;
vector<float, QD3D12_NPOM_MAX_OUTPUT> dy = 2.0 * lossW * d;
vector<float, QD3D12_NPOM_MAX_OUTPUT> dz3 = dy * OutputActivationDeriv(y);
float loss = 0.0;
[unroll]
for (uint lo = 0; lo < NPOM_O; ++lo)
loss += lossW[lo] * d[lo] * d[lo];
AtomicAddF32(gBatchStats, 0u, loss);
AtomicAddF32(gBatchStats, 4u, 1.0);
AtomicAddF32(gBatchStats, 8u, 1.0);
gBatchStats.Store(16u, NPOM_STATUS_MAGIC);
W3TTy w3t = W3TTy::Load<MatrixLayout::ColMajor>(gWeights, gW3OffsetBytes, NPOM_MAX_H * 4u);
vector<float, QD3D12_NPOM_MAX_HIDDEN> dh2 = Multiply<float>(w3t, dz3);
vector<float, QD3D12_NPOM_MAX_HIDDEN> dz2 = dh2 * (1.0 - h2 * h2);
W2TTy w2t = W2TTy::Load<MatrixLayout::ColMajor>(gWeights, gW2OffsetBytes, NPOM_MAX_H * 4u);
vector<float, QD3D12_NPOM_MAX_HIDDEN> dh1 = Multiply<float>(w2t, dz2);
vector<float, QD3D12_NPOM_MAX_HIDDEN> dz1 = dh1 * (1.0 - h1 * h1);
W1TTy w1t = W1TTy::Load<MatrixLayout::ColMajor>(gWeights, gW1OffsetBytes, NPOM_MAX_I * 4u);
vector<float, QD3D12_NPOM_MAX_INPUT> dx = Multiply<float>(w1t, dz1);
// Do not use MatrixScope::Thread Matrix::InterlockedAccumulate here. In the
// SM 6.10 dx/linalg.h API that overload has only two arguments because
// thread-scope accumulator matrices are written in OuterProductOptimal layout.
// This trainer's Adam pass expects gGradWeights to have the same padded
// row-major layout as gWeightsRW, so accumulate rows explicitly.
[loop]
for (uint wo = 0; wo < NPOM_MAX_O; ++wo)
{
if (wo < NPOM_O)
AtomicAccumulateVectorWeights(h2 * dz3[wo], gW3OffsetBytes + wo * NPOM_MAX_H * 4u, gHiddenCount);
}
[loop]
for (uint wh2 = 0; wh2 < NPOM_MAX_H; ++wh2)
{
if (wh2 < gHiddenCount)
AtomicAccumulateVectorWeights(h1 * dz2[wh2], gW2OffsetBytes + wh2 * NPOM_MAX_H * 4u, gHiddenCount);
}
[loop]
for (uint wh1 = 0; wh1 < NPOM_MAX_H; ++wh1)
{
if (wh1 < gHiddenCount)
AtomicAccumulateInputRow(input * dz1[wh1], gW1OffsetBytes + wh1 * NPOM_MAX_I * 4u, gInputCount);
}
[unroll]
for (uint bo = 0; bo < NPOM_MAX_O; ++bo)
{
if (bo < NPOM_O)
AtomicAddF32(gGradWeights, gB3OffsetBytes + bo * 4u, dz3[bo]);
}
[unroll]
for (uint bh = 0; bh < NPOM_MAX_H; ++bh)
{
if (bh < gHiddenCount)
{
AtomicAddF32(gGradWeights, gB2OffsetBytes + bh * 4u, dz2[bh]);
AtomicAddF32(gGradWeights, gB1OffsetBytes + bh * 4u, dz1[bh]);
}
}
vector<float, QD3D12_NPOM_MAX_LATENT_CHANNELS> dLatent = (vector<float, QD3D12_NPOM_MAX_LATENT_CHANNELS>)0;
[unroll]
for (uint lc = 0; lc < NPOM_MAX_C; ++lc)
{
if (lc < gLatentChannels)
dLatent[lc] = dx[lc];
}
AtomicAccumulateLatent(dLatent * lw0, latentIdx0 * 4u, gLatentChannels);
AtomicAccumulateLatent(dLatent * lw1, latentIdx1 * 4u, gLatentChannels);
AtomicAccumulateLatent(dLatent * lw2, latentIdx2 * 4u, gLatentChannels);
AtomicAccumulateLatent(dLatent * lw3, latentIdx3 * 4u, gLatentChannels);
}
float AdamUpdateOne(float p, float g, inout float m, inout float v, float lr, float l2, float clampLo, float clampHi)
{
g = g * gGradientScale + p * l2;
m = gBeta1 * m + (1.0 - gBeta1) * g;
v = gBeta2 * v + (1.0 - gBeta2) * g * g;
float mh = m * gInvBiasCorrection1;
float vh = v * gInvBiasCorrection2;
p -= lr * mh / (sqrt(max(vh, 0.0)) + gAdamEps);
return clamp(p, clampLo, clampHi);
}
[numthreads(256, 1, 1)]
void CSAdamUpdate(uint3 dispatchThreadId : SV_DispatchThreadID)
{
uint id = dispatchThreadId.x;
if (id < gWeightFloatCount)
{
uint byteOffset = id * 4u;
float p = LoadF32RW(gWeightsRW, byteOffset);
float g = LoadF32RW(gGradWeights, byteOffset);
float m = LoadF32RW(gAdamMWeights, byteOffset);
float v = LoadF32RW(gAdamVWeights, byteOffset);
p = AdamUpdateOne(p, g, m, v, gLearningRate, gWeightL2, -8.0, 8.0);
StoreF32(gWeightsRW, byteOffset, p);
StoreF32(gAdamMWeights, byteOffset, m);
StoreF32(gAdamVWeights, byteOffset, v);
}
if (id < gLatentFloatCount)
{
uint byteOffsetL = id * 4u;
float pL = LoadF32RW(gLatentRW, byteOffsetL);
float gL = LoadF32RW(gGradLatent, byteOffsetL);
float mL = LoadF32RW(gAdamMLatent, byteOffsetL);
float vL = LoadF32RW(gAdamVLatent, byteOffsetL);
pL = AdamUpdateOne(pL, gL, mL, vL, gLatentLearningRate, gLatentL2, -4.0, 4.0);
StoreF32(gLatentRW, byteOffsetL, pL);
StoreF32(gAdamMLatent, byteOffsetL, mL);
StoreF32(gAdamVLatent, byteOffsetL, vL);
}
}