From f4e91305471a3727d52886ef6d410eb570cd484f Mon Sep 17 00:00:00 2001 From: Jaret Burkett Date: Fri, 7 Aug 2026 14:52:47 -0600 Subject: [PATCH] Fix race condition that can corrupt grads under certain conditions. --- extensions_built_in/sd_trainer/SDTrainer.py | 4 ++ toolkit/memory_management/__init__.py | 3 +- toolkit/memory_management/manager_modules.py | 61 +++++++++++++++----- 3 files changed, 53 insertions(+), 15 deletions(-) diff --git a/extensions_built_in/sd_trainer/SDTrainer.py b/extensions_built_in/sd_trainer/SDTrainer.py index a5197a47..9552e88f 100644 --- a/extensions_built_in/sd_trainer/SDTrainer.py +++ b/extensions_built_in/sd_trainer/SDTrainer.py @@ -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): diff --git a/toolkit/memory_management/__init__.py b/toolkit/memory_management/__init__.py index 2eeef37d..83fa980f 100644 --- a/toolkit/memory_management/__init__.py +++ b/toolkit/memory_management/__init__.py @@ -1 +1,2 @@ -from .manager import MemoryManager \ No newline at end of file +from .manager import MemoryManager +from .manager_modules import sync_grad_transfers diff --git a/toolkit/memory_management/manager_modules.py b/toolkit/memory_management/manager_modules.py index 44f3b288..9680fdc4 100644 --- a/toolkit/memory_management/manager_modules.py +++ b/toolkit/memory_management/manager_modules.py @@ -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 (