Skip to content
Merged
Show file tree
Hide file tree
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
13 changes: 0 additions & 13 deletions megatron/core/models/gpt/fine_grained_callables.py
Original file line number Diff line number Diff line change
Expand Up @@ -316,12 +316,6 @@ def backward_impl(self, outputs, output_grad):
detached_grad = tuple([e.grad for e in self.detached])
grads = output_grad + detached_grad
self.default_backward_func(outputs + self.before_detached, grads)
# release the output grad memory after backward finishes,
# except when delay_wgrad_comptue is enabled, the grad should be
# kept until all modules' backward_dw has been invoked.
if self.delay_wgrad_compute:
self.output_grads = grads
self.delay_grads_release = len(self.bwd_dw_callables) > 0

# return grads for record stream
return grads
Expand All @@ -339,13 +333,6 @@ def backward_dw(self):
module.backward_dw()
nvtx_range_pop(nvtx_msg)

# the output grad memory is last used in wgrad compute, should be safe to release.
assert self.delay_grads_release, "output grad memory should be valid before wgrad."
if self.manual_release_grads:
for tensor in self.output_grads:
tensor.untyped_storage().resize_(0)
self.output_grads = None

self.bwd_dw_callables = None

def __del__(self):
Expand Down
8 changes: 0 additions & 8 deletions megatron/core/pipeline_parallel/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -182,8 +182,6 @@ def __init__(
self.free_input = free_input
self.inputs = None
self.outputs = None
self.delay_grads_release = False
self.manual_release_grads = False

def default_backward_func(self, outputs, output_grad):
"""Default backward function"""
Expand Down Expand Up @@ -263,12 +261,6 @@ def _backward(self, *output_grad):
for g in output_grad:
if g is not None:
g.record_stream(self.stream)
# Manually trigger the memory release of dgrad tensor
# to avoid delayed garbage collection. If
# delay_grads_release is True, dgrad is last used in
# wgrad compute and skip the release here.
if self.manual_release_grads and not self.delay_grads_release:
g.untyped_storage().resize_(0)

grads = self.get_grad()
self._release_state()
Expand Down
Loading