Fix race condition that can corrupt grads under certain conditions.
This commit is contained in:
@@ -20,6 +20,7 @@ from toolkit.guidance import get_targeted_guidance_loss, get_guidance_loss, Guid
|
|||||||
from toolkit.image_utils import show_tensors, show_latents
|
from toolkit.image_utils import show_tensors, show_latents
|
||||||
from toolkit.ip_adapter import IPAdapter
|
from toolkit.ip_adapter import IPAdapter
|
||||||
from toolkit.custom_adapter import CustomAdapter
|
from toolkit.custom_adapter import CustomAdapter
|
||||||
|
from toolkit.memory_management import sync_grad_transfers
|
||||||
from toolkit.print import print_acc
|
from toolkit.print import print_acc
|
||||||
from toolkit.prompt_utils import PromptEmbeds, concat_prompt_embeds
|
from toolkit.prompt_utils import PromptEmbeds, concat_prompt_embeds
|
||||||
from toolkit.reference_adapter import ReferenceAdapter
|
from toolkit.reference_adapter import ReferenceAdapter
|
||||||
@@ -2218,6 +2219,9 @@ class SDTrainer(BaseSDTrainProcess):
|
|||||||
|
|
||||||
|
|
||||||
if not self.is_grad_accumulation_step:
|
if not self.is_grad_accumulation_step:
|
||||||
|
# grads of memory-managed (offloaded) params are async D2H copies into
|
||||||
|
# pinned tensors; join them before anything on the CPU reads .grad
|
||||||
|
sync_grad_transfers()
|
||||||
# fix this for multi params
|
# fix this for multi params
|
||||||
if self.train_config.optimizer != 'adafactor':
|
if self.train_config.optimizer != 'adafactor':
|
||||||
if isinstance(self.params[0], dict):
|
if isinstance(self.params[0], dict):
|
||||||
|
|||||||
@@ -1 +1,2 @@
|
|||||||
from .manager import MemoryManager
|
from .manager import MemoryManager
|
||||||
|
from .manager_modules import sync_grad_transfers
|
||||||
|
|||||||
@@ -118,9 +118,17 @@ def _release_backward_weight_slot(state, idx):
|
|||||||
state["bwd_slot_free"][idx].record()
|
state["bwd_slot_free"][idx].record()
|
||||||
|
|
||||||
|
|
||||||
def _stage_grads_to_cpu(state, idx, grad_w_gpu, grad_b_gpu):
|
def _stage_grads_to_cpu(state, idx, grad_w_gpu, grad_b_gpu, weight_cpu, bias_cpu):
|
||||||
"""Copy freshly-computed device grads (in staging slot idx) to CPU on the
|
"""Copy freshly-computed device grads (in staging slot idx) to CPU on the
|
||||||
grad stream, overlapping the next H2D. Returns (grad_w_cpu, grad_b_cpu)."""
|
grad stream, overlapping the next H2D. Returns (grad_w_cpu, grad_b_cpu).
|
||||||
|
|
||||||
|
The returned tensors are pinned-memory destinations of an ASYNC copy: their
|
||||||
|
contents are undefined until the grad stream reaches grad_xfer_done. GPU
|
||||||
|
consumers are ordered by that event; host consumers must join it first —
|
||||||
|
the optimizer/clip path does so via sync_grad_transfers(). The one host
|
||||||
|
read we can't defer is grad accumulation: when the param already holds a
|
||||||
|
.grad, AccumulateGrad does `grad += returned` on the engine thread the
|
||||||
|
moment backward returns, so block here until the copy has landed."""
|
||||||
gs = state["transfer_grad_stream"]
|
gs = state["transfer_grad_stream"]
|
||||||
state["grad_compute_done"][idx].record() # on the compute stream
|
state["grad_compute_done"][idx].record() # on the compute stream
|
||||||
grad_w_cpu = grad_b_cpu = None
|
grad_w_cpu = grad_b_cpu = None
|
||||||
@@ -131,9 +139,27 @@ def _stage_grads_to_cpu(state, idx, grad_w_gpu, grad_b_gpu):
|
|||||||
if grad_b_gpu is not None:
|
if grad_b_gpu is not None:
|
||||||
grad_b_cpu = grad_b_gpu.to("cpu", non_blocking=True)
|
grad_b_cpu = grad_b_gpu.to("cpu", non_blocking=True)
|
||||||
state["grad_xfer_done"][idx].record()
|
state["grad_xfer_done"][idx].record()
|
||||||
|
if (grad_w_cpu is not None and weight_cpu.grad is not None) or (
|
||||||
|
grad_b_cpu is not None and bias_cpu.grad is not None
|
||||||
|
):
|
||||||
|
state["grad_xfer_done"][idx].synchronize()
|
||||||
return grad_w_cpu, grad_b_cpu
|
return grad_w_cpu, grad_b_cpu
|
||||||
|
|
||||||
|
|
||||||
|
def sync_grad_transfers():
|
||||||
|
"""Host-join every device's grad D2H stream.
|
||||||
|
|
||||||
|
Staged weight/bias grads of memory-managed layers are async copies into
|
||||||
|
pinned CPU tensors; nothing else orders those copies against the host.
|
||||||
|
Call this after backward and before anything on the CPU reads .grad of a
|
||||||
|
memory-managed parameter (grad clipping, optimizer step). No-op when no
|
||||||
|
offloading is active."""
|
||||||
|
for state in _DEVICE_STATE.values():
|
||||||
|
stream = state.get("transfer_grad_stream")
|
||||||
|
if stream is not None:
|
||||||
|
stream.synchronize()
|
||||||
|
|
||||||
|
|
||||||
# (ADD) detect torchao wrapper tensors
|
# (ADD) detect torchao wrapper tensors
|
||||||
def _is_ao_quantized_tensor(t: Optional[torch.Tensor]) -> bool:
|
def _is_ao_quantized_tensor(t: Optional[torch.Tensor]) -> bool:
|
||||||
if t is None:
|
if t is None:
|
||||||
@@ -287,11 +313,16 @@ class _BouncingLinearFn(torch.autograd.Function):
|
|||||||
return out.to(x.device)
|
return out.to(x.device)
|
||||||
|
|
||||||
state = _get_device_state(device)
|
state = _get_device_state(device)
|
||||||
idx, w_gpu, b_gpu = _stage_forward_weight(
|
# the guard makes current_stream() (used by the staging helpers' event
|
||||||
state, device, _materialize_linear_weight, weight_cpu, bias_cpu
|
# waits/records) resolve to the process device; without it they hit
|
||||||
)
|
# device 0's streams when training on another gpu and nothing orders
|
||||||
out = F.linear(x, w_gpu, b_gpu)
|
# the H2D against the compute
|
||||||
_release_forward_slot(state, idx)
|
with torch.cuda.device(device):
|
||||||
|
idx, w_gpu, b_gpu = _stage_forward_weight(
|
||||||
|
state, device, _materialize_linear_weight, weight_cpu, bias_cpu
|
||||||
|
)
|
||||||
|
out = F.linear(x, w_gpu, b_gpu)
|
||||||
|
_release_forward_slot(state, idx)
|
||||||
|
|
||||||
ctx.save_for_backward(x, weight_cpu, bias_cpu)
|
ctx.save_for_backward(x, weight_cpu, bias_cpu)
|
||||||
ctx.device = device
|
ctx.device = device
|
||||||
@@ -376,7 +407,7 @@ class _BouncingLinearFn(torch.autograd.Function):
|
|||||||
b_grad_gpu = grad_out.sum(dim=tuple(range(grad_out.ndim - 1)))
|
b_grad_gpu = grad_out.sum(dim=tuple(range(grad_out.ndim - 1)))
|
||||||
state["b_grad_buffers"][idx] = b_grad_gpu
|
state["b_grad_buffers"][idx] = b_grad_gpu
|
||||||
grad_weight, grad_bias = _stage_grads_to_cpu(
|
grad_weight, grad_bias = _stage_grads_to_cpu(
|
||||||
state, idx, w_grad_gpu, b_grad_gpu
|
state, idx, w_grad_gpu, b_grad_gpu, weight_cpu, bias_cpu
|
||||||
)
|
)
|
||||||
|
|
||||||
return grad_input.to(dtype=grad_out.dtype), grad_weight, grad_bias, None
|
return grad_input.to(dtype=grad_out.dtype), grad_weight, grad_bias, None
|
||||||
@@ -431,11 +462,13 @@ class _BouncingConv2dFn(torch.autograd.Function):
|
|||||||
return out.to(x.device)
|
return out.to(x.device)
|
||||||
|
|
||||||
state = _get_device_state(device)
|
state = _get_device_state(device)
|
||||||
idx, w_gpu, b_gpu = _stage_forward_weight(
|
# device guard: see _BouncingLinearFn.forward
|
||||||
state, device, _materialize_conv_weight, weight_cpu, bias_cpu
|
with torch.cuda.device(device):
|
||||||
)
|
idx, w_gpu, b_gpu = _stage_forward_weight(
|
||||||
out = F.conv2d(x, w_gpu, b_gpu, stride, padding, dilation, groups)
|
state, device, _materialize_conv_weight, weight_cpu, bias_cpu
|
||||||
_release_forward_slot(state, idx)
|
)
|
||||||
|
out = F.conv2d(x, w_gpu, b_gpu, stride, padding, dilation, groups)
|
||||||
|
_release_forward_slot(state, idx)
|
||||||
|
|
||||||
ctx.save_for_backward(x, weight_cpu, bias_cpu)
|
ctx.save_for_backward(x, weight_cpu, bias_cpu)
|
||||||
ctx.meta = (device, stride, padding, dilation, groups, target_dtype)
|
ctx.meta = (device, stride, padding, dilation, groups, target_dtype)
|
||||||
@@ -563,7 +596,7 @@ class _BouncingConv2dFn(torch.autograd.Function):
|
|||||||
b_grad_gpu = grad_out.sum(dim=(0, 2, 3))
|
b_grad_gpu = grad_out.sum(dim=(0, 2, 3))
|
||||||
state["b_grad_buffers"][idx] = b_grad_gpu
|
state["b_grad_buffers"][idx] = b_grad_gpu
|
||||||
grad_weight, grad_bias = _stage_grads_to_cpu(
|
grad_weight, grad_bias = _stage_grads_to_cpu(
|
||||||
state, idx, w_grad_gpu, b_grad_gpu
|
state, idx, w_grad_gpu, b_grad_gpu, weight_cpu, bias_cpu
|
||||||
)
|
)
|
||||||
|
|
||||||
return (
|
return (
|
||||||
|
|||||||
Reference in New Issue
Block a user