瀏覽代碼

gradient_checkpointing_enable()

Kai Wu 7 月之前
父節點
當前提交
50dff0b78e
共有 1 個文件被更改,包括 1 次插入1 次删除
  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():