From 21a4c42da0bcb65274802bc7635b9bdf5e625ba9 Mon Sep 17 00:00:00 2001 From: mayukhdeb Date: Sun, 17 Mar 2024 05:14:06 -0700 Subject: [PATCH] dtype selection ofload --- trainer/utils/dtype.py | 7 +++++++ trainer_pti.py | 8 ++------ 2 files changed, 9 insertions(+), 6 deletions(-) create mode 100644 trainer/utils/dtype.py diff --git a/trainer/utils/dtype.py b/trainer/utils/dtype.py new file mode 100644 index 0000000..c7d0fa4 --- /dev/null +++ b/trainer/utils/dtype.py @@ -0,0 +1,7 @@ +import torch + +dtype_map = { + "fp16": torch.float16, + "bf16": torch.bfloat16, + "fp32": torch.float32 +} \ No newline at end of file diff --git a/trainer_pti.py b/trainer_pti.py index c78ff54..7b746e1 100755 --- a/trainer_pti.py +++ b/trainer_pti.py @@ -23,7 +23,7 @@ from dataset_and_utils import * from lora_utils import * from io_utils import make_validation_img_grid import matplotlib.pyplot as plt - +from trainer.utils.dtype import dtype_map def print_trainable_parameters(model, name = ''): trainable_params = 0 @@ -313,11 +313,7 @@ def main( print("Using seed", seed) torch.manual_seed(seed) - weight_dtype = torch.float32 - if mixed_precision == "fp16": - weight_dtype = torch.float16 - elif mixed_precision == "bf16": - weight_dtype = torch.bfloat16 + weight_dtype = dtype_map[mixed_precision] print(f"Loading models with weight_dtype: {weight_dtype}")