yum-mirror/slang
Making it easier to work with shaders
git clone https://git.yummers.dev/yum-mirror/slang
f28f67d98
master
1// In this example, we implement a simple multi-layer perceptron (MLP) training loop on 2// Vulkan (through slang-rhi). See also the mlp-training-coopvec example, which 3// implements the same MLP training loop using cooperative vector intrinsics for better 4// performance. 5// 6// The simple MLP is trained to approximate a polynomial expression. 7// The network contains one hidden layer with 16 neurons. It takes 4 inputs and produces 4 8// outputs. 9 10#include "core/slang-basic.h" 11#include "examples/example-base/example-base.h" 12#include "external/slang-rhi/include/slang-rhi.h" 13#include "slang-com-ptr.h" 14#include "slang.h" 15 16#include <string> 17 18using Slang ::ComPtr ; 19 20static const ExampleResources resourceBase ("mlp-training" ); 21 22typedef uint16_t NFloat ; 23 24static const int kLayerSizes []= {4 ,16 ,4 }; 25static const int kLayerCount = sizeof (kLayerSizes ) /sizeof (int )- 1 ; 26 27int getNetworkLayerWeightStride (int i ) 28{ 29return kLayerSizes [i ]* sizeof (NFloat ); 30} 31 32int getNetworkLayerWeightCount (int i ) 33{ 34return kLayerSizes [i ]* kLayerSizes [i + 1 ]; 35} 36 37int getNetworkLayerBiasCount (int i ) 38{ 39return kLayerSizes [i + 1 ]; 40} 41 42struct Kernel 43{ 44ComPtr < rhi::IShaderProgram > program ; 45ComPtr < rhi::IComputePipeline > pipeline ; 46 operatorbool () {return program && pipeline ; } 47}; 48 49struct ClearBufferParams 50{ 51 rhi::DeviceAddress buffer ; 52uint32_t count ; 53}; 54 55struct LearnGradParams 56{ 57 rhi::DeviceAddress networkBuffer ; 58 rhi::DeviceAddress lossBuffer ; 59 rhi::DeviceAddress inputs ; 60uint32_t count ; 61}; 62 63struct AdjustParamsParams 64{ 65 rhi::DeviceAddress adamStates ; 66 rhi::DeviceAddress params ; 67 rhi::DeviceAddress gradients ; 68uint32_t count ; 69}; 70 71struct ExampleProgram :public TestBase 72{ 73ComPtr < rhi::IDevice > gDevice ; 74 75ComPtr < slang::ISession > gSlangSession ; 76ComPtr < slang::IModule > gSlangModule ; 77Kernel gLearnGradProgram ; 78Kernel gAdjustParamProgram ; 79 80// Sub-allocated buffer range for each network layer's parameters (weights, biases, gradients). 81// 82struct NetworkParameterAllocation 83 { 84size_t weightsOffset ; 85size_t weightsSize ; 86size_t biasOffset ; 87size_t biasSize ; 88size_t weightsGradOffset ; 89size_t biasGradOffset ; 90 }; 91 92SlangResult execute (int argc ,char * argv []) 93 { 94parseOption (argc ,argv ); 95 96 rhi::DeviceDesc deviceDesc ; 97deviceDesc .slang .targetProfile = "spirv_1_6" ; 98deviceDesc .deviceType = rhi::DeviceType ::Vulkan ; 99 100gDevice = rhi::getRHI ()-> createDevice (deviceDesc ); 101if (!gDevice ) 102return SLANG_FAIL ; 103 104SLANG_RETURN_ON_FAIL (loadShaderKernels ()); 105 106// Create a buffer to hold all network parameters (weights, biases, gradients). 107// This buffer is arranged as following: 108// (segment 1): | weights0 | bias0 | weights1 | bias1 | ... | weightsN | biasN | 109// (segment 2): | weightsGrad0 | biasGrad0 | weightsGrad1 | biasGrad1 | ... | 110// 111// Where the first segment contains all weights and biases for each layer in row-major 112// layout. The second segment contains gradients for weights and biases in row-major layout. 113 114// Total size of all network parameters. 115size_t paramBufferSize ; 116 117// Offset for the second segment, where gradients for weights and biases in row-major layout 118// start. 119size_t gradientOffset ; 120 121// Sub-allocated weight/Bias offsets for each layer. 122 std::vector < NetworkParameterAllocation > layerAllocations ; 123allocateNetworkParameterStorage (layerAllocations ,paramBufferSize ,gradientOffset ); 124 125 std::vector < uint16_t > initParams ; 126srand (1072 ); 127for (int i = 0 ;i < paramBufferSize /sizeof (NFloat );i ++ ) 128 { 129if (i < gradientOffset /sizeof (NFloat )) 130 { 131float v = rand () / (float )RAND_MAX ; 132v = v * 2.0f - 1.0f ;// Normalize to [-1, 1] 133initParams .push_back (floatToHalf (v )); 134 } 135else 136 { 137// Initialize gradients to zero. 138initParams .push_back (0 ); 139 } 140 } 141auto networkParamsBuffer = createBuffer (paramBufferSize ,initParams .data ()); 142 143static const size_t kAdamStateSize = sizeof (NFloat )* 2 + sizeof (int32_t ); 144auto adamStateBuffer = createBuffer (initParams .size ()* kAdamStateSize ); 145clearBuffer (adamStateBuffer ); 146 147 std::vector < uint64_t > networkConstantBufferData ; 148auto paramBufferAddr = networkParamsBuffer -> getDeviceAddress (); 149for (int i = 0 ;i < kLayerCount ;i ++ ) 150 { 151networkConstantBufferData .push_back ( 152paramBufferAddr + layerAllocations [i ].weightsOffset ); 153networkConstantBufferData .push_back ( 154paramBufferAddr + layerAllocations [i ].weightsGradOffset ); 155networkConstantBufferData .push_back (paramBufferAddr + layerAllocations [i ].biasOffset ); 156networkConstantBufferData .push_back ( 157paramBufferAddr + layerAllocations [i ].biasGradOffset ); 158 } 159auto networkConstantBuffer = createBuffer ( 160networkConstantBufferData .size ()* sizeof (uint64_t ), 161networkConstantBufferData .data ()); 162 163static const int inputCount = 32 ; 164 std::vector < float > inputBufferData ; 165for (int i = 0 ;i < inputCount ;i ++ ) 166 { 167inputBufferData .push_back ((float )rand () /RAND_MAX ); 168 } 169auto inputBuffer = createBuffer (inputCount * sizeof (float ),inputBufferData .data ()); 170 171// Create buffer for receiving current loss value. 172auto lossBuffer = createBuffer (sizeof (uint64_t )); 173 174auto queue = gDevice -> getQueue (rhi::QueueType ::Graphics ); 175 176for (int k = 0 ;k < 1000 ;k ++ ) 177 { 178clearBuffer (lossBuffer ); 179 180// Compute gradients. 181 { 182LearnGradParams entryPointParams = {}; 183entryPointParams .inputs = inputBuffer -> getDeviceAddress (); 184entryPointParams .count = inputCount /2 ; 185entryPointParams .lossBuffer = lossBuffer -> getDeviceAddress (); 186entryPointParams .networkBuffer = networkConstantBuffer -> getDeviceAddress (); 187dispatchKernel ( 188gLearnGradProgram , 189entryPointParams , 190 (entryPointParams .count + 255 ) /256 ); 191 } 192// Adjust parameters in row-major buffer (adam optimize). 193 { 194AdjustParamsParams entryPointParams = {}; 195entryPointParams .adamStates = adamStateBuffer -> getDeviceAddress (); 196entryPointParams .params = networkParamsBuffer -> getDeviceAddress (); 197entryPointParams .count = (paramBufferSize - gradientOffset ) /sizeof (NFloat ); 198entryPointParams .gradients = 199networkParamsBuffer -> getDeviceAddress ()+ gradientOffset ; 200dispatchKernel ( 201gAdjustParamProgram , 202entryPointParams , 203 (entryPointParams .count + 255 ) /256 ); 204 } 205if ((k + 1 ) %10 == 0 ) 206 { 207queue -> waitOnHost (); 208ComPtr < ISlangBlob > blob ; 209gDevice -> readBuffer (lossBuffer ,0 ,sizeof (float ),blob .writeRef ()); 210printf ("Loss after %d iterations: %f\n" ,k + 1 ,* (float * )blob -> getBufferPointer ()); 211 } 212 } 213return SLANG_OK ; 214 } 215 216// Allocate storage for network parameters, including weights, biases, and gradients. 217void allocateNetworkParameterStorage ( 218 std::vector < NetworkParameterAllocation >& paramStorage , 219size_t & outParamBufferSize , 220size_t & outGradientOffset ) 221 { 222outParamBufferSize = 0 ; 223 224auto allocRowMajorStorage = [& ](size_t size ) 225 { 226size = (size + 63 ) /64 * 64 ; 227size_t offset = outParamBufferSize ; 228outParamBufferSize += size ; 229return offset ; 230 }; 231 232for (int i = 0 ;i < kLayerCount ;i ++ ) 233 { 234size_t biasSize = getNetworkLayerBiasCount (i )* sizeof (NFloat ); 235NetworkParameterAllocation layer = {}; 236layer .weightsSize = getNetworkLayerWeightCount (i )* sizeof (NFloat ); 237layer .weightsOffset = allocRowMajorStorage (layer .weightsSize ); 238layer .biasSize = biasSize ; 239layer .biasOffset = allocRowMajorStorage (biasSize ); 240paramStorage .push_back (layer ); 241 } 242 243// Alloc storage for gradients. 244outGradientOffset = outParamBufferSize ; 245for (int i = 0 ;i < kLayerCount ;i ++ ) 246 { 247paramStorage [i ].weightsGradOffset = allocRowMajorStorage (paramStorage [i ].weightsSize ); 248paramStorage [i ].biasGradOffset = allocRowMajorStorage (paramStorage [i ].biasSize ); 249 } 250 } 251 252template < typename Args > 253void dispatchKernel (Kernel & kernel ,Args & args ,size_t numWorkGroups ) 254 { 255auto queue = gDevice -> getQueue (rhi::QueueType ::Graphics ); 256ComPtr < rhi::ICommandEncoder > encoder ; 257queue -> createCommandEncoder (encoder .writeRef ()); 258 { 259auto computeEncoder = encoder -> beginComputePass (); 260auto rootShaderObject = computeEncoder -> bindPipeline (kernel .pipeline .get ()); 261rootShaderObject -> getEntryPoint (0 )-> setData (rhi::ShaderOffset (),& args ,sizeof (args )); 262computeEncoder -> dispatchCompute (numWorkGroups ,1 ,1 ); 263computeEncoder -> end (); 264 } 265ComPtr < rhi::ICommandBuffer > commandBuffer ; 266encoder -> finish (commandBuffer .writeRef ()); 267queue -> submit (commandBuffer ); 268 } 269 270// Create a buffer with the specified size and optional initial data. 271ComPtr < rhi::IBuffer > createBuffer (size_t size ,void * initData = nullptr ) 272 { 273 rhi::BufferDesc bufferDesc = {}; 274bufferDesc .size = size ; 275bufferDesc .defaultState = rhi::ResourceState ::UnorderedAccess ; 276bufferDesc .usage = rhi::BufferUsage ::CopySource | rhi::BufferUsage ::CopyDestination | 277 rhi::BufferUsage ::UnorderedAccess ; 278bufferDesc .memoryType = rhi::MemoryType ::DeviceLocal ; 279return gDevice -> createBuffer (bufferDesc ,initData ); 280 } 281 282void clearBuffer (rhi::IBuffer * buffer ) 283 { 284auto queue = gDevice -> getQueue (rhi::QueueType ::Graphics ); 285auto encoder = queue -> createCommandEncoder (); 286encoder -> clearBuffer (buffer ); 287auto cmdBuffer = encoder -> finish (); 288queue -> submit (cmdBuffer ); 289 } 290 291Kernel loadComputeProgram (slang::IModule * slangModule ,char const * entryPointName ) 292 { 293ComPtr < slang::IEntryPoint > entryPoint ; 294slangModule -> findEntryPointByName (entryPointName ,entryPoint .writeRef ()); 295 296ComPtr < slang::IComponentType > linkedProgram ; 297entryPoint -> link (linkedProgram .writeRef ()); 298 299if (isTestMode ()) 300 { 301printEntrypointHashes (1 ,1 ,linkedProgram ); 302 } 303 304Kernel result ; 305 306 rhi::ComputePipelineDesc desc ; 307auto program = gDevice -> createShaderProgram (linkedProgram ); 308desc .program = program .get (); 309result .program = program ; 310result .pipeline = gDevice -> createComputePipeline (desc ); 311return result ; 312 } 313 314inline unsigned short floatToHalf (float val ) 315 { 316uint32_t x = 0 ; 317memcpy (& x ,& val ,sizeof (float )); 318 319unsigned short bits = (x >>16 )& 0x8000 ; 320unsigned short m = (x >>12 )& 0x07ff ; 321unsigned int e = (x >>23 )& 0xff ; 322if (e < 103 ) 323return bits ; 324if (e > 142 ) 325 { 326bits |=0x7c00u ; 327bits |=e == 255 && (x & 0x007fffffu ); 328return bits ; 329 } 330if (e < 113 ) 331 { 332m |=0x0800u ; 333bits |= (m >> (114 - e ))+ ((m >> (113 - e ))& 1 ); 334return bits ; 335 } 336bits |= ((e - 112 ) <<10 ) | (m >>1 ); 337bits += m & 1 ; 338return bits ; 339 } 340 341ComPtr < slang::ISession > createSlangSession (rhi::IDevice * device ) 342 { 343ComPtr < slang::ISession > slangSession = device -> getSlangSession (); 344return slangSession ; 345 } 346 347ComPtr < slang::IModule > compileShaderModuleFromFile ( 348 slang::ISession * slangSession , 349char const * filePath ) 350 { 351ComPtr < slang::IModule > slangModule ; 352ComPtr < slang::IBlob > diagnosticBlob ; 353Slang ::String path = resourceBase .resolveResource (filePath ); 354slangModule = slangSession -> loadModule (path .getBuffer (),diagnosticBlob .writeRef ()); 355diagnoseIfNeeded (diagnosticBlob ); 356 357return slangModule ; 358 } 359 360SlangResult loadShaderKernels () 361 { 362Slang ::String path = resourceBase .resolveResource ("kernels.slang" ); 363 364gSlangSession = createSlangSession (gDevice ); 365gSlangModule = compileShaderModuleFromFile (gSlangSession ,path .getBuffer ()); 366if (!gSlangModule ) 367return SLANG_FAIL ; 368 369gLearnGradProgram = loadComputeProgram (gSlangModule ,"learnGradient" ); 370if (!gLearnGradProgram ) 371return SLANG_FAIL ; 372 373gAdjustParamProgram = loadComputeProgram (gSlangModule ,"adjustParameters" ); 374if (!gAdjustParamProgram ) 375return SLANG_FAIL ; 376 377return SLANG_OK ; 378 } 379}; 380 381int exampleMain (int argc ,char ** argv ) 382{ 383ExampleProgram app ; 384if (SLANG_FAILED (app .execute (argc ,argv ))) 385 { 386return -1 ; 387 } 388return 0 ; 389}