Skip to content
8 changes: 6 additions & 2 deletions core/common/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -1143,7 +1143,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 +1155,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 +1188,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
83 changes: 66 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,12 @@ 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)
# 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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is probably a defensive check -- a new version is copying concepts from HEAD which should already be indexed/vectorised.

if instance.released:
instance.index_resources_for_self_as_latest_released()
else:
Expand Down Expand Up @@ -696,33 +703,75 @@ 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(ignore_result=True)
def sync_source_concept_vectors(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).
"""
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:
source.sync_concept_vectors_async(recheck=False, countdown=settings.VECTOR_SYNC_RECHECK_SECONDS)


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).
"""
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]))
batch_index_with_summary(
update_concepts_locale_fields, queryset.filter(~in_semantic_version), single_batch, parallel)
with_vectors = queryset.filter(in_semantic_version)
if with_vectors.exists():
batch_index_with_summary(
source.batch_index, with_vectors, ConceptDocument, single_batch=single_batch, parallel=parallel,
**get_batch_index_relations(Concept))


def get_concepts_to_index(source, locales=None, exclude_locale=None):
from core.concepts.models import ConceptName
queryset = source.concepts
Expand Down
32 changes: 31 additions & 1 deletion 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)
from core.concepts.documents import ConceptDocument
from core.concepts.models import Concept
from core.mappings.documents import MappingDocument
Expand Down Expand Up @@ -1291,6 +1291,36 @@ 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_encode_texts_ci_env_returns_none_for_each_text(self, settings_mock):
settings_mock.ENV = 'ci'
self.assertEqual(encode_texts(['a', 'b']), [None, None])
self.assertEqual(encode_texts([]), [])

@patch('core.common.utils.settings')
def test_encode_texts_encodes_in_one_batched_call(self, settings_mock):
settings_mock.ENV = 'production'
settings_mock.LM_ENCODE_BATCH_SIZE = 32
settings_mock.LM.encode.return_value = [[0.1], [0.2]]

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

settings_mock.LM.encode.assert_called_once_with(['a', '2'], batch_size=32)

@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.ENV = 'production'
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
15 changes: 15 additions & 0 deletions core/common/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1019,3 +1019,18 @@ def get_embeddings(txt):
from sentence_transformers import SentenceTransformer
model = SentenceTransformer(settings.LM_MODEL_NAME)
return model.encode(str(txt))


def encode_texts(texts):
"""Embeddings for several texts, from one batched model call (OpenConceptLab/ocl_online#247)."""
texts = [str(text) for text in texts]
if not texts:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

this not texts check should be before texts = [str(text) for text in texts]

return []
if settings.ENV == 'ci':

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

this check should be before texts = [str(text) for text in texts]

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

should we also add "demo" env check here -- it was overridden in settings.py by making it default constant

return [None] * len(texts)

model = settings.LM

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

1032-1035 are common with get_embeddings method above -- can be extracted as a separate method get_LM_model

if not model:
from sentence_transformers import SentenceTransformer
model = SentenceTransformer(settings.LM_MODEL_NAME)
return list(model.encode(texts, batch_size=settings.LM_ENCODE_BATCH_SIZE))

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

wouldn't want to mutate texts input value -- not sure how the caller wants to reuse it

def get_lm_model():
    model = settings.LM
    if not model:
        from sentence_transformers import SentenceTransformer
        model = SentenceTransformer(settings.LM_MODEL_NAME)
   return model

def encode_texts(texts):
    """Embeddings for several texts, from one batched model call (OpenConceptLab/ocl_online#247)."""
    if not texts or not isinstance(texts, (set, list, str)):
        return []
    if settings.ENV == 'ci':
        return [None] * len(texts)

    model = get_lm_model()
    return list(model.encode([str(text) for text in texts], batch_size=settings.LM_ENCODE_BATCH_SIZE))

94 changes: 74 additions & 20 deletions core/concepts/documents.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,23 @@
from itertools import islice

from django.conf import settings
from django_elasticsearch_dsl import Document, fields
from django_elasticsearch_dsl.registries import registry
from pydash import compact, get

from core.common.utils import jsonify_safe, flatten_extras, get_embeddings, drop_version
from core.common.utils import jsonify_safe, flatten_extras, drop_version
from core.concepts.embeddings import ConceptVectors, needs_vectors
from core.concepts.models import Concept

# The text a vector encoded, kept in _source only: reuse reads it back, nothing searches it (ocl_online#247)
EMBEDDING_TEXT = {"type": "keyword", "index": False, "doc_values": False}
# What ocl_online#247 adds to an existing concepts index's mapping (the concept_vector_mapping command)
VECTOR_PROVENANCE_MAPPING = {
'_embeddings': {'type': 'nested', 'properties': {'text': EMBEDDING_TEXT}},
'_synonyms_embeddings': {'type': 'nested', 'properties': {'text': EMBEDDING_TEXT}},
'_embeddings_model': {'type': 'keyword'},
}


@registry.register_document
class ConceptDocument(Document):
Expand Down Expand Up @@ -65,7 +78,8 @@ class Index:
},
"type": {
"type": "text"
}
},
"text": EMBEDDING_TEXT,
}
)
_synonyms_embeddings = fields.NestedField(
Expand All @@ -75,9 +89,17 @@ class Index:
},
"type": {
"type": "text"
}
},
"text": EMBEDDING_TEXT,
}
)
_embeddings_model = fields.KeywordField()

VECTOR_CHUNK_SIZE = 100
_vectors = None # the ConceptVectors of the chunk being prepared, see _get_actions
vectors_reused = 0 # over this document instance's _get_actions calls
texts_encoded = 0
_source_versions = None # (version, match_algorithms) of each repo version the row being prepared belongs to

class Django:
model = Concept
Expand Down Expand Up @@ -168,9 +190,9 @@ def prepare_numeric_id(instance):
def prepare_locale(instance):
return compact(set(instance.active_names.values_list('locale', flat=True)))

@staticmethod
def prepare_source_version(instance):
return list(instance.sources.values_list('version', flat=True))
def prepare_source_version(self, instance):
versions = self._source_versions or instance.sources.values_list('version', 'match_algorithms')
return [version for version, _ in versions]

@staticmethod
def prepare_extras(instance):
Expand Down Expand Up @@ -209,8 +231,33 @@ def prepare_description_types(instance):
def prepare_description(instance):
return '. '.join(compact(set(instance.active_descriptions.values_list('name', flat=True))))

def _get_actions(self, object_list, action):
"""
As django_elasticsearch_dsl's, but prepares 100 docs at a time, so that their vectors are resolved together
(ConceptVectors): reused where the stored docs already hold them, and the rest encoded in one batched call.
"""
if action == 'delete':
yield from super()._get_actions(object_list, action)
return
objects = iter(object_list)
while chunk := list(islice(objects, self.VECTOR_CHUNK_SIZE)):
self._vectors = ConceptVectors(self._index._name)
try:
actions = [self._prepare_action(obj, action) for obj in chunk if self.should_index_object(obj)]

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

can use actions = self.get_actions(chunk, action)

self._vectors.resolve()
self.vectors_reused += self._vectors.reused
self.texts_encoded += self._vectors.encoded
finally:
self._vectors = None
yield from actions

def prepare(self, instance):
data = super().prepare(instance)
self._source_versions = list(instance.sources.values_list('version', 'match_algorithms'))
try:
data = super().prepare(instance)
finally:
versions_match_algorithms = [match_algorithms for _, match_algorithms in self._source_versions]
self._source_versions = None
same_as_mapped_codes, other_mapped_codes, verbose_info = self.get_mapped_codes(instance)
data['same_as_map_codes'] = same_as_mapped_codes
data['other_map_codes'] = other_mapped_codes
Expand All @@ -224,19 +271,8 @@ def prepare(self, instance):
data['synonyms'] = compact(set(n.name for n in synonyms))
data['_synonyms'] = data['synonyms']

if instance.parent.has_semantic_match_algorithm:
data['_embeddings'] = {
'vector': get_embeddings(name),
'type': get(preferred_locale, 'type'),
'locale': get(preferred_locale, 'locale')
}
data['_synonyms_embeddings'] = [
{
'vector': get_embeddings(s.name),
'type': get(s, 'type'),
'locale': get(s, 'locale')
} for s in synonyms
]
if needs_vectors(instance.parent, versions_match_algorithms):
self.add_vectors(data, instance, name, preferred_locale, synonyms)

expansions = list(instance.expansion_set.only('mnemonic', 'uri'))
data['expansion'] = [e.mnemonic for e in expansions]
Expand All @@ -248,6 +284,24 @@ def prepare(self, instance):

return data

def add_vectors(self, data, instance, name, preferred_locale, synonyms): # pylint: disable=too-many-arguments
"""
Vector entries for the display name and each synonym, each with the text it encodes, and the model on the doc.
Their vectors come from the chunk's ConceptVectors once every doc in it is prepared -- or, for a doc prepared
on its own, right away.
"""
data['_embeddings'] = {
'vector': None, 'text': name, 'type': get(preferred_locale, 'type'), 'locale': get(preferred_locale, 'locale')
}
data['_synonyms_embeddings'] = [
{'vector': None, 'text': s.name, 'type': get(s, 'type'), 'locale': get(s, 'locale')} for s in synonyms
]
data['_embeddings_model'] = settings.LM_MODEL_NAME
vectors = self._vectors or ConceptVectors(self._index._name)
vectors.add(instance, [data['_embeddings'], *data['_synonyms_embeddings']])
if self._vectors is None:
vectors.resolve()

@staticmethod
def get_mapped_codes(instance):
mappings = instance.get_unidirectional_mappings()
Expand Down
Loading
Loading