From 33fb95980b0120cdd4d4f2d51f5f116e808dd4aa Mon Sep 17 00:00:00 2001 From: Yong He Date: Fri, 6 Jan 2023 13:39:06 -0800 Subject: Split bwd_diff op into separate ops for primal and propagate func. (#2582) * Split bwd_diff op into separate ops for primal and propagate func. * Fix. * Download swiftshader with github actions instead of curl on linux. * Fix github action. Co-authored-by: Yong He --- source/slang/slang-ir-autodiff-transcriber-base.cpp | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) (limited to 'source/slang/slang-ir-autodiff-transcriber-base.cpp') diff --git a/source/slang/slang-ir-autodiff-transcriber-base.cpp b/source/slang/slang-ir-autodiff-transcriber-base.cpp index 69cef941c..4aab0f835 100644 --- a/source/slang/slang-ir-autodiff-transcriber-base.cpp +++ b/source/slang/slang-ir-autodiff-transcriber-base.cpp @@ -259,7 +259,7 @@ IRType* AutoDiffTranscriberBase::_differentiateTypeImpl(IRBuilder* builder, IRTy } case kIROp_FuncType: - return differentiateFunctionType(builder, as(primalType)); + return differentiateFunctionType(builder, nullptr, as(primalType)); case kIROp_OutType: if (auto diffValueType = differentiateType(builder, as(primalType)->getValueType())) @@ -436,7 +436,7 @@ InstPair AutoDiffTranscriberBase::transcribeParam(IRBuilder* builder, IRParam* o { auto primalDataType = findOrTranscribePrimalInst(builder, origParam->getDataType()); // Do not differentiate generic type (and witness table) parameters - if (as(primalDataType) || as(primalDataType)) + if (isGenericParam(origParam)) { return InstPair( cloneInst(&cloneEnv, builder, origParam), -- cgit v1.2.3