diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index 257d79945a4..e8da712e2d2 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -2650,8 +2650,8 @@ def get_cpu_offload_context( retain_pinned_cpu_buffers, ): """Get CPU offload context and sync function.""" - if is_te_min_version("2.5.0"): - # Enables the additional double buffering switch for activations during LLM training + if is_te_min_version("2.10.0"): + # TE 2.10+ supports retain_pinned_cpu_buffers context, sync_func = _get_cpu_offload_context( enabled, num_layers, @@ -2661,6 +2661,16 @@ def get_cpu_offload_context( double_buffering, retain_pinned_cpu_buffers=retain_pinned_cpu_buffers, ) + elif is_te_min_version("2.5.0"): + # TE 2.5-2.9 supports double_buffering but not retain_pinned_cpu_buffers + context, sync_func = _get_cpu_offload_context( + enabled, + num_layers, + model_layers, + activation_offloading, + weight_offloading, + double_buffering, + ) elif is_te_min_version("1.10.0.dev0"): context, sync_func = _get_cpu_offload_context( enabled, num_layers, model_layers, activation_offloading, weight_offloading