From f017815083894102150e77b9bc8d0e244fa83614 Mon Sep 17 00:00:00 2001 From: ZhikangNiu Date: Sat, 15 Mar 2025 15:48:47 +0800 Subject: [PATCH] fix #858 and pass use_reentrant explicitly in checkpoint_activation mode --- src/f5_tts/model/backbones/dit.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/src/f5_tts/model/backbones/dit.py b/src/f5_tts/model/backbones/dit.py index c462528..2ff9670 100644 --- a/src/f5_tts/model/backbones/dit.py +++ b/src/f5_tts/model/backbones/dit.py @@ -219,7 +219,9 @@ class DiT(nn.Module): for block in self.transformer_blocks: if self.checkpoint_activations: - x = torch.utils.checkpoint.checkpoint(self.ckpt_wrapper(block), x, t, mask, rope) + # if you have question, please check: https://pytorch.org/docs/stable/checkpoint.html#torch.utils.checkpoint.checkpoint + # After PyTorch 2.4, we must pass the use_reentrant explicitly + x = torch.utils.checkpoint.checkpoint(self.ckpt_wrapper(block), x, t, mask, rope, use_reentrant=False) else: x = block(x, t, mask=mask, rope=rope)