Skip to content
Merged
2 changes: 2 additions & 0 deletions core/common/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -146,3 +146,5 @@
ALL = '*'
CANONICAL_URL_REQUEST_PARAM = 'canonicalUrl'
SAME_STANDARD_CHECKSUM_ERROR = 'No changes detected. Standard checksum is same as last version.'
# Docs per bulk request when (re)indexing or syncing a repo's resources (batch_index_full, sync_concept_vectors)
INDEX_BATCH_SIZE = 500
20 changes: 13 additions & 7 deletions core/common/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,8 @@
ACCESS_TYPE_VIEW, SUPER_ADMIN_USER_ID,
HEAD, PERSIST_NEW_ERROR_MESSAGE, SOURCE_PARENT_CANNOT_BE_NONE, PARENT_RESOURCE_CANNOT_BE_NONE,
CREATOR_CANNOT_BE_NONE, CANNOT_DELETE_ONLY_VERSION, OPENMRS_VALIDATION_SCHEMA, VALIDATION_SCHEMAS,
DEFAULT_VALIDATION_SCHEMA, ES_REQUEST_TIMEOUT, UPDATED_BY_USERNAME_PARAM, ES_RETRYABLE_ERROR_TYPES)
DEFAULT_VALIDATION_SCHEMA, ES_REQUEST_TIMEOUT, UPDATED_BY_USERNAME_PARAM, ES_RETRYABLE_ERROR_TYPES,
INDEX_BATCH_SIZE)
from .es import ESScript
from .exceptions import Http400, BatchIndexingError
from .fields import URIField
Expand Down Expand Up @@ -408,8 +409,9 @@ def get_actions(batch_ids):
def batch_index_full( # pylint: disable=too-many-arguments
single_batch: bool, queryset, document, prefetch, select_related, parallel=True, refresh=None):
"""
Full (re)index, 500 docs a batch (single_batch: all in one). Every batch is attempted; if any failed, raises
BatchIndexingError with the counts once they all have (see BatchIndexRun). Returns the run's summary.
Full (re)index, INDEX_BATCH_SIZE docs a batch (single_batch: all in one). Every batch is attempted; if any
failed, raises BatchIndexingError with the counts once they all have (see BatchIndexRun). Returns the run's
summary.
refresh=False doesn't make ES refresh after each batch (the document's auto_refresh does by default).
"""
if get(settings, 'TEST_MODE', False):
Expand All @@ -435,7 +437,7 @@ def index_batch(objects):
if single_batch:
run.attempt(0, list(queryset.all()), index_batch)
else:
batch_size = 500
batch_size = INDEX_BATCH_SIZE
start = 0
while True:
batch = list(queryset.order_by('-id')[start:start+batch_size])
Expand Down Expand Up @@ -479,7 +481,7 @@ def index_batch(ids):
if single_batch:
run.attempt(0, list(queryset.all().values_list('id', flat=True)), index_batch)
else:
batch_size = 500
batch_size = INDEX_BATCH_SIZE
start = 0
id_qs = queryset.order_by('-id').values_list('id', flat=True)
while True:
Expand Down Expand Up @@ -1143,7 +1145,7 @@ def index_resources_for_self_as_latest_released(self, only_update=False): # pyl
pass

@classmethod
def persist_changes(cls, obj, updated_by, original_schema, **kwargs): # pylint: disable=too-many-locals
def persist_changes(cls, obj, updated_by, original_schema, **kwargs): # pylint: disable=too-many-locals,too-many-branches
errors = {}
parent_resource = kwargs.pop('parent_resource', obj.parent)
if not parent_resource:
Expand All @@ -1155,6 +1157,8 @@ def persist_changes(cls, obj, updated_by, original_schema, **kwargs): # pylint:
is_source = cls.__name__ == 'Source'
should_reindex_resources = is_source and obj.released != original_repo.released
concepts_reindex_filters = obj.get_concepts_reindex_filters(original_repo) if is_source else None
should_sync_vectors = is_source and (
bool(obj.has_semantic_match_algorithm) != bool(original_repo.has_semantic_match_algorithm))

obj._should_update_public_access = is_source and obj.public_access != original_repo.public_access # pylint: disable=protected-access
obj._should_update_is_active = is_source and obj.is_active != original_repo.is_active # pylint: disable=protected-access
Expand Down Expand Up @@ -1186,8 +1190,10 @@ def persist_changes(cls, obj, updated_by, original_schema, **kwargs): # pylint:
obj.index_resources_for_self_as_latest_released(only_update=True)
else:
obj.index_resources_for_self_as_unreleased()
elif concepts_reindex_filters is not None:
if concepts_reindex_filters is not None:
obj.index_concepts_async(obj.updated_by, **concepts_reindex_filters)
if should_sync_vectors:
obj.sync_concept_vectors_async(obj.updated_by)

except IntegrityError as ex:
errors.update({'__all__': ex.args})
Expand Down
125 changes: 108 additions & 17 deletions core/common/tasks.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@
from core.common.exceptions import BatchIndexingError
from core.common.utils import write_export_file, web_url, get_resource_class_from_resource_name, get_export_service, \
get_date_range_label
from core.concepts.embeddings import sync_concept_vectors
from core.reports.models import ResourceUsageReport
from core.tasks.models import QueueOnceCustomTask, Task

Expand Down Expand Up @@ -411,6 +412,15 @@ def seed_children_to_new_version(self, resource, obj_id, export=True, sync=False
instance.seed_concepts(index=False)
instance.seed_mappings(index=False)
instance.update_children_counts(sync)
# A semantic version's vectors come from this sync. Seeding copies HEAD's membership only, and the
# indexing below only appends this version to its members' docs, a partial update that never
# computes vectors: with a lexical HEAD, those docs have none, so this sync is what embeds the
# version; with a semantic HEAD it's a cheap, read-only check. The flag is read now, not as it was
# when the task started: an opt-in while seeding ran synced before any of this version's members
# existed. Whether or not it's still the latest release, and however its indexing below goes.
instance.refresh_from_db(fields=['match_algorithms'])
if instance.has_semantic_match_algorithm:
instance.sync_concept_vectors_async(instance.created_by)
Comment thread
snyaggarwal marked this conversation as resolved.
if instance.released:
instance.index_resources_for_self_as_latest_released()
else:
Expand Down Expand Up @@ -696,33 +706,114 @@ def index_source_concepts( # pylint: disable=too-many-arguments,too-many-locals
from core.sources.models import Source
source = Source.objects.filter(id=source_id).first()
if source:
from core.concepts.documents import ConceptDocument
prefetch = ['sources', 'names', 'descriptions'] if should_prefetch else []
select_related = [
'parent', 'parent__organization', 'parent__user', 'created_by', 'updated_by'
] if should_select_related else []
narrowed = bool(locales or exclude_locale)
queryset = get_concepts_to_index(source, locales, exclude_locale)
if (locales or exclude_locale) and not partial_doc and not source.has_semantic_match_algorithm:
batch_index_with_summary(update_concepts_locale_fields, queryset, single_batch, parallel)
source.clear_concepts_cache()
return
try:
kwargs = {'partial_doc': partial_doc} if partial_doc else {
'prefetch': prefetch, 'select_related': select_related}
kwargs['single_batch'] = single_batch
kwargs['parallel'] = parallel
batch_index_with_summary(source.batch_index, queryset, ConceptDocument, **kwargs)
except Exception as ex: # pragma: no cover
if not partial_doc or (isinstance(ex, BatchIndexingError) and ex.rejected):
raise
logger.exception('Falling back to full concept reindex for source %s', source_id)
batch_index_with_summary(
source.batch_index, queryset, ConceptDocument, prefetch=prefetch, select_related=select_related,
parallel=parallel)
if narrowed and not partial_doc and not source.has_semantic_match_algorithm:
index_concepts_locale_change(source, queryset, single_batch, parallel)
else:
index_concepts(source, queryset, partial_doc, single_batch, parallel, prefetch, select_related)
finally:
source.clear_concepts_cache()


@app.task(bind=True, ignore_result=True)
def sync_source_concept_vectors(self, source_id, recheck=True):
"""
Gives a repo version's concept docs the vectors they need and nothing else (sync_concept_vectors). Unless this
is the recheck, then queues one more sync VECTOR_SYNC_RECHECK_SECONDS later, even when this one failed: a doc
written from the flags as they were just before the change that queued this sync can land after this sync has
checked it, and the recheck sees it (OpenConceptLab/ocl_online#247). The recheck is this task's child, so
revoking or rerunning this task revokes it too. Not a chained task, which wouldn't run after a failure.
"""
from core.sources.models import Source
source = Source.objects.filter(id=source_id).first()
if source:
try:
batch_index_with_summary(sync_concept_vectors, source)
finally:
if recheck:
recheck_task = source.sync_concept_vectors_async(
recheck=False, countdown=settings.VECTOR_SYNC_RECHECK_SECONDS)
task = Task.objects.filter(id=self.request.id).first() if recheck_task else None
if task:
task.children = [*(task.children or []), recheck_task.id]
task.save(update_fields=['children'])


def index_concepts(source, queryset, partial_doc, single_batch, parallel, prefetch, select_related): # pylint: disable=too-many-arguments
from core.concepts.documents import ConceptDocument
try:
kwargs = {'partial_doc': partial_doc} if partial_doc else {
'prefetch': prefetch, 'select_related': select_related}
kwargs['single_batch'] = single_batch
kwargs['parallel'] = parallel
batch_index_with_summary(source.batch_index, queryset, ConceptDocument, **kwargs)
except Exception as ex: # pragma: no cover
if not partial_doc or (isinstance(ex, BatchIndexingError) and ex.rejected):
raise
logger.exception('Falling back to full concept reindex for source %s', source.id)
batch_index_with_summary(
source.batch_index, queryset, ConceptDocument, prefetch=prefetch, select_related=select_related,
parallel=parallel)


def index_concepts_locale_change(source, queryset, single_batch, parallel):
"""
A locale change on a repo whose HEAD isn't semantic only updates the docs' name fields -- except for the docs of
rows in a semantic version, which carry vectors: they're rebuilt, so their display-name vector follows the new
display name (reusing the vectors their names already have). Each set is attempted even when the other fails, and
the task's summary counts both (index_in_parts). Returns it.
"""
from core.concepts.documents import ConceptDocument
from core.concepts.models import Concept
in_semantic_version = Exists(Concept.sources.through.objects.filter(
concept_id=OuterRef('id'), source__match_algorithms__contains=[source.SEMANTIC_MATCH_ALGORITHM]))
with_vectors = queryset.filter(in_semantic_version)

def rebuild_with_vectors():
if not with_vectors.exists():
return None
return source.batch_index(
with_vectors, ConceptDocument, single_batch=single_batch, parallel=parallel,
**get_batch_index_relations(Concept))

return batch_index_with_summary(
index_in_parts,
lambda: update_concepts_locale_fields(queryset.filter(~in_semantic_version), single_batch, parallel),
rebuild_with_vectors)


def index_in_parts(*parts):
"""
Runs each part, a batch indexing call that returns its run's summary, even after an earlier part failed, and
returns their summaries summed. If any part failed, raises once they all have run: BatchIndexingError with the
summed summary when every failure was one, else the first other error.
"""
summaries, failures = [], []
for part in parts:
try:
summaries.append(part())
except Exception as ex: # pylint: disable=broad-except
failures.append(ex)
summaries.append(getattr(ex, 'summary', None))
summaries = [summary for summary in summaries if isinstance(summary, dict)]
summary = {
key: sum(get(part_summary, key, 0) for part_summary in summaries)
for key in dict.fromkeys(key for part_summary in summaries for key in part_summary)
} if summaries else None
if not failures:
return summary
if all(isinstance(ex, BatchIndexingError) for ex in failures):
raise BatchIndexingError(
'; '.join(str(ex) for ex in failures), summary, any(ex.rejected for ex in failures)) from failures[0]
raise next(ex for ex in failures if not isinstance(ex, BatchIndexingError))


def get_concepts_to_index(source, locales=None, exclude_locale=None):
from core.concepts.models import ConceptName
queryset = source.concepts
Expand Down
56 changes: 53 additions & 3 deletions core/common/tests.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,7 +62,7 @@
split_list_by_condition, is_zip_file, get_date_range_label, get_prev_month, from_string_to_date, get_end_of_month,
get_start_of_month, es_id_in, web_url, get_queue_task_names, get_resource_class_from_resource_uri, encode_string,
to_parent_kwargs_from_uri, reverse_resource, reverse_resource_version, write_export_file, queue_bulk_import,
get_bulk_import_celery_once_lock_key, generic_sort, get_embeddings)
get_bulk_import_celery_once_lock_key, generic_sort, get_embeddings, encode_texts, get_lm_model)
from core.concepts.documents import ConceptDocument
from core.concepts.models import Concept
from core.mappings.documents import MappingDocument
Expand Down Expand Up @@ -1272,13 +1272,14 @@ def test_generic_sort(self):

@patch('core.common.utils.settings')
def test_get_embeddings_ci_env_returns_none(self, settings_mock):
settings_mock.ENV = 'ci'
settings_mock.LM_DISABLED = True
self.assertIsNone(get_embeddings('some text'))
settings_mock.LM.encode.assert_not_called()

@patch('sentence_transformers.SentenceTransformer')
@patch('core.common.utils.settings')
def test_get_embeddings_loads_model_when_not_ci(self, settings_mock, sentence_transformer_mock):
settings_mock.ENV = 'production'
settings_mock.LM_DISABLED = False
settings_mock.LM = None
settings_mock.LM_MODEL_NAME = 'some-model'
model_instance_mock = Mock()
Expand All @@ -1291,6 +1292,55 @@ def test_get_embeddings_loads_model_when_not_ci(self, settings_mock, sentence_tr
model_instance_mock.encode.assert_called_once_with('some text')
self.assertEqual(result, [0.1, 0.2])

@patch('core.common.utils.settings')
def test_get_lm_model_is_the_loaded_model(self, settings_mock):
self.assertIs(get_lm_model(), settings_mock.LM)

@patch('sentence_transformers.SentenceTransformer')
@patch('core.common.utils.settings')
def test_get_lm_model_loads_it_when_not_loaded(self, settings_mock, sentence_transformer_mock):
settings_mock.LM = None
settings_mock.LM_MODEL_NAME = 'some-model'

self.assertIs(get_lm_model(), sentence_transformer_mock.return_value)

sentence_transformer_mock.assert_called_once_with('some-model')

@patch('core.common.utils.settings')
def test_encode_texts_returns_none_for_each_text_when_the_model_is_disabled(self, settings_mock):
settings_mock.LM_DISABLED = True
self.assertEqual(encode_texts(['a', 'b']), [None, None])
self.assertEqual(encode_texts([]), [])
settings_mock.LM.encode.assert_not_called()

@patch('core.common.utils.settings')
def test_encode_texts_encodes_in_one_batched_call(self, settings_mock):
settings_mock.LM_DISABLED = False
settings_mock.LM_ENCODE_BATCH_SIZE = 32
settings_mock.LM.encode.return_value = [[0.1], [0.2]]
texts = ['a', 2]

self.assertEqual(encode_texts(texts), [[0.1], [0.2]])

settings_mock.LM.encode.assert_called_once_with(['a', '2'], batch_size=32)
self.assertEqual(texts, ['a', 2])
self.assertEqual(encode_texts([]), [])
settings_mock.LM.encode.assert_called_once() # not for no texts

@patch('sentence_transformers.SentenceTransformer')
@patch('core.common.utils.settings')
def test_encode_texts_loads_model_when_not_loaded(self, settings_mock, sentence_transformer_mock):
settings_mock.LM_DISABLED = False
settings_mock.LM = None
settings_mock.LM_MODEL_NAME = 'some-model'
settings_mock.LM_ENCODE_BATCH_SIZE = 64
sentence_transformer_mock.return_value.encode.return_value = [[0.1]]

self.assertEqual(encode_texts(['a']), [[0.1]])

sentence_transformer_mock.assert_called_once_with('some-model')
sentence_transformer_mock.return_value.encode.assert_called_once_with(['a'], batch_size=64)


class BaseModelTest(OCLTestCase):
def test_model_name(self):
Expand Down
26 changes: 21 additions & 5 deletions core/common/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1010,12 +1010,28 @@ def clean_term(term):
return term.lower().replace(' ', '').replace('-', '').replace('_', '')


def get_embeddings(txt):
if settings.ENV == 'ci':
return None

def get_lm_model():
"""The language model: the one loaded at startup, or else loaded now."""
model = settings.LM
if not model:
from sentence_transformers import SentenceTransformer
model = SentenceTransformer(settings.LM_MODEL_NAME)
return model.encode(str(txt))
return model


def get_embeddings(txt):
if settings.LM_DISABLED:
return None
return get_lm_model().encode(str(txt))


def encode_texts(texts):
"""
Embeddings for a list of texts, from one batched model call (OpenConceptLab/ocl_online#247). Where the model is
disabled, None for each text.
"""
if not texts:
Comment thread
snyaggarwal marked this conversation as resolved.
return []
if settings.LM_DISABLED:
return [None] * len(texts)
return list(get_lm_model().encode([str(text) for text in texts], batch_size=settings.LM_ENCODE_BATCH_SIZE))
Loading
Loading