From 78517dc392f0d2ebba25f0ac3f4d4e004b0f0ab0 Mon Sep 17 00:00:00 2001 From: Sai Praveen Bangaru <31557731+saipraveenb25@users.noreply.github.com> Date: Fri, 14 Mar 2025 17:15:36 -0700 Subject: Fix lowering of associated types in generic interfaces (#6600) * Fix lowering of associated types in generic interfaces. * Update diff-assoctype-generic-interface.slang * Fix-up lowering of differentiable witnesses for implicit ops * Update slang-ir-autodiff-transcriber-base.cpp * Fix issue with differentiating type-packs --- source/slang/slang-lower-to-ir.cpp | 41 ++++++++++++++++++++++++++++---------- 1 file changed, 30 insertions(+), 11 deletions(-) (limited to 'source/slang/slang-lower-to-ir.cpp') diff --git a/source/slang/slang-lower-to-ir.cpp b/source/slang/slang-lower-to-ir.cpp index 775986a9a..decfe4a91 100644 --- a/source/slang/slang-lower-to-ir.cpp +++ b/source/slang/slang-lower-to-ir.cpp @@ -1919,6 +1919,28 @@ struct ValLoweringVisitor : ValVisitorgetIntValue(type, val->getValue())); } + IRType* visitDifferentialPairType(DifferentialPairType* pairType) + { + IRType* primalType = lowerType(context, pairType->getPrimalType()); + if (as(primalType) || as(primalType)) + { + List operands; + SubstitutionSet(pairType->getDeclRef()) + .forEachSubstitutionArg( + [&](Val* arg) + { + auto argVal = lowerVal(context, arg).val; + SLANG_ASSERT(argVal); + operands.add(argVal); + }); + + auto undefined = getBuilder()->emitUndefined(operands[1]->getFullType()); + return getBuilder()->getDifferentialPairUserCodeType(primalType, undefined); + } + else + return lowerSimpleIntrinsicType(pairType); + } + IRFuncType* visitFuncType(FuncType* type) { IRType* resultType = lowerType(context, type->getResultType()); @@ -10195,15 +10217,17 @@ struct DeclLoweringVisitor : DeclVisitor // If our function is differentiable, register a callback so the derivative // annotations for types can be lowered. // - if (auto diffAttr = decl->findModifier()) + if (decl->findModifier() && !isInterfaceRequirement(decl)) { + auto diffAttr = decl->findModifier(); + auto diffTypeWitnessMap = diffAttr->getMapTypeToIDifferentiableWitness(); - OrderedDictionary resolveddiffTypeWitnessMap; + OrderedDictionary resolveddiffTypeWitnessMap; // Go through each entry in the map and resolve the key. for (auto& entry : diffTypeWitnessMap) { - auto resolvedKey = as(entry.key->resolve()); + auto resolvedKey = as(entry.key->resolve()); resolveddiffTypeWitnessMap[resolvedKey] = as(as(entry.value)->resolve()); } @@ -10211,14 +10235,9 @@ struct DeclLoweringVisitor : DeclVisitor subContext->registerTypeCallback( [=](IRGenContext* context, Type* type, IRType* irType) { - if (!as(type)) - return irType; - - DeclRefBase* declRefBase = as(type)->getDeclRefBase(); - if (resolveddiffTypeWitnessMap.containsKey(declRefBase)) + if (resolveddiffTypeWitnessMap.containsKey(type)) { - auto irWitness = - lowerVal(subContext, resolveddiffTypeWitnessMap[declRefBase]).val; + auto irWitness = lowerVal(subContext, resolveddiffTypeWitnessMap[type]).val; if (irWitness) { IRInst* args[] = {irType, irWitness}; @@ -11328,7 +11347,7 @@ LoweredValInfo emitDeclRef(IRGenContext* context, Decl* decl, DeclRefBase* subst // interface definitions. return emitDeclRef( context, - createDefaultSpecializedDeclRef(context, nullptr, decl), + decl->getDefaultDeclRef(), context->irBuilder->getTypeKind()); } -- cgit v1.2.3