This commit is contained in:
sayakpaul
2024-12-25 21:34:56 +05:30
parent 5aef4a6c80
commit 6b54867c3c
4 changed files with 45 additions and 20 deletions
+1 -1
View File
@@ -14,7 +14,7 @@ machine_rank: 0
main_training_function: main main_training_function: main
mixed_precision: bf16 mixed_precision: bf16
num_machines: 1 num_machines: 1
num_processes: 2 num_processes: 1
rdzv_backend: static rdzv_backend: static
same_network: true same_network: true
tpu_env: [] tpu_env: []
+3 -3
View File
@@ -1,14 +1,14 @@
compute_environment: LOCAL_MACHINE compute_environment: LOCAL_MACHINE
debug: false debug: false
distributed_type: MULTI_GPU distributed_type: 'NO'
downcast_bf16: 'no' downcast_bf16: 'no'
enable_cpu_affinity: false enable_cpu_affinity: false
gpu_ids: 0,1 gpu_ids: 0
machine_rank: 0 machine_rank: 0
main_training_function: main main_training_function: main
mixed_precision: bf16 mixed_precision: bf16
num_machines: 1 num_machines: 1
num_processes: 2 num_processes: 1
rdzv_backend: static rdzv_backend: static
same_network: true same_network: true
tpu_env: [] tpu_env: []
+1
View File
@@ -675,6 +675,7 @@ _DTYPE_MAP = {
"fp16": torch.float16, "fp16": torch.float16,
"fp32": torch.float32, "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: def _map_to_args_type(args: Dict[str, Any]) -> Args:
+40 -16
View File
@@ -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 peft import LoraConfig, get_peft_model_state_dict, set_peft_model_state_dict
from tqdm import tqdm from tqdm import tqdm
from .args import Args, validate_args from .args import Args, validate_args, _INVERSE_DTYPE_MAP
from .constants import ( from .constants import (
FINETRAINERS_LOG_LEVEL, FINETRAINERS_LOG_LEVEL,
PRECOMPUTED_DIR_NAME, PRECOMPUTED_DIR_NAME,
@@ -340,6 +340,27 @@ class Trainer:
pin_memory=self.args.pin_memory, 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: def prepare_trainable_parameters(self) -> None:
logger.info("Initializing trainable parameters") 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 # 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. # as these weights are only used for inference, keeping weights in full precision is not required.
weight_dtype = torch.float32 weight_dtype = self.sort_out_weight_dtype(accelerator=self.state.accelerator)
if self.state.accelerator.mixed_precision == "fp16":
weight_dtype = torch.float16
elif self.state.accelerator.mixed_precision == "bf16":
weight_dtype = torch.bfloat16
if torch.backends.mps.is_available() and weight_dtype == torch.bfloat16: if torch.backends.mps.is_available() and weight_dtype == torch.bfloat16:
# due to pytorch#99272, MPS does not yet support 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 # TODO(aryan): handle torch dtype from accelerator vs model dtype; refactor
self.state.weight_dtype = weight_dtype 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.transformer.to(dtype=weight_dtype)
self._move_components_to_device() self._move_components_to_device()
@@ -492,16 +514,20 @@ class Trainer:
params_to_optimize = [transformer_parameters_with_lr] params_to_optimize = [transformer_parameters_with_lr]
self.state.num_trainable_parameters = sum(p.numel() for p in transformer_lora_parameters) 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( optimizer = get_optimizer(
params_to_optimize=params_to_optimize, params_to_optimize=params_to_optimize,
optimizer_name=self.args.optimizer, optimizer_name=self.args.optimizer,
learning_rate=self.args.lr, learning_rate=self.state.learning_rate,
beta1=self.args.beta1, beta1=self.args.beta1,
beta2=self.args.beta2, beta2=self.args.beta2,
beta3=self.args.beta3, beta3=self.args.beta3,
epsilon=self.args.epsilon, epsilon=self.args.epsilon,
weight_decay=self.args.weight_decay, 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) 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.train_steps = self.state.train_epochs * num_update_steps_per_epoch
self.state.overwrote_max_train_steps = True 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( lr_scheduler = get_scheduler(
name=self.args.lr_scheduler, name=self.args.lr_scheduler,
optimizer=optimizer, optimizer=optimizer,
@@ -516,6 +546,7 @@ class Trainer:
num_training_steps=self.state.train_steps * self.state.accelerator.num_processes, num_training_steps=self.state.train_steps * self.state.accelerator.num_processes,
num_cycles=self.args.lr_num_cycles, num_cycles=self.args.lr_num_cycles,
power=self.args.lr_power, power=self.args.lr_power,
use_deepspeed=use_deepspeed_lr_schd
) )
self.optimizer = optimizer self.optimizer = optimizer
@@ -692,6 +723,7 @@ class Trainer:
loss = loss.mean(list(range(1, loss.ndim))) loss = loss.mean(list(range(1, loss.ndim)))
# Average loss across batch dimension # Average loss across batch dimension
loss = loss.mean() loss = loss.mean()
print(f"{self.transformer.dtype=}, {accelerator.mixed_precision=}")
accelerator.backward(loss) accelerator.backward(loss)
if accelerator.sync_gradients: if accelerator.sync_gradients:
@@ -756,14 +788,6 @@ class Trainer:
if accelerator.is_main_process: if accelerator.is_main_process:
# TODO: consider factoring this out when supporting other types of training algos. # TODO: consider factoring this out when supporting other types of training algos.
self.transformer = unwrap_model(accelerator, self.transformer) 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) transformer_lora_layers = get_peft_model_state_dict(self.transformer)
self.model_config["pipeline_cls"].save_lora_weights( self.model_config["pipeline_cls"].save_lora_weights(