Skip to content
Merged
Show file tree
Hide file tree
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
4 changes: 4 additions & 0 deletions medcat-trainer/webapp/api/api/admin/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,6 +96,10 @@ def _set_proj_from_group(self, proj: ProjectAnnotateEntities, group: ProjectGrou
proj.project_status = group.project_status
proj.concept_db = group.concept_db
proj.vocab = group.vocab
proj.model_pack = group.model_pack
proj.deid_model_annotation = group.deid_model_annotation
proj.use_model_service = group.use_model_service
proj.model_service_url = group.model_service_url
proj.require_entity_validation = group.require_entity_validation
proj.train_model_on_submit = group.train_model_on_submit
proj.add_new_entities = group.add_new_entities
Expand Down
9 changes: 9 additions & 0 deletions medcat-trainer/webapp/api/api/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -499,7 +499,16 @@ class Meta:
model_service_url = models.CharField(max_length=500, blank=True, null=True,
help_text='URL of the remote MedCAT service API (e.g., http://medcat-service:8000)')

def _normalize_model_config(self):
"""ModelPack and CDB/Vocab are mutually exclusive configuration options."""
if self.model_pack_id:
self.concept_db = None
self.vocab = None
elif self.concept_db_id and self.vocab_id:
self.model_pack = None

def save(self, *args, **kwargs):
self._normalize_model_config()
# If using remote model service, skip local model validation
if not self.use_model_service:
if self.model_pack is None and (self.concept_db is None or self.vocab is None):
Expand Down
89 changes: 89 additions & 0 deletions medcat-trainer/webapp/api/api/tests/test_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,9 @@
MetaAnnotation,
MetaTask,
MetaTaskValue,
ModelPack,
ProjectAnnotateEntities,
ProjectGroup,
Relation,
Vocabulary,
cdb_name_validator,
Expand Down Expand Up @@ -103,6 +105,93 @@ def test_use_model_service_with_url_skips_model_validation(self):
proj.save()
self.assertIsNotNone(proj.id)

def test_save_with_model_pack_clears_stale_cdb_vocab(self):
mp = ModelPack(name='normalize-mp')
mp.save(skip_load=True)
proj = self._new_project(concept_db=self.cdb, vocab=self.vocab, model_pack=mp)
proj.save()
proj.refresh_from_db()
self.assertEqual(proj.model_pack_id, mp.id)
self.assertIsNone(proj.concept_db_id)
self.assertIsNone(proj.vocab_id)

def test_save_amended_model_pack_with_stale_cdb_vocab_succeeds(self):
mp_a = ModelPack(name='normalize-mp-a')
mp_a.save(skip_load=True)
mp_b = ModelPack(name='normalize-mp-b')
mp_b.save(skip_load=True)
proj = self._new_project(concept_db=self.cdb, vocab=self.vocab, model_pack=mp_a)
proj.save()
proj.model_pack = mp_b
proj.save()
proj.refresh_from_db()
self.assertEqual(proj.model_pack_id, mp_b.id)
self.assertIsNone(proj.concept_db_id)
self.assertIsNone(proj.vocab_id)

def test_save_prefers_model_pack_when_cdb_vocab_also_set(self):
# When both are present (e.g. stale CDB/Vocab after a ModelPack change),
# ModelPack wins and CDB/Vocab are cleared.
mp = ModelPack(name='normalize-mp-clear')
mp.save(skip_load=True)
proj = self._new_project(model_pack=mp)
proj.save()
proj.concept_db = self.cdb
proj.vocab = self.vocab
proj.save()
proj.refresh_from_db()
self.assertEqual(proj.model_pack_id, mp.id)
self.assertIsNone(proj.concept_db_id)
self.assertIsNone(proj.vocab_id)

def test_save_with_cdb_vocab_after_clearing_model_pack(self):
mp = ModelPack(name='normalize-mp-switch')
mp.save(skip_load=True)
proj = self._new_project(model_pack=mp)
proj.save()
proj.model_pack = None
proj.concept_db = self.cdb
proj.vocab = self.vocab
proj.save()
proj.refresh_from_db()
self.assertIsNone(proj.model_pack_id)
self.assertEqual(proj.concept_db_id, self.cdb.id)
self.assertEqual(proj.vocab_id, self.vocab.id)


@override_settings(MEDIA_ROOT='/tmp/mct-tests-models')
class ProjectGroupModelConfigValidationTests(TestCase):
@classmethod
def setUpTestData(cls):
cdb = ConceptDB(name='pg_val_cdb', cdb_file='pg_val_cdb.dat')
cdb.save(skip_load=True)
vocab = Vocabulary(name='pg_val_vocab', vocab_file='pg_val_vocab.dat')
vocab.save(skip_load=True)
cls.cdb = cdb
cls.vocab = vocab
cls.dataset = create_dataset(name='pg_val_ds', file_name='pg_val_ds.csv')

def test_save_amended_model_pack_clears_stale_cdb_vocab(self):
mp_a = ModelPack(name='pg-mp-a')
mp_a.save(skip_load=True)
mp_b = ModelPack(name='pg-mp-b')
mp_b.save(skip_load=True)
group = ProjectGroup(
name='pg-switch-model-pack',
dataset=self.dataset,
concept_db=self.cdb,
vocab=self.vocab,
model_pack=mp_a,
cuis='',
)
group.save()
group.model_pack = mp_b
group.save()
group.refresh_from_db()
self.assertEqual(group.model_pack_id, mp_b.id)
self.assertIsNone(group.concept_db_id)
self.assertIsNone(group.vocab_id)


@override_settings(MEDIA_ROOT='/tmp/mct-tests-models')
class AnnotatedEntitySaveUpdatesProjectTests(TestCase):
Expand Down
Loading