Explorar o código

gradient_checkpointing_enable()

Kai Wu hai 7 meses
pai
achega
50dff0b78e
Modificáronse 1 ficheiros con 1 adicións e 1 borrados
  1. 1 1
      src/llama_recipes/finetuning.py

+ 1 - 1
src/llama_recipes/finetuning.py

@@ -208,7 +208,7 @@ def main(**kwargs):
         )
         if fsdp_config.fsdp_activation_checkpointing:            
             model.enable_input_require_grads()
-            #model.gradient_checkpointing_enable()
+            model.gradient_checkpointing_enable()
             apply_fsdp_checkpointing(model)                      
     elif not train_config.quantization and not train_config.enable_fsdp:
         if is_xpu_available():