fix: 🐛 keep the last model match instead of first

See #184 for details
This commit is contained in:
Mel Massadian
2024-06-21 19:53:25 +02:00
parent fa3199be2b
commit 1edc2cd10d
2 changed files with 38 additions and 26 deletions
+1 -15
View File
@@ -5,7 +5,7 @@ import torch
import tqdm
from ..log import log
from ..utils import StringConvertibleEnum
from ..utils import Operation, Precision
from ..utils import output_dir as comfy_out_dir
PRUNE_DATA = {
@@ -27,20 +27,6 @@ dtypes_to_bf16 = {torch.float32, torch.float64, torch.float16}
dtypes_to_fp8 = {torch.float32, torch.float64, torch.bfloat16, torch.float16}
class Precision(StringConvertibleEnum):
FULL = "full"
FP32 = "fp32"
FP16 = "fp16"
BF16 = "bf16"
FP8 = "fp8"
class Operation(StringConvertibleEnum):
COPY = "copy"
CONVERT = "convert"
DELETE = "delete"
class MTB_ModelPruner:
@classmethod
def INPUT_TYPES(cls):
+37 -11
View File
@@ -11,7 +11,7 @@ import sys
import uuid
from enum import Enum
from pathlib import Path
from typing import List, Optional, TypeVar, Union
from typing import TypeVar
import folder_paths
import numpy as np
@@ -211,6 +211,7 @@ def get_server_info():
# region MISC Utilities
# TODO: use mtb.core directly instead of copying parts here
T = TypeVar("T", bound="StringConvertibleEnum")
@@ -292,12 +293,37 @@ class StringConvertibleEnum(Enum):
return self.value
class Precision(StringConvertibleEnum):
FULL = "full"
FP32 = "fp32"
FP16 = "fp16"
BF16 = "bf16"
FP8 = "fp8"
def to_dtype(self):
match self:
case Precision.FP32 | Precision.FULL:
return torch.float32
case Precision.FP16:
return torch.float16
case Precision.BF16:
return torch.bfloat16
case Precision.FP8:
return torch.float8_e4m3fn
class Operation(StringConvertibleEnum):
COPY = "copy"
CONVERT = "convert"
DELETE = "delete"
def backup_file(
fp: Path,
target: Optional[Path] = None,
target: Path | None = None,
backup_dir: str = ".bak",
suffix: Optional[str] = None,
prefix: Optional[str] = None,
suffix: str | None = None,
prefix: str | None = None,
):
if not fp.exists():
raise FileNotFoundError(f"No file found at {fp}")
@@ -475,7 +501,7 @@ PIL_FILTER_MAP = {
# region TENSOR Utilities
def tensor2pil(image: torch.Tensor) -> List[Image.Image]:
def tensor2pil(image: torch.Tensor) -> list[Image.Image]:
batch_count = image.size(0) if len(image.shape) > 3 else 1
if batch_count > 1:
out = []
@@ -492,7 +518,7 @@ def tensor2pil(image: torch.Tensor) -> List[Image.Image]:
]
def pil2tensor(image: Union[Image.Image, List[Image.Image]]) -> torch.Tensor:
def pil2tensor(image: Image.Image | list[Image.Image]) -> torch.Tensor:
if isinstance(image, list):
return torch.cat([pil2tensor(img) for img in image], dim=0)
@@ -501,14 +527,14 @@ def pil2tensor(image: Union[Image.Image, List[Image.Image]]) -> torch.Tensor:
).unsqueeze(0)
def np2tensor(img_np: Union[np.ndarray, List[np.ndarray]]) -> torch.Tensor:
def np2tensor(img_np: np.ndarray | list[np.ndarray]) -> torch.Tensor:
if isinstance(img_np, list):
return torch.cat([np2tensor(img) for img in img_np], dim=0)
return torch.from_numpy(img_np.astype(np.float32) / 255.0).unsqueeze(0)
def tensor2np(tensor: torch.Tensor) -> List[np.ndarray]:
def tensor2np(tensor: torch.Tensor) -> list[np.ndarray]:
batch_count = tensor.size(0) if len(tensor.shape) > 3 else 1
if batch_count > 1:
out = []
@@ -770,11 +796,11 @@ def get_model_path(fam, model=None):
if res:
if isinstance(res, list):
if len(res) > 1:
warn_msg = f"Found multiple match, we will pick the first {res[0]}\n{res}"
warn_msg = f"Found multiple match, we will pick the last {res[-1]}\n{res}"
if warn_msg not in warned_messages:
log.warning(warn_msg)
log.info(warn_msg)
warned_messages.add(warn_msg)
res = res[0]
res = res[-1]
res = Path(res)
log.debug(f"Resolved model path from folder_paths: {res}")
else: