Fix race condition that can corrupt grads under certain conditions.
This commit is contained in:
parent
817f3dcbcb
commit
f4e9130547
|
|
@ -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.ip_adapter import IPAdapter
|
||||
from toolkit.custom_adapter import CustomAdapter
|
||||
from toolkit.memory_management import sync_grad_transfers
|
||||
from toolkit.print import print_acc
|
||||
from toolkit.prompt_utils import PromptEmbeds, concat_prompt_embeds
|
||||
from toolkit.reference_adapter import ReferenceAdapter
|
||||
|
|
@ -2218,6 +2219,9 @@ class SDTrainer(BaseSDTrainProcess):
|
|||
|
||||
|
||||
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
|
||||
if self.train_config.optimizer != 'adafactor':
|
||||
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()
|
||||
|
||||
|
||||
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
|
||||
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"]
|
||||
state["grad_compute_done"][idx].record() # on the compute stream
|
||||
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:
|
||||
grad_b_cpu = grad_b_gpu.to("cpu", non_blocking=True)
|
||||
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
|
||||
|
||||
|
||||
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
|
||||
def _is_ao_quantized_tensor(t: Optional[torch.Tensor]) -> bool:
|
||||
if t is None:
|
||||
|
|
@ -287,11 +313,16 @@ class _BouncingLinearFn(torch.autograd.Function):
|
|||
return out.to(x.device)
|
||||
|
||||
state = _get_device_state(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)
|
||||
# the guard makes current_stream() (used by the staging helpers' event
|
||||
# waits/records) resolve to the process device; without it they hit
|
||||
# device 0's streams when training on another gpu and nothing orders
|
||||
# the H2D against the compute
|
||||
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.device = device
|
||||
|
|
@ -376,7 +407,7 @@ class _BouncingLinearFn(torch.autograd.Function):
|
|||
b_grad_gpu = grad_out.sum(dim=tuple(range(grad_out.ndim - 1)))
|
||||
state["b_grad_buffers"][idx] = b_grad_gpu
|
||||
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
|
||||
|
|
@ -431,11 +462,13 @@ class _BouncingConv2dFn(torch.autograd.Function):
|
|||
return out.to(x.device)
|
||||
|
||||
state = _get_device_state(device)
|
||||
idx, w_gpu, b_gpu = _stage_forward_weight(
|
||||
state, device, _materialize_conv_weight, weight_cpu, bias_cpu
|
||||
)
|
||||
out = F.conv2d(x, w_gpu, b_gpu, stride, padding, dilation, groups)
|
||||
_release_forward_slot(state, idx)
|
||||
# device guard: see _BouncingLinearFn.forward
|
||||
with torch.cuda.device(device):
|
||||
idx, w_gpu, b_gpu = _stage_forward_weight(
|
||||
state, device, _materialize_conv_weight, weight_cpu, bias_cpu
|
||||
)
|
||||
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.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))
|
||||
state["b_grad_buffers"][idx] = b_grad_gpu
|
||||
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 (
|
||||
|
|
|
|||
Loading…
Reference in New Issue