diff --git a/astrai/parallel/executor.py b/astrai/parallel/executor.py index 84f3c4c..af0a70e 100644 --- a/astrai/parallel/executor.py +++ b/astrai/parallel/executor.py @@ -231,13 +231,6 @@ class DDPExecutor(BaseExecutor): return model.module.state_dict() return model.state_dict() - def _gather_state_dict(self, model: nn.Module): - if not self.use_distributed: - return self.unwrap_model(model) - if get_rank() != 0: - return None - return self.unwrap_model(model) - @ExecutorFactory.register("fsdp") class FSDPExecutor(BaseExecutor): @@ -298,12 +291,7 @@ class FSDPExecutor(BaseExecutor): def clip_grad_norm(self, model: nn.Module, max_norm: Optional[float]) -> float: if max_norm is None: - total_norm = torch.norm( - torch.stack( - [p.grad.norm(2) for p in model.parameters() if p.grad is not None] - ) - ) - return total_norm.item() + return super().clip_grad_norm(model, max_norm) if isinstance(model, FSDP) and self.use_distributed: total_norm = model.clip_grad_norm_(max_norm) if isinstance(total_norm, torch.Tensor): diff --git a/astrai/trainer/strategy.py b/astrai/trainer/strategy.py index e4acc89..2dde1c9 100644 --- a/astrai/trainer/strategy.py +++ b/astrai/trainer/strategy.py @@ -98,7 +98,6 @@ class BaseStrategy(ABC): self.model = model self.device = device self.executor = kwargs.pop("executor", None) - self.model_fn = kwargs.pop("model_fn", None) self.extra_kwargs = kwargs @abstractmethod diff --git a/astrai/trainer/train_context.py b/astrai/trainer/train_context.py index 3071f9e..dc7d411 100644 --- a/astrai/trainer/train_context.py +++ b/astrai/trainer/train_context.py @@ -211,7 +211,6 @@ class TrainContextBuilder: model=context.model, device=device, executor=executor, - model_fn=cfg.model_fn, **strategy_kwargs, )