Fix race condition that can corrupt grads under certain conditions.

This commit is contained in:
Jaret Burkett
2026-08-07 14:52:47 -06:00
parent 817f3dcbcb
commit f4e9130547
3 changed files with 53 additions and 15 deletions

View File

@@ -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):

View File

@@ -1 +1,2 @@
from .manager import MemoryManager from .manager import MemoryManager
from .manager_modules import sync_grad_transfers

View File

@@ -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 (