diff options
| author | Sai Praveen Bangaru <31557731+saipraveenb25@users.noreply.github.com> | 2022-11-29 20:01:41 -0500 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2022-11-29 17:01:41 -0800 |
| commit | f5581786a1891cedb165adb1afe71fe34f26e030 (patch) | |
| tree | 86da2f1acbaec920ac0c38349897b293b405c021 /source/slang/slang-ir-autodiff-rev.cpp | |
| parent | af7f40063dfed1c651d33b93956c7623a7d2c050 (diff) | |
Refactored reverse-mode implementation to use 4 separate passes. (#2531)
* Added partial implementation for reverse-mode
* Fixing several compile and runtime errors.
* Fixed several issues with reverse-mode passes.
* Fixed more issues. Basic reverse-mode tests passing
Co-authored-by: Edward Liu <shiqiu1105@gmail.com>
Diffstat (limited to 'source/slang/slang-ir-autodiff-rev.cpp')
| -rw-r--r-- | source/slang/slang-ir-autodiff-rev.cpp | 196 |
1 files changed, 185 insertions, 11 deletions
diff --git a/source/slang/slang-ir-autodiff-rev.cpp b/source/slang/slang-ir-autodiff-rev.cpp index 52567e887..522c995b0 100644 --- a/source/slang/slang-ir-autodiff-rev.cpp +++ b/source/slang/slang-ir-autodiff-rev.cpp @@ -6,6 +6,11 @@ #include "slang-ir-util.h" #include "slang-ir-inst-pass-base.h" +#include "slang-ir-autodiff-fwd.h" +#include "slang-ir-autodiff-propagate.h" +#include "slang-ir-autodiff-unzip.h" +#include "slang-ir-autodiff-transpose.h" + namespace Slang { @@ -44,12 +49,33 @@ struct BackwardDiffTranscriber IRWitnessTable* differentialBottomWitness = nullptr; Dictionary<InstPair, IRInst*> differentialPairTypes; + // References to other passes that for reverse-mode transcription. + ForwardDerivativeTranscriber *fwdDiffTranscriber; + DiffTransposePass *diffTransposePass; + DiffPropagationPass *diffPropagationPass; + DiffUnzipPass *diffUnzipPass; + + // Allocate space for the passes. + ForwardDerivativeTranscriber fwdDiffTranscriberStorage; + DiffTransposePass diffTransposePassStorage; + DiffPropagationPass diffPropagationPassStorage; + DiffUnzipPass diffUnzipPassStorage; + + BackwardDiffTranscriber(AutoDiffSharedContext* shared, SharedIRBuilder* inSharedBuilder, DiagnosticSink* inSink) : autoDiffSharedContext(shared) , sink(inSink) , differentiableTypeConformanceContext(shared) , sharedBuilder(inSharedBuilder) - {} + , fwdDiffTranscriberStorage(shared, inSharedBuilder) + , diffTransposePassStorage(shared) + , diffPropagationPassStorage(shared) + , diffUnzipPassStorage(shared) + , fwdDiffTranscriber(&fwdDiffTranscriberStorage) + , diffTransposePass(&diffTransposePassStorage) + , diffPropagationPass(&diffPropagationPassStorage) + , diffUnzipPass(&diffUnzipPassStorage) + { } DiagnosticSink* getSink() { @@ -413,20 +439,166 @@ struct BackwardDiffTranscriber return result; } - // Transcribe a function definition. - InstPair transcribeFunc(IRBuilder* inBuilder, IRFunc* primalFunc, IRFunc* diffFunc) + // Puts parameters into their own block. + void makeParameterBlock(IRBuilder* inBuilder, IRFunc* func) { IRBuilder builder(inBuilder->getSharedBuilder()); - builder.setInsertInto(diffFunc); - differentiableTypeConformanceContext.setFunc(primalFunc); - // Transcribe children from origFunc into diffFunc - for (auto block = primalFunc->getFirstBlock(); block; block = block->getNextBlock()) - this->transcribeBlock(&builder, block); + auto firstBlock = func->getFirstBlock(); + IRParam* param = func->getFirstParam(); + + builder.setInsertBefore(firstBlock); + + // Note: It looks like emitBlock() doesn't use the current + // builder position, so we're going to manually move the new block + // to before the existing block. + auto paramBlock = builder.emitBlock(); + paramBlock->insertBefore(firstBlock); + builder.setInsertInto(paramBlock); + + while(param) + { + IRParam* nextParam = param->getNextParam(); + + // Copy inst into the new parameter block. + auto clonedParam = cloneInst(&cloneEnv, &builder, param); + param->replaceUsesWith(clonedParam); + param->removeAndDeallocate(); + + param = nextParam; + } + + // Replace this block as the first block. + firstBlock->replaceUsesWith(paramBlock); + + // Add terminator inst. + builder.emitBranch(firstBlock); + } + + // Transcribe a function definition. + InstPair transcribeFunc(IRBuilder* builder, IRFunc* primalFunc, IRFunc* diffFunc) + { + SLANG_ASSERT(primalFunc); + SLANG_ASSERT(diffFunc); + // Reverse-mode transcription uses 4 separate steps: + // TODO(sai): Fill in documentation. + + // Generate a temporary forward derivative function as an intermediate step. + IRFunc* fwdDiffFunc = as<IRFunc>(fwdDiffTranscriber->transcribeFuncHeader(builder, (IRFunc*)primalFunc).differential); + SLANG_ASSERT(fwdDiffFunc); + + // Transcribe the body of the primal function into it's linear (fwd-diff) form. + // TODO(sai): Handle the case when we already have a user-defined fwd-derivative function. + fwdDiffTranscriber->transcribeFunc(builder, primalFunc, as<IRFunc>(fwdDiffFunc)); + + // Split first block into a paramter block. + this->makeParameterBlock(builder, as<IRFunc>(fwdDiffFunc)); + + // This steps adds a decoration to instructions that are computing the differential. + diffPropagationPass->propagateDiffInstDecoration(builder, fwdDiffFunc); + + // Copy primal insts to the first block of the unzipped function, copy diff insts to the + // second block of the unzipped function. + // + IRFunc* unzippedFwdDiffFunc = diffUnzipPass->unzipDiffInsts(fwdDiffFunc); + + // Clone the primal blocks from unzippedFwdDiffFunc + // to the reverse-mode function. + // TODO: This is the spot where we can make a decision to split + // the primal and differential into two different funcitons + // instead of two blocks in the same function. + // + // Special care needs to be taken for the first block since it holds the parameters + + // Clone all blocks into a temporary diff func. + // We're using a temporary sice we don't want to clone decorations, + // only blocks, and right now there's no provision in slang-ir-clone.h + // for that. + // + builder->setInsertInto(diffFunc->getParent()); + auto tempDiffFunc = as<IRFunc>(cloneInst(&cloneEnv, builder, unzippedFwdDiffFunc)); + + // Move blocks to the diffFunc shell. + { + List<IRBlock*> workList; + for (auto block = tempDiffFunc->getFirstBlock(); block; block = block->getNextBlock()) + workList.add(block); + + for (auto block : workList) + block->insertAtEnd(diffFunc); + } + + // Transpose the first block (parameter block) + transcribeParameterBlock(builder, diffFunc); + + builder->setInsertInto(diffFunc); + + auto dOutParameter = diffFunc->getLastParam(); + + // Transpose differential blocks from unzippedFwdDiffFunc into diffFunc (with dOutParameter) representing the + DiffTransposePass::FuncTranspositionInfo info = {dOutParameter, nullptr}; + diffTransposePass->transposeDiffBlocksInFunc(diffFunc, info); + + // Clean up by deallocating intermediate steps. + tempDiffFunc->removeAndDeallocate(); + unzippedFwdDiffFunc->removeAndDeallocate(); + fwdDiffFunc->removeAndDeallocate(); return InstPair(primalFunc, diffFunc); } + void transcribeParameterBlock(IRBuilder* builder, IRFunc* diffFunc) + { + IRBlock* fwdDiffParameterBlock = diffFunc->getFirstBlock(); + + // Find the 'next' block using the terminator inst of the parameter block. + auto fwdParamBlockBranch = as<IRUnconditionalBranch>(fwdDiffParameterBlock->getTerminator()); + auto nextBlock = fwdParamBlockBranch->getTargetBlock(); + + builder->setInsertInto(fwdDiffParameterBlock); + + // 1. Turn fwd-diff versions of the parameters into reverse-diff versions by wrapping them as InOutType<> + for (auto child = fwdDiffParameterBlock->getFirstParam(); child;) + { + IRParam* nextChild = child->getNextParam(); + + auto fwdParam = as<IRParam>(child); + SLANG_ASSERT(fwdParam); + + // TODO: Handle ptr<pair> types. + if (auto diffPairType = as<IRDifferentialPairType>(fwdParam->getDataType())) + { + // Create inout version. + auto inoutDiffPairType = builder->getInOutType(diffPairType); + auto newParam = builder->emitParam(inoutDiffPairType); + + // Map the _load_ of the new parameter as the clone of the old one. + auto newParamLoad = builder->emitLoad(newParam); + newParamLoad->insertAtStart(nextBlock); // Move to first block _after_ the parameter block. + fwdParam->replaceUsesWith(newParamLoad); + fwdParam->removeAndDeallocate(); + } + else + { + // Default case (parameter has nothing to do with differentiation) + // Do nothing. + } + + child = nextChild; + } + + auto paramCount = as<IRFuncType>(diffFunc->getDataType())->getParamCount(); + + // 2. Add a parameter for 'derivative of the output' (d_out). + // The type is the last parameter type of the function. + // + auto dOutParamType = as<IRFuncType>(diffFunc->getDataType())->getParamType(paramCount - 1); + + SLANG_ASSERT(dOutParamType); + + builder->emitParam(dOutParamType); + } + IRInst* copyParam(IRBuilder* builder, IRParam* origParam) { auto primalDataType = origParam->getDataType(); @@ -652,7 +824,6 @@ struct BackwardDiffTranscriber struct ReverseDerivativePass : public InstPassBase { - DiagnosticSink* getSink() { return sink; @@ -664,7 +835,7 @@ struct ReverseDerivativePass : public InstPassBase IRBuilder builderStorage(autodiffContext->sharedBuilder); IRBuilder* builder = &builderStorage; - // Process all ForwardDifferentiate instructions (kIROp_ForwardDifferentiate), by + // Process all ForwardDifferentiate instructions (kIROp_ForwardDifferentiate), by // generating derivative code for the referenced function. // bool modified = processReferencedFunctions(builder); @@ -800,6 +971,9 @@ struct ReverseDerivativePass : public InstPassBase pairBuilderStorage(context) { backwardTranscriberStorage.pairBuilder = &pairBuilderStorage; + backwardTranscriberStorage.fwdDiffTranscriberStorage.sink = sink; + backwardTranscriberStorage.fwdDiffTranscriberStorage.autoDiffSharedContext = context; + backwardTranscriberStorage.fwdDiffTranscriberStorage.pairBuilder = &(pairBuilderStorage); } protected: @@ -829,4 +1003,4 @@ bool processReverseDerivativeCalls( return changed; } -}
\ No newline at end of file +} |
