Some corrections and simplifications to device management

This commit is contained in:
Yahweasel
2026-01-15 10:34:32 -05:00
parent 956d9ab1c4
commit 0b33d60842
2 changed files with 6 additions and 2 deletions
+3
View File
@@ -65,7 +65,10 @@ class HFTLoadPipeline:
model_kwargs["quantization_config"] = kwargs["quantization_config"]
del kwargs["quantization_config"]
if to_device:
print("TO ", to_device)
kwargs["device"] = to_device
kwargs["dtype"] = kwargs["torch_dtype"]
del kwargs["torch_dtype"]
if pipeline_class != "":
kwargs["pipeline_class"] = getattr(transformers, pipeline_class)
+3 -2
View File
@@ -30,7 +30,8 @@ if torch.cuda.is_available():
DEVICES.append("cuda")
for i in range(torch.cuda.device_count()):
DEVICES.append(f"cuda:{i}")
DEFAULT_DEVICE = "cuda"
if torch.cuda.device_count() > 0:
DEFAULT_DEVICE = "cuda:0"
DTYPES = ("default", "float32", "bfloat16", "float16", "bitsandbytes_8bit", "bitsandbytes_4bit")
def get_device(device):
@@ -65,7 +66,7 @@ def apply_device(kwargs, device, dtype, enable_model_cpu_offload=False, quant="p
if ":" in device:
to_device = device
else:
kwargs["device_map"] = get_device(device)
kwargs["device_map"] = device
kwargs["torch_dtype"] = torch.bfloat16
if dtype[0:13] == "bitsandbytes_":