Skip to content

Commit aa66d5c

Browse files
authored
Merge pull request #4168 from alejoe91/fix-pydantic-validator
Fix pydantic model after validator
2 parents 64fff45 + 7aa5cc0 commit aa66d5c

1 file changed

Lines changed: 13 additions & 15 deletions

File tree

src/spikeinterface/curation/curation_model.py

Lines changed: 13 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -378,19 +378,17 @@ def validate_fields(cls, values):
378378
return values
379379

380380
@model_validator(mode="after")
381-
def validate_curation_dict(cls, values):
382-
if values.format_version not in values.supported_versions:
381+
def validate_curation_dict(self):
382+
if self.format_version not in self.supported_versions:
383383
raise ValueError(
384-
f"Format version {values.format_version} not supported. Only {values.supported_versions} are valid"
384+
f"Format version {self.format_version} not supported. Only {self.supported_versions} are valid"
385385
)
386386

387-
labeled_unit_set = set([lbl.unit_id for lbl in values.manual_labels]) if values.manual_labels else set()
388-
merged_units_set = (
389-
set(chain.from_iterable(merge.unit_ids for merge in values.merges)) if values.merges else set()
390-
)
391-
split_units_set = set(split.unit_id for split in values.splits) if values.splits else set()
392-
removed_set = set(values.removed) if values.removed else set()
393-
unit_ids = values.unit_ids
387+
labeled_unit_set = set([lbl.unit_id for lbl in self.manual_labels]) if self.manual_labels else set()
388+
merged_units_set = set(chain.from_iterable(merge.unit_ids for merge in self.merges)) if self.merges else set()
389+
split_units_set = set(split.unit_id for split in self.splits) if self.splits else set()
390+
removed_set = set(self.removed) if self.removed else set()
391+
unit_ids = self.unit_ids
394392

395393
unit_set = set(unit_ids)
396394
if not labeled_unit_set.issubset(unit_set):
@@ -403,7 +401,7 @@ def validate_curation_dict(cls, values):
403401
raise ValueError("Curation format: some removed units are not in the unit list")
404402

405403
# Check for units being merged multiple times
406-
all_merging_groups = [set(merge.unit_ids) for merge in values.merges] if values.merges else []
404+
all_merging_groups = [set(merge.unit_ids) for merge in self.merges] if self.merges else []
407405
for gp_1, gp_2 in combinations(all_merging_groups, 2):
408406
if len(gp_1.intersection(gp_2)) != 0:
409407
raise ValueError("Curation format: some units belong to multiple merge groups")
@@ -416,19 +414,19 @@ def validate_curation_dict(cls, values):
416414
if len(merged_units_set.intersection(split_units_set)) != 0:
417415
raise ValueError("Curation format: some units were both merged and split")
418416

419-
for manual_label in values.manual_labels:
420-
for label_key in values.label_definitions.keys():
417+
for manual_label in self.manual_labels:
418+
for label_key in self.label_definitions.keys():
421419
if label_key in manual_label.labels:
422420
unit_id = manual_label.unit_id
423421
label_value = manual_label.labels[label_key]
424422
if not isinstance(label_value, list):
425423
raise ValueError(f"Curation format: manual_labels {unit_id} is invalid should be a list")
426424

427-
is_exclusive = values.label_definitions[label_key].exclusive
425+
is_exclusive = self.label_definitions[label_key].exclusive
428426

429427
if is_exclusive and not len(label_value) <= 1:
430428
raise ValueError(
431429
f"Curation format: manual_labels {unit_id} {label_key} are exclusive labels. {label_value} is invalid"
432430
)
433431

434-
return values
432+
return self

0 commit comments

Comments
 (0)