@@ -331,13 +331,41 @@ FlatSymbolRefAttr MEnzymeLogic::CreateSplitModeDiff(
331331
332332 SymbolTable symbolTable (SymbolTable::getNearestSymbolTable (fn));
333333
334+ SmallVector<mlir::Attribute> argAttrs;
335+ if (auto prevArgAttrs = fn.getAllArgAttrs ())
336+ argAttrs.assign (prevArgAttrs.begin (), prevArgAttrs.end ());
337+
338+ SmallVector<Attribute> argActivityAttrs;
339+ for (auto [i, act] : llvm::enumerate (constants)) {
340+ argActivityAttrs.push_back (activityFromDiffeType (fn.getContext (), act));
341+
342+ if (!argAttrs.empty () && act == DIFFE_TYPE ::DUP_ARG )
343+ argAttrs.insert (argAttrs.begin () + i + 1 -
344+ (argAttrs.size () - fn.getNumArguments ()),
345+ nullptr );
346+ }
347+
348+ SmallVector<Attribute> retActivityAttrs;
349+ for (auto act : retType)
350+ retActivityAttrs.push_back (activityFromDiffeType (fn.getContext (), act));
351+
334352 if (auto existingCustomRule =
335353 fn->getAttrOfType <FlatSymbolRefAttr>(" enzyme.custom_rule" )) {
336354 auto CR = symbolTable.lookup <enzyme::CustomReverseRuleOp>(
337355 existingCustomRule.getValue ());
338356
339357 if (CR ) {
340- return existingCustomRule;
358+ auto getAttrActivity = [](auto attr) {
359+ return cast<ActivityAttr>(attr).getValue ();
360+ };
361+
362+ SmallVector<Activity> ArgActivity =
363+ llvm::map_to_vector (argActivityAttrs, getAttrActivity);
364+ SmallVector<Activity> RetActivity =
365+ llvm::map_to_vector (retActivityAttrs, getAttrActivity);
366+
367+ if (!failed (CR .activityMatch (ArgActivity, RetActivity)))
368+ return existingCustomRule;
341369 }
342370 }
343371
@@ -360,24 +388,6 @@ FlatSymbolRefAttr MEnzymeLogic::CreateSplitModeDiff(
360388
361389 auto name = fn.getName ();
362390
363- SmallVector<mlir::Attribute> argAttrs;
364- if (auto prevArgAttrs = fn.getAllArgAttrs ())
365- argAttrs.assign (prevArgAttrs.begin (), prevArgAttrs.end ());
366-
367- SmallVector<Attribute> argActivityAttrs;
368- for (auto [i, act] : llvm::enumerate (constants)) {
369- argActivityAttrs.push_back (activityFromDiffeType (fn.getContext (), act));
370-
371- if (!argAttrs.empty () && act == DIFFE_TYPE ::DUP_ARG )
372- argAttrs.insert (argAttrs.begin () + i + 1 -
373- (argAttrs.size () - fn.getNumArguments ()),
374- nullptr );
375- }
376-
377- SmallVector<Attribute> retActivityAttrs;
378- for (auto act : retType)
379- retActivityAttrs.push_back (activityFromDiffeType (fn.getContext (), act));
380-
381391 auto argActivityAttr = ArrayAttr::get (fn.getContext (), argActivityAttrs);
382392 auto retActivityAttr = ArrayAttr::get (fn.getContext (), retActivityAttrs);
383393
0 commit comments