fix vitmatte model loader
This commit is contained in:
+14
-5
@@ -91,6 +91,14 @@ def load_light_leak_images() -> list:
|
||||
file = os.path.join(folder_paths.models_dir, "layerstyle", "light_leak.pkl")
|
||||
return load_pickle(file)
|
||||
|
||||
def check_and_download_model(model_path, repo_id):
|
||||
model_path = os.path.join(folder_paths.models_dir, model_path)
|
||||
|
||||
if not os.path.exists(model_path):
|
||||
print(f"Downloading {repo_id} model...")
|
||||
from huggingface_hub import snapshot_download
|
||||
snapshot_download(repo_id=repo_id, local_dir=model_path, ignore_patterns=["*.md", "*.txt", "onnx", ".git"])
|
||||
|
||||
'''Converter'''
|
||||
|
||||
def cv22ski(cv2_image:np.ndarray) -> np.array:
|
||||
@@ -1518,12 +1526,13 @@ class VITMatteModel:
|
||||
self.processor = processor
|
||||
|
||||
def load_VITMatte_model(model_name:str, local_files_only:bool=False) -> object:
|
||||
if local_files_only:
|
||||
model_name = Path(os.path.join(folder_paths.models_dir, "vitmatte"))
|
||||
# model_name = Path(os.path.join(folder_paths.models_dir, "vitmatte"))
|
||||
model_name = "vitmatte"
|
||||
model_path = os.path.join(folder_paths.models_dir, model_name)
|
||||
model_repo = "hustvl/vitmatte-small-composition-1k"
|
||||
check_and_download_model(model_name, model_repo)
|
||||
from transformers import VitMatteImageProcessor, VitMatteForImageMatting
|
||||
model = VitMatteForImageMatting.from_pretrained(model_name, local_files_only=local_files_only)
|
||||
processor = VitMatteImageProcessor.from_pretrained(model_name, local_files_only=local_files_only)
|
||||
model = VitMatteForImageMatting.from_pretrained(model_path, local_files_only=local_files_only)
|
||||
processor = VitMatteImageProcessor.from_pretrained(model_path, local_files_only=local_files_only)
|
||||
vitmatte = VITMatteModel(model, processor)
|
||||
return vitmatte
|
||||
|
||||
|
||||
Reference in New Issue
Block a user