Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 6 additions & 2 deletions medcat-den/src/medcat_den/wrappers.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,11 @@
from typing import Union, Optional
from typing import Union, Optional, Type

from medcat.cat import CAT
from medcat.utils.defaults import DEFAULT_PACK_NAME
from medcat.storage.serialisers import AvailableSerialisers
from medcat.trainer import Trainer
from medcat.data.mctexport import MedCATTrainerExport
from medcat.components.addons.addons import AddonComponent

from medcat_den.base import ModelInfo
from medcat_den.config import DenConfig, RemoteDenConfig
Expand Down Expand Up @@ -90,6 +91,7 @@ def trainer(self) -> Trainer:
def load_model_pack(cls, model_pack_path: str,
config_dict: Optional[dict] = None,
addon_config_dict: Optional[dict[str, dict]] = None,
keep_addons_of_types: Optional[list[Type[AddonComponent]]] = None,
model_info: Optional[ModelInfo] = None,
den_cnf: Optional[DenConfig] = None,
) -> 'CAT':
Expand Down Expand Up @@ -121,7 +123,9 @@ def load_model_pack(cls, model_pack_path: str,
CAT: The loaded model pack.
"""
_cat = super().load_model_pack(
model_pack_path, config_dict, addon_config_dict)
model_pack_path, config_dict, addon_config_dict,
keep_addons_of_types=keep_addons_of_types,
)
cat = cls(_cat)
if model_info is None:
raise CannotWrapModel("Model info must be provided")
Expand Down
Loading