Add Wall-X serving and Turtle2 TCP WebSocket bridge
Pre-commit / pre-commit (push) Canceled after 0s
Pre-commit / pre-commit (push) Canceled after 0s
This commit is contained in:
@@ -95,6 +95,18 @@ class FSDPTrainer(DistributedTrainer):
|
||||
self.load_processor()
|
||||
self.load_model()
|
||||
|
||||
# Single-file pretrained checkpoints use native model parameter names.
|
||||
# Load them before PEFT or other parameter-name-changing wrappers.
|
||||
if self._resume_from_single_file():
|
||||
self.load_state_dict(
|
||||
self.model,
|
||||
{"ckpt": self.cfg.checkpoint.resume_from},
|
||||
)
|
||||
|
||||
self.model = self.adapter.finalize_model_after_weight_load(
|
||||
self.model, self.model_config
|
||||
)
|
||||
|
||||
# DDP needs to see correct requires_grad at wrap time, so freeze first.
|
||||
self._freeze_params_if_needed(self.model)
|
||||
|
||||
@@ -107,12 +119,6 @@ class FSDPTrainer(DistributedTrainer):
|
||||
name: p.shape for name, p in self.model.named_parameters()
|
||||
}
|
||||
|
||||
if self._resume_from_single_file():
|
||||
self.load_state_dict(
|
||||
self.model,
|
||||
{"ckpt": self.cfg.checkpoint.resume_from},
|
||||
)
|
||||
|
||||
self._wrap_model(self.model)
|
||||
self._create_optimizer()
|
||||
self._create_scheduler()
|
||||
|
||||
Reference in New Issue
Block a user