Files
ComfyUI/comfy/storage.py
T

112 lines
3.3 KiB
Python

import functools
import logging
import os
import platform
import re
import comfy_aimdo.storage
from comfy.cli_args import args
_NVME_NAMESPACE = re.compile(r"^(nvme\d+)n\d+$")
def _read(path):
try:
with open(path, encoding="utf-8") as f:
return f.read().strip()
except OSError:
return None
def _physical_block_devices(name):
partition = f"/sys/class/block/{name}/partition"
if os.path.exists(partition):
name = os.path.basename(os.path.dirname(os.path.realpath(f"/sys/class/block/{name}")))
slaves = f"/sys/class/block/{name}/slaves"
try:
children = os.listdir(slaves)
except OSError:
children = []
if children:
devices = []
for child in children:
devices.extend(_physical_block_devices(child))
return devices
return [name]
def _fast_nvme(name):
match = _NVME_NAMESPACE.match(name)
if match is None:
return False
controller = match.group(1)
speed = _read(f"/sys/class/nvme/{controller}/device/current_link_speed")
width = _read(f"/sys/class/nvme/{controller}/device/current_link_width")
if speed is None or width is None:
return None
try:
speed_gts = float(speed.split()[0])
width = int(width)
except ValueError:
return None
return (speed_gts >= 8.0 and width >= 4) or (speed_gts >= 32.0 and width >= 2)
@functools.lru_cache(maxsize=None)
def _linux_fast_storage(device):
sys_device = f"/sys/dev/block/{os.major(device)}:{os.minor(device)}"
if not os.path.exists(sys_device):
return None
name = os.path.basename(os.path.realpath(sys_device))
devices = _physical_block_devices(name)
results = [_fast_nvme(x) for x in devices]
if any(x is None for x in results):
return None
return all(results)
def fast_storage(path):
system = platform.system()
if system == "Linux":
try:
device = os.stat(os.path.realpath(path)).st_dev
except OSError:
return None
return _linux_fast_storage(device)
if system == "Windows":
return comfy_aimdo.storage.fast_disk(path)
return None
def annotate_state_dict(state_dict, path):
path = os.path.realpath(path)
for value in state_dict.values():
untyped_storage = getattr(value, "untyped_storage", None)
if untyped_storage is not None:
untyped_storage()._comfy_source_path = path
def state_dict_fast_disk(state_dict):
state_dicts = state_dict if isinstance(state_dict, (list, tuple)) else (state_dict,)
paths = set()
for sd in state_dicts:
for value in sd.values():
untyped_storage = getattr(value, "untyped_storage", None)
if untyped_storage is not None:
path = getattr(untyped_storage(), "_comfy_source_path", None)
if path is not None:
paths.add(path)
return model_fast_disk(sorted(paths))
def model_fast_disk(paths):
if args.fast_disk or args.disable_fast_disk:
fast = not args.disable_fast_disk
else:
results = [fast_storage(path) for path in paths]
fast = bool(results) and all(result is True for result in results)
logging.info("Model storage policy: fast_disk=%s paths=%s", fast, [os.path.realpath(path) for path in paths])
return fast