From 8b25bc42de2ed1ad809aafa846c3e6af95bc6025 Mon Sep 17 00:00:00 2001 From: Lucia Quirke Date: Wed, 29 Jul 2026 14:38:19 +0000 Subject: [PATCH] fix: import load_from_optimizer lazily to break a spawn-time cycle Importing it at module scope pulls in bergson.gradients, which re-enters the package __init__ and reaches magic.cli -> validate -> magic.trainer. In-process that resolves, but a spawned worker unpickling through magic.cli hits magic.trainer half-initialized and fails on TrainerState, taking out test_grad_accum_matches_full_batch. --- bergson/magic/trainer.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/bergson/magic/trainer.py b/bergson/magic/trainer.py index 6a92c33a..99ab3d12 100644 --- a/bergson/magic/trainer.py +++ b/bergson/magic/trainer.py @@ -22,7 +22,6 @@ from ..config.config import TrainingConfig from ..data import sorted_checkpoints -from ..utils.load_from_optimizer import save_second_moments_as_optimizer_pt from ..utils.utils import get_device from ..utils.worker_utils import setup_model_and_peft from .config import MagicSaveMode @@ -623,6 +622,12 @@ def train( # Write before the DCP save starts, into the same directory. if optimizer_cfg is not None: + # Local import: at module scope this cycles back into + # magic.trainer via the package __init__. + from ..utils.load_from_optimizer import ( + save_second_moments_as_optimizer_pt, + ) + os.makedirs(p, exist_ok=True) save_second_moments_as_optimizer_pt( self.model,