@@ -20,6 +20,42 @@ def _autograd_graph_contains(fn, needle: str) -> bool:
2020
2121
2222class DynamoCudagraphWorkaroundTests (unittest .TestCase ):
23+ def test_activation_checkpointing_downgrades_reduce_overhead_to_default (self ):
24+ from simpletuner .helpers .training import dynamo
25+
26+ config = SimpleNamespace (
27+ dynamo_backend = "inductor" ,
28+ dynamo_mode = "reduce-overhead" ,
29+ dynamo_use_regional_compilation = True ,
30+ gradient_checkpointing = True ,
31+ )
32+ with (
33+ unittest .mock .patch .object (dynamo .logger , "warning" ) as warning ,
34+ unittest .mock .patch .dict (os .environ , {}, clear = True ),
35+ ):
36+ self .assertTrue (dynamo .apply_checkpointing_cudagraph_compatibility (config ))
37+ self .assertEqual (os .environ ["TRAINING_DYNAMO_MODE" ], "default" )
38+ self .assertEqual (os .environ ["ACCELERATE_DYNAMO_MODE" ], "default" )
39+
40+ self .assertEqual (config .dynamo_mode , "default" )
41+ self .assertTrue (config .dynamo_use_regional_compilation )
42+ warning .assert_called_once ()
43+ self .assertEqual (warning .call_args .args [2 ], "https://github.com/pytorch/pytorch/issues/154306" )
44+
45+ def test_reduce_overhead_is_preserved_without_activation_checkpointing (self ):
46+ from simpletuner .helpers .training import dynamo
47+
48+ config = SimpleNamespace (
49+ dynamo_backend = "inductor" ,
50+ dynamo_mode = "reduce-overhead" ,
51+ gradient_checkpointing = False ,
52+ )
53+ with unittest .mock .patch .dict (os .environ , {}, clear = True ):
54+ self .assertFalse (dynamo .apply_checkpointing_cudagraph_compatibility (config ))
55+ self .assertNotIn ("TRAINING_DYNAMO_MODE" , os .environ )
56+
57+ self .assertEqual (config .dynamo_mode , "reduce-overhead" )
58+
2359 def test_peft_lora_cudagraph_patch_clones_base_result (self ):
2460 from peft import LoraConfig
2561 from peft .tuners .lora .layer import Linear
@@ -123,6 +159,25 @@ def test_inductor_cudagraph_tree_mode_enabled_from_accelerate_env(self):
123159 ):
124160 self .assertTrue (dynamo ._inductor_cudagraphs_enabled (SimpleNamespace (dynamo_backend = None )))
125161
162+ def test_configured_inductor_overrides_outer_accelerate_no_backend (self ):
163+ import torch ._inductor .config as inductor_config
164+
165+ from simpletuner .helpers .training import dynamo
166+
167+ config = SimpleNamespace (dynamo_backend = "inductor" , dynamo_mode = "reduce-overhead" )
168+ with (
169+ unittest .mock .patch .dict (
170+ os .environ ,
171+ {
172+ "ACCELERATE_DYNAMO_BACKEND" : "NO" ,
173+ "ACCELERATE_DYNAMO_MODE" : "default" ,
174+ },
175+ ),
176+ unittest .mock .patch .object (inductor_config .triton , "cudagraphs" , False ),
177+ unittest .mock .patch .object (inductor_config .triton , "cudagraph_trees" , True ),
178+ ):
179+ self .assertTrue (dynamo ._inductor_cudagraphs_enabled (config ))
180+
126181
127182if __name__ == "__main__" :
128183 unittest .main ()
0 commit comments