From f9e1178de76937558003cbedfd136786e38e0445 Mon Sep 17 00:00:00 2001 From: City <125218114+city96@users.noreply.github.com> Date: Sun, 21 Apr 2024 20:41:25 +0200 Subject: [PATCH] T5 offload w/ bnb Mentioned in #22 --- README.md | 2 +- T5/nodes.py | 2 +- T5/t5v11.py | 6 +++--- 3 files changed, 5 insertions(+), 5 deletions(-) diff --git a/README.md b/README.md index 340297f..f41e509 100644 --- a/README.md +++ b/README.md @@ -140,7 +140,7 @@ If you have a second GPU, selecting "cuda:1" as the device will allow you to use Loaded in bnb4bit mode, it only takes around 6GB VRAM, making it work with 12GB cards. The only drawback is that it'll constantly stay in VRAM since BitsAndBytes doesn't allow moving the weights to the system RAM temporarily. Switching to a different workflow *should* still release the VRAM as expected. Pascal cards (1080ti, P40) seem to struggle with 4bit. Select "cpu" if you encounter issues. -On windows, you may need a newer version of bitsandbytes for 4bit. Try `python -m pip install bitsandbytes --prefer-binary --extra-index-url=https://jllllll.github.io/bitsandbytes-windows-webui` +On windows, you may need a newer version of bitsandbytes for 4bit. Try `python -m pip install bitsandbytes` > [!IMPORTANT] > You may also need to upgrade transformers and install spiece for the tokenizer. `pip install -r requirements.txt` diff --git a/T5/nodes.py b/T5/nodes.py index d1153f7..8b9e5cc 100644 --- a/T5/nodes.py +++ b/T5/nodes.py @@ -53,7 +53,7 @@ class T5v11Loader: def load_model(self, t5v11_name, t5v11_ver, path_type, device, dtype): if "bnb" in dtype: - assert device == "gpu", "BitsAndBytes only works on CUDA! Set device to 'gpu'." + assert device == "gpu" or device.startswith("cuda"), "BitsAndBytes only works on CUDA! Set device to 'gpu'." dtype = string_to_dtype(dtype, "text_encoder") if device == "cpu": assert dtype in [None, torch.float32], f"Can't use dtype '{dtype}' with CPU! Set dtype to 'default'." diff --git a/T5/t5v11.py b/T5/t5v11.py index f43e8d6..76cec9c 100644 --- a/T5/t5v11.py +++ b/T5/t5v11.py @@ -33,9 +33,9 @@ class T5v11Model(torch.nn.Module): else: if dtype: model_args["torch_dtype"] = dtype self.bnb = False - # second GPU offload hack part 2 - if device.startswith("cuda"): - model_args["device_map"] = device + # second GPU offload hack part 2 + if device.startswith("cuda"): + model_args["device_map"] = device print(f"Loading T5 from '{textmodel_path}'") self.transformer = T5EncoderModel.from_pretrained(textmodel_path, **model_args) else: