diff --git a/deepspeed/runtime/activation_checkpointing/checkpointing.py b/deepspeed/runtime/activation_checkpointing/checkpointing.py index f48b79b20521..d203ae92f25b 100644 --- a/deepspeed/runtime/activation_checkpointing/checkpointing.py +++ b/deepspeed/runtime/activation_checkpointing/checkpointing.py @@ -368,11 +368,15 @@ def partition_activations(args, cpu_checkpoint, contiguous_checkpoint): global contiguous_data_buffers, data_offsets inputs = [] - for i, item in enumerate(args): + num_non_fp_tensors = 0 + + for arg_index, item in enumerate(args): if not is_activation_to_checkpoint(item): inputs.append(item) + num_non_fp_tensors += 1 continue + i = arg_index - num_non_fp_tensors partition_size = get_partition_size(item) partition = item.detach().contiguous().view(-1).narrow( 0, @@ -387,7 +391,7 @@ def partition_activations(args, cpu_checkpoint, contiguous_checkpoint): torch.tensor(()).new_empty([partition_size], dtype=partition.dtype, device=buffer_device) - for i in range(num_layers) + for _ in range(num_layers) ] contiguous_data_buffers.append(tensor_list) data_offsets.append(0) @@ -396,7 +400,7 @@ def partition_activations(args, cpu_checkpoint, contiguous_checkpoint): torch.tensor(()).new_empty([partition_size], dtype=partition.dtype, device=buffer_device) - for i in range(num_layers) + for _ in range(num_layers) ] contiguous_data_buffers[i] = tensor_list data_offsets[i] = 0 @@ -430,15 +434,19 @@ def get_partitioned_activations_for_backward(args, inputs, contiguous_checkpoint global contiguous_size_buffers, size_offsets new_args = [] - for i, (arg, inp) in enumerate(zip(args, inputs)): + num_non_fp_tensors = 0 + + for arg_index, (arg, inp) in enumerate(zip(args, inputs)): size = torch.tensor(arg.size()) if torch.is_tensor(arg) else None if not is_activation_to_checkpoint(arg): new_args.append(arg) new_args.append(size) + num_non_fp_tensors += 1 continue arg.data = inp.data new_args.append(arg) + i = arg_index - num_non_fp_tensors if contiguous_checkpoint: numel = size.numel()