diff --git a/src/lmflow/pipeline/finetuner.py b/src/lmflow/pipeline/finetuner.py index 704f8c111..77f109992 100644 --- a/src/lmflow/pipeline/finetuner.py +++ b/src/lmflow/pipeline/finetuner.py @@ -311,8 +311,7 @@ def __init__(self, n_layers, interval_steps, model): self.layers_attribute = 'model.transformer.h' # General access path self.total_layers = len(eval('self.' + self.layers_attribute)) # Dynamically execute to get the number of layers - # Freeze all layers upon initialization - self.freeze_all_layers() + self.switch_active_layers() self.active_layers_indices = [] def freeze_all_layers(self): @@ -323,7 +322,7 @@ def freeze_all_layers(self): def on_step_begin(self, args, state, control, **kwargs): # Check if it's time to switch active layers, including at step 0 - if state.global_step % self.interval_steps == 0 or state.global_step == 1: + if state.global_step % self.interval_steps == 0 : self.switch_active_layers() def switch_active_layers(self):