#include "slang-compile.h" #include #include #include "Containers/Array.h" #include #include #define CHECK_RESULT(x) {SlangResult r = x; if(r != 0) {throw std::runtime_error(fmt::format("Error: {0}", r));}} #define CHECK_DIAGNOSTICS() {if(diagnostics) {std::cout << (const char*)diagnostics->getBufferPointer() << std::endl; assert(false);}} Slang::ComPtr Seele::generateShader(const ShaderCreateInfo& createInfo, SlangCompileTarget target, Map& paramMapping) { thread_local Slang::ComPtr globalSession; if(!globalSession) { slang::createGlobalSession(globalSession.writeRef()); } slang::SessionDesc sessionDesc; sessionDesc.flags = 0; sessionDesc.defaultMatrixLayoutMode = SLANG_MATRIX_LAYOUT_COLUMN_MAJOR; Array macros; for(const auto& [key, val] : createInfo.defines) { macros.add(slang::PreprocessorMacroDesc{ .name = key, .value = val, }); } sessionDesc.preprocessorMacroCount = macros.size(); sessionDesc.preprocessorMacros = macros.data(); slang::TargetDesc targetDesc; targetDesc.profile = globalSession->findProfile("sm_6_6"); targetDesc.format = target; sessionDesc.targetCount = 1; sessionDesc.targets = &targetDesc; StaticArray searchPaths = {"shaders/", "shaders/lib/", "shaders/generated/"}; sessionDesc.searchPaths = searchPaths.data(); sessionDesc.searchPathCount = searchPaths.size(); Slang::ComPtr session; CHECK_RESULT(globalSession->createSession(sessionDesc, session.writeRef())); Slang::ComPtr diagnostics; Array modules; Slang::ComPtr entrypoint; slang::IModule* mainModule = nullptr; for (const auto& moduleName : createInfo.additionalModules) { modules.add(session->loadModule(moduleName.c_str(), diagnostics.writeRef())); if (moduleName == createInfo.mainModule) { mainModule = (slang::IModule*)modules.back(); } CHECK_DIAGNOSTICS(); } CHECK_DIAGNOSTICS(); mainModule->findEntryPointByName(createInfo.entryPoint.c_str(), entrypoint.writeRef()); modules.add(entrypoint); Slang::ComPtr moduleComposition; session->createCompositeComponentType(modules.data(), modules.size(), moduleComposition.writeRef(), diagnostics.writeRef()); CHECK_DIAGNOSTICS(); Slang::ComPtr linkedProgram; moduleComposition->link(linkedProgram.writeRef(), diagnostics.writeRef()); CHECK_DIAGNOSTICS(); slang::ProgramLayout* reflection = linkedProgram->getLayout(0, diagnostics.writeRef()); CHECK_DIAGNOSTICS(); Array specialization; for(const auto& [key, value] : createInfo.typeParameter) { specialization.add(slang::SpecializationArg::fromType(reflection->findTypeByName(value))); } Slang::ComPtr specializedComponent; linkedProgram->specialize(specialization.data(), specialization.size(), specializedComponent.writeRef(), diagnostics.writeRef()); CHECK_DIAGNOSTICS(); Slang::ComPtr kernelBlob; specializedComponent->getEntryPointCode( 0, 0, kernelBlob.writeRef(), diagnostics.writeRef() ); CHECK_DIAGNOSTICS(); slang::ProgramLayout* signature = specializedComponent->getLayout(0, diagnostics.writeRef()); CHECK_DIAGNOSTICS(); auto entry = signature->findEntryPointByName(createInfo.entryPoint.c_str()); uint32 offset = 0; if(target == SLANG_DXIL) { offset = 1;// idk why } for(size_t i = 0; i < signature->getParameterCount(); ++i) { paramMapping[param->getName()] = offset++; } return kernelBlob; }