T5 second GPU offload

This commit is contained in:
City
2024-04-10 01:54:08 +02:00
parent 5e3886c094
commit 92323c62e7
4 changed files with 16 additions and 2 deletions
+6
View File
@@ -35,6 +35,12 @@ class EXM_T5v11:
self.load_device = "cpu"
self.offload_device = "cpu"
self.init_device="cpu"
elif device.startswith("cuda"):
print("Direct CUDA device override!\nVRAM will not be freed by default.")
size = 0
self.load_device = device
self.offload_device = device
self.init_device = device
else:
size = 0
self.load_device = model_management.get_torch_device()
+5 -1
View File
@@ -33,12 +33,16 @@ else: dtypes += ["FP8 E4M3", "FP8 E5M2"]
class T5v11Loader:
@classmethod
def INPUT_TYPES(s):
devices = ["auto", "cpu", "gpu"]
# hack for using second GPU as offload
for k in range(1, torch.cuda.device_count()):
devices.append(f"cuda:{k}")
return {
"required": {
"t5v11_name": (folder_paths.get_filename_list("t5"),),
"t5v11_ver": (["xxl"],),
"path_type": (["folder", "file"],),
"device": (["auto", "cpu", "gpu"],{"default":"cpu"}),
"device": (devices, {"default":"cpu"}),
"dtype": (dtypes,),
}
}
+3 -1
View File
@@ -33,7 +33,9 @@ class T5v11Model(torch.nn.Module):
else:
if dtype: model_args["torch_dtype"] = dtype
self.bnb = False
# TODO: custom device map?
# 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: