From b9e1f62afe3e22faff243dc0e7f050bfcc3822bf Mon Sep 17 00:00:00 2001 From: hongbinl Date: Tue, 9 Jun 2026 03:25:53 -0700 Subject: [PATCH 01/12] Add TE op CPU offload opt-out API Signed-off-by: hongbinl --- tests/pytorch/test_fusible_ops.py | 42 +++++++++++++++++ .../pytorch/ops/basic/activation.py | 7 +-- .../pytorch/ops/basic/basic_linear.py | 10 ++-- .../pytorch/ops/basic/dropout.py | 4 +- .../pytorch/ops/basic/grouped_linear.py | 33 ++++++------- .../pytorch/ops/basic/l2normalization.py | 4 +- .../pytorch/ops/basic/layer_norm.py | 4 +- .../pytorch/ops/basic/rmsnorm.py | 4 +- .../pytorch/ops/basic/swiglu.py | 10 ++-- .../pytorch/ops/fused/forward_grouped_mlp.py | 14 ++---- .../fused/forward_linear_bias_activation.py | 4 +- .../ops/fused/forward_linear_bias_add.py | 4 +- .../ops/fused/forward_linear_scale_add.py | 4 +- .../ops/fused/userbuffers_forward_linear.py | 4 +- transformer_engine/pytorch/ops/op.py | 46 +++++++++++++++++++ 15 files changed, 122 insertions(+), 72 deletions(-) diff --git a/tests/pytorch/test_fusible_ops.py b/tests/pytorch/test_fusible_ops.py index 7c75d11e3b..dc5294e13c 100644 --- a/tests/pytorch/test_fusible_ops.py +++ b/tests/pytorch/test_fusible_ops.py @@ -71,6 +71,48 @@ # Supported devices _devices: list[torch.device] = [torch.device("cpu"), torch.device("cuda")] + +def test_basic_operation_cpu_offloading_control(monkeypatch): + """BasicOperation should expose a public opt-out for activation CPU offload.""" + import transformer_engine.pytorch.cpu_offload as cpu_offload + + calls = [] + tensor = torch.empty(1) + tensor_id = id(tensor) + op = te_ops.Identity() + + monkeypatch.setattr(cpu_offload, "is_cpu_offload_enabled", lambda: True) + monkeypatch.setattr( + cpu_offload, + "start_offload", + lambda *tensors: calls.append(("start", [id(t) for t in tensors])), + ) + monkeypatch.setattr( + cpu_offload, + "mark_activation_offload", + lambda *tensors: calls.append(("mark", [id(t) for t in tensors])), + ) + monkeypatch.setattr( + cpu_offload, + "mark_not_offload", + lambda *tensors: calls.append(("skip", [id(t) for t in tensors])), + ) + + op.maybe_mark_and_start_activation_offload(tensor, None, start=True) + assert calls == [("start", [tensor_id]), ("mark", [tensor_id])] + + calls.clear() + op.disable_cpu_offloading() + op.maybe_mark_and_start_activation_offload(tensor, start=True) + assert calls == [("skip", [tensor_id])] + + calls.clear() + op.enable_cpu_offloading() + monkeypatch.setattr(cpu_offload, "is_cpu_offload_enabled", lambda: False) + op.maybe_mark_and_start_activation_offload(tensor, start=True) + assert calls == [] + + # Supported quantization recipes _quantization_list: list[Optional[str]] = [None] if fp8_available: diff --git a/transformer_engine/pytorch/ops/basic/activation.py b/transformer_engine/pytorch/ops/basic/activation.py index f4beffe90c..64a4de55e6 100644 --- a/transformer_engine/pytorch/ops/basic/activation.py +++ b/transformer_engine/pytorch/ops/basic/activation.py @@ -13,7 +13,6 @@ import transformer_engine_torch as tex from ...constants import DType -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...tensor.float8_tensor import Float8CurrentScalingQuantizer, Quantizer from ...utils import clear_tensor_data from ..op import BasicOperation, OperationContext @@ -114,8 +113,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - if is_cpu_offload_enabled(): - mark_activation_offload(x) + self.maybe_mark_and_start_activation_offload(x) ctx.save_for_backward(x) ctx.dtype = dtype ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer @@ -414,8 +412,7 @@ def fuser_forward( ctx = basic_op_ctxs[0] if ctx.requires_grad: - if is_cpu_offload_enabled(): - mark_activation_offload(x) + self.maybe_mark_and_start_activation_offload(x) ctx.input_requires_grad = True ctx.extra_input_requires_grad = extra_input.requires_grad ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/basic/basic_linear.py b/transformer_engine/pytorch/ops/basic/basic_linear.py index 6b17d66fcd..5b59a7be2c 100644 --- a/transformer_engine/pytorch/ops/basic/basic_linear.py +++ b/transformer_engine/pytorch/ops/basic/basic_linear.py @@ -13,7 +13,6 @@ import torch from ...cpp_extensions import general_gemm -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...distributed import ( CudaRNGStatesTracker, gather_along_first_dim, @@ -1049,11 +1048,10 @@ def op_forward( else: saved_input = x_local saved_weight = w - if is_cpu_offload_enabled(): - # No special CPU offloading logic is needed for weights. saved_weight is - # either self.weight (nn.Parameter, auto-excluded from offload) or a - # workspace freshly created each forward pass. - mark_activation_offload(saved_input) + # No special CPU offloading logic is needed for weights. saved_weight is + # either self.weight (nn.Parameter, auto-excluded from offload) or a + # workspace freshly created each forward pass. + self.maybe_mark_and_start_activation_offload(saved_input) ctx.save_for_backward(saved_input, saved_weight) ctx.with_quantized_compute = with_quantized_compute and backward_override is None ctx.backward_override = backward_override diff --git a/transformer_engine/pytorch/ops/basic/dropout.py b/transformer_engine/pytorch/ops/basic/dropout.py index 8850604aad..7a6cd9ad8f 100644 --- a/transformer_engine/pytorch/ops/basic/dropout.py +++ b/transformer_engine/pytorch/ops/basic/dropout.py @@ -9,7 +9,6 @@ import torch import transformer_engine_torch as tex -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...tensor import Quantizer from ...tensor.storage.float8_tensor_storage import Float8TensorStorage from .._common import maybe_autocast_dtype, maybe_dequantize @@ -71,8 +70,7 @@ def op_forward( # Save context for backward if ctx.requires_grad: - if is_cpu_offload_enabled(): - mark_activation_offload(mask) + self.maybe_mark_and_start_activation_offload(mask) ctx.save_for_backward(mask) ctx.impl = impl ctx.dropout_probability = self.dropout_probability diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index 73e328c9d1..3c05cdcde5 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -23,7 +23,6 @@ _2X_ACC_DGRAD, _2X_ACC_WGRAD, ) -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload, start_offload from ...quantization import FP8GlobalStateManager, QuantizerRole, Recipe from ...quantized_tensor import QuantizedTensorStorage from ...tensor import MXFP8Quantizer, MXFP8Tensor, Quantizer @@ -1032,19 +1031,16 @@ def fuser_forward_save_ctx( # Note: No special logic is needed for weights. They are # either nn.Parameter (auto-excluded from offload) or are # temporary workspaces freshly created in each forward pass. - if is_cpu_offload_enabled(): - saved = tensors_to_save[0] - offset = 4 if self._scale_bias else 3 - if use_grouped_tensor_path: - # Layout: [split_sizes, base_split_offsets, split_points, (scales?), grouped_x, *weights] - grouped_x = saved[offset] - if grouped_x is not None: - mark_activation_offload(grouped_x) - else: - # Layout: [split_sizes, None, None, (scales?), *xs, *ws] - live_xs = [t for t in saved[offset : offset + self.num_groups] if t is not None] - if live_xs: - mark_activation_offload(*live_xs) + saved = tensors_to_save[0] + offset = 4 if self._scale_bias else 3 + if use_grouped_tensor_path: + # Layout: [split_sizes, base_split_offsets, split_points, (scales?), grouped_x, *weights] + grouped_x = saved[offset] + self.maybe_mark_and_start_activation_offload(grouped_x) + else: + # Layout: [split_sizes, None, None, (scales?), *xs, *ws] + live_xs = [t for t in saved[offset : offset + self.num_groups] if t is not None] + self.maybe_mark_and_start_activation_offload(*live_xs) ctx.save_for_backward(*tensors_to_save[0]) @@ -1130,10 +1126,8 @@ def _fuser_forward_split_quantize( xs = tex.split_quantize(x, split_sizes_int, input_quantizers) else: xs = torch.split(x, split_sizes_int) - if is_cpu_offload_enabled(): - live_xs = [t for t in xs if t is not None] - if live_xs: - start_offload(*live_xs) + live_xs = [t for t in xs if t is not None] + self.maybe_mark_and_start_activation_offload(*live_xs, start=True) # Allocate output tensor in_shape = list(input_.size()) @@ -1239,8 +1233,7 @@ def _fuser_forward_grouped_tensor( tensor_offsets=base_split_offsets * self.in_features, ) - if is_cpu_offload_enabled() and grouped_x is not None: - start_offload(grouped_x) + self.maybe_mark_and_start_activation_offload(grouped_x, start=True) # Build the weight GroupedTensor / list. if self.single_grouped_weight: diff --git a/transformer_engine/pytorch/ops/basic/l2normalization.py b/transformer_engine/pytorch/ops/basic/l2normalization.py index be155c9356..4e1e88bf20 100644 --- a/transformer_engine/pytorch/ops/basic/l2normalization.py +++ b/transformer_engine/pytorch/ops/basic/l2normalization.py @@ -11,7 +11,6 @@ import torch from ...torch_version import torch_version -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...jit import ( l2normalization_fused, l2normalization_fwd_fused, @@ -102,8 +101,7 @@ def op_forward( # Save state for backward pass if requires_grad: - if is_cpu_offload_enabled(): - mark_activation_offload(x, rsqrt_norm) + self.maybe_mark_and_start_activation_offload(x, rsqrt_norm) ctx.save_for_backward(x, rsqrt_norm) return y diff --git a/transformer_engine/pytorch/ops/basic/layer_norm.py b/transformer_engine/pytorch/ops/basic/layer_norm.py index 3fda5145c6..4b17813bce 100644 --- a/transformer_engine/pytorch/ops/basic/layer_norm.py +++ b/transformer_engine/pytorch/ops/basic/layer_norm.py @@ -14,7 +14,6 @@ from transformer_engine_torch import layernorm_bwd, layernorm_fwd from ...constants import TE_DType -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...export import is_in_onnx_export_mode from ...tensor import Quantizer from ...utils import ( @@ -216,8 +215,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - if is_cpu_offload_enabled(): - mark_activation_offload(x, means, rstdevs) + self.maybe_mark_and_start_activation_offload(x, means, rstdevs) ctx.save_for_backward(x, means, rstdevs) ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/basic/rmsnorm.py b/transformer_engine/pytorch/ops/basic/rmsnorm.py index 1d8d8be971..d679ceeb31 100644 --- a/transformer_engine/pytorch/ops/basic/rmsnorm.py +++ b/transformer_engine/pytorch/ops/basic/rmsnorm.py @@ -14,7 +14,6 @@ from transformer_engine_torch import rmsnorm_bwd, rmsnorm_fwd from ...constants import TE_DType -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...export import is_in_onnx_export_mode from ...tensor import Quantizer from ...utils import ( @@ -197,8 +196,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - if is_cpu_offload_enabled(): - mark_activation_offload(x, rstdevs) + self.maybe_mark_and_start_activation_offload(x, rstdevs) ctx.save_for_backward(x, rstdevs) ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/basic/swiglu.py b/transformer_engine/pytorch/ops/basic/swiglu.py index 02f330ede3..3bfc6cd5ad 100644 --- a/transformer_engine/pytorch/ops/basic/swiglu.py +++ b/transformer_engine/pytorch/ops/basic/swiglu.py @@ -12,7 +12,6 @@ import transformer_engine_torch as tex from ...constants import DType -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...tensor import Float8CurrentScalingQuantizer, Quantizer from ...utils import clear_tensor_data from ..op import BasicOperation, OperationContext @@ -127,8 +126,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - if is_cpu_offload_enabled(): - mark_activation_offload(input_) + self.maybe_mark_and_start_activation_offload(input_) ctx.save_for_backward(input_) ctx.dtype = dtype ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer @@ -311,8 +309,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - if is_cpu_offload_enabled(): - mark_activation_offload(x) + self.maybe_mark_and_start_activation_offload(x) ctx.save_for_backward(x) ctx.dtype = dtype ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer @@ -462,8 +459,7 @@ def fuser_forward( # Save state for backward pass ctx = basic_op_ctxs[0] if ctx.requires_grad: - if is_cpu_offload_enabled(): - mark_activation_offload(input_) + self.maybe_mark_and_start_activation_offload(input_) ctx.input_requires_grad = True ctx.extra_input_requires_grad = extra_input.requires_grad ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py b/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py index f03ccc15b5..5767923925 100644 --- a/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py +++ b/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py @@ -13,7 +13,6 @@ import torch import transformer_engine_torch as tex -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload, start_offload from ...cpp_extensions import general_gemm, general_grouped_gemm_for_grouped_tensor from ...quantization import Recipe from ...tensor import NVFP4Quantizer, NVFP4Tensor, Quantizer @@ -170,7 +169,7 @@ def fuser_forward( basic_op_kwargs: list[dict[str, Any]], ) -> tuple[torch.Tensor, Iterable[Iterable[torch.Tensor]]]: # Get basic operations - fc1_op, _, fc2_op = self.basic_ops + fc1_op, activation_op, fc2_op = self.basic_ops fc1_ctx, activation_ctx, fc2_ctx = basic_op_ctxs # Tensor properties @@ -726,8 +725,6 @@ def fuser_forward( # Save state for backward pass if requires_grad: mark_grouped_tensor(grouped_fc1_x, activation_in, scales, grouped_fc2_x) - activation_op = self.basic_ops[1] - cpu_offloading = is_cpu_offload_enabled() activation_is_srelu = isinstance(activation_op, ScaledSReLU) activation_recompute_in_mlp = bool( getattr(activation_op, "activation_recompute_in_mlp", False) @@ -749,12 +746,9 @@ def fuser_forward( grouped_fc_x.rowwise_data = None grouped_fc_x.scale_inv = None - if cpu_offloading: - activation_tensors = [ - t for t in (grouped_fc1_x, activation_in, saved_grouped_fc2_x) if t is not None - ] - start_offload(*activation_tensors) - mark_activation_offload(*activation_tensors) + fc1_op.maybe_mark_and_start_activation_offload(grouped_fc1_x, start=True) + activation_op.maybe_mark_and_start_activation_offload(activation_in, start=True) + fc2_op.maybe_mark_and_start_activation_offload(saved_grouped_fc2_x, start=True) # FC1 saved-tensor layout. # [split_sizes, base_split_offsets, split_points, diff --git a/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py b/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py index 8df929f799..b3e0cd7497 100644 --- a/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py +++ b/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py @@ -10,7 +10,6 @@ import torch -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...quantization import FP8GlobalStateManager from ...tensor import Quantizer from ..basic import BasicLinear, Bias @@ -129,8 +128,7 @@ def fuser_forward( else: saved_input = x_local saved_weight = w - if is_cpu_offload_enabled(): - mark_activation_offload(saved_input) + linear_op.maybe_mark_and_start_activation_offload(saved_input) linear_op_ctx.save_for_backward(saved_input, saved_weight) linear_op_ctx.with_quantized_compute = ( with_quantized_compute and backward_override is None diff --git a/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py b/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py index 5376a7d264..1fedb34de9 100644 --- a/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py +++ b/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py @@ -10,7 +10,6 @@ import torch -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...quantization import FP8GlobalStateManager from ...tensor import Quantizer from ..basic import AddExtraInput, BasicLinear, Bias @@ -126,8 +125,7 @@ def fuser_forward( else: saved_input = x_local saved_weight = w - if is_cpu_offload_enabled(): - mark_activation_offload(saved_input) + linear_op.maybe_mark_and_start_activation_offload(saved_input) linear_op_ctx.save_for_backward(saved_input, saved_weight) linear_op_ctx.with_quantized_compute = ( with_quantized_compute and backward_override is None diff --git a/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py b/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py index abeb39adfa..3f3ddb4057 100644 --- a/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py +++ b/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py @@ -10,7 +10,6 @@ import torch -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...quantization import FP8GlobalStateManager from ...tensor import Quantizer from ..basic import AddExtraInput, BasicLinear, ConstantScale @@ -107,8 +106,7 @@ def fuser_forward( else: saved_input = x_local saved_weight = w - if is_cpu_offload_enabled(): - mark_activation_offload(saved_input) + linear_op.maybe_mark_and_start_activation_offload(saved_input) linear_op_ctx.save_for_backward(saved_input, saved_weight) linear_op_ctx.with_quantized_compute = ( with_quantized_compute and backward_override is None diff --git a/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py b/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py index 3a8ff5438d..6fbf815d0d 100644 --- a/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py +++ b/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py @@ -12,7 +12,6 @@ from transformer_engine_torch import CommOverlapType from ...cpp_extensions import general_gemm -from ...cpu_offload import is_cpu_offload_enabled, mark_activation_offload from ...distributed import get_distributed_world_size from ...quantization import FP8GlobalStateManager from ...module.base import ( @@ -354,8 +353,7 @@ def fuser_forward( # Save state for backward pass if linear_op_ctx.requires_grad: - if is_cpu_offload_enabled(): - mark_activation_offload(x_local) + linear_op.maybe_mark_and_start_activation_offload(x_local) linear_op_ctx.save_for_backward(x_local, w) linear_op_ctx.with_quantized_compute = with_quantized_compute linear_op_ctx.input_quantizer = input_quantizer diff --git a/transformer_engine/pytorch/ops/op.py b/transformer_engine/pytorch/ops/op.py index 86bd60ed9c..f7364558ef 100644 --- a/transformer_engine/pytorch/ops/op.py +++ b/transformer_engine/pytorch/ops/op.py @@ -189,11 +189,57 @@ def __init__(self) -> None: # Objects for quantization self._fp8_metas: Optional[dict[str, dict[str, Any]]] = None self._quantizers: Optional[dict[str, list[Quantizer]]] = None + self._cpu_offloading_disabled: bool = False @property def is_fused_op(self) -> bool: return False + def disable_cpu_offloading(self, disabled: bool = True) -> None: + """Disable CPU offloading for activation tensors saved by this op. + + CPU offloading is controlled by the surrounding offload context. This + setting only opts this operation's saved activation tensors out of that + context. + """ + self._cpu_offloading_disabled = disabled + + def enable_cpu_offloading(self) -> None: + """Re-enable CPU offloading for activation tensors saved by this op.""" + self.disable_cpu_offloading(False) + + def maybe_mark_and_start_activation_offload( + self, + *tensors: Any, + start: bool = False, + ) -> None: + """Mark saved activation tensors for CPU offloading when enabled. + + If CPU offloading has been disabled for this op, mark the tensors so + the active offload context skips them. + """ + from ..cpu_offload import ( # pylint: disable=import-outside-toplevel + is_cpu_offload_enabled, + mark_activation_offload, + mark_not_offload, + start_offload, + ) + + if not is_cpu_offload_enabled(): + return + + tensors = tuple(tensor for tensor in tensors if tensor is not None) + if not tensors: + return + + if self._cpu_offloading_disabled: + mark_not_offload(*tensors) + return + + if start: + start_offload(*tensors) + mark_activation_offload(*tensors) + def num_quantizers( self, mode: str, # pylint: disable=unused-argument From 3c1178457f8ed4da273e163c9886d7c3ae463a69 Mon Sep 17 00:00:00 2001 From: hongbinl Date: Tue, 9 Jun 2026 05:42:36 -0700 Subject: [PATCH 02/12] Rename op API to activation offloading Signed-off-by: hongbinl --- tests/pytorch/test_fusible_ops.py | 6 +++--- transformer_engine/pytorch/ops/op.py | 25 ++++++++++++------------- 2 files changed, 15 insertions(+), 16 deletions(-) diff --git a/tests/pytorch/test_fusible_ops.py b/tests/pytorch/test_fusible_ops.py index dc5294e13c..0bebe7fb73 100644 --- a/tests/pytorch/test_fusible_ops.py +++ b/tests/pytorch/test_fusible_ops.py @@ -72,7 +72,7 @@ _devices: list[torch.device] = [torch.device("cpu"), torch.device("cuda")] -def test_basic_operation_cpu_offloading_control(monkeypatch): +def test_basic_operation_activation_offloading_control(monkeypatch): """BasicOperation should expose a public opt-out for activation CPU offload.""" import transformer_engine.pytorch.cpu_offload as cpu_offload @@ -102,12 +102,12 @@ def test_basic_operation_cpu_offloading_control(monkeypatch): assert calls == [("start", [tensor_id]), ("mark", [tensor_id])] calls.clear() - op.disable_cpu_offloading() + op.disable_activation_offloading() op.maybe_mark_and_start_activation_offload(tensor, start=True) assert calls == [("skip", [tensor_id])] calls.clear() - op.enable_cpu_offloading() + op.enable_activation_offloading() monkeypatch.setattr(cpu_offload, "is_cpu_offload_enabled", lambda: False) op.maybe_mark_and_start_activation_offload(tensor, start=True) assert calls == [] diff --git a/transformer_engine/pytorch/ops/op.py b/transformer_engine/pytorch/ops/op.py index f7364558ef..b98a8deda5 100644 --- a/transformer_engine/pytorch/ops/op.py +++ b/transformer_engine/pytorch/ops/op.py @@ -189,24 +189,23 @@ def __init__(self) -> None: # Objects for quantization self._fp8_metas: Optional[dict[str, dict[str, Any]]] = None self._quantizers: Optional[dict[str, list[Quantizer]]] = None - self._cpu_offloading_disabled: bool = False + self.activation_offloading: bool = True @property def is_fused_op(self) -> bool: return False - def disable_cpu_offloading(self, disabled: bool = True) -> None: - """Disable CPU offloading for activation tensors saved by this op. + def disable_activation_offloading(self, disabled: bool = True) -> None: + """Disable activation CPU offloading for tensors saved by this op. - CPU offloading is controlled by the surrounding offload context. This - setting only opts this operation's saved activation tensors out of that - context. + CPU offloading is controlled by the surrounding offload context. This setting only + opts this operation's saved activation tensors out of that context. """ - self._cpu_offloading_disabled = disabled + self.activation_offloading = not disabled - def enable_cpu_offloading(self) -> None: - """Re-enable CPU offloading for activation tensors saved by this op.""" - self.disable_cpu_offloading(False) + def enable_activation_offloading(self) -> None: + """Re-enable activation CPU offloading for tensors saved by this op.""" + self.disable_activation_offloading(False) def maybe_mark_and_start_activation_offload( self, @@ -215,8 +214,8 @@ def maybe_mark_and_start_activation_offload( ) -> None: """Mark saved activation tensors for CPU offloading when enabled. - If CPU offloading has been disabled for this op, mark the tensors so - the active offload context skips them. + If activation offloading has been disabled for this op, mark the tensors so the + active offload context skips them. """ from ..cpu_offload import ( # pylint: disable=import-outside-toplevel is_cpu_offload_enabled, @@ -232,7 +231,7 @@ def maybe_mark_and_start_activation_offload( if not tensors: return - if self._cpu_offloading_disabled: + if not self.activation_offloading: mark_not_offload(*tensors) return From 02dc7709c940651b01e6904e645725ec1746d49b Mon Sep 17 00:00:00 2001 From: hongbinl Date: Tue, 9 Jun 2026 05:53:39 -0700 Subject: [PATCH 03/12] Move CPU offload gating to TE op call sites Signed-off-by: hongbinl --- tests/pytorch/test_fusible_ops.py | 8 +++---- .../pytorch/ops/basic/activation.py | 7 ++++-- .../pytorch/ops/basic/basic_linear.py | 4 +++- .../pytorch/ops/basic/dropout.py | 4 +++- .../pytorch/ops/basic/grouped_linear.py | 24 +++++++++++-------- .../pytorch/ops/basic/l2normalization.py | 4 +++- .../pytorch/ops/basic/layer_norm.py | 4 +++- .../pytorch/ops/basic/rmsnorm.py | 4 +++- .../pytorch/ops/basic/swiglu.py | 10 +++++--- .../pytorch/ops/fused/forward_grouped_mlp.py | 9 ++++--- .../fused/forward_linear_bias_activation.py | 4 +++- .../ops/fused/forward_linear_bias_add.py | 4 +++- .../ops/fused/forward_linear_scale_add.py | 4 +++- .../ops/fused/userbuffers_forward_linear.py | 4 +++- transformer_engine/pytorch/ops/op.py | 6 +---- 15 files changed, 63 insertions(+), 37 deletions(-) diff --git a/tests/pytorch/test_fusible_ops.py b/tests/pytorch/test_fusible_ops.py index 0bebe7fb73..799834bf27 100644 --- a/tests/pytorch/test_fusible_ops.py +++ b/tests/pytorch/test_fusible_ops.py @@ -72,8 +72,8 @@ _devices: list[torch.device] = [torch.device("cpu"), torch.device("cuda")] -def test_basic_operation_activation_offloading_control(monkeypatch): - """BasicOperation should expose a public opt-out for activation CPU offload.""" +def test_basic_operation_activation_offloading_policy(monkeypatch): + """BasicOperation should expose a public opt-out for saved activation CPU offload.""" import transformer_engine.pytorch.cpu_offload as cpu_offload calls = [] @@ -81,7 +81,6 @@ def test_basic_operation_activation_offloading_control(monkeypatch): tensor_id = id(tensor) op = te_ops.Identity() - monkeypatch.setattr(cpu_offload, "is_cpu_offload_enabled", lambda: True) monkeypatch.setattr( cpu_offload, "start_offload", @@ -108,9 +107,8 @@ def test_basic_operation_activation_offloading_control(monkeypatch): calls.clear() op.enable_activation_offloading() - monkeypatch.setattr(cpu_offload, "is_cpu_offload_enabled", lambda: False) op.maybe_mark_and_start_activation_offload(tensor, start=True) - assert calls == [] + assert calls == [("start", [tensor_id]), ("mark", [tensor_id])] # Supported quantization recipes diff --git a/transformer_engine/pytorch/ops/basic/activation.py b/transformer_engine/pytorch/ops/basic/activation.py index 64a4de55e6..394473cded 100644 --- a/transformer_engine/pytorch/ops/basic/activation.py +++ b/transformer_engine/pytorch/ops/basic/activation.py @@ -13,6 +13,7 @@ import transformer_engine_torch as tex from ...constants import DType +from ...cpu_offload import is_cpu_offload_enabled from ...tensor.float8_tensor import Float8CurrentScalingQuantizer, Quantizer from ...utils import clear_tensor_data from ..op import BasicOperation, OperationContext @@ -113,7 +114,8 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - self.maybe_mark_and_start_activation_offload(x) + if is_cpu_offload_enabled(): + self.maybe_mark_and_start_activation_offload(x) ctx.save_for_backward(x) ctx.dtype = dtype ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer @@ -412,7 +414,8 @@ def fuser_forward( ctx = basic_op_ctxs[0] if ctx.requires_grad: - self.maybe_mark_and_start_activation_offload(x) + if is_cpu_offload_enabled(): + self.maybe_mark_and_start_activation_offload(x) ctx.input_requires_grad = True ctx.extra_input_requires_grad = extra_input.requires_grad ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/basic/basic_linear.py b/transformer_engine/pytorch/ops/basic/basic_linear.py index 5b59a7be2c..03814d2949 100644 --- a/transformer_engine/pytorch/ops/basic/basic_linear.py +++ b/transformer_engine/pytorch/ops/basic/basic_linear.py @@ -13,6 +13,7 @@ import torch from ...cpp_extensions import general_gemm +from ...cpu_offload import is_cpu_offload_enabled from ...distributed import ( CudaRNGStatesTracker, gather_along_first_dim, @@ -1051,7 +1052,8 @@ def op_forward( # No special CPU offloading logic is needed for weights. saved_weight is # either self.weight (nn.Parameter, auto-excluded from offload) or a # workspace freshly created each forward pass. - self.maybe_mark_and_start_activation_offload(saved_input) + if is_cpu_offload_enabled(): + self.maybe_mark_and_start_activation_offload(saved_input) ctx.save_for_backward(saved_input, saved_weight) ctx.with_quantized_compute = with_quantized_compute and backward_override is None ctx.backward_override = backward_override diff --git a/transformer_engine/pytorch/ops/basic/dropout.py b/transformer_engine/pytorch/ops/basic/dropout.py index 7a6cd9ad8f..43899b66a4 100644 --- a/transformer_engine/pytorch/ops/basic/dropout.py +++ b/transformer_engine/pytorch/ops/basic/dropout.py @@ -9,6 +9,7 @@ import torch import transformer_engine_torch as tex +from ...cpu_offload import is_cpu_offload_enabled from ...tensor import Quantizer from ...tensor.storage.float8_tensor_storage import Float8TensorStorage from .._common import maybe_autocast_dtype, maybe_dequantize @@ -70,7 +71,8 @@ def op_forward( # Save context for backward if ctx.requires_grad: - self.maybe_mark_and_start_activation_offload(mask) + if is_cpu_offload_enabled(): + self.maybe_mark_and_start_activation_offload(mask) ctx.save_for_backward(mask) ctx.impl = impl ctx.dropout_probability = self.dropout_probability diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index 3c05cdcde5..e0aa33a69d 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -16,6 +16,7 @@ import transformer_engine_torch as tex from ...constants import DType from ...cpp_extensions import general_grouped_gemm, general_grouped_gemm_for_grouped_tensor +from ...cpu_offload import is_cpu_offload_enabled from ...distributed import CudaRNGStatesTracker from ...module._common import WeightGradStore from ...module.base import ( @@ -1033,14 +1034,15 @@ def fuser_forward_save_ctx( # temporary workspaces freshly created in each forward pass. saved = tensors_to_save[0] offset = 4 if self._scale_bias else 3 - if use_grouped_tensor_path: - # Layout: [split_sizes, base_split_offsets, split_points, (scales?), grouped_x, *weights] - grouped_x = saved[offset] - self.maybe_mark_and_start_activation_offload(grouped_x) - else: - # Layout: [split_sizes, None, None, (scales?), *xs, *ws] - live_xs = [t for t in saved[offset : offset + self.num_groups] if t is not None] - self.maybe_mark_and_start_activation_offload(*live_xs) + if is_cpu_offload_enabled(): + if use_grouped_tensor_path: + # Layout: [split_sizes, base_split_offsets, split_points, (scales?), grouped_x, *weights] + grouped_x = saved[offset] + self.maybe_mark_and_start_activation_offload(grouped_x) + else: + # Layout: [split_sizes, None, None, (scales?), *xs, *ws] + live_xs = [t for t in saved[offset : offset + self.num_groups] if t is not None] + self.maybe_mark_and_start_activation_offload(*live_xs) ctx.save_for_backward(*tensors_to_save[0]) @@ -1127,7 +1129,8 @@ def _fuser_forward_split_quantize( else: xs = torch.split(x, split_sizes_int) live_xs = [t for t in xs if t is not None] - self.maybe_mark_and_start_activation_offload(*live_xs, start=True) + if is_cpu_offload_enabled(): + self.maybe_mark_and_start_activation_offload(*live_xs, start=True) # Allocate output tensor in_shape = list(input_.size()) @@ -1233,7 +1236,8 @@ def _fuser_forward_grouped_tensor( tensor_offsets=base_split_offsets * self.in_features, ) - self.maybe_mark_and_start_activation_offload(grouped_x, start=True) + if is_cpu_offload_enabled(): + self.maybe_mark_and_start_activation_offload(grouped_x, start=True) # Build the weight GroupedTensor / list. if self.single_grouped_weight: diff --git a/transformer_engine/pytorch/ops/basic/l2normalization.py b/transformer_engine/pytorch/ops/basic/l2normalization.py index 4e1e88bf20..85b2bba56e 100644 --- a/transformer_engine/pytorch/ops/basic/l2normalization.py +++ b/transformer_engine/pytorch/ops/basic/l2normalization.py @@ -10,6 +10,7 @@ import torch +from ...cpu_offload import is_cpu_offload_enabled from ...torch_version import torch_version from ...jit import ( l2normalization_fused, @@ -101,7 +102,8 @@ def op_forward( # Save state for backward pass if requires_grad: - self.maybe_mark_and_start_activation_offload(x, rsqrt_norm) + if is_cpu_offload_enabled(): + self.maybe_mark_and_start_activation_offload(x, rsqrt_norm) ctx.save_for_backward(x, rsqrt_norm) return y diff --git a/transformer_engine/pytorch/ops/basic/layer_norm.py b/transformer_engine/pytorch/ops/basic/layer_norm.py index 4b17813bce..ce3b89023f 100644 --- a/transformer_engine/pytorch/ops/basic/layer_norm.py +++ b/transformer_engine/pytorch/ops/basic/layer_norm.py @@ -14,6 +14,7 @@ from transformer_engine_torch import layernorm_bwd, layernorm_fwd from ...constants import TE_DType +from ...cpu_offload import is_cpu_offload_enabled from ...export import is_in_onnx_export_mode from ...tensor import Quantizer from ...utils import ( @@ -215,7 +216,8 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - self.maybe_mark_and_start_activation_offload(x, means, rstdevs) + if is_cpu_offload_enabled(): + self.maybe_mark_and_start_activation_offload(x, means, rstdevs) ctx.save_for_backward(x, means, rstdevs) ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/basic/rmsnorm.py b/transformer_engine/pytorch/ops/basic/rmsnorm.py index d679ceeb31..f2b0b4fb66 100644 --- a/transformer_engine/pytorch/ops/basic/rmsnorm.py +++ b/transformer_engine/pytorch/ops/basic/rmsnorm.py @@ -14,6 +14,7 @@ from transformer_engine_torch import rmsnorm_bwd, rmsnorm_fwd from ...constants import TE_DType +from ...cpu_offload import is_cpu_offload_enabled from ...export import is_in_onnx_export_mode from ...tensor import Quantizer from ...utils import ( @@ -196,7 +197,8 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - self.maybe_mark_and_start_activation_offload(x, rstdevs) + if is_cpu_offload_enabled(): + self.maybe_mark_and_start_activation_offload(x, rstdevs) ctx.save_for_backward(x, rstdevs) ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/basic/swiglu.py b/transformer_engine/pytorch/ops/basic/swiglu.py index 3bfc6cd5ad..3e579da6a0 100644 --- a/transformer_engine/pytorch/ops/basic/swiglu.py +++ b/transformer_engine/pytorch/ops/basic/swiglu.py @@ -11,6 +11,7 @@ import torch import transformer_engine_torch as tex +from ...cpu_offload import is_cpu_offload_enabled from ...constants import DType from ...tensor import Float8CurrentScalingQuantizer, Quantizer from ...utils import clear_tensor_data @@ -126,7 +127,8 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - self.maybe_mark_and_start_activation_offload(input_) + if is_cpu_offload_enabled(): + self.maybe_mark_and_start_activation_offload(input_) ctx.save_for_backward(input_) ctx.dtype = dtype ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer @@ -309,7 +311,8 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - self.maybe_mark_and_start_activation_offload(x) + if is_cpu_offload_enabled(): + self.maybe_mark_and_start_activation_offload(x) ctx.save_for_backward(x) ctx.dtype = dtype ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer @@ -459,7 +462,8 @@ def fuser_forward( # Save state for backward pass ctx = basic_op_ctxs[0] if ctx.requires_grad: - self.maybe_mark_and_start_activation_offload(input_) + if is_cpu_offload_enabled(): + self.maybe_mark_and_start_activation_offload(input_) ctx.input_requires_grad = True ctx.extra_input_requires_grad = extra_input.requires_grad ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py b/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py index 5767923925..2d8cec4683 100644 --- a/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py +++ b/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py @@ -14,6 +14,7 @@ import transformer_engine_torch as tex from ...cpp_extensions import general_gemm, general_grouped_gemm_for_grouped_tensor +from ...cpu_offload import is_cpu_offload_enabled from ...quantization import Recipe from ...tensor import NVFP4Quantizer, NVFP4Tensor, Quantizer from ...utils import ( @@ -725,6 +726,7 @@ def fuser_forward( # Save state for backward pass if requires_grad: mark_grouped_tensor(grouped_fc1_x, activation_in, scales, grouped_fc2_x) + cpu_offloading = is_cpu_offload_enabled() activation_is_srelu = isinstance(activation_op, ScaledSReLU) activation_recompute_in_mlp = bool( getattr(activation_op, "activation_recompute_in_mlp", False) @@ -746,9 +748,10 @@ def fuser_forward( grouped_fc_x.rowwise_data = None grouped_fc_x.scale_inv = None - fc1_op.maybe_mark_and_start_activation_offload(grouped_fc1_x, start=True) - activation_op.maybe_mark_and_start_activation_offload(activation_in, start=True) - fc2_op.maybe_mark_and_start_activation_offload(saved_grouped_fc2_x, start=True) + if cpu_offloading: + fc1_op.maybe_mark_and_start_activation_offload(grouped_fc1_x, start=True) + activation_op.maybe_mark_and_start_activation_offload(activation_in, start=True) + fc2_op.maybe_mark_and_start_activation_offload(saved_grouped_fc2_x, start=True) # FC1 saved-tensor layout. # [split_sizes, base_split_offsets, split_points, diff --git a/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py b/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py index b3e0cd7497..eb91326b97 100644 --- a/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py +++ b/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py @@ -10,6 +10,7 @@ import torch +from ...cpu_offload import is_cpu_offload_enabled from ...quantization import FP8GlobalStateManager from ...tensor import Quantizer from ..basic import BasicLinear, Bias @@ -128,7 +129,8 @@ def fuser_forward( else: saved_input = x_local saved_weight = w - linear_op.maybe_mark_and_start_activation_offload(saved_input) + if is_cpu_offload_enabled(): + linear_op.maybe_mark_and_start_activation_offload(saved_input) linear_op_ctx.save_for_backward(saved_input, saved_weight) linear_op_ctx.with_quantized_compute = ( with_quantized_compute and backward_override is None diff --git a/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py b/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py index 1fedb34de9..4a8e9ae634 100644 --- a/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py +++ b/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py @@ -10,6 +10,7 @@ import torch +from ...cpu_offload import is_cpu_offload_enabled from ...quantization import FP8GlobalStateManager from ...tensor import Quantizer from ..basic import AddExtraInput, BasicLinear, Bias @@ -125,7 +126,8 @@ def fuser_forward( else: saved_input = x_local saved_weight = w - linear_op.maybe_mark_and_start_activation_offload(saved_input) + if is_cpu_offload_enabled(): + linear_op.maybe_mark_and_start_activation_offload(saved_input) linear_op_ctx.save_for_backward(saved_input, saved_weight) linear_op_ctx.with_quantized_compute = ( with_quantized_compute and backward_override is None diff --git a/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py b/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py index 3f3ddb4057..74fdddd7d1 100644 --- a/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py +++ b/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py @@ -10,6 +10,7 @@ import torch +from ...cpu_offload import is_cpu_offload_enabled from ...quantization import FP8GlobalStateManager from ...tensor import Quantizer from ..basic import AddExtraInput, BasicLinear, ConstantScale @@ -106,7 +107,8 @@ def fuser_forward( else: saved_input = x_local saved_weight = w - linear_op.maybe_mark_and_start_activation_offload(saved_input) + if is_cpu_offload_enabled(): + linear_op.maybe_mark_and_start_activation_offload(saved_input) linear_op_ctx.save_for_backward(saved_input, saved_weight) linear_op_ctx.with_quantized_compute = ( with_quantized_compute and backward_override is None diff --git a/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py b/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py index 6fbf815d0d..9eb76d20dc 100644 --- a/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py +++ b/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py @@ -12,6 +12,7 @@ from transformer_engine_torch import CommOverlapType from ...cpp_extensions import general_gemm +from ...cpu_offload import is_cpu_offload_enabled from ...distributed import get_distributed_world_size from ...quantization import FP8GlobalStateManager from ...module.base import ( @@ -353,7 +354,8 @@ def fuser_forward( # Save state for backward pass if linear_op_ctx.requires_grad: - linear_op.maybe_mark_and_start_activation_offload(x_local) + if is_cpu_offload_enabled(): + linear_op.maybe_mark_and_start_activation_offload(x_local) linear_op_ctx.save_for_backward(x_local, w) linear_op_ctx.with_quantized_compute = with_quantized_compute linear_op_ctx.input_quantizer = input_quantizer diff --git a/transformer_engine/pytorch/ops/op.py b/transformer_engine/pytorch/ops/op.py index b98a8deda5..c6175d4fea 100644 --- a/transformer_engine/pytorch/ops/op.py +++ b/transformer_engine/pytorch/ops/op.py @@ -212,21 +212,17 @@ def maybe_mark_and_start_activation_offload( *tensors: Any, start: bool = False, ) -> None: - """Mark saved activation tensors for CPU offloading when enabled. + """Mark saved activation tensors for CPU offloading in an active offload context. If activation offloading has been disabled for this op, mark the tensors so the active offload context skips them. """ from ..cpu_offload import ( # pylint: disable=import-outside-toplevel - is_cpu_offload_enabled, mark_activation_offload, mark_not_offload, start_offload, ) - if not is_cpu_offload_enabled(): - return - tensors = tuple(tensor for tensor in tensors if tensor is not None) if not tensors: return From b93a9af8c8febb5c6b88be61f9714d331de4fe47 Mon Sep 17 00:00:00 2001 From: hongbinl Date: Tue, 9 Jun 2026 05:58:19 -0700 Subject: [PATCH 04/12] Preserve grouped linear offload start semantics Signed-off-by: hongbinl --- tests/pytorch/test_fusible_ops.py | 4 ++-- .../pytorch/ops/basic/grouped_linear.py | 19 +++++++++++-------- transformer_engine/pytorch/ops/op.py | 7 +++++-- 3 files changed, 18 insertions(+), 12 deletions(-) diff --git a/tests/pytorch/test_fusible_ops.py b/tests/pytorch/test_fusible_ops.py index 799834bf27..ebda9ac6ac 100644 --- a/tests/pytorch/test_fusible_ops.py +++ b/tests/pytorch/test_fusible_ops.py @@ -107,8 +107,8 @@ def test_basic_operation_activation_offloading_policy(monkeypatch): calls.clear() op.enable_activation_offloading() - op.maybe_mark_and_start_activation_offload(tensor, start=True) - assert calls == [("start", [tensor_id]), ("mark", [tensor_id])] + op.maybe_mark_and_start_activation_offload(tensor, start=True, mark=False) + assert calls == [("start", [tensor_id])] # Supported quantization recipes diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index e0aa33a69d..0292fb24d2 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -1032,17 +1032,19 @@ def fuser_forward_save_ctx( # Note: No special logic is needed for weights. They are # either nn.Parameter (auto-excluded from offload) or are # temporary workspaces freshly created in each forward pass. - saved = tensors_to_save[0] - offset = 4 if self._scale_bias else 3 if is_cpu_offload_enabled(): + saved = tensors_to_save[0] + offset = 4 if self._scale_bias else 3 if use_grouped_tensor_path: # Layout: [split_sizes, base_split_offsets, split_points, (scales?), grouped_x, *weights] grouped_x = saved[offset] - self.maybe_mark_and_start_activation_offload(grouped_x) + if grouped_x is not None: + self.maybe_mark_and_start_activation_offload(grouped_x) else: # Layout: [split_sizes, None, None, (scales?), *xs, *ws] live_xs = [t for t in saved[offset : offset + self.num_groups] if t is not None] - self.maybe_mark_and_start_activation_offload(*live_xs) + if live_xs: + self.maybe_mark_and_start_activation_offload(*live_xs) ctx.save_for_backward(*tensors_to_save[0]) @@ -1128,9 +1130,10 @@ def _fuser_forward_split_quantize( xs = tex.split_quantize(x, split_sizes_int, input_quantizers) else: xs = torch.split(x, split_sizes_int) - live_xs = [t for t in xs if t is not None] if is_cpu_offload_enabled(): - self.maybe_mark_and_start_activation_offload(*live_xs, start=True) + live_xs = [t for t in xs if t is not None] + if live_xs: + self.maybe_mark_and_start_activation_offload(*live_xs, start=True, mark=False) # Allocate output tensor in_shape = list(input_.size()) @@ -1236,8 +1239,8 @@ def _fuser_forward_grouped_tensor( tensor_offsets=base_split_offsets * self.in_features, ) - if is_cpu_offload_enabled(): - self.maybe_mark_and_start_activation_offload(grouped_x, start=True) + if is_cpu_offload_enabled() and grouped_x is not None: + self.maybe_mark_and_start_activation_offload(grouped_x, start=True, mark=False) # Build the weight GroupedTensor / list. if self.single_grouped_weight: diff --git a/transformer_engine/pytorch/ops/op.py b/transformer_engine/pytorch/ops/op.py index c6175d4fea..24aac99ff2 100644 --- a/transformer_engine/pytorch/ops/op.py +++ b/transformer_engine/pytorch/ops/op.py @@ -211,11 +211,13 @@ def maybe_mark_and_start_activation_offload( self, *tensors: Any, start: bool = False, + mark: bool = True, ) -> None: """Mark saved activation tensors for CPU offloading in an active offload context. If activation offloading has been disabled for this op, mark the tensors so the - active offload context skips them. + active offload context skips them. If mark is False, only start offload for tensors + that were already selected by the active offload context. """ from ..cpu_offload import ( # pylint: disable=import-outside-toplevel mark_activation_offload, @@ -233,7 +235,8 @@ def maybe_mark_and_start_activation_offload( if start: start_offload(*tensors) - mark_activation_offload(*tensors) + if mark: + mark_activation_offload(*tensors) def num_quantizers( self, From 0b4130259795a82e833cd43bb909db1d6e8f1387 Mon Sep 17 00:00:00 2001 From: hongbinl Date: Tue, 9 Jun 2026 06:01:27 -0700 Subject: [PATCH 05/12] Use setter for activation offload policy Signed-off-by: hongbinl --- tests/pytorch/test_fusible_ops.py | 4 ++-- transformer_engine/pytorch/ops/op.py | 12 ++++-------- 2 files changed, 6 insertions(+), 10 deletions(-) diff --git a/tests/pytorch/test_fusible_ops.py b/tests/pytorch/test_fusible_ops.py index ebda9ac6ac..fff3b3bb1e 100644 --- a/tests/pytorch/test_fusible_ops.py +++ b/tests/pytorch/test_fusible_ops.py @@ -101,12 +101,12 @@ def test_basic_operation_activation_offloading_policy(monkeypatch): assert calls == [("start", [tensor_id]), ("mark", [tensor_id])] calls.clear() - op.disable_activation_offloading() + op.set_activation_offloading(False) op.maybe_mark_and_start_activation_offload(tensor, start=True) assert calls == [("skip", [tensor_id])] calls.clear() - op.enable_activation_offloading() + op.set_activation_offloading(True) op.maybe_mark_and_start_activation_offload(tensor, start=True, mark=False) assert calls == [("start", [tensor_id])] diff --git a/transformer_engine/pytorch/ops/op.py b/transformer_engine/pytorch/ops/op.py index 24aac99ff2..4a278b8cc1 100644 --- a/transformer_engine/pytorch/ops/op.py +++ b/transformer_engine/pytorch/ops/op.py @@ -195,17 +195,13 @@ def __init__(self) -> None: def is_fused_op(self) -> bool: return False - def disable_activation_offloading(self, disabled: bool = True) -> None: - """Disable activation CPU offloading for tensors saved by this op. + def set_activation_offloading(self, enabled: bool) -> None: + """Enable or disable activation CPU offloading for tensors saved by this op. CPU offloading is controlled by the surrounding offload context. This setting only - opts this operation's saved activation tensors out of that context. + opts this operation's saved activation tensors in or out of that context. """ - self.activation_offloading = not disabled - - def enable_activation_offloading(self) -> None: - """Re-enable activation CPU offloading for tensors saved by this op.""" - self.disable_activation_offloading(False) + self.activation_offloading = enabled def maybe_mark_and_start_activation_offload( self, From 1b8178697b1f9278b04bec30a165e76cb42d665d Mon Sep 17 00:00:00 2001 From: hongbinl Date: Tue, 9 Jun 2026 06:18:15 -0700 Subject: [PATCH 06/12] Limit activation offload policy helper to marking Signed-off-by: hongbinl --- tests/pytorch/test_fusible_ops.py | 15 +++++---------- .../pytorch/ops/basic/activation.py | 4 ++-- .../pytorch/ops/basic/basic_linear.py | 2 +- transformer_engine/pytorch/ops/basic/dropout.py | 2 +- .../pytorch/ops/basic/grouped_linear.py | 10 +++++----- .../pytorch/ops/basic/l2normalization.py | 2 +- .../pytorch/ops/basic/layer_norm.py | 2 +- transformer_engine/pytorch/ops/basic/rmsnorm.py | 2 +- transformer_engine/pytorch/ops/basic/swiglu.py | 6 +++--- .../pytorch/ops/fused/forward_grouped_mlp.py | 12 ++++++++---- .../ops/fused/forward_linear_bias_activation.py | 2 +- .../pytorch/ops/fused/forward_linear_bias_add.py | 2 +- .../ops/fused/forward_linear_scale_add.py | 2 +- .../ops/fused/userbuffers_forward_linear.py | 2 +- transformer_engine/pytorch/ops/op.py | 16 +++------------- 15 files changed, 35 insertions(+), 46 deletions(-) diff --git a/tests/pytorch/test_fusible_ops.py b/tests/pytorch/test_fusible_ops.py index fff3b3bb1e..6296d9ca91 100644 --- a/tests/pytorch/test_fusible_ops.py +++ b/tests/pytorch/test_fusible_ops.py @@ -81,11 +81,6 @@ def test_basic_operation_activation_offloading_policy(monkeypatch): tensor_id = id(tensor) op = te_ops.Identity() - monkeypatch.setattr( - cpu_offload, - "start_offload", - lambda *tensors: calls.append(("start", [id(t) for t in tensors])), - ) monkeypatch.setattr( cpu_offload, "mark_activation_offload", @@ -97,18 +92,18 @@ def test_basic_operation_activation_offloading_policy(monkeypatch): lambda *tensors: calls.append(("skip", [id(t) for t in tensors])), ) - op.maybe_mark_and_start_activation_offload(tensor, None, start=True) - assert calls == [("start", [tensor_id]), ("mark", [tensor_id])] + op.maybe_mark_activation_offload(tensor, None) + assert calls == [("mark", [tensor_id])] calls.clear() op.set_activation_offloading(False) - op.maybe_mark_and_start_activation_offload(tensor, start=True) + op.maybe_mark_activation_offload(tensor) assert calls == [("skip", [tensor_id])] calls.clear() op.set_activation_offloading(True) - op.maybe_mark_and_start_activation_offload(tensor, start=True, mark=False) - assert calls == [("start", [tensor_id])] + op.maybe_mark_activation_offload(tensor) + assert calls == [("mark", [tensor_id])] # Supported quantization recipes diff --git a/transformer_engine/pytorch/ops/basic/activation.py b/transformer_engine/pytorch/ops/basic/activation.py index 394473cded..709ea9f5c7 100644 --- a/transformer_engine/pytorch/ops/basic/activation.py +++ b/transformer_engine/pytorch/ops/basic/activation.py @@ -115,7 +115,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: if is_cpu_offload_enabled(): - self.maybe_mark_and_start_activation_offload(x) + self.maybe_mark_activation_offload(x) ctx.save_for_backward(x) ctx.dtype = dtype ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer @@ -415,7 +415,7 @@ def fuser_forward( ctx = basic_op_ctxs[0] if ctx.requires_grad: if is_cpu_offload_enabled(): - self.maybe_mark_and_start_activation_offload(x) + self.maybe_mark_activation_offload(x) ctx.input_requires_grad = True ctx.extra_input_requires_grad = extra_input.requires_grad ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/basic/basic_linear.py b/transformer_engine/pytorch/ops/basic/basic_linear.py index 03814d2949..9e241a5c12 100644 --- a/transformer_engine/pytorch/ops/basic/basic_linear.py +++ b/transformer_engine/pytorch/ops/basic/basic_linear.py @@ -1053,7 +1053,7 @@ def op_forward( # either self.weight (nn.Parameter, auto-excluded from offload) or a # workspace freshly created each forward pass. if is_cpu_offload_enabled(): - self.maybe_mark_and_start_activation_offload(saved_input) + self.maybe_mark_activation_offload(saved_input) ctx.save_for_backward(saved_input, saved_weight) ctx.with_quantized_compute = with_quantized_compute and backward_override is None ctx.backward_override = backward_override diff --git a/transformer_engine/pytorch/ops/basic/dropout.py b/transformer_engine/pytorch/ops/basic/dropout.py index 43899b66a4..519d29ede2 100644 --- a/transformer_engine/pytorch/ops/basic/dropout.py +++ b/transformer_engine/pytorch/ops/basic/dropout.py @@ -72,7 +72,7 @@ def op_forward( # Save context for backward if ctx.requires_grad: if is_cpu_offload_enabled(): - self.maybe_mark_and_start_activation_offload(mask) + self.maybe_mark_activation_offload(mask) ctx.save_for_backward(mask) ctx.impl = impl ctx.dropout_probability = self.dropout_probability diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index 0292fb24d2..c7c1f33b5e 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -16,7 +16,6 @@ import transformer_engine_torch as tex from ...constants import DType from ...cpp_extensions import general_grouped_gemm, general_grouped_gemm_for_grouped_tensor -from ...cpu_offload import is_cpu_offload_enabled from ...distributed import CudaRNGStatesTracker from ...module._common import WeightGradStore from ...module.base import ( @@ -24,6 +23,7 @@ _2X_ACC_DGRAD, _2X_ACC_WGRAD, ) +from ...cpu_offload import is_cpu_offload_enabled, start_offload from ...quantization import FP8GlobalStateManager, QuantizerRole, Recipe from ...quantized_tensor import QuantizedTensorStorage from ...tensor import MXFP8Quantizer, MXFP8Tensor, Quantizer @@ -1039,12 +1039,12 @@ def fuser_forward_save_ctx( # Layout: [split_sizes, base_split_offsets, split_points, (scales?), grouped_x, *weights] grouped_x = saved[offset] if grouped_x is not None: - self.maybe_mark_and_start_activation_offload(grouped_x) + self.maybe_mark_activation_offload(grouped_x) else: # Layout: [split_sizes, None, None, (scales?), *xs, *ws] live_xs = [t for t in saved[offset : offset + self.num_groups] if t is not None] if live_xs: - self.maybe_mark_and_start_activation_offload(*live_xs) + self.maybe_mark_activation_offload(*live_xs) ctx.save_for_backward(*tensors_to_save[0]) @@ -1133,7 +1133,7 @@ def _fuser_forward_split_quantize( if is_cpu_offload_enabled(): live_xs = [t for t in xs if t is not None] if live_xs: - self.maybe_mark_and_start_activation_offload(*live_xs, start=True, mark=False) + start_offload(*live_xs) # Allocate output tensor in_shape = list(input_.size()) @@ -1240,7 +1240,7 @@ def _fuser_forward_grouped_tensor( ) if is_cpu_offload_enabled() and grouped_x is not None: - self.maybe_mark_and_start_activation_offload(grouped_x, start=True, mark=False) + start_offload(grouped_x) # Build the weight GroupedTensor / list. if self.single_grouped_weight: diff --git a/transformer_engine/pytorch/ops/basic/l2normalization.py b/transformer_engine/pytorch/ops/basic/l2normalization.py index 85b2bba56e..eee1794439 100644 --- a/transformer_engine/pytorch/ops/basic/l2normalization.py +++ b/transformer_engine/pytorch/ops/basic/l2normalization.py @@ -103,7 +103,7 @@ def op_forward( # Save state for backward pass if requires_grad: if is_cpu_offload_enabled(): - self.maybe_mark_and_start_activation_offload(x, rsqrt_norm) + self.maybe_mark_activation_offload(x, rsqrt_norm) ctx.save_for_backward(x, rsqrt_norm) return y diff --git a/transformer_engine/pytorch/ops/basic/layer_norm.py b/transformer_engine/pytorch/ops/basic/layer_norm.py index ce3b89023f..c1a8255132 100644 --- a/transformer_engine/pytorch/ops/basic/layer_norm.py +++ b/transformer_engine/pytorch/ops/basic/layer_norm.py @@ -217,7 +217,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: if is_cpu_offload_enabled(): - self.maybe_mark_and_start_activation_offload(x, means, rstdevs) + self.maybe_mark_activation_offload(x, means, rstdevs) ctx.save_for_backward(x, means, rstdevs) ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/basic/rmsnorm.py b/transformer_engine/pytorch/ops/basic/rmsnorm.py index f2b0b4fb66..a9079b6159 100644 --- a/transformer_engine/pytorch/ops/basic/rmsnorm.py +++ b/transformer_engine/pytorch/ops/basic/rmsnorm.py @@ -198,7 +198,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: if is_cpu_offload_enabled(): - self.maybe_mark_and_start_activation_offload(x, rstdevs) + self.maybe_mark_activation_offload(x, rstdevs) ctx.save_for_backward(x, rstdevs) ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/basic/swiglu.py b/transformer_engine/pytorch/ops/basic/swiglu.py index 3e579da6a0..72c0286fff 100644 --- a/transformer_engine/pytorch/ops/basic/swiglu.py +++ b/transformer_engine/pytorch/ops/basic/swiglu.py @@ -128,7 +128,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: if is_cpu_offload_enabled(): - self.maybe_mark_and_start_activation_offload(input_) + self.maybe_mark_activation_offload(input_) ctx.save_for_backward(input_) ctx.dtype = dtype ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer @@ -312,7 +312,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: if is_cpu_offload_enabled(): - self.maybe_mark_and_start_activation_offload(x) + self.maybe_mark_activation_offload(x) ctx.save_for_backward(x) ctx.dtype = dtype ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer @@ -463,7 +463,7 @@ def fuser_forward( ctx = basic_op_ctxs[0] if ctx.requires_grad: if is_cpu_offload_enabled(): - self.maybe_mark_and_start_activation_offload(input_) + self.maybe_mark_activation_offload(input_) ctx.input_requires_grad = True ctx.extra_input_requires_grad = extra_input.requires_grad ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py b/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py index 2d8cec4683..fa7d56a78d 100644 --- a/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py +++ b/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py @@ -13,8 +13,8 @@ import torch import transformer_engine_torch as tex +from ...cpu_offload import is_cpu_offload_enabled, start_offload from ...cpp_extensions import general_gemm, general_grouped_gemm_for_grouped_tensor -from ...cpu_offload import is_cpu_offload_enabled from ...quantization import Recipe from ...tensor import NVFP4Quantizer, NVFP4Tensor, Quantizer from ...utils import ( @@ -749,9 +749,13 @@ def fuser_forward( grouped_fc_x.scale_inv = None if cpu_offloading: - fc1_op.maybe_mark_and_start_activation_offload(grouped_fc1_x, start=True) - activation_op.maybe_mark_and_start_activation_offload(activation_in, start=True) - fc2_op.maybe_mark_and_start_activation_offload(saved_grouped_fc2_x, start=True) + activation_tensors = [ + t for t in (grouped_fc1_x, activation_in, saved_grouped_fc2_x) if t is not None + ] + start_offload(*activation_tensors) + fc1_op.maybe_mark_activation_offload(grouped_fc1_x) + activation_op.maybe_mark_activation_offload(activation_in) + fc2_op.maybe_mark_activation_offload(saved_grouped_fc2_x) # FC1 saved-tensor layout. # [split_sizes, base_split_offsets, split_points, diff --git a/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py b/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py index eb91326b97..8e5c5f841b 100644 --- a/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py +++ b/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py @@ -130,7 +130,7 @@ def fuser_forward( saved_input = x_local saved_weight = w if is_cpu_offload_enabled(): - linear_op.maybe_mark_and_start_activation_offload(saved_input) + linear_op.maybe_mark_activation_offload(saved_input) linear_op_ctx.save_for_backward(saved_input, saved_weight) linear_op_ctx.with_quantized_compute = ( with_quantized_compute and backward_override is None diff --git a/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py b/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py index 4a8e9ae634..d2e085e41b 100644 --- a/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py +++ b/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py @@ -127,7 +127,7 @@ def fuser_forward( saved_input = x_local saved_weight = w if is_cpu_offload_enabled(): - linear_op.maybe_mark_and_start_activation_offload(saved_input) + linear_op.maybe_mark_activation_offload(saved_input) linear_op_ctx.save_for_backward(saved_input, saved_weight) linear_op_ctx.with_quantized_compute = ( with_quantized_compute and backward_override is None diff --git a/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py b/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py index 74fdddd7d1..8852a60e92 100644 --- a/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py +++ b/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py @@ -108,7 +108,7 @@ def fuser_forward( saved_input = x_local saved_weight = w if is_cpu_offload_enabled(): - linear_op.maybe_mark_and_start_activation_offload(saved_input) + linear_op.maybe_mark_activation_offload(saved_input) linear_op_ctx.save_for_backward(saved_input, saved_weight) linear_op_ctx.with_quantized_compute = ( with_quantized_compute and backward_override is None diff --git a/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py b/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py index 9eb76d20dc..07d66a3b3a 100644 --- a/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py +++ b/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py @@ -355,7 +355,7 @@ def fuser_forward( # Save state for backward pass if linear_op_ctx.requires_grad: if is_cpu_offload_enabled(): - linear_op.maybe_mark_and_start_activation_offload(x_local) + linear_op.maybe_mark_activation_offload(x_local) linear_op_ctx.save_for_backward(x_local, w) linear_op_ctx.with_quantized_compute = with_quantized_compute linear_op_ctx.input_quantizer = input_quantizer diff --git a/transformer_engine/pytorch/ops/op.py b/transformer_engine/pytorch/ops/op.py index 4a278b8cc1..615f269572 100644 --- a/transformer_engine/pytorch/ops/op.py +++ b/transformer_engine/pytorch/ops/op.py @@ -203,22 +203,15 @@ def set_activation_offloading(self, enabled: bool) -> None: """ self.activation_offloading = enabled - def maybe_mark_and_start_activation_offload( - self, - *tensors: Any, - start: bool = False, - mark: bool = True, - ) -> None: + def maybe_mark_activation_offload(self, *tensors: Any) -> None: """Mark saved activation tensors for CPU offloading in an active offload context. If activation offloading has been disabled for this op, mark the tensors so the - active offload context skips them. If mark is False, only start offload for tensors - that were already selected by the active offload context. + active offload context skips them. """ from ..cpu_offload import ( # pylint: disable=import-outside-toplevel mark_activation_offload, mark_not_offload, - start_offload, ) tensors = tuple(tensor for tensor in tensors if tensor is not None) @@ -229,10 +222,7 @@ def maybe_mark_and_start_activation_offload( mark_not_offload(*tensors) return - if start: - start_offload(*tensors) - if mark: - mark_activation_offload(*tensors) + mark_activation_offload(*tensors) def num_quantizers( self, From df5e715e6f6729bcede61ffe586e45582de9b084 Mon Sep 17 00:00:00 2001 From: hongbinl Date: Tue, 9 Jun 2026 06:21:13 -0700 Subject: [PATCH 07/12] Move CPU offload imports to op module scope Signed-off-by: hongbinl --- transformer_engine/pytorch/ops/op.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/transformer_engine/pytorch/ops/op.py b/transformer_engine/pytorch/ops/op.py index 615f269572..9e2c442ff5 100644 --- a/transformer_engine/pytorch/ops/op.py +++ b/transformer_engine/pytorch/ops/op.py @@ -14,6 +14,7 @@ import torch from transformer_engine.common.recipe import Recipe +from ..cpu_offload import mark_activation_offload, mark_not_offload from ..quantization import ( FP8GlobalStateManager, QuantizerRole, @@ -209,11 +210,6 @@ def maybe_mark_activation_offload(self, *tensors: Any) -> None: If activation offloading has been disabled for this op, mark the tensors so the active offload context skips them. """ - from ..cpu_offload import ( # pylint: disable=import-outside-toplevel - mark_activation_offload, - mark_not_offload, - ) - tensors = tuple(tensor for tensor in tensors if tensor is not None) if not tensors: return From f26534d8b2c88f95b92c7ff82e64832f555b3276 Mon Sep 17 00:00:00 2001 From: hongbinl Date: Tue, 9 Jun 2026 06:29:28 -0700 Subject: [PATCH 08/12] Patch activation offload test bound symbols Signed-off-by: hongbinl --- tests/pytorch/test_fusible_ops.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/pytorch/test_fusible_ops.py b/tests/pytorch/test_fusible_ops.py index 6296d9ca91..b97bccd1b2 100644 --- a/tests/pytorch/test_fusible_ops.py +++ b/tests/pytorch/test_fusible_ops.py @@ -74,7 +74,7 @@ def test_basic_operation_activation_offloading_policy(monkeypatch): """BasicOperation should expose a public opt-out for saved activation CPU offload.""" - import transformer_engine.pytorch.cpu_offload as cpu_offload + import transformer_engine.pytorch.ops.op as op_module calls = [] tensor = torch.empty(1) @@ -82,12 +82,12 @@ def test_basic_operation_activation_offloading_policy(monkeypatch): op = te_ops.Identity() monkeypatch.setattr( - cpu_offload, + op_module, "mark_activation_offload", lambda *tensors: calls.append(("mark", [id(t) for t in tensors])), ) monkeypatch.setattr( - cpu_offload, + op_module, "mark_not_offload", lambda *tensors: calls.append(("skip", [id(t) for t in tensors])), ) From 7d9128cf214e0dbbb1758c5683e25fc34e24e2c6 Mon Sep 17 00:00:00 2001 From: Tim Moon Date: Thu, 11 Jun 2026 00:36:26 +0000 Subject: [PATCH 09/12] Refactor base class offloading infrastructure Handle inclusion and exclusion in same function. Check whether CPU offloading is enabled internally. Tweak documentation and style. Signed-off-by: Tim Moon --- transformer_engine/pytorch/ops/op.py | 118 +++++++++++++++++++++------ 1 file changed, 91 insertions(+), 27 deletions(-) diff --git a/transformer_engine/pytorch/ops/op.py b/transformer_engine/pytorch/ops/op.py index 9e2c442ff5..a506192284 100644 --- a/transformer_engine/pytorch/ops/op.py +++ b/transformer_engine/pytorch/ops/op.py @@ -14,14 +14,22 @@ import torch from transformer_engine.common.recipe import Recipe -from ..cpu_offload import mark_activation_offload, mark_not_offload +from ..cpu_offload import is_cpu_offload_enabled, mark_activation_offload, mark_not_offload from ..quantization import ( FP8GlobalStateManager, QuantizerRole, RecipeState, autocast, ) -from ..tensor import Quantizer +from ..tensor import ( + GroupedTensorStorage, + QuantizedTensorStorage, + Quantizer, +) + + +# Tensor class supported by fusible operation +TensorLike = torch.Tensor | QuantizedTensorStorage | GroupedTensorStorage @dataclasses.dataclass @@ -190,36 +198,14 @@ def __init__(self) -> None: # Objects for quantization self._fp8_metas: Optional[dict[str, dict[str, Any]]] = None self._quantizers: Optional[dict[str, list[Quantizer]]] = None - self.activation_offloading: bool = True + + # Whether to participate when activation CPU offloading is enabled + self._activation_offloading_enabled: bool = True @property def is_fused_op(self) -> bool: return False - def set_activation_offloading(self, enabled: bool) -> None: - """Enable or disable activation CPU offloading for tensors saved by this op. - - CPU offloading is controlled by the surrounding offload context. This setting only - opts this operation's saved activation tensors in or out of that context. - """ - self.activation_offloading = enabled - - def maybe_mark_activation_offload(self, *tensors: Any) -> None: - """Mark saved activation tensors for CPU offloading in an active offload context. - - If activation offloading has been disabled for this op, mark the tensors so the - active offload context skips them. - """ - tensors = tuple(tensor for tensor in tensors if tensor is not None) - if not tensors: - return - - if not self.activation_offloading: - mark_not_offload(*tensors) - return - - mark_activation_offload(*tensors) - def num_quantizers( self, mode: str, # pylint: disable=unused-argument @@ -455,6 +441,72 @@ def _load_fp8_metas(self, fp8_metas: Optional[dict[str, Any]]) -> None: self._fp8_metas[mode][fp8_meta_key].scale.copy_(scale) self._fp8_metas[mode][fp8_meta_key].amax_history.copy_(amax_history) + def set_activation_offloading(self, enabled: bool) -> None: + """Set whether to participate when activation CPU offloading is enabled. + + Offloading is controlled by an offloading context (see + ``get_cpu_offload_context``). Disabling this setting allows + this operation to opt out when the offloading context is + active, but enabling does not activate offloading outside that + context. + """ + self._activation_offloading_enabled = enabled + + def mark_for_cpu_offload_if_needed( + self, + *tensors: TensorLike | Iterable[Optional[TensorLike]] | None, + exclude_tensors: TensorLike | Iterable[Optional[TensorLike]] | None = None, + ) -> None: + """Mark tensors to include and exclude from activation CPU offloading. + + Call in op forward implementation in order to control what + tensors participate in CPU offloading. It does nothing if + offloading is not enabled (see ``get_cpu_offload_context``), + and it excludes all tensors if the op is not participating in + offloading. + """ + if not tensors and not exclude_tensors: + return + if not is_cpu_offload_enabled(): + return + + supported_classes = (torch.Tensor, QuantizedTensorStorage, GroupedTensorStorage) + + def filter_supported_and_extend( + out: list[TensorLike], + ts: TensorLike | Iterable[Optional[TensorLike]] | None, + ) -> None: + """Extend a list with objects that support CPU offloading. + + Filters out ``None`` and checks for unsupported classes. + """ + if ts is None: + return + if isinstance(ts, supported_classes): + ts = (ts,) + for t in ts: + if t is None: + continue + if not isinstance(t, supported_classes): + raise TypeError(f"{t.__class__.__name__} does not support CPU offloading.") + out.append(t) + + # Choose tensors to include and exclude from CPU offloading + include = [] + exclude = [] + if self._activation_offloading_enabled: + filter_supported_and_extend(include, tensors) + filter_supported_and_extend(exclude, exclude_tensors) + else: + filter_supported_and_extend(exclude, tensors) + filter_supported_and_extend(exclude, exclude_tensors) + + # Mark tensors + if include: + mark_activation_offload(*include) + if exclude: + mark_not_offload(*exclude) + @abc.abstractmethod def op_forward( self, @@ -774,6 +826,18 @@ def get_input_quantizer(self) -> Optional[Quantizer]: def get_grad_output_quantizer(self) -> Optional[Quantizer]: return self.basic_ops[-1].get_grad_output_quantizer() + def set_activation_offloading(self, enabled: bool) -> None: + """Whether to participate when activation CPU offloading is enabled globally. + + Offloading is controlled by an offloading context (see + ``get_cpu_offload_context``). Disabling this setting allows + this operation to opt out when the offloading context is + active, but enabling does not activate offloading outside that + context. + """ + for op in self.basic_ops: + op.set_activation_offloading(enabled) + def pre_first_fuser_forward(self) -> None: for op in self.basic_ops: op.pre_first_fuser_forward() From 9c3d262a1aed21487bba2bfa21c2a0210686a196 Mon Sep 17 00:00:00 2001 From: Tim Moon Date: Thu, 11 Jun 2026 00:59:28 +0000 Subject: [PATCH 10/12] Propagate activation offload policy helper Use BasicOperation.mark_for_cpu_offload_if_needed at op call sites and keep explicit offload synchronization checks where needed. Co-authored-by: OpenAI Codex Signed-off-by: Tim Moon --- tests/pytorch/test_fusible_ops.py | 7 ++++--- .../pytorch/ops/basic/activation.py | 7 ++----- .../pytorch/ops/basic/basic_linear.py | 9 ++++---- .../pytorch/ops/basic/dropout.py | 4 +--- .../pytorch/ops/basic/grouped_linear.py | 21 +++++++------------ .../pytorch/ops/basic/l2normalization.py | 4 +--- .../pytorch/ops/basic/layer_norm.py | 4 +--- .../pytorch/ops/basic/rmsnorm.py | 4 +--- .../pytorch/ops/basic/swiglu.py | 10 +++------ .../pytorch/ops/fused/forward_grouped_mlp.py | 6 +++--- .../fused/forward_linear_bias_activation.py | 4 +--- .../ops/fused/forward_linear_bias_add.py | 4 +--- .../ops/fused/forward_linear_scale_add.py | 4 +--- .../ops/fused/userbuffers_forward_linear.py | 4 +--- 14 files changed, 33 insertions(+), 59 deletions(-) diff --git a/tests/pytorch/test_fusible_ops.py b/tests/pytorch/test_fusible_ops.py index b97bccd1b2..fccfe8ffb3 100644 --- a/tests/pytorch/test_fusible_ops.py +++ b/tests/pytorch/test_fusible_ops.py @@ -91,18 +91,19 @@ def test_basic_operation_activation_offloading_policy(monkeypatch): "mark_not_offload", lambda *tensors: calls.append(("skip", [id(t) for t in tensors])), ) + monkeypatch.setattr(op_module, "is_cpu_offload_enabled", lambda: True) - op.maybe_mark_activation_offload(tensor, None) + op.mark_for_cpu_offload_if_needed(tensor, None) assert calls == [("mark", [tensor_id])] calls.clear() op.set_activation_offloading(False) - op.maybe_mark_activation_offload(tensor) + op.mark_for_cpu_offload_if_needed(tensor) assert calls == [("skip", [tensor_id])] calls.clear() op.set_activation_offloading(True) - op.maybe_mark_activation_offload(tensor) + op.mark_for_cpu_offload_if_needed(tensor) assert calls == [("mark", [tensor_id])] diff --git a/transformer_engine/pytorch/ops/basic/activation.py b/transformer_engine/pytorch/ops/basic/activation.py index 709ea9f5c7..8d458331e5 100644 --- a/transformer_engine/pytorch/ops/basic/activation.py +++ b/transformer_engine/pytorch/ops/basic/activation.py @@ -13,7 +13,6 @@ import transformer_engine_torch as tex from ...constants import DType -from ...cpu_offload import is_cpu_offload_enabled from ...tensor.float8_tensor import Float8CurrentScalingQuantizer, Quantizer from ...utils import clear_tensor_data from ..op import BasicOperation, OperationContext @@ -114,8 +113,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - if is_cpu_offload_enabled(): - self.maybe_mark_activation_offload(x) + self.mark_for_cpu_offload_if_needed(x) ctx.save_for_backward(x) ctx.dtype = dtype ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer @@ -414,8 +412,7 @@ def fuser_forward( ctx = basic_op_ctxs[0] if ctx.requires_grad: - if is_cpu_offload_enabled(): - self.maybe_mark_activation_offload(x) + self.mark_for_cpu_offload_if_needed(x) ctx.input_requires_grad = True ctx.extra_input_requires_grad = extra_input.requires_grad ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/basic/basic_linear.py b/transformer_engine/pytorch/ops/basic/basic_linear.py index 9e241a5c12..0b9296b7ef 100644 --- a/transformer_engine/pytorch/ops/basic/basic_linear.py +++ b/transformer_engine/pytorch/ops/basic/basic_linear.py @@ -13,7 +13,6 @@ import torch from ...cpp_extensions import general_gemm -from ...cpu_offload import is_cpu_offload_enabled from ...distributed import ( CudaRNGStatesTracker, gather_along_first_dim, @@ -1049,11 +1048,13 @@ def op_forward( else: saved_input = x_local saved_weight = w - # No special CPU offloading logic is needed for weights. saved_weight is + + # Activation CPU offloading + # Note: No special CPU offloading logic is needed for weights. saved_weight is # either self.weight (nn.Parameter, auto-excluded from offload) or a # workspace freshly created each forward pass. - if is_cpu_offload_enabled(): - self.maybe_mark_activation_offload(saved_input) + self.mark_for_cpu_offload_if_needed(saved_input) + ctx.save_for_backward(saved_input, saved_weight) ctx.with_quantized_compute = with_quantized_compute and backward_override is None ctx.backward_override = backward_override diff --git a/transformer_engine/pytorch/ops/basic/dropout.py b/transformer_engine/pytorch/ops/basic/dropout.py index 519d29ede2..aa95a1d363 100644 --- a/transformer_engine/pytorch/ops/basic/dropout.py +++ b/transformer_engine/pytorch/ops/basic/dropout.py @@ -9,7 +9,6 @@ import torch import transformer_engine_torch as tex -from ...cpu_offload import is_cpu_offload_enabled from ...tensor import Quantizer from ...tensor.storage.float8_tensor_storage import Float8TensorStorage from .._common import maybe_autocast_dtype, maybe_dequantize @@ -71,8 +70,7 @@ def op_forward( # Save context for backward if ctx.requires_grad: - if is_cpu_offload_enabled(): - self.maybe_mark_activation_offload(mask) + self.mark_for_cpu_offload_if_needed(mask) ctx.save_for_backward(mask) ctx.impl = impl ctx.dropout_probability = self.dropout_probability diff --git a/transformer_engine/pytorch/ops/basic/grouped_linear.py b/transformer_engine/pytorch/ops/basic/grouped_linear.py index c7c1f33b5e..7b5b4a263a 100644 --- a/transformer_engine/pytorch/ops/basic/grouped_linear.py +++ b/transformer_engine/pytorch/ops/basic/grouped_linear.py @@ -1032,19 +1032,14 @@ def fuser_forward_save_ctx( # Note: No special logic is needed for weights. They are # either nn.Parameter (auto-excluded from offload) or are # temporary workspaces freshly created in each forward pass. - if is_cpu_offload_enabled(): - saved = tensors_to_save[0] - offset = 4 if self._scale_bias else 3 - if use_grouped_tensor_path: - # Layout: [split_sizes, base_split_offsets, split_points, (scales?), grouped_x, *weights] - grouped_x = saved[offset] - if grouped_x is not None: - self.maybe_mark_activation_offload(grouped_x) - else: - # Layout: [split_sizes, None, None, (scales?), *xs, *ws] - live_xs = [t for t in saved[offset : offset + self.num_groups] if t is not None] - if live_xs: - self.maybe_mark_activation_offload(*live_xs) + saved = tensors_to_save[0] + offset = 4 if self._scale_bias else 3 + if use_grouped_tensor_path: + # Layout: [split_sizes, base_split_offsets, split_points, (scales?), grouped_x, *weights] + self.mark_for_cpu_offload_if_needed(saved[offset]) + else: + # Layout: [split_sizes, None, None, (scales?), *xs, *ws] + self.mark_for_cpu_offload_if_needed(saved[offset : offset + self.num_groups]) ctx.save_for_backward(*tensors_to_save[0]) diff --git a/transformer_engine/pytorch/ops/basic/l2normalization.py b/transformer_engine/pytorch/ops/basic/l2normalization.py index eee1794439..fb9a81180a 100644 --- a/transformer_engine/pytorch/ops/basic/l2normalization.py +++ b/transformer_engine/pytorch/ops/basic/l2normalization.py @@ -10,7 +10,6 @@ import torch -from ...cpu_offload import is_cpu_offload_enabled from ...torch_version import torch_version from ...jit import ( l2normalization_fused, @@ -102,8 +101,7 @@ def op_forward( # Save state for backward pass if requires_grad: - if is_cpu_offload_enabled(): - self.maybe_mark_activation_offload(x, rsqrt_norm) + self.mark_for_cpu_offload_if_needed(x, rsqrt_norm) ctx.save_for_backward(x, rsqrt_norm) return y diff --git a/transformer_engine/pytorch/ops/basic/layer_norm.py b/transformer_engine/pytorch/ops/basic/layer_norm.py index c1a8255132..2e00740967 100644 --- a/transformer_engine/pytorch/ops/basic/layer_norm.py +++ b/transformer_engine/pytorch/ops/basic/layer_norm.py @@ -14,7 +14,6 @@ from transformer_engine_torch import layernorm_bwd, layernorm_fwd from ...constants import TE_DType -from ...cpu_offload import is_cpu_offload_enabled from ...export import is_in_onnx_export_mode from ...tensor import Quantizer from ...utils import ( @@ -216,8 +215,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - if is_cpu_offload_enabled(): - self.maybe_mark_activation_offload(x, means, rstdevs) + self.mark_for_cpu_offload_if_needed(x, means, rstdevs) ctx.save_for_backward(x, means, rstdevs) ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/basic/rmsnorm.py b/transformer_engine/pytorch/ops/basic/rmsnorm.py index a9079b6159..4df086a4e1 100644 --- a/transformer_engine/pytorch/ops/basic/rmsnorm.py +++ b/transformer_engine/pytorch/ops/basic/rmsnorm.py @@ -14,7 +14,6 @@ from transformer_engine_torch import rmsnorm_bwd, rmsnorm_fwd from ...constants import TE_DType -from ...cpu_offload import is_cpu_offload_enabled from ...export import is_in_onnx_export_mode from ...tensor import Quantizer from ...utils import ( @@ -197,8 +196,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - if is_cpu_offload_enabled(): - self.maybe_mark_activation_offload(x, rstdevs) + self.mark_for_cpu_offload_if_needed(x, rstdevs) ctx.save_for_backward(x, rstdevs) ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/basic/swiglu.py b/transformer_engine/pytorch/ops/basic/swiglu.py index 72c0286fff..131786a9c9 100644 --- a/transformer_engine/pytorch/ops/basic/swiglu.py +++ b/transformer_engine/pytorch/ops/basic/swiglu.py @@ -11,7 +11,6 @@ import torch import transformer_engine_torch as tex -from ...cpu_offload import is_cpu_offload_enabled from ...constants import DType from ...tensor import Float8CurrentScalingQuantizer, Quantizer from ...utils import clear_tensor_data @@ -127,8 +126,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - if is_cpu_offload_enabled(): - self.maybe_mark_activation_offload(input_) + self.mark_for_cpu_offload_if_needed(input_) ctx.save_for_backward(input_) ctx.dtype = dtype ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer @@ -311,8 +309,7 @@ def op_forward( # Save state for backward pass if ctx.requires_grad: - if is_cpu_offload_enabled(): - self.maybe_mark_activation_offload(x) + self.mark_for_cpu_offload_if_needed(x) ctx.save_for_backward(x) ctx.dtype = dtype ctx.prev_op_grad_output_quantizer = prev_op_grad_output_quantizer @@ -462,8 +459,7 @@ def fuser_forward( # Save state for backward pass ctx = basic_op_ctxs[0] if ctx.requires_grad: - if is_cpu_offload_enabled(): - self.maybe_mark_activation_offload(input_) + self.mark_for_cpu_offload_if_needed(input_) ctx.input_requires_grad = True ctx.extra_input_requires_grad = extra_input.requires_grad ctx.dtype = dtype diff --git a/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py b/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py index fa7d56a78d..1df4bf164e 100644 --- a/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py +++ b/transformer_engine/pytorch/ops/fused/forward_grouped_mlp.py @@ -753,9 +753,9 @@ def fuser_forward( t for t in (grouped_fc1_x, activation_in, saved_grouped_fc2_x) if t is not None ] start_offload(*activation_tensors) - fc1_op.maybe_mark_activation_offload(grouped_fc1_x) - activation_op.maybe_mark_activation_offload(activation_in) - fc2_op.maybe_mark_activation_offload(saved_grouped_fc2_x) + fc1_op.mark_for_cpu_offload_if_needed(grouped_fc1_x) + activation_op.mark_for_cpu_offload_if_needed(activation_in) + fc2_op.mark_for_cpu_offload_if_needed(saved_grouped_fc2_x) # FC1 saved-tensor layout. # [split_sizes, base_split_offsets, split_points, diff --git a/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py b/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py index 8e5c5f841b..b0911b8ae1 100644 --- a/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py +++ b/transformer_engine/pytorch/ops/fused/forward_linear_bias_activation.py @@ -10,7 +10,6 @@ import torch -from ...cpu_offload import is_cpu_offload_enabled from ...quantization import FP8GlobalStateManager from ...tensor import Quantizer from ..basic import BasicLinear, Bias @@ -129,8 +128,7 @@ def fuser_forward( else: saved_input = x_local saved_weight = w - if is_cpu_offload_enabled(): - linear_op.maybe_mark_activation_offload(saved_input) + linear_op.mark_for_cpu_offload_if_needed(saved_input) linear_op_ctx.save_for_backward(saved_input, saved_weight) linear_op_ctx.with_quantized_compute = ( with_quantized_compute and backward_override is None diff --git a/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py b/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py index d2e085e41b..5aed52546b 100644 --- a/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py +++ b/transformer_engine/pytorch/ops/fused/forward_linear_bias_add.py @@ -10,7 +10,6 @@ import torch -from ...cpu_offload import is_cpu_offload_enabled from ...quantization import FP8GlobalStateManager from ...tensor import Quantizer from ..basic import AddExtraInput, BasicLinear, Bias @@ -126,8 +125,7 @@ def fuser_forward( else: saved_input = x_local saved_weight = w - if is_cpu_offload_enabled(): - linear_op.maybe_mark_activation_offload(saved_input) + linear_op.mark_for_cpu_offload_if_needed(saved_input) linear_op_ctx.save_for_backward(saved_input, saved_weight) linear_op_ctx.with_quantized_compute = ( with_quantized_compute and backward_override is None diff --git a/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py b/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py index 8852a60e92..286dc58e8e 100644 --- a/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py +++ b/transformer_engine/pytorch/ops/fused/forward_linear_scale_add.py @@ -10,7 +10,6 @@ import torch -from ...cpu_offload import is_cpu_offload_enabled from ...quantization import FP8GlobalStateManager from ...tensor import Quantizer from ..basic import AddExtraInput, BasicLinear, ConstantScale @@ -107,8 +106,7 @@ def fuser_forward( else: saved_input = x_local saved_weight = w - if is_cpu_offload_enabled(): - linear_op.maybe_mark_activation_offload(saved_input) + linear_op.mark_for_cpu_offload_if_needed(saved_input) linear_op_ctx.save_for_backward(saved_input, saved_weight) linear_op_ctx.with_quantized_compute = ( with_quantized_compute and backward_override is None diff --git a/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py b/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py index 07d66a3b3a..291fdb9773 100644 --- a/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py +++ b/transformer_engine/pytorch/ops/fused/userbuffers_forward_linear.py @@ -12,7 +12,6 @@ from transformer_engine_torch import CommOverlapType from ...cpp_extensions import general_gemm -from ...cpu_offload import is_cpu_offload_enabled from ...distributed import get_distributed_world_size from ...quantization import FP8GlobalStateManager from ...module.base import ( @@ -354,8 +353,7 @@ def fuser_forward( # Save state for backward pass if linear_op_ctx.requires_grad: - if is_cpu_offload_enabled(): - linear_op.maybe_mark_activation_offload(x_local) + linear_op.mark_for_cpu_offload_if_needed(x_local) linear_op_ctx.save_for_backward(x_local, w) linear_op_ctx.with_quantized_compute = with_quantized_compute linear_op_ctx.input_quantizer = input_quantizer From 62b3da9c7335c13392590299d6e02f2401e64ac6 Mon Sep 17 00:00:00 2001 From: Tim Moon Date: Thu, 11 Jun 2026 01:21:10 +0000 Subject: [PATCH 11/12] Move test into TestFuser suite Signed-off-by: Tim Moon --- tests/pytorch/test_fusible_ops.py | 70 +++++++++++++++---------------- 1 file changed, 34 insertions(+), 36 deletions(-) diff --git a/tests/pytorch/test_fusible_ops.py b/tests/pytorch/test_fusible_ops.py index fccfe8ffb3..2f7e19792e 100644 --- a/tests/pytorch/test_fusible_ops.py +++ b/tests/pytorch/test_fusible_ops.py @@ -71,42 +71,6 @@ # Supported devices _devices: list[torch.device] = [torch.device("cpu"), torch.device("cuda")] - -def test_basic_operation_activation_offloading_policy(monkeypatch): - """BasicOperation should expose a public opt-out for saved activation CPU offload.""" - import transformer_engine.pytorch.ops.op as op_module - - calls = [] - tensor = torch.empty(1) - tensor_id = id(tensor) - op = te_ops.Identity() - - monkeypatch.setattr( - op_module, - "mark_activation_offload", - lambda *tensors: calls.append(("mark", [id(t) for t in tensors])), - ) - monkeypatch.setattr( - op_module, - "mark_not_offload", - lambda *tensors: calls.append(("skip", [id(t) for t in tensors])), - ) - monkeypatch.setattr(op_module, "is_cpu_offload_enabled", lambda: True) - - op.mark_for_cpu_offload_if_needed(tensor, None) - assert calls == [("mark", [tensor_id])] - - calls.clear() - op.set_activation_offloading(False) - op.mark_for_cpu_offload_if_needed(tensor) - assert calls == [("skip", [tensor_id])] - - calls.clear() - op.set_activation_offloading(True) - op.mark_for_cpu_offload_if_needed(tensor) - assert calls == [("mark", [tensor_id])] - - # Supported quantization recipes _quantization_list: list[Optional[str]] = [None] if fp8_available: @@ -670,6 +634,40 @@ def test_pyt_autocast( assert x.grad.dtype == model_dtype assert op.weight.grad.dtype == model_dtype + def test_activation_offloading_policy(self, monkeypatch): + """Test opt-out API for activation CPU offloading.""" + import transformer_engine.pytorch.ops.op as op_module + + calls = [] + tensor = torch.empty(1) + tensor_id = id(tensor) + op = te_ops.Identity() + + monkeypatch.setattr( + op_module, + "mark_activation_offload", + lambda *tensors: calls.append(("mark", [id(t) for t in tensors])), + ) + monkeypatch.setattr( + op_module, + "mark_not_offload", + lambda *tensors: calls.append(("skip", [id(t) for t in tensors])), + ) + monkeypatch.setattr(op_module, "is_cpu_offload_enabled", lambda: True) + + op.mark_for_cpu_offload_if_needed(tensor, None) + assert calls == [("mark", [tensor_id])] + + calls.clear() + op.set_activation_offloading(False) + op.mark_for_cpu_offload_if_needed(tensor) + assert calls == [("skip", [tensor_id])] + + calls.clear() + op.set_activation_offloading(True) + op.mark_for_cpu_offload_if_needed(tensor) + assert calls == [("mark", [tensor_id])] + class TestBasicOps: """Tests for individual operations""" From d858ba244c7aa41f386b02990353f97497216774 Mon Sep 17 00:00:00 2001 From: Tim Moon Date: Thu, 11 Jun 2026 01:36:48 +0000 Subject: [PATCH 12/12] Debug failure with grouped linear Signed-off-by: Tim Moon --- transformer_engine/pytorch/ops/op.py | 33 ++++++++++++++-------------- 1 file changed, 17 insertions(+), 16 deletions(-) diff --git a/transformer_engine/pytorch/ops/op.py b/transformer_engine/pytorch/ops/op.py index a506192284..21924ce08b 100644 --- a/transformer_engine/pytorch/ops/op.py +++ b/transformer_engine/pytorch/ops/op.py @@ -454,8 +454,8 @@ def set_activation_offloading(self, enabled: bool) -> None: def mark_for_cpu_offload_if_needed( self, - *tensors: TensorLike | Iterable[Optional[TensorLike]] | None, - exclude_tensors: TensorLike | Iterable[Optional[TensorLike]] | None = None, + *tensors: Iterable[Optional[TensorLike]] | TensorLike | None, + exclude: Iterable[Optional[TensorLike]] | TensorLike | None = None, ) -> None: """Mark tensors to include and exclude from activation CPU offloading. @@ -465,7 +465,7 @@ def mark_for_cpu_offload_if_needed( and it excludes all tensors if the op is not participating in offloading. """ - if not tensors and not exclude_tensors: + if not tensors and not exclude: return if not is_cpu_offload_enabled(): return @@ -474,7 +474,7 @@ def mark_for_cpu_offload_if_needed( def filter_supported_and_extend( out: list[TensorLike], - ts: TensorLike | Iterable[Optional[TensorLike]] | None, + ts: Iterable[Optional[TensorLike]] | TensorLike | None, ) -> None: """Extend a list with objects that support CPU offloading. @@ -492,20 +492,21 @@ def filter_supported_and_extend( out.append(t) # Choose tensors to include and exclude from CPU offloading - include = [] - exclude = [] - if self._activation_offloading_enabled: - filter_supported_and_extend(include, tensors) - filter_supported_and_extend(exclude, exclude_tensors) - else: - filter_supported_and_extend(exclude, tensors) - filter_supported_and_extend(exclude, exclude_tensors) + include_tensors = [] + exclude_tensors = [] + is_enabled = self._activation_offloading_enabled + for t in tensors: + filter_supported_and_extend( + include_tensors if is_enabled else exclude_tensors, + t, + ) + filter_supported_and_extend(exclude_tensors, exclude) # Mark tensors - if include: - mark_activation_offload(*include) - if exclude: - mark_not_offload(*exclude) + if include_tensors: + mark_activation_offload(*include_tensors) + if exclude_tensors: + mark_not_offload(*exclude_tensors) @abc.abstractmethod def op_forward(