summaryrefslogtreecommitdiff
path: root/source/slang/slang-ir-autodiff-transcriber-base.cpp
diff options
context:
space:
mode:
authorYong He <yonghe@outlook.com>2023-01-11 15:33:28 -0800
committerGitHub <noreply@github.com>2023-01-11 15:33:28 -0800
commita3ac6e71cbc922b7c941c45f23ee18a9fc274d1f (patch)
treeacf8c18601f124e9290494f8b379d2420369fc35 /source/slang/slang-ir-autodiff-transcriber-base.cpp
parent20262684bcbb707d16669b2670039df870b65ca8 (diff)
Make backward differentiation work with generics. (#2586)
* Make backward differentiation work with generics. * Fix. * Another fix. * More fix. Co-authored-by: Yong He <yhe@nvidia.com>
Diffstat (limited to 'source/slang/slang-ir-autodiff-transcriber-base.cpp')
-rw-r--r--source/slang/slang-ir-autodiff-transcriber-base.cpp62
1 files changed, 43 insertions, 19 deletions
diff --git a/source/slang/slang-ir-autodiff-transcriber-base.cpp b/source/slang/slang-ir-autodiff-transcriber-base.cpp
index c0404e036..deb1b2da9 100644
--- a/source/slang/slang-ir-autodiff-transcriber-base.cpp
+++ b/source/slang/slang-ir-autodiff-transcriber-base.cpp
@@ -75,36 +75,38 @@ bool AutoDiffTranscriberBase::hasDifferentialInst(IRInst* origInst)
return instMapD.ContainsKey(origInst);
}
-bool AutoDiffTranscriberBase::shouldUseOriginalAsPrimal(IRInst* origInst)
+bool AutoDiffTranscriberBase::shouldUseOriginalAsPrimal(IRInst* currentParent, IRInst* origInst)
{
if (as<IRGlobalValueWithCode>(origInst))
return true;
if (origInst->parent && origInst->parent->getOp() == kIROp_Module)
return true;
+ if (isChildInstOf(currentParent, origInst->getParent()))
+ return true;
return false;
}
-IRInst* AutoDiffTranscriberBase::lookupPrimalInst(IRInst* origInst)
+IRInst* AutoDiffTranscriberBase::lookupPrimalInstImpl(IRInst* currentParent, IRInst* origInst)
{
if (!origInst)
return nullptr;
- if (shouldUseOriginalAsPrimal(origInst))
+ if (shouldUseOriginalAsPrimal(currentParent, origInst))
return origInst;
return cloneEnv.mapOldValToNew[origInst];
}
-IRInst* AutoDiffTranscriberBase::lookupPrimalInst(IRInst* origInst, IRInst* defaultInst)
+IRInst* AutoDiffTranscriberBase::lookupPrimalInst(IRInst* currentParent, IRInst* origInst, IRInst* defaultInst)
{
if (!origInst)
return nullptr;
- return (hasPrimalInst(origInst)) ? lookupPrimalInst(origInst) : defaultInst;
+ return (hasPrimalInst(currentParent, origInst)) ? lookupPrimalInstImpl(currentParent, origInst) : defaultInst;
}
-bool AutoDiffTranscriberBase::hasPrimalInst(IRInst* origInst)
+bool AutoDiffTranscriberBase::hasPrimalInst(IRInst* currentParent, IRInst* origInst)
{
if (!origInst)
return false;
- if (shouldUseOriginalAsPrimal(origInst))
+ if (shouldUseOriginalAsPrimal(currentParent, origInst))
return true;
return cloneEnv.mapOldValToNew.ContainsKey(origInst);
}
@@ -124,26 +126,48 @@ IRInst* AutoDiffTranscriberBase::findOrTranscribePrimalInst(IRBuilder* builder,
{
if (!origInst)
return origInst;
+
+ auto currentParent = builder->getInsertLoc().getParent();
- if (shouldUseOriginalAsPrimal(origInst))
+ if (shouldUseOriginalAsPrimal(currentParent, origInst))
return origInst;
- if (!hasPrimalInst(origInst))
+ if (!hasPrimalInst(currentParent, origInst))
{
transcribe(builder, origInst);
- SLANG_ASSERT(hasPrimalInst(origInst));
+ SLANG_ASSERT(hasPrimalInst(currentParent, origInst));
}
- return lookupPrimalInst(origInst);
+ return lookupPrimalInstImpl(currentParent, origInst);
}
IRInst* AutoDiffTranscriberBase::maybeCloneForPrimalInst(IRBuilder* builder, IRInst* inst)
{
- IRInst* primal = lookupPrimalInst(inst, inst);
-
- if (primal == inst &&
- !isChildInstOf(builder->getInsertLoc().getParent(), inst->getParent()))
- primal = cloneInst(&cloneEnv, builder, inst);
+ IRInst* primal = lookupPrimalInst(builder, inst, nullptr);
+ if (!primal)
+ {
+ IRInst* type = inst->getFullType();
+ if (type)
+ {
+ type = maybeCloneForPrimalInst(builder, type);
+ }
+ List<IRInst*> operands;
+ for (UInt i = 0; i < inst->getOperandCount(); i++)
+ {
+ auto operand = maybeCloneForPrimalInst(builder, inst->getOperand(i));
+ operands.add(operand);
+ }
+ auto cloneResult = builder->emitIntrinsicInst(
+ (IRType*)type, inst->getOp(), operands.getCount(), operands.getBuffer());
+ IRBuilder subBuilder = *builder;
+ subBuilder.setInsertInto(cloneResult);
+ for (auto child : inst->getDecorationsAndChildren())
+ {
+ maybeCloneForPrimalInst(&subBuilder, child);
+ }
+ cloneEnv.mapOldValToNew[inst] = cloneResult;
+ return cloneResult;
+ }
return primal;
}
@@ -223,7 +247,7 @@ IRType* AutoDiffTranscriberBase::_differentiateTypeImpl(IRBuilder* builder, IRTy
// If there is an explicit primal version of this type in the local scope, load that
// otherwise use the original type.
//
- IRInst* primalType = lookupPrimalInst(origType, origType);
+ IRInst* primalType = lookupPrimalInst(builder, origType, origType);
// Special case certain compound types (PtrType, FuncType, etc..)
// otherwise try to lookup a differential definition for the given type.
@@ -390,7 +414,7 @@ IRType* AutoDiffTranscriberBase::differentiateExtractExistentialType(IRBuilder*
if (lookupKeyPath.getCount())
{
// `interfaceType` does conform to `IDifferentiable`.
- outWitnessTable = builder->emitExtractExistentialWitnessTable(lookupPrimalInstIfExists(origType->getOperand(0)));
+ outWitnessTable = builder->emitExtractExistentialWitnessTable(lookupPrimalInstIfExists(builder, origType->getOperand(0)));
for (auto node : lookupKeyPath)
{
outWitnessTable = builder->emitLookupInterfaceMethodInst((IRType*)node->getRequirementVal(), outWitnessTable, node->getRequirementKey());
@@ -731,7 +755,7 @@ IRInst* AutoDiffTranscriberBase::transcribe(IRBuilder* builder, IRInst* origInst
//
if (auto diffInst = lookupDiffInst(origInst, nullptr))
{
- SLANG_ASSERT(lookupPrimalInst(origInst)); // Consistency check.
+ SLANG_ASSERT(lookupPrimalInst(builder, origInst)); // Consistency check.
return diffInst;
}