Fix single gpu limit code overriding the wrong cuda gpu id via env (#228)

This commit is contained in:
Qubitium 2024-03-15 21:12:16 +08:00 committed by GitHub
commit 39713e66ed

View file

@ -17,22 +17,16 @@ import importlib
# Currently only supports 1 GPU, or else seg faults will occur. # Currently only supports 1 GPU, or else seg faults will occur.
if "CUDA_VISIBLE_DEVICES" in os.environ: if "CUDA_VISIBLE_DEVICES" in os.environ:
device = os.environ["CUDA_VISIBLE_DEVICES"] devices = os.environ["CUDA_VISIBLE_DEVICES"]
if not device.isdigit(): # check if there are multiple cuda devices set in env
if not devices.isdigit():
first_id = devices.split(',')[0]
warnings.warn( warnings.warn(
f"Unsloth: 'CUDA_VISIBLE_DEVICES' is currently {device} "\ f"Unsloth: 'CUDA_VISIBLE_DEVICES' is currently {devices} \n"\
"but we require 'CUDA_VISIBLE_DEVICES=0'\n"\ "Multiple CUDA devices detected but we require a single device.\n"\
"We shall set it ourselves." f"We will override CUDA_VISIBLE_DEVICES to first device: {first_id}."
) )
os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" os.environ["CUDA_VISIBLE_DEVICES"] = str(first_id)
os.environ["CUDA_VISIBLE_DEVICES"] = "0"
elif "CUDA_DEVICE_ORDER" not in os.environ:
warnings.warn(
f"Unsloth: 'CUDA_DEVICE_ORDER' is not set "\
"but we require 'CUDA_DEVICE_ORDER=PCI_BUS_ID'\n"\
"We shall set it ourselves."
)
os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"
else: else:
# warnings.warn("Unsloth: 'CUDA_VISIBLE_DEVICES' is not set. We shall set it ourselves.") # warnings.warn("Unsloth: 'CUDA_VISIBLE_DEVICES' is not set. We shall set it ourselves.")
os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID" os.environ["CUDA_DEVICE_ORDER"] = "PCI_BUS_ID"