From 26db8262a81d7418e79e8b3a70caf7c94bee6ecb Mon Sep 17 00:00:00 2001 From: Jim Lee <21117946+jimlee2048@users.noreply.github.com> Date: Thu, 14 Nov 2024 03:30:48 +0800 Subject: [PATCH] refactor: improve model loading logic for existing BiRefNet-General model --- py/birefnet_ultra_v2.py | 23 ++++++++++++----------- 1 file changed, 12 insertions(+), 11 deletions(-) diff --git a/py/birefnet_ultra_v2.py b/py/birefnet_ultra_v2.py index 04d1f76..0602978 100755 --- a/py/birefnet_ultra_v2.py +++ b/py/birefnet_ultra_v2.py @@ -80,18 +80,19 @@ class LS_LoadBiRefNetModelV2: os.makedirs(birefnet_path, exist_ok=True) model_path = os.path.join(birefnet_path, version) - old_model_path = os.path.join(birefnet_path, 'pth') - if version == "BiRefNet-General" and os.path.exists(old_model_path): - model = "BiRefNet-general-epoch_244.pth" - from .BiRefNet_v2.models.birefnet import BiRefNet - from .BiRefNet_v2.utils import check_state_dict - model_dict = get_models() - self.birefnet = BiRefNet(bb_pretrained=False) - self.state_dict = torch.load(model_dict[model], map_location='cpu', weights_only=True) - self.state_dict = check_state_dict(self.state_dict) - self.birefnet.load_state_dict(self.state_dict) - return (self.birefnet,) + if version == "BiRefNet-General": + old_birefnet_path = os.path.join(birefnet_path, 'pth') + old_model = "BiRefNet-general-epoch_244.pth" + old_model_path = os.path.join(old_birefnet_path, old_model) + if os.path.exists(old_model_path): + from .BiRefNet_v2.models.birefnet import BiRefNet + from .BiRefNet_v2.utils import check_state_dict + self.birefnet = BiRefNet(bb_pretrained=False) + self.state_dict = torch.load(old_model_path, map_location='cpu', weights_only=True) + self.state_dict = check_state_dict(self.state_dict) + self.birefnet.load_state_dict(self.state_dict) + return (self.birefnet,) elif not os.path.exists(model_path): log(f"Downloading {version} model...") repo_id = self.birefnet_model_repos[version]