[pre-commit.ci] auto fixes from pre-commit.com hooks
for more information, see https://pre-commit.ci
This commit is contained in:
parent
9ad9d23349
commit
ea0545170a
2 changed files with 8 additions and 4 deletions
|
|
@ -51,6 +51,7 @@ def _multi_gpu_device_map_kwargs() -> dict:
|
|||
return {}
|
||||
try:
|
||||
from utils.hardware import get_device_map, get_parent_visible_gpu_ids
|
||||
|
||||
visible = get_parent_visible_gpu_ids()
|
||||
if len(visible) > 1:
|
||||
device_map = get_device_map(visible)
|
||||
|
|
|
|||
|
|
@ -38,9 +38,7 @@ class _FakeLogger:
|
|||
def _load_helpers(fake_torch, fake_logger):
|
||||
tree = ast.parse(_SAVE_PY.read_text(encoding = "utf-8"))
|
||||
keep = [
|
||||
node
|
||||
for node in tree.body
|
||||
if isinstance(node, ast.FunctionDef) and node.name in _WANTED
|
||||
node for node in tree.body if isinstance(node, ast.FunctionDef) and node.name in _WANTED
|
||||
]
|
||||
assert len(keep) == len(_WANTED), "release helpers missing from save.py"
|
||||
namespace = {"torch": fake_torch, "logger": fake_logger}
|
||||
|
|
@ -58,7 +56,12 @@ def _fake_torch(cuda_available = True):
|
|||
|
||||
|
||||
class _FakeModel:
|
||||
def __init__(self, device_map = None, devices = ("cuda:0",), quantized = False):
|
||||
def __init__(
|
||||
self,
|
||||
device_map = None,
|
||||
devices = ("cuda:0",),
|
||||
quantized = False,
|
||||
):
|
||||
if device_map is not None:
|
||||
self.hf_device_map = device_map
|
||||
self._devices = [types.SimpleNamespace(device = d) for d in devices]
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue