walidsobhie-code commited on
Commit
98e3329
·
1 Parent(s): c2013aa

fix: set expandable_segments=False to fix PyTorch #124807 gradient checkpointing bug

Browse files
Files changed (1) hide show
  1. train_simple_nobnb.py +3 -2
train_simple_nobnb.py CHANGED
@@ -144,7 +144,8 @@ def train(config: dict):
144
  use_8bit = hardware_config.get("use_8bit", False)
145
 
146
  # Set environment variables for better CUDA memory management
147
- os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True,max_split_size_mb:512"
 
148
 
149
  # Clear CUDA cache before loading
150
  if torch.cuda.is_available():
@@ -229,7 +230,7 @@ def train(config: dict):
229
  report_to="none",
230
  dataloader_num_workers=0,
231
  remove_unused_columns=False,
232
- optim="paged_adamw_32bit" if (use_4bit or use_8bit) else "adamw_torch",
233
  )
234
 
235
  data_collator = DataCollatorForLanguageModeling(
 
144
  use_8bit = hardware_config.get("use_8bit", False)
145
 
146
  # Set environment variables for better CUDA memory management
147
+ # expandable_segments:False fixes a known PyTorch bug (#124807, #128829) with gradient checkpointing
148
+ os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:False"
149
 
150
  # Clear CUDA cache before loading
151
  if torch.cuda.is_available():
 
230
  report_to="none",
231
  dataloader_num_workers=0,
232
  remove_unused_columns=False,
233
+ optim="paged_adamw_32bit" if (use_4bit or use_8bit) else "adamw_torch_fused",
234
  )
235
 
236
  data_collator = DataCollatorForLanguageModeling(