Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 12 additions & 4 deletions deepspeed/runtime/activation_checkpointing/checkpointing.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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)
Expand All @@ -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
Expand Down Expand Up @@ -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()
Expand Down