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
27 changes: 25 additions & 2 deletions sciencebeam_parser/config/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
import os
import copy
from pathlib import Path
from typing import Any, Optional, Union
from typing import Any, Optional, Tuple, Union

import yaml

Expand All @@ -27,6 +27,29 @@ def _deep_merge(base: dict, overlay: dict) -> dict:
return result


def _resolve_sequence_model_profile(
seq_profiles: dict,
name: str,
_seen: Tuple[str, ...] = ()
) -> dict:
if name not in seq_profiles:
raise ValueError(
f'Unknown sequence_model_profile {name!r}. Available: {sorted(seq_profiles)}'
)
if name in _seen:
raise ValueError(
f'Circular extends detected for sequence_model_profile {name!r} '
f'(chain: {" -> ".join([*_seen, name])})'
)
profile = seq_profiles[name]
base_name = profile.get('extends')
overlay = {key: value for key, value in profile.items() if key != 'extends'}
if not base_name:
return overlay
base = _resolve_sequence_model_profile(seq_profiles, base_name, _seen + (name,))
return _deep_merge(base, overlay)


class AppConfig:
def __init__(self, props: dict):
self.props = props
Expand Down Expand Up @@ -87,7 +110,7 @@ def resolve_profile(self, profile_name: Optional[str] = None) -> 'AppConfig':
f'Profile {resolved!r} references unknown sequence_model_profile '
f'{seq_name!r}. Available: {sorted(seq_profiles)}'
)
overlay['models'] = seq_profiles[seq_name]
overlay['models'] = _resolve_sequence_model_profile(seq_profiles, seq_name)

for key, value in profile.items():
if key != 'sequence_models':
Expand Down
7 changes: 7 additions & 0 deletions sciencebeam_parser/resources/default_config/config.yml
Original file line number Diff line number Diff line change
Expand Up @@ -179,6 +179,8 @@ sequence_model_profiles:
citation:
path: 'https://github.com/kermitt2/grobid/raw/0.9.0/grobid-home/models/citation'
engine: 'wapiti'
grobid_custom_hybrid:
extends: grobid_crf_0_9_0

profiles:
biorxiv_elife:
Expand All @@ -188,6 +190,11 @@ profiles:
processors:
fulltext:
noise_filter_enabled: false
grobid_custom_hybrid:
sequence_models: grobid_custom_hybrid
processors:
fulltext:
noise_filter_enabled: false

profile_aliases:
grobid_crf: grobid_crf_0_9_0
Expand Down
89 changes: 88 additions & 1 deletion tests/config/config_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,12 @@
import pytest
import yaml

from sciencebeam_parser.config.config import AppConfig, _deep_merge
from sciencebeam_parser.config.config import (
AppConfig,
_deep_merge,
_resolve_sequence_model_profile
)
from sciencebeam_parser.resources.default_config import DEFAULT_CONFIG_FILE


MINIMAL_PROFILE_CONFIG = {
Expand All @@ -18,10 +23,15 @@
'segmentation': {'path': 'path_b/segmentation', 'engine': 'wapiti'},
'header': {'path': 'path_b/header', 'engine': 'wapiti'},
},
'profile_b_extended': {
'extends': 'profile_b',
'header': {'path': 'path_b_extended/header', 'engine': 'wapiti'},
},
},
'profiles': {
'profile_a': {'sequence_models': 'profile_a'},
'profile_b': {'sequence_models': 'profile_b'},
'profile_b_extended': {'sequence_models': 'profile_b_extended'},
'profile_with_extra': {
'sequence_models': 'profile_a',
'processors': {'fulltext': {'use_cv_model': True}},
Expand Down Expand Up @@ -67,6 +77,69 @@ def test_does_not_mutate_base(self):
assert base['a']['b'] == 1


class TestResolveSequenceModelProfile:
def test_returns_profile_without_extends_unchanged(self):
seq_profiles = {
'base': {'segmentation': {'path': 'base/segmentation'}},
}
result = _resolve_sequence_model_profile(seq_profiles, 'base')
assert result == {'segmentation': {'path': 'base/segmentation'}}

def test_merges_extended_profile(self):
seq_profiles = {
'base': {
'segmentation': {'path': 'base/segmentation', 'engine': 'wapiti'},
'header': {'path': 'base/header', 'engine': 'wapiti'},
},
'child': {
'extends': 'base',
'header': {'path': 'child/header'},
},
}
result = _resolve_sequence_model_profile(seq_profiles, 'child')
assert result['segmentation'] == {'path': 'base/segmentation', 'engine': 'wapiti'}
assert result['header'] == {'path': 'child/header', 'engine': 'wapiti'}
assert 'extends' not in result

def test_supports_chained_extends(self):
seq_profiles = {
'grandparent': {'segmentation': {'path': 'gp/segmentation'}},
'parent': {'extends': 'grandparent', 'header': {'path': 'p/header'}},
'child': {'extends': 'parent', 'table': {'path': 'c/table'}},
}
result = _resolve_sequence_model_profile(seq_profiles, 'child')
assert result['segmentation'] == {'path': 'gp/segmentation'}
assert result['header'] == {'path': 'p/header'}
assert result['table'] == {'path': 'c/table'}

def test_raises_on_unknown_profile(self):
with pytest.raises(ValueError, match='Unknown sequence_model_profile'):
_resolve_sequence_model_profile({}, 'missing')

def test_raises_on_unknown_extends_target(self):
seq_profiles = {'child': {'extends': 'missing'}}
with pytest.raises(ValueError, match='Unknown sequence_model_profile'):
_resolve_sequence_model_profile(seq_profiles, 'child')

def test_raises_on_circular_extends(self):
seq_profiles = {
'a': {'extends': 'b'},
'b': {'extends': 'a'},
}
with pytest.raises(ValueError, match='Circular extends'):
_resolve_sequence_model_profile(seq_profiles, 'a')

def test_reports_the_circular_chain_in_the_order_it_was_followed(self):
seq_profiles = {
'alpha': {'extends': 'beta'},
'beta': {'extends': 'gamma'},
'gamma': {'extends': 'alpha'},
}
with pytest.raises(ValueError) as exc_info:
_resolve_sequence_model_profile(seq_profiles, 'alpha')
assert 'chain: alpha -> beta -> gamma -> alpha' in str(exc_info.value)


class TestAppConfigResolveProfile:
def _make_config(self, extra: Optional[dict] = None) -> AppConfig:
props = dict(MINIMAL_PROFILE_CONFIG)
Expand All @@ -80,6 +153,11 @@ def test_applies_sequence_model_profile(self):
assert config['models']['segmentation']['engine'] == 'wapiti'
assert config['models']['header']['path'] == 'path_b/header'

def test_applies_extended_sequence_model_profile(self):
config = self._make_config().resolve_profile('profile_b_extended')
assert config['models']['segmentation']['path'] == 'path_b/segmentation'
assert config['models']['header']['path'] == 'path_b_extended/header'

def test_inherits_base_model_keys_not_in_profile(self):
config = self._make_config().resolve_profile('profile_a')
assert config['models']['segmentation']['path'] == 'path_a/segmentation'
Expand Down Expand Up @@ -215,3 +293,12 @@ def test_should_override_bool_value_with_env_var(
config = AppConfig.load_yaml(str(config_path))
config = config.apply_environment_variables()
assert config.props['key1'] is False


class TestDefaultConfigProfiles:
def _resolve_models(self, profile_name: str) -> dict:
config = AppConfig.load_yaml(DEFAULT_CONFIG_FILE)
return config.resolve_profile(profile_name)['models']

def test_grobid_custom_hybrid_inherits_every_grobid_crf_model(self):
assert self._resolve_models('grobid_custom_hybrid') == self._resolve_models('grobid_crf')
Loading