From 85c1569308793cc2408088e539a3ed1da5f9d235 Mon Sep 17 00:00:00 2001 From: Yong He Date: Fri, 24 Feb 2023 14:33:32 -0800 Subject: Support dynamic dispatch a backward differentiable function. (#2678) Co-authored-by: Yong He --- source/slang/slang-check-decl.cpp | 69 +-------------------------------------- 1 file changed, 1 insertion(+), 68 deletions(-) (limited to 'source/slang/slang-check-decl.cpp') diff --git a/source/slang/slang-check-decl.cpp b/source/slang/slang-check-decl.cpp index 142842e12..a1d5acfb0 100644 --- a/source/slang/slang-check-decl.cpp +++ b/source/slang/slang-check-decl.cpp @@ -2677,24 +2677,6 @@ namespace Slang val->func = satisfyingMemberDeclRef; witnessTable->add(bwdReq, RequirementWitness(val)); } - else if (auto primalReq = as(reqRefDecl->referencedDecl)) - { - DifferentiateVal* val = m_astBuilder->create(); - val->func = satisfyingMemberDeclRef; - witnessTable->add(primalReq, RequirementWitness(val)); - } - else if (auto propReq = as(reqRefDecl->referencedDecl)) - { - DifferentiateVal* val = m_astBuilder->create(); - val->func = satisfyingMemberDeclRef; - witnessTable->add(propReq, RequirementWitness(val)); - } - else if (auto itypeReq = as(reqRefDecl->referencedDecl)) - { - DifferentiateVal* val = m_astBuilder->create(); - val->func = satisfyingMemberDeclRef; - witnessTable->add(itypeReq, RequirementWitness(val)); - } } witnessTable->add(requiredMemberDeclRef, RequirementWitness(satisfyingMemberDeclRef)); } @@ -5920,7 +5902,7 @@ namespace Slang if (auto interfaceDecl = findParentInterfaceDecl(decl)) { bool isDiffFunc = false; - if (decl->hasModifier()) + if (decl->hasModifier() || decl->hasModifier()) { auto reqDecl = m_astBuilder->create(); cloneModifiers(reqDecl, decl); @@ -5954,55 +5936,6 @@ namespace Slang reqRef->parentDecl = decl; decl->members.add(reqRef); } - // Requirement for backward derivative intermediate type. - auto intermediateTypeReqDecl = m_astBuilder->create(); - auto intermediateType = m_astBuilder->getOrCreateDeclRefType( - intermediateTypeReqDecl, createDefaultSubstitutions(m_astBuilder, this, decl)); - { - cloneModifiers(intermediateTypeReqDecl, decl); - interfaceDecl->members.add(intermediateTypeReqDecl); - intermediateTypeReqDecl->parentDecl = interfaceDecl; - - auto reqRef = m_astBuilder->create(); - reqRef->referencedDecl = intermediateTypeReqDecl; - reqRef->parentDecl = decl; - decl->members.add(reqRef); - } - // Requirement for backward derivative primal func. - { - auto reqDecl = m_astBuilder->create(); - cloneModifiers(reqDecl, decl); - FuncType* primalFuncType = m_astBuilder->create(); - primalFuncType->resultType = originalFuncType->resultType; - primalFuncType->paramTypes.addRange(originalFuncType->paramTypes); - auto outType = m_astBuilder->getOutType(intermediateType); - primalFuncType->paramTypes.add(outType); - setFuncTypeIntoRequirementDecl(reqDecl, primalFuncType); - interfaceDecl->members.add(reqDecl); - reqDecl->parentDecl = interfaceDecl; - - auto reqRef = m_astBuilder->create(); - reqRef->referencedDecl = reqDecl; - reqRef->parentDecl = decl; - decl->members.add(reqRef); - } - // Requirement for backward derivative propagate func. - { - auto reqDecl = m_astBuilder->create(); - cloneModifiers(reqDecl, decl); - interfaceDecl->members.add(reqDecl); - reqDecl->parentDecl = interfaceDecl; - FuncType* propagateFuncType = m_astBuilder->create(); - propagateFuncType->resultType = diffFuncType->resultType; - propagateFuncType->paramTypes.addRange(diffFuncType->paramTypes); - propagateFuncType->paramTypes.add(intermediateType); - setFuncTypeIntoRequirementDecl(reqDecl, propagateFuncType); - auto reqRef = m_astBuilder->create(); - reqRef->referencedDecl = reqDecl; - reqRef->parentDecl = decl; - decl->members.add(reqRef); - } - isDiffFunc = true; } if (isDiffFunc) -- cgit v1.2.3