Some corrections and simplifications to device management
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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_":
|
||||
|
||||
Reference in New Issue
Block a user