diff --git a/accelerate_configs/deepspeed.yaml b/accelerate_configs/deepspeed.yaml index 2827648..efbbf6f 100644 --- a/accelerate_configs/deepspeed.yaml +++ b/accelerate_configs/deepspeed.yaml @@ -14,7 +14,7 @@ machine_rank: 0 main_training_function: main mixed_precision: bf16 num_machines: 1 -num_processes: 2 +num_processes: 1 rdzv_backend: static same_network: true tpu_env: [] diff --git a/accelerate_configs/uncompiled_2.yaml b/accelerate_configs/uncompiled_2.yaml index c5216da..0c99039 100644 --- a/accelerate_configs/uncompiled_2.yaml +++ b/accelerate_configs/uncompiled_2.yaml @@ -1,14 +1,14 @@ compute_environment: LOCAL_MACHINE debug: false -distributed_type: MULTI_GPU +distributed_type: 'NO' downcast_bf16: 'no' enable_cpu_affinity: false -gpu_ids: 0,1 +gpu_ids: 0 machine_rank: 0 main_training_function: main mixed_precision: bf16 num_machines: 1 -num_processes: 2 +num_processes: 1 rdzv_backend: static same_network: true tpu_env: [] diff --git a/finetrainers/args.py b/finetrainers/args.py index 31c2a76..4c095db 100644 --- a/finetrainers/args.py +++ b/finetrainers/args.py @@ -675,6 +675,7 @@ _DTYPE_MAP = { "fp16": torch.float16, "fp32": torch.float32, } +_INVERSE_DTYPE_MAP = {v: k for k, v in _DTYPE_MAP.items()} def _map_to_args_type(args: Dict[str, Any]) -> Args: diff --git a/finetrainers/trainer.py b/finetrainers/trainer.py index b3bb4db..777d37e 100644 --- a/finetrainers/trainer.py +++ b/finetrainers/trainer.py @@ -33,7 +33,7 @@ from huggingface_hub import create_repo, upload_folder from peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict from tqdm import tqdm -from .args import Args, validate_args +from .args import Args, validate_args, _INVERSE_DTYPE_MAP from .constants import ( FINETRAINERS_LOG_LEVEL, PRECOMPUTED_DIR_NAME, @@ -340,6 +340,27 @@ class Trainer: pin_memory=self.args.pin_memory, ) + def sort_out_weight_dtype(self, accelerator): + weight_dtype = torch.float32 + if accelerator.state.deepspeed_plugin: + # DeepSpeed is handling precision, use what's in the DeepSpeed config + if ( + "fp16" in accelerator.state.deepspeed_plugin.deepspeed_config + and accelerator.state.deepspeed_plugin.deepspeed_config["fp16"]["enabled"] + ): + weight_dtype = torch.float16 + if ( + "bf16" in accelerator.state.deepspeed_plugin.deepspeed_config + and accelerator.state.deepspeed_plugin.deepspeed_config["bf16"]["enabled"] + ): + weight_dtype = torch.bfloat16 + else: + if self.state.accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + elif self.state.accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + return weight_dtype + def prepare_trainable_parameters(self) -> None: logger.info("Initializing trainable parameters") @@ -360,11 +381,7 @@ class Trainer: # For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision # as these weights are only used for inference, keeping weights in full precision is not required. - weight_dtype = torch.float32 - if self.state.accelerator.mixed_precision == "fp16": - weight_dtype = torch.float16 - elif self.state.accelerator.mixed_precision == "bf16": - weight_dtype = torch.bfloat16 + weight_dtype = self.sort_out_weight_dtype(accelerator=self.state.accelerator) if torch.backends.mps.is_available() and weight_dtype == torch.bfloat16: # due to pytorch#99272, MPS does not yet support bfloat16. @@ -374,6 +391,11 @@ class Trainer: # TODO(aryan): handle torch dtype from accelerator vs model dtype; refactor self.state.weight_dtype = weight_dtype + if self.args.mixed_precision != _INVERSE_DTYPE_MAP[weight_dtype]: + logger.info( + f"`mixed_precision` was set to {_INVERSE_DTYPE_MAP[weight_dtype]} which different from what was initially passed ({self.args.mixed_precision})." + ) + self.args.mixed_precision = _INVERSE_DTYPE_MAP[weight_dtype] self.transformer.to(dtype=weight_dtype) self._move_components_to_device() @@ -492,16 +514,20 @@ class Trainer: params_to_optimize = [transformer_parameters_with_lr] self.state.num_trainable_parameters = sum(p.numel() for p in transformer_lora_parameters) - # TODO(aryan): add deepspeed support + use_deepspeed_opt = ( + self.state.accelerator.state.deepspeed_plugin is not None + and "optimizer" in self.state.accelerator.state.deepspeed_plugin.deepspeed_config + ) optimizer = get_optimizer( params_to_optimize=params_to_optimize, optimizer_name=self.args.optimizer, - learning_rate=self.args.lr, + learning_rate=self.state.learning_rate, beta1=self.args.beta1, beta2=self.args.beta2, beta3=self.args.beta3, epsilon=self.args.epsilon, weight_decay=self.args.weight_decay, + use_deepspeed=use_deepspeed_opt ) num_update_steps_per_epoch = math.ceil(len(self.dataloader) / self.args.gradient_accumulation_steps) @@ -509,6 +535,10 @@ class Trainer: self.state.train_steps = self.state.train_epochs * num_update_steps_per_epoch self.state.overwrote_max_train_steps = True + use_deepspeed_lr_schd = ( + self.state.accelerator.state.deepspeed_plugin is not None + and "scheduler" in self.state.accelerator.state.deepspeed_plugin.deepspeed_config + ) lr_scheduler = get_scheduler( name=self.args.lr_scheduler, optimizer=optimizer, @@ -516,6 +546,7 @@ class Trainer: num_training_steps=self.state.train_steps * self.state.accelerator.num_processes, num_cycles=self.args.lr_num_cycles, power=self.args.lr_power, + use_deepspeed=use_deepspeed_lr_schd ) self.optimizer = optimizer @@ -692,6 +723,7 @@ class Trainer: loss = loss.mean(list(range(1, loss.ndim))) # Average loss across batch dimension loss = loss.mean() + print(f"{self.transformer.dtype=}, {accelerator.mixed_precision=}") accelerator.backward(loss) if accelerator.sync_gradients: @@ -756,14 +788,6 @@ class Trainer: if accelerator.is_main_process: # TODO: consider factoring this out when supporting other types of training algos. self.transformer = unwrap_model(accelerator, self.transformer) - dtype = ( - torch.float16 - if self.args.mixed_precision == "fp16" - else torch.bfloat16 - if self.args.mixed_precision == "bf16" - else torch.float32 - ) - self.transformer = self.transformer.to(dtype) transformer_lora_layers = get_peft_model_state_dict(self.transformer) self.model_config["pipeline_cls"].save_lora_weights(