diff --git a/transformers_nodes.py b/transformers_nodes.py index a3d58e2..2d555e2 100644 --- a/transformers_nodes.py +++ b/transformers_nodes.py @@ -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) diff --git a/util.py b/util.py index 7b98bf2..dcc524b 100644 --- a/util.py +++ b/util.py @@ -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_":