[Inductor][1/n]Split cat customization (#123045)

Summary: Change the config and revise the group batch fusion in order not to reuse the exsiting pre_grad and post_grad fusion options

Test Plan:
# unit test
```
buck2 test @mode/dev-nosan //caffe2/test/inductor:split_cat_fx_passes
```
Test UI: https://www.internalfb.com/intern/testinfra/testrun/17732923560510096
Network: Up: 15MiB  Down: 155MiB  (reSessionID-6a577a14-1772-42d9-9ae8-bfdc62f406a3)
Jobs completed: 267487. Time elapsed: 2:39.7s.
Cache hits: 99%. Commands: 104465 (cached: 104457, remote: 8, local: 0)
Tests finished: Pass 11. Fail 0. Fatal 0. Skip 0. Build failure 0

```
buck2 test @mode/dev-nosan //caffe2/test/inductor/fb:split_cat_fx_passes_fb
```
Test UI: https://www.internalfb.com/intern/testinfra/testrun/9007199283031382
Network: Up: 28MiB  Down: 177MiB  (reSessionID-a3081518-7cba-4c83-b442-c16655ecb2cd)
Jobs completed: 183164. Time elapsed: 1:41.4s.
Cache hits: 99%. Commands: 75875 (cached: 75862, remote: 12, local: 1)
Tests finished: Pass 1. Fail 0. Fatal 0. Skip 0. Build failure 0

```
buck2 test @mode/dev-nosan //caffe2/test/inductor:group_batch_fusion
```
Test UI: https://www.internalfb.com/intern/testinfra/testrun/10133099189612276
Network: Up: 1.3MiB           Down: 3.1MiB           (reSessionID-0d312a2d-e19e-4ba6-9f96-7eb5863734e7)
Discovered 9. Pass 0. Fail 0. Fatal 0. Skip 0. Timeout 0
Network: Up: 1.4MiB  Down: 3.2MiB  (reSessionID-0d312a2d-e19e-4ba6-9f96-7eb5863734e7)
Jobs completed: 68. Time elapsed: 2:19.9s.
Cache hits: 0%. Commands: 13 (cached: 0, remote: 1, local: 12)
Tests finished: Pass 9. Fail 0. Fatal 0. Skip 0. Build failure 0
```
buck2 test @mode/dev-nosan //caffe2/test/inductor:perf
```
Test UI: https://www.internalfb.com/intern/testinfra/testrun/5066549804623287
Network: Up: 1.5MiB  Down: 1.1MiB  (reSessionID-8d912a20-fceb-4698-89c3-d28e0708831f)
Jobs completed: 164. Time elapsed: 1:42.2s.
Cache hits: 0%. Commands: 13 (cached: 0, remote: 1, local: 12)
Tests finished: Pass 57. Fail 0. Fatal 0. Skip 0. Build failure 0

# local reproduce
case 1: with split cat
```
buck2 run @mode/opt //scripts/jackiexu0313/pt2:local_model_with_pt2 -- --test_mode batch-split --model_type "cmf" --flow_id 524546542
```
optimus parameter sent to the scuba:
```
{'before_recompile_pre_grad': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GLL6RBZb-ssXJYcBAMzw0oaKtp80br0LAAAz', 'BatchLayernormFusion': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GH1LAxcxv0Ae_BkFAHVav3K3oosDbr0LAAAz', 'BatchSigmoidPreGradFusion': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GNb0jwR-Ukkqns4CAGRmOqucfedDbr0LAAAz', 'normalization_pass_pre_grad': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GHsIQxm-hn3SPrgCAKq1E-HBsoZHbr0LAAAz', 'remove_split_with_size_one_pass_pre_grad': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GOrJORmbMTV_xlQDAOwolqclPsIAbr0LAAAz', 'merge_getitem_cat_pass_pre_grad': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GCqkmRblvVKybGUDACVxkwVIrWxLbr0LAAAz', 'merge_splits_pass_pre_grad': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GCB1QBfko_kVN0wFAKGjSZv4DJULbr0LAAAz', 'after_recompile_pre_grad': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GMwJPRmu4ry88swDAO1gdA5RCKIXbr0LAAAz', 'before_recompile_post_grad': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GLXCORnNiKeQFmoDABR93CRKmP8Sbr0LAAAz', 'BatchMulPostGradFusion': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GBMIPRnlwQyjSD4BANPuaMhV7MUjbr0LAAAz', 'after_recompile_post_grad': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GJ9KPxkOv4LL8_0DAA65D4kh4JYDbr0LAAAz', 'inductor': Counter({'pattern_matcher_nodes': 2844, 'pattern_matcher_count': 2604, 'normalization_pass': 886, 'remove_split_with_size_one_pass': 748, 'merge_splits_pass': 82, 'merge_getitem_cat_pass': 11, 'scmerge_split_sections_removed': 4, 'batch_aten_mul': 4, 'batch_sigmoid': 2, 'batch_aten_sub': 2, 'batch_layernorm': 1, 'scmerge_split_added': 1, 'scmerge_cat_added': 1, 'scmerge_split_removed': 1, 'scmerge_cat_removed': 1, 'batch_aten_add': 1}), 'BatchAddPostGradFusion': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GEcvPxmxBj-pd8gCABE1QgB-d6N6br0LAAAz', 'BatchSubPostGradFusion': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GEvQxhYomJGj2FMBAEXXAI8Vgzhmbr0LAAAz'}
```
P1202819405

case 2: without split cat
```
buck2 run @mode/opt //scripts/jackiexu0313/pt2:local_model_with_pt2 -- --test_mode batch --model_type "cmf" --flow_id 524546542
```
optimus parameter sent to the scuba:
```
{'before_recompile_pre_grad': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GAY7PxmGthuyjSwEAHF_A767YbMkbr0LAAAz', 'BatchLayernormFusion': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GLDPtBacXyybEOICAKaGCPatq5oabr0LAAAz', 'BatchSigmoidPreGradFusion': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GBu7ORkiDJu42QAEAGmlVTgO_Mpbbr0LAAAz', 'after_recompile_pre_grad': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GC893BZNl99ftY4BAHm5Z8sM4ptSbr0LAAAz', 'before_recompile_post_grad': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GCAeuRYgzPO5RcsCAPO3Z7tdMNMKbr0LAAAz', 'BatchMulPostGradFusion': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GHBIQxm1jlU-xhsFAONkzhh2mgknbr0LAAAz', 'after_recompile_post_grad': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GDoUPhmZ0noiaGMDAJHYuuiwHEAUbr0LAAAz', 'inductor': Counter({'pattern_matcher_nodes': 1189, 'pattern_matcher_count': 757, 'batch_aten_mul': 9, 'batch_aten_sub': 3, 'batch_sigmoid': 2, 'batch_aten_add': 2, 'batch_layernorm': 1, 'batch_linear_post_grad': 1}), 'BatchAddPostGradFusion': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GAluthYxi8uxpI4BAIQDzn3OyywUbr0LAAAz', 'BatchSubPostGradFusion': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GDjsJhTK5VAcot4CADIcAixghrYibr0LAAAz', 'PostGradBatchLinearFusion': 'https://www.internalfb.com/intern/everpaste/?color=0&handle=GEPfJxfJwktC7wsEAA0QbkqYNuVAbr0LAAAz'}
```
P1202823734

# e2e
training_platform:fd4f02cd855f5cc0ccb49317a5a6c8bb

with split cat
f546646358

without split cat
f546647159

Pull Request resolved: https://github.com/pytorch/pytorch/pull/123045
Approved by: https://github.com/jackiexu1992
diff --git a/test/inductor/test_perf.py b/test/inductor/test_perf.py
index 1650458..22ccb63 100644
--- a/test/inductor/test_perf.py
+++ b/test/inductor/test_perf.py
@@ -270,6 +270,19 @@
         self.assertExpectedInline(count_numel(f, *inp), """400""")
 
     @patch.object(config, "split_cat_fx_passes", False)
+    @patch.object(
+        config,
+        "pre_grad_fusion_options",
+        {
+            "batch_linear": {},
+            "batch_linear_lhs": {},
+            "batch_layernorm": {},
+            "batch_tanh": {},
+            "batch_relu": {},
+            "batch_sigmoid": {},
+        },
+    )
+    @patch.object(config, "post_grad_fusion_options", {})
     def test_cat_pointwise_many_complex_inputs(self):
         def f(*inputs):
             input = [torch.nn.functional.gelu(val) for val in inputs]
@@ -279,6 +292,19 @@
         self.assertExpectedInline(count_numel(f, *inp), """6400""")
 
     @patch.object(config, "split_cat_fx_passes", False)
+    @patch.object(
+        config,
+        "pre_grad_fusion_options",
+        {
+            "batch_linear": {},
+            "batch_linear_lhs": {},
+            "batch_layernorm": {},
+            "batch_tanh": {},
+            "batch_relu": {},
+            "batch_sigmoid": {},
+        },
+    )
+    @patch.object(config, "post_grad_fusion_options", {})
     def test_cat_pointwise_many_simple_inputs(self):
         def f(*inputs):
             input = [torch.nn.functional.relu(val) for val in inputs]
diff --git a/torch/_inductor/compile_fx.py b/torch/_inductor/compile_fx.py
index 193ed66..d83798a 100644
--- a/torch/_inductor/compile_fx.py
+++ b/torch/_inductor/compile_fx.py
@@ -662,9 +662,10 @@
         optimus_scuba_log["inductor"] = counters["inductor"]
         signpost_event(
             "optimus",
-            "compile_fx.post_grad_passes",
+            "compile_fx",
             optimus_scuba_log,
         )
+        log.debug("optimus parameter sent to the scuba: %s", optimus_scuba_log)
 
     with V.set_fake_mode(fake_mode):
         const_output_index = None
diff --git a/torch/_inductor/config.py b/torch/_inductor/config.py
index 02bf820..90bfe9a 100644
--- a/torch/_inductor/config.py
+++ b/torch/_inductor/config.py
@@ -98,7 +98,7 @@
 # use post-grad passes.
 pre_grad_custom_pass: Optional[Callable[[torch.fx.graph.Graph], None]] = None
 
-# Optimize away split cat patterns (Experimental)
+# Deprecated
 split_cat_fx_passes = True
 
 # Optimize conv-batchnorm if batchnorm is in eval mode. Slightly reduces numerical stability.
@@ -113,8 +113,24 @@
 # Deprecated
 batch_fusion = True
 
-# Pre grad group/batch fusion and options in order, set to empty dict to disable fusion.
+# Pre grad fusion and options in order, set to empty dict to disable fusion.
 # Call `torch._inductor.fx_passes.group_batch_fusion.list_group_batch_fusions()` to see available fusions.
+# batch fusion options:
+# batch_linear
+# batch_linear_lhs
+# batch_layernorm
+# batch_tanh
+# batch_relu
+# batch_sigmoid
+
+# split cat fusion options:
+# normalization_pass
+# remove_split_with_size_one_pass
+# merge_getitem_cat_pass
+# merge_stack_tahn_unbind
+# merge_splits_pass
+# mutate_cat_pass
+# split_cat_pass
 pre_grad_fusion_options: Dict[str, Dict[str, Any]] = {
     "batch_linear": {},
     "batch_linear_lhs": {},
@@ -122,9 +138,16 @@
     "batch_tanh": {},
     "batch_relu": {},
     "batch_sigmoid": {},
+    "normalization_pass": {},
+    "remove_split_with_size_one_pass": {},
+    "merge_getitem_cat_pass": {},
+    "merge_stack_tahn_unbind_pass": {},
+    "merge_splits_pass": {},
+    "mutate_cat_pass": {},
+    "split_cat_pass": {},
 }
 
-# Post grad group/batch fusion and options, set to empty dict to disable fusion.
+# Post grad fusion and options, set to empty dict to disable fusion.
 # Call `torch._inductor.fx_passes.group_batch_fusion.list_group_batch_fusions(False)` to see available fusions.
 post_grad_fusion_options: Dict[str, Dict[str, Any]] = {}
 
diff --git a/torch/_inductor/fx_passes/group_batch_fusion.py b/torch/_inductor/fx_passes/group_batch_fusion.py
index caf2748..9ad1e39 100644
--- a/torch/_inductor/fx_passes/group_batch_fusion.py
+++ b/torch/_inductor/fx_passes/group_batch_fusion.py
@@ -1028,6 +1028,9 @@
 def generate_fusion_from_config(config_options: Dict[str, Any], pre_grad=True):
     fusions: List[GroupBatchFusionBase] = []
     for name, options in config_options.items():
+        # we skip all patterns from pattern_matcher passes (e.g., split_cat)
+        if name not in PRE_GRAD_FUSIONS and name not in POST_GRAD_FUSIONS:
+            continue
         fusion_cls = PRE_GRAD_FUSIONS[name] if pre_grad else POST_GRAD_FUSIONS[name]
         _options = graph_search_options.copy()
         _options.update(options)
diff --git a/torch/_inductor/fx_passes/pre_grad.py b/torch/_inductor/fx_passes/pre_grad.py
index 2142b36..dbec853 100644
--- a/torch/_inductor/fx_passes/pre_grad.py
+++ b/torch/_inductor/fx_passes/pre_grad.py
@@ -1,6 +1,6 @@
 import copy
 import logging
-from typing import List, Optional
+from typing import Dict, List, Optional
 
 import torch
 import torch.nn as nn
@@ -28,33 +28,9 @@
 
 log = logging.getLogger(__name__)
 
-normalization_pass = PatternMatcherPass(
-    prevent_match_across_mutations=True, pass_name="normalization_pass"
-)
-merge_splits_pass = PatternMatcherPass(
-    prevent_match_across_mutations=True, pass_name="merge_splits_pass"
-)
-split_cat_pass = PatternMatcherPass(
-    prevent_match_across_mutations=True, pass_name="split_cat_pass"
-)
-unbind_stack_pass = PatternMatcherPass(
-    prevent_match_across_mutations=True, pass_name="unbind_stack_pass"
-)
 efficient_conv_bn_eval_pass = PatternMatcherPass(
     prevent_match_across_mutations=True, pass_name="efficient_conv_bn_eval_pass"
 )
-merge_getitem_cat_pass = PatternMatcherPass(
-    prevent_match_across_mutations=True, pass_name="merge_getitem_cat_pass"
-)
-merge_stack_tahn_unbind_pass = PatternMatcherPass(
-    prevent_match_across_mutations=True, pass_name="merge_stack_tahn_unbind_pass"
-)
-mutate_cat_pass = PatternMatcherPass(
-    prevent_match_across_mutations=True, pass_name="mutate_cat_pass"
-)
-remove_split_with_size_one_pass = PatternMatcherPass(
-    prevent_match_across_mutations=True, pass_name="remove_split_with_size_one_pass"
-)
 
 fuse_split_linear_add_pass = PatternMatcherPass(
     prevent_match_across_mutations=True,
@@ -99,6 +75,10 @@
     return True
 
 
+def construct_pattern_matcher_pass(pass_name):
+    return PatternMatcherPass(prevent_match_across_mutations=True, pass_name=pass_name)
+
+
 def fuse_parallel_linear_pass(graph):
     return None
 
@@ -107,17 +87,26 @@
     return None
 
 
-pattern_matcher_passes: List[PatternMatcherPass] = [
-    normalization_pass,
-    remove_split_with_size_one_pass,
-    merge_getitem_cat_pass,
-    merge_stack_tahn_unbind_pass,
-    merge_splits_pass,
-    mutate_cat_pass,
-    split_cat_pass,
-    unbind_stack_pass,
-    efficient_conv_bn_eval_pass,
+PRE_GRAD_PATTERNS: Dict[str, PatternMatcherPass] = dict()
+pass_names = [
+    "normalization_pass",
+    "remove_split_with_size_one_pass",
+    "merge_getitem_cat_pass",
+    "merge_stack_tahn_unbind_pass",
+    "merge_splits_pass",
+    "mutate_cat_pass",
+    "split_cat_pass",
+    "unbind_stack_pass",
 ]
+
+for pass_name in pass_names:
+    PRE_GRAD_PATTERNS[pass_name] = construct_pattern_matcher_pass(pass_name)
+# split_cat related fusions
+pattern_matcher_passes = list(PRE_GRAD_PATTERNS.values())
+# non-split_cat related fusions
+# TODO: move them to the fusions dict too.
+pattern_matcher_passes.append(efficient_conv_bn_eval_pass)
+
 pattern_matcher_passes_aten: List[PatternMatcherPass] = [
     remove_split_with_size_one_pass_aten,
     merge_getitem_cat_pass_aten,
diff --git a/torch/_inductor/fx_passes/split_cat.py b/torch/_inductor/fx_passes/split_cat.py
index 2b204b6..b2c31d7 100644
--- a/torch/_inductor/fx_passes/split_cat.py
+++ b/torch/_inductor/fx_passes/split_cat.py
@@ -27,16 +27,7 @@
     RepeatedExpr,
 )
 from .group_batch_fusion import is_node_meta_valid
-from .pre_grad import (
-    merge_getitem_cat_pass,
-    merge_splits_pass,
-    merge_stack_tahn_unbind_pass,
-    mutate_cat_pass,
-    normalization_pass,
-    remove_split_with_size_one_pass,
-    split_cat_pass,
-    unbind_stack_pass,
-)
+from .pre_grad import PRE_GRAD_PATTERNS
 
 log = logging.getLogger(__name__)
 
@@ -157,12 +148,12 @@
 
 @register_graph_pattern(
     CallFunctionVarArgs(torch.split, users=MULTIPLE),
-    pass_dict=normalization_pass,
+    pass_dict=PRE_GRAD_PATTERNS["normalization_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 @register_graph_pattern(
     CallMethodVarArgs("split", users=MULTIPLE),
-    pass_dict=normalization_pass,
+    pass_dict=PRE_GRAD_PATTERNS["normalization_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 def normalize_split_default(match: Match, *args, **kwargs):
@@ -171,12 +162,12 @@
 
 @register_graph_pattern(
     CallFunctionVarArgs(torch.split, users=MULTIPLE),
-    pass_dict=remove_split_with_size_one_pass,
+    pass_dict=PRE_GRAD_PATTERNS["remove_split_with_size_one_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 @register_graph_pattern(
     CallMethodVarArgs("split", users=MULTIPLE),
-    pass_dict=remove_split_with_size_one_pass,
+    pass_dict=PRE_GRAD_PATTERNS["remove_split_with_size_one_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 def remove_split_with_size_one(match: Match, *args, **kwargs):
@@ -211,12 +202,12 @@
 
 @register_graph_pattern(
     CallFunctionVarArgs(torch.unbind, users=MULTIPLE),
-    pass_dict=normalization_pass,
+    pass_dict=PRE_GRAD_PATTERNS["normalization_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 @register_graph_pattern(
     CallMethodVarArgs("unbind", users=MULTIPLE),
-    pass_dict=normalization_pass,
+    pass_dict=PRE_GRAD_PATTERNS["normalization_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 def normalize_unbind_default(match: Match, *args, **kwargs):
@@ -253,7 +244,7 @@
 
 @register_graph_pattern(
     CallFunctionVarArgs(torch.cat, users=MULTIPLE),
-    pass_dict=normalization_pass,
+    pass_dict=PRE_GRAD_PATTERNS["normalization_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 def normalize_cat_default(match: Match, *args, **kwargs):
@@ -315,7 +306,7 @@
 
 @register_graph_pattern(
     CallFunctionVarArgs(torch.stack, users=MULTIPLE),
-    pass_dict=normalization_pass,
+    pass_dict=PRE_GRAD_PATTERNS["normalization_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 def normalize_stack_default(match: Match, *args, **kwargs):
@@ -361,7 +352,7 @@
 
 @register_graph_pattern(
     CallMethodVarArgs("squeeze", users=MULTIPLE),
-    pass_dict=normalization_pass,
+    pass_dict=PRE_GRAD_PATTERNS["normalization_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 def normalize_squeeze_default(match: Match, *args, **kwargs):
@@ -444,7 +435,7 @@
         ),
         KeywordArg("next_split_sections"),
     ),
-    pass_dict=merge_splits_pass,
+    pass_dict=PRE_GRAD_PATTERNS["merge_splits_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 def merge_splits(
@@ -1050,7 +1041,7 @@
             _users=MULTIPLE,
         ),
     ),
-    pass_dict=split_cat_pass,
+    pass_dict=PRE_GRAD_PATTERNS["split_cat_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 @register_graph_pattern(
@@ -1068,7 +1059,7 @@
             _users=MULTIPLE,
         )
     ),
-    pass_dict=split_cat_pass,
+    pass_dict=PRE_GRAD_PATTERNS["split_cat_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 def merge_split_squeeze(
@@ -1122,21 +1113,21 @@
 
 @register_graph_pattern(
     CallFunction([torch.stack, torch.cat], getitem_unbind, Ignored(), _users=MULTIPLE),
-    pass_dict=unbind_stack_pass,
+    pass_dict=PRE_GRAD_PATTERNS["unbind_stack_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 @register_graph_pattern(
     CallFunction(
         [torch.stack, torch.cat], getitem_unbind, dim=Ignored(), _users=MULTIPLE
     ),
-    pass_dict=unbind_stack_pass,
+    pass_dict=PRE_GRAD_PATTERNS["unbind_stack_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 @register_graph_pattern(
     CallFunction(
         [torch.stack, torch.cat], tensors=getitem_unbind, dim=Ignored(), _users=MULTIPLE
     ),
-    pass_dict=unbind_stack_pass,
+    pass_dict=PRE_GRAD_PATTERNS["unbind_stack_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 def merge_unbind_stack(match: Match, unbind_input: torch.fx.Node, dim: int):
@@ -1165,7 +1156,7 @@
         dim=Ignored(),
         _users=MULTIPLE,
     ),
-    pass_dict=split_cat_pass,
+    pass_dict=PRE_GRAD_PATTERNS["split_cat_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 @register_graph_pattern(
@@ -1175,7 +1166,7 @@
         dim=Ignored(),
         _users=MULTIPLE,
     ),
-    pass_dict=split_cat_pass,
+    pass_dict=PRE_GRAD_PATTERNS["split_cat_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 @register_graph_pattern(
@@ -1185,7 +1176,7 @@
         Ignored(),
         _users=MULTIPLE,
     ),
-    pass_dict=split_cat_pass,
+    pass_dict=PRE_GRAD_PATTERNS["split_cat_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 def simplify_split_cat(match: Match, split_sections: List[int], dim: int):
@@ -1270,7 +1261,7 @@
         dim=Ignored(),
         _users=MULTIPLE,
     ),
-    pass_dict=merge_getitem_cat_pass,
+    pass_dict=PRE_GRAD_PATTERNS["merge_getitem_cat_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 def merge_getitem_cat(match: Match, split_sections: List[int], dim: int):
@@ -1378,7 +1369,7 @@
         dim=Ignored(),
         _users=MULTIPLE,
     ),
-    pass_dict=mutate_cat_pass,
+    pass_dict=PRE_GRAD_PATTERNS["mutate_cat_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 def mutate_cat_node(match: Match, split_sections: List[int], dim: int):
@@ -1474,7 +1465,7 @@
             dim=Ignored(),
         ),
     ),
-    pass_dict=merge_getitem_cat_pass,
+    pass_dict=PRE_GRAD_PATTERNS["merge_getitem_cat_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 @register_graph_pattern(
@@ -1486,7 +1477,7 @@
             dim=Ignored(),
         ),
     ),
-    pass_dict=merge_getitem_cat_pass,
+    pass_dict=PRE_GRAD_PATTERNS["merge_getitem_cat_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 @register_graph_pattern(
@@ -1498,7 +1489,7 @@
             Ignored(),
         ),
     ),
-    pass_dict=merge_stack_tahn_unbind_pass,
+    pass_dict=PRE_GRAD_PATTERNS["merge_stack_tahn_unbind_pass"],
     extra_check=config_flag("split_cat_fx_passes"),
 )
 def merge_stack_tahn_unbind(match: Match, split_sections: List[int], dim: int):