diff --git a/core/common/models.py b/core/common/models.py index 50b8d5b6..a237e874 100644 --- a/core/common/models.py +++ b/core/common/models.py @@ -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: @@ -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 @@ -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}) diff --git a/core/common/tasks.py b/core/common/tasks.py index 1abae664..bb19c202 100644 --- a/core/common/tasks.py +++ b/core/common/tasks.py @@ -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 @@ -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) if instance.released: instance.index_resources_for_self_as_latest_released() else: @@ -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 diff --git a/core/common/tests.py b/core/common/tests.py index f676098a..315f367b 100644 --- a/core/common/tests.py +++ b/core/common/tests.py @@ -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 @@ -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): diff --git a/core/common/utils.py b/core/common/utils.py index 6a5aa446..f656e84b 100644 --- a/core/common/utils.py +++ b/core/common/utils.py @@ -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: + return [] + if settings.ENV == 'ci': + return [None] * len(texts) + + model = settings.LM + 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)) diff --git a/core/concepts/documents.py b/core/concepts/documents.py index f4fc4d7d..8d79acbc 100644 --- a/core/concepts/documents.py +++ b/core/concepts/documents.py @@ -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): @@ -65,7 +78,8 @@ class Index: }, "type": { "type": "text" - } + }, + "text": EMBEDDING_TEXT, } ) _synonyms_embeddings = fields.NestedField( @@ -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 @@ -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): @@ -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)] + 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 @@ -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] @@ -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() diff --git a/core/concepts/embeddings.py b/core/concepts/embeddings.py new file mode 100644 index 00000000..895457b4 --- /dev/null +++ b/core/concepts/embeddings.py @@ -0,0 +1,181 @@ +""" +Concept vectors (OpenConceptLab/ocl_online#247). + +A concept doc carries vectors when any repo version the row belongs to is semantic, HEAD included: exactly the docs a +semantic $match can return, since it always searches one repo version's members. Each vector +records the exact text it encoded, and each doc the model, so a rebuild can reuse a vector rather than re-encode it: +only when the doc itself, or its versioned object's doc, already holds a vector for the same text from the same model. +The other texts are encoded together, in one batched call per chunk of docs. Docs written before #247 record neither, +so their vectors are re-encoded the next time they're rebuilt. +""" +from django.conf import settings +from django.db.models import Exists, OuterRef +from elasticsearch_dsl.connections import connections +from pydash import get + +from core.common.utils import encode_texts + +SYNC_BATCH_SIZE = 500 + + +def needs_vectors(parent, versions_match_algorithms): + """ + Whether a concept row's doc carries vectors: any repo version it belongs to (HEAD included) is semantic, given + each one's match_algorithms. A row in no version yet counts as its HEAD's (`parent`): a new concept is saved, and + indexed, before it's added to HEAD, so this keeps those writes in step with the one after, whichever lands last. + get_concept_ids_needing_vectors is the same rule in SQL. + """ + if not versions_match_algorithms: + return bool(parent.has_semantic_match_algorithm) + return any(parent.SEMANTIC_MATCH_ALGORITHM in (algorithms or []) for algorithms in versions_match_algorithms) + + +def get_concept_ids_needing_vectors(ids): + """The ids, among these concept rows, whose docs carry vectors (needs_vectors, in SQL).""" + from core.concepts.models import Concept + from core.sources.models import Source + semantic = [Source.SEMANTIC_MATCH_ALGORITHM] + through = Concept.sources.through.objects + in_semantic_version = through.filter( + concept_id__in=ids, source__match_algorithms__contains=semantic).values_list('concept_id', flat=True) + in_no_version_of_semantic_head = Concept.objects.filter( + id__in=ids, parent__match_algorithms__contains=semantic + ).exclude(Exists(through.filter(concept_id=OuterRef('id')))).values_list('id', flat=True) + return set(in_semantic_version) | set(in_no_version_of_semantic_head) + + +def get_ids_with_vectors(index_name, ids, synonyms_too=False): + """ + The ids, among these concept docs, that carry a display-name vector -- or with synonyms_too, any vector at all. + Searches, so a doc written in the last refresh interval may not count yet. + """ + paths = ['_embeddings', '_synonyms_embeddings'] if synonyms_too else ['_embeddings'] + response = connections.get_connection().search( + index=index_name, size=len(ids), source=False, query={'bool': { + 'filter': [{'ids': {'values': [str(_id) for _id in ids]}}], + 'should': [ + {'nested': {'path': path, 'query': {'exists': {'field': f'{path}.vector'}}}} for path in paths + ], + 'minimum_should_match': 1, + }}) + return {int(hit['_id']) for hit in response['hits']['hits']} + + +class ConceptVectors: + """ + The vectors for one chunk of concept docs being prepared. add() takes each doc's vector entries ({'text', ...}, + still without a 'vector'), and resolve() fills in every one: reused from the stored docs where it can, the rest + encoded in one batched call. + """ + STORED_FIELDS = ['_embeddings', '_synonyms_embeddings', '_embeddings_model'] + + def __init__(self, index_name): + self.index_name = index_name + self.model = settings.LM_MODEL_NAME + self.pending = [] + self.reused = 0 + self.encoded = 0 + + def add(self, concept, entries): + """Queues a concept row's entries. Its own doc and its versioned object's are where vectors are reused from.""" + doc_ids = [str(_id) for _id in (concept.versioned_object_id, concept.id) if _id] + self.pending.append((doc_ids, entries)) + + def resolve(self): + if not self.pending: + return + stored = self.get_stored_vectors({doc_id for doc_ids, _ in self.pending for doc_id in doc_ids}) + missing = {} + for doc_ids, entries in self.pending: + reusable = {} + for doc_id in doc_ids: # the row's own doc last, so it wins + reusable.update(stored.get(doc_id, {})) + for entry in entries: + vector = reusable.get(entry['text']) + if vector is None: + missing.setdefault(entry['text'], []).append(entry) + else: + entry['vector'] = vector + self.reused += 1 + if missing: + texts = list(missing) + for text, vector in zip(texts, encode_texts(texts)): + for entry in missing[text]: + entry['vector'] = vector + self.encoded += len(texts) + self.pending = [] + + def get_stored_vectors(self, doc_ids): + """ + {doc id: {text: vector}} from the stored docs that recorded the text of each vector and this model. A doc that + isn't there, or an index that isn't (ES answers each id with an error), gives nothing. + """ + response = connections.get_connection().mget( + index=self.index_name, ids=sorted(doc_ids), source_includes=self.STORED_FIELDS) + stored = {} + for doc in response['docs']: + source = doc.get('_source') or {} + if not doc.get('found') or source.get('_embeddings_model') != self.model: + continue + entries = [*self.as_list(source.get('_embeddings')), *self.as_list(source.get('_synonyms_embeddings'))] + stored[doc['_id']] = { + entry['text']: entry['vector'] for entry in entries + if isinstance(entry, dict) and isinstance(entry.get('text'), str) and entry.get('vector') + } + return stored + + @staticmethod + def as_list(value): + if isinstance(value, list): + return value + return [value] if value else [] + + +def sync_concept_vectors(version, parallel=True): + """ + Gives a repo version's concept docs the vectors they need, and leaves every other doc alone: + - a doc that needs vectors and has none is rebuilt, which embeds it (reusing what it can); + - a doc that has vectors, though no version it belongs to is semantic any more, is rebuilt without them; + - the rest aren't touched. So opting a version in embeds only what's missing, and opting it out never strips the + vectors another semantic version (HEAD included) still uses. + Checks 500 rows at a time, from the flags as they are at that moment, after refreshing the index so that what an + earlier sync wrote counts. Every batch is attempted, and a failed one fails the run (BatchIndexRun). Returns the + run's summary, plus the docs filled and stripped, and the vectors reused and texts encoded. + + A doc written from flags read just before a change can land after that change's own sync has checked it, so + every sync that a change queues runs once more a few minutes later (sync_source_concept_vectors). + """ + if get(settings, 'TEST_MODE', False): + return None + + from core.common.models import BatchIndexRun + from core.concepts.documents import ConceptDocument + from core.concepts.models import Concept + + doc = ConceptDocument() + index_name = doc._index._name # pylint: disable=protected-access + run = BatchIndexRun(ConceptDocument) + counts = {'filled': 0, 'stripped': 0} + ids = sorted(Concept.sources.through.objects.filter( + source_id=version.id).values_list('concept_id', flat=True), reverse=True) + + def sync_batch(batch): + needing = get_concept_ids_needing_vectors(batch) + with_vectors = get_ids_with_vectors(index_name, batch) + fill = [_id for _id in batch if _id in needing and _id not in with_vectors] + not_needing = [_id for _id in batch if _id not in needing] + strip = list(get_ids_with_vectors(index_name, not_needing, synonyms_too=True)) if not_needing else [] + if fill or strip: + concepts = list(Concept.objects.filter(id__in=fill + strip).select_related( + 'parent', 'parent__organization', 'parent__user', 'created_by', 'updated_by' + ).prefetch_related('names')) + run.retry_rejected(lambda: BatchIndexRun.bulk( + doc, doc._get_actions(concepts, 'index'), parallel, refresh=False)) # pylint: disable=protected-access + counts['filled'] += len(fill) + counts['stripped'] += len(strip) + + connections.get_connection().indices.refresh(index=index_name) + for start in range(0, len(ids), SYNC_BATCH_SIZE): + run.attempt(start, ids[start:start + SYNC_BATCH_SIZE], sync_batch) + + return {**run.finish(), **counts, 'vectors_reused': doc.vectors_reused, 'texts_encoded': doc.texts_encoded} diff --git a/core/concepts/management/__init__.py b/core/concepts/management/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/core/concepts/management/commands/__init__.py b/core/concepts/management/commands/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/core/concepts/management/commands/concept_vector_mapping.py b/core/concepts/management/commands/concept_vector_mapping.py new file mode 100644 index 00000000..42a03e27 --- /dev/null +++ b/core/concepts/management/commands/concept_vector_mapping.py @@ -0,0 +1,45 @@ +from django.core.management import BaseCommand, CommandError +from elasticsearch import ApiError +from elasticsearch_dsl.connections import connections + +from core.concepts.documents import ConceptDocument, VECTOR_PROVENANCE_MAPPING + + +class Command(BaseCommand): + help = ( + 'Adds the vector provenance fields (OpenConceptLab/ocl_online#247: the text each vector encoded, and the model) ' + 'to an existing concepts index, then checks them. Additive: nothing is reindexed. Run it before deploying the ' + 'code that writes them: once a doc carries them, ES has mapped them dynamically and this can no longer apply.' + ) + + def add_arguments(self, parser): + parser.add_argument('--check', action='store_true', help='Only check the mapping, changing nothing.') + parser.add_argument('--index', help='The index (default: the concepts index).') + + def handle(self, *args, **options): + index = options['index'] or ConceptDocument._index._name # pylint: disable=protected-access + es = connections.get_connection() + if not options['check']: + try: + es.indices.put_mapping(index=index, properties=VECTOR_PROVENANCE_MAPPING) + except ApiError as ex: + raise CommandError(f'{index}: the mapping was not changed: {ex}') from ex + + problems = [] + for path, expected in self.get_expected_fields(): + found = es.indices.get_field_mapping(index=index, fields=path)[index]['mappings'].get(path) + actual = found['mapping'][path.split('.')[-1]] if found else None + if actual != expected: + problems.append(f'{path} is {actual or "missing"}, expected {expected}') + if problems: + raise CommandError(f'{index}: ' + '; '.join(problems)) + self.stdout.write(f'{index}: vector provenance fields ok') + + @staticmethod + def get_expected_fields(): + for field, mapping in VECTOR_PROVENANCE_MAPPING.items(): + if mapping.get('type') == 'nested': + for sub_field, sub_mapping in mapping['properties'].items(): + yield f'{field}.{sub_field}', sub_mapping + else: + yield field, mapping diff --git a/core/concepts/tests/tests_vectorization.py b/core/concepts/tests/tests_vectorization.py new file mode 100644 index 00000000..2dd89068 --- /dev/null +++ b/core/concepts/tests/tests_vectorization.py @@ -0,0 +1,840 @@ +""" +Release-level vectorization (OpenConceptLab/ocl_online#247): which concept docs carry vectors, reusing a doc's vectors +instead of re-encoding them, opting a version in or out, new versions inheriting vectorization, and the reindex paths +that must not strip vectors a semantic version still uses. + +These run against the test Elasticsearch, with a deterministic stand-in for the language model. +""" +import hashlib +import uuid +from io import StringIO +from unittest.mock import patch, Mock + +import numpy +from django.core.management import call_command, CommandError +from django.test import override_settings +from elasticsearch.helpers import streaming_bulk +from elasticsearch_dsl.connections import connections + +from core.common.models import BatchIndexRun +from core.common.exceptions import BatchIndexingError +from core.common.tasks import index_source_concepts, batch_index_resources, index_concepts_mapped_codes, \ + handle_save, sync_source_concept_vectors, seed_children_to_new_version +from core.common.tests import OCLTestCase, OCLAPITestCase +from core.concepts.documents import ConceptDocument +from core.concepts.embeddings import sync_concept_vectors, ConceptVectors, get_concept_ids_needing_vectors +from core.concepts.models import Concept +from core.concepts.tests.factories import ConceptFactory, ConceptNameFactory +from core.importers.models import BulkImportInline +from core.orgs.models import Organization +from core.orgs.tests.factories import OrganizationFactory +from core.sources.models import Source +from core.sources.tests.factories import OrganizationSourceFactory +from core.tasks.models import Task +from core.users.models import UserProfile +from core.users.tests.factories import UserProfileFactory + +DIMS = 384 +MODEL = 'all-MiniLM-L6-v2' +LLM = ['es', 'llm'] + + +def fake_vector(text): + seed = int(hashlib.sha256(str(text).encode()).hexdigest()[:8], 16) + return numpy.random.default_rng(seed).random(DIMS, dtype=numpy.float32) + + +class FakeEncoder: + """Stands in for the language model: a deterministic vector per text, and a record of what it was asked.""" + def __init__(self): + self.calls = [] + + def __call__(self, texts): + texts = list(texts) + self.calls.append(texts) + return [fake_vector(text) for text in texts] + + @property + def texts(self): + return sorted(text for call in self.calls for text in call) + + def reset(self): + self.calls = [] + + +class VectorTestMixin: + """ + Concepts in a source whose docs are written for real, and helpers to read them back from Elasticsearch. Tests + index explicitly: an earlier test's pause_indexing() can leave save signals off for the rest of the run. + """ + def setUp(self): + super().setUp() + self.encoder = FakeEncoder() + patcher = patch('core.concepts.embeddings.encode_texts', self.encoder) + patcher.start() + self.addCleanup(patcher.stop) + # parallel_bulk builds the documents on pool threads, whose DB connections can't see this test's rows + patcher = patch('core.common.models.parallel_bulk', streaming_bulk) + patcher.start() + self.addCleanup(patcher.stop) + self.org = OrganizationFactory(mnemonic=f'Vec{uuid.uuid4().hex[:8]}') + self.source = OrganizationSourceFactory(organization=self.org, default_locale='en', supported_locales=['en']) + + @staticmethod + def es(): + return connections.get_connection() + + @staticmethod + def index_name(): + return ConceptDocument._index._name # pylint: disable=protected-access + + def get_doc(self, concept): + return self.es().get(index=self.index_name(), id=str(concept.id)) + + def get_source(self, concept): + return self.get_doc(concept)['_source'] + + def seq_no(self, concept): + return self.get_doc(concept)['_seq_no'] + + def refresh(self): + self.es().indices.refresh(index=self.index_name()) + + def create_concept(self, *names, source=None): + """A concept with these English names (the first preferred); returns its latest version row.""" + concept = ConceptFactory( + parent=source or self.source, + names=[ + ConceptNameFactory.build(name=name, locale='en', locale_preferred=index == 0) + for index, name in enumerate(names) + ] + ) + return concept.get_latest_version() + + def create_version(self, version, match_algorithms=None, concepts=(), released=True): + repo_version = OrganizationSourceFactory( + mnemonic=self.source.mnemonic, organization=self.org, version=version, released=released, + default_locale='en', supported_locales=['en'], match_algorithms=match_algorithms or ['es'] + ) + repo_version.concepts.add(*concepts) + return repo_version + + def index(self, *concepts): + ConceptDocument().update(list(concepts), refresh=True, parallel=False) + + def assert_vectors(self, concept, names): + """The doc carries a vector for each name, recorded with the text and the model it came from.""" + source = self.get_source(concept) + display = source['_embeddings'] + self.assertEqual(display['text'], names[0]) + self.assertEqual(source['_embeddings_model'], MODEL) + numpy.testing.assert_allclose(display['vector'], fake_vector(names[0]), rtol=1e-6) + synonyms = source['_synonyms_embeddings'] + self.assertEqual(sorted(entry['text'] for entry in synonyms), sorted(names[1:])) + for entry in synonyms: + numpy.testing.assert_allclose(entry['vector'], fake_vector(entry['text']), rtol=1e-6) + + def assert_no_vectors(self, concept): + source = self.get_source(concept) + self.assertFalse(source.get('_embeddings')) + self.assertFalse(source.get('_synonyms_embeddings')) + self.assertFalse(source.get('_embeddings_model')) + + def strip_provenance(self, concept): + """Makes the doc look like one written before #247: vectors with no record of their text or model.""" + self.es().update(index=self.index_name(), id=str(concept.id), refresh=True, script={'source': ( + "ctx._source._embeddings.remove('text'); ctx._source.remove('_embeddings_model');" + "for (entry in ctx._source._synonyms_embeddings) { entry.remove('text'); }" + )}) + + +@override_settings(LM_MODEL_NAME=MODEL) +class ConceptDocumentVectorsTest(VectorTestMixin, OCLTestCase): + def test_vectors_when_head_is_semantic(self): + self.source.match_algorithms = LLM + self.source.save() + concept = self.create_concept('Malaria', 'Paludism') + + self.index(concept) + + self.assert_vectors(concept, ['Malaria', 'Paludism']) + + def test_vectors_when_a_version_it_belongs_to_is_semantic(self): + concept = self.create_concept('Malaria', 'Paludism') + self.create_version('v1', LLM, [concept]) + + self.index(concept) + + self.assert_vectors(concept, ['Malaria', 'Paludism']) + + def test_no_vectors_when_no_version_it_belongs_to_is_semantic(self): + concept = self.create_concept('Malaria', 'Paludism') + self.create_version('v1', ['es'], [concept]) + other = self.create_concept('Fever') + self.create_version('v2', LLM, [other]) + + self.index(concept) + + self.assert_no_vectors(concept) + + def test_superseded_row_under_a_semantic_head_has_vectors_only_if_a_semantic_version_holds_it(self): + self.source.match_algorithms = LLM + self.source.save() + superseded = self.create_concept('Malaria', 'Paludism') + superseded.sources.remove(self.source) # as unmark_latest_version does once a newer version replaces it + self.create_version('v0', ['es'], [superseded]) + kept = self.create_concept('Fever') + kept.sources.remove(self.source) + self.create_version('v1', LLM, [kept]) + + self.index(superseded, kept) + + self.assert_no_vectors(superseded) # only in a lexical release: no semantic $match can return it + self.assert_vectors(kept, ['Fever']) + + def test_a_row_in_no_version_yet_counts_as_its_heads(self): + concept = self.create_concept('Malaria', 'Paludism') + concept.sources.clear() # as while it's being created, before it joins HEAD + self.create_version('v1', LLM, []) + + self.index(concept) + self.assert_no_vectors(concept) # HEAD isn't semantic + + self.source.match_algorithms = LLM + self.source.save() + self.index(Concept.objects.get(id=concept.id)) + self.assert_vectors(concept, ['Malaria', 'Paludism']) # so a write that races its HEAD membership agrees + self.assertEqual(get_concept_ids_needing_vectors([concept.id]), {concept.id}) + + def test_no_stored_vectors_to_reuse_without_an_index(self): + self.assertEqual(ConceptVectors(f'concepts-missing-{uuid.uuid4().hex[:8]}').get_stored_vectors({'1'}), {}) + + def test_prepare_outside_a_batch_resolves_vectors_for_that_doc(self): + concept = self.create_concept('Malaria', 'Paludism') + self.create_version('v1', LLM, [concept]) + + data = ConceptDocument().prepare(concept) + + self.assertEqual(data['_embeddings']['text'], 'Malaria') + numpy.testing.assert_allclose(data['_embeddings']['vector'], fake_vector('Malaria')) + self.assertEqual([entry['text'] for entry in data['_synonyms_embeddings']], ['Paludism']) + self.assertEqual(data['_embeddings_model'], MODEL) + self.assertIn('v1', data['source_version']) + + +@override_settings(LM_MODEL_NAME=MODEL) +class VectorReuseTest(VectorTestMixin, OCLTestCase): + def setUp(self): + super().setUp() + self.source.match_algorithms = LLM + self.source.save() + + def test_reindex_reuses_vectors_when_text_and_model_are_unchanged(self): + concept = self.create_concept('Malaria', 'Paludism', 'Marsh fever') + self.index(concept) + self.encoder.reset() + + self.index(Concept.objects.get(id=concept.id)) + + self.assertEqual(self.encoder.texts, []) + self.assert_vectors(concept, ['Malaria', 'Paludism', 'Marsh fever']) + + def test_reindex_encodes_only_the_changed_text(self): + concept = self.create_concept('Malaria', 'Paludism') + self.index(concept) + self.encoder.reset() + name = concept.names.get(name='Paludism') + name.name = 'Marsh fever' + name.save() + + self.index(Concept.objects.get(id=concept.id)) + + self.assertEqual(self.encoder.texts, ['Marsh fever']) + self.assert_vectors(concept, ['Malaria', 'Marsh fever']) + + def test_reindex_does_not_reuse_vectors_from_another_model(self): + concept = self.create_concept('Malaria', 'Paludism') + with override_settings(LM_MODEL_NAME='some-older-model'): + self.index(concept.versioned_object) + self.index(concept) + self.encoder.reset() + + self.index(Concept.objects.get(id=concept.id)) + + self.assertEqual(self.encoder.texts, ['Malaria', 'Paludism']) + self.assert_vectors(concept, ['Malaria', 'Paludism']) + + def test_reindex_does_not_reuse_vectors_that_record_no_text(self): + concept = self.create_concept('Malaria', 'Paludism') + self.index(concept.versioned_object, concept) + self.strip_provenance(concept.versioned_object) + self.strip_provenance(concept) + self.encoder.reset() + + self.index(Concept.objects.get(id=concept.id)) + + self.assertEqual(self.encoder.texts, ['Malaria', 'Paludism']) + self.assert_vectors(concept, ['Malaria', 'Paludism']) + + def test_new_version_row_reuses_its_versioned_objects_vectors(self): + concept = self.create_concept('Malaria', 'Paludism') + versioned_object = concept.versioned_object + self.index(versioned_object) + new_row = versioned_object.clone() + new_row.save() + new_row.version = new_row.id + new_row.save() + new_row.set_locales(versioned_object.names.all(), type(versioned_object.names.first())) + new_row.sources.add(self.source) # an edit's new version joins HEAD + self.encoder.reset() + + self.index(Concept.objects.get(id=new_row.id)) + + self.assertEqual(self.encoder.texts, []) + self.assert_vectors(new_row, ['Malaria', 'Paludism']) + + def test_encodes_a_chunk_of_docs_in_one_batched_call(self): + self.source.match_algorithms = ['es'] + self.source.save() # so that creating them writes no vectors + concepts = [self.create_concept(f'Name {index}', f'Synonym {index}', 'Shared') for index in range(5)] + self.source.match_algorithms = LLM + self.source.save() + concepts = list(Concept.objects.filter(id__in=[concept.id for concept in concepts]).order_by('id')) + self.encoder.reset() + + self.index(*concepts) + + self.assertEqual(len(self.encoder.calls), 1) + self.assertEqual( + sorted(self.encoder.calls[0]), + sorted(['Shared', *[f'Name {index}' for index in range(5)], *[f'Synonym {index}' for index in range(5)]]) + ) + for index, concept in enumerate(concepts): + self.assert_vectors(concept, [f'Name {index}', f'Synonym {index}', 'Shared']) + + +@override_settings(LM_MODEL_NAME=MODEL) +class VersionVectorSyncTest(VectorTestMixin, OCLTestCase): + """HEAD isn't semantic: only a version's own flag decides whether its docs carry vectors.""" + def setUp(self): + super().setUp() + self.shared = [self.create_concept('Malaria', 'Paludism'), self.create_concept('Fever', 'Pyrexia')] + self.own = self.create_concept('Cholera', 'Asiatic cholera') + + def sync(self, version): + self.refresh() + self.encoder.reset() + with override_settings(TEST_MODE=False): + summary = sync_concept_vectors(version, parallel=False) + self.refresh() + self.assertEqual(summary['failed_docs'], 0) + self.assertEqual(summary['texts_encoded'], len(set(self.encoder.texts))) + return {key: summary[key] for key in ('docs', 'filled', 'stripped')} + + def test_opting_a_version_in_embeds_only_its_docs_without_vectors(self): + self.create_version('v1', LLM, self.shared) + v2 = self.create_version('v2', ['es'], [*self.shared, self.own]) + self.index(*self.shared, self.own) + self.assert_no_vectors(self.own) + seq_nos = [self.seq_no(concept) for concept in self.shared] + v2.match_algorithms = LLM + v2.save() + + summary = self.sync(v2) + + self.assertEqual(summary, {'docs': 3, 'filled': 1, 'stripped': 0}) + self.assertEqual(self.encoder.texts, ['Asiatic cholera', 'Cholera']) + self.assert_vectors(self.own, ['Cholera', 'Asiatic cholera']) + self.assertEqual([self.seq_no(concept) for concept in self.shared], seq_nos) # not rewritten + + def test_opting_a_version_out_keeps_the_vectors_another_semantic_version_uses(self): + self.create_version('v1', LLM, self.shared) + v2 = self.create_version('v2', LLM, [*self.shared, self.own]) + self.index(*self.shared, self.own) + self.assert_vectors(self.own, ['Cholera', 'Asiatic cholera']) + seq_nos = [self.seq_no(concept) for concept in self.shared] + v2.match_algorithms = ['es'] + v2.save() + + summary = self.sync(v2) + + self.assertEqual(summary, {'docs': 3, 'filled': 0, 'stripped': 1}) + self.assertEqual(self.encoder.texts, []) + self.assert_vectors(self.shared[0], ['Malaria', 'Paludism']) + self.assert_vectors(self.shared[1], ['Fever', 'Pyrexia']) + self.assertEqual([self.seq_no(concept) for concept in self.shared], seq_nos) # not rewritten + self.assert_no_vectors(self.own) + + def test_opting_a_version_out_keeps_vectors_when_head_is_semantic(self): + v2 = self.create_version('v2', LLM, [*self.shared, self.own]) + self.index(*self.shared, self.own) + self.source.match_algorithms = LLM + self.source.save() + v2.match_algorithms = ['es'] + v2.save() + + summary = self.sync(v2) + + self.assertEqual(summary, {'docs': 3, 'filled': 0, 'stripped': 0}) + self.assert_vectors(self.own, ['Cholera', 'Asiatic cholera']) + + def test_opting_head_in_fills_its_members_only(self): + superseded = self.create_concept('Ague') + superseded.sources.remove(self.source) + self.create_version('v0', ['es'], [superseded]) # superseded: only in an old lexical release + self.index(*self.shared, self.own, superseded) + self.source.match_algorithms = LLM + self.source.save() + + summary = self.sync(self.source) + + self.assertEqual(summary['filled'], summary['docs']) # every HEAD member: latest rows and versioned objects + self.assert_vectors(self.own, ['Cholera', 'Asiatic cholera']) + self.assert_no_vectors(superseded) # in no version, so no semantic $match can return it + + def test_sync_is_a_no_op_when_docs_already_match(self): + v1 = self.create_version('v1', LLM, self.shared) + self.index(*self.shared, self.own) + + summary = self.sync(v1) + + self.assertEqual(summary, {'docs': 2, 'filled': 0, 'stripped': 0}) + self.assertEqual(self.encoder.calls, []) + + def during_the_next_write(self, change): + """Patches the sync's write: after its docs are prepared, and before they're sent, runs change().""" + real_bulk = BatchIndexRun.bulk + pending = [change] + + def bulk(doc, actions, *args, **kwargs): + actions = list(actions) # prepared from the flags as they were + if pending: + pending.pop()() + return real_bulk(doc, iter(actions), *args, **kwargs) + return patch('core.common.models.BatchIndexRun.bulk', side_effect=bulk) + + def opt_in_and_sync(self, version): + """A version opts in, and the sync that queues runs at once: it finds the doc's vectors still there.""" + def change(): + Source.objects.filter(id=version.id).update(match_algorithms=LLM) + with override_settings(TEST_MODE=False): + self.assertEqual(sync_concept_vectors(version, parallel=False)['filled'], 0) + return change + + def test_an_opt_in_during_another_versions_opt_out_is_repaired_by_its_recheck(self): + v1 = self.create_version('v1', LLM, [*self.shared, self.own]) + v2 = self.create_version('v2', ['es'], [self.own]) + self.index(*self.shared, self.own) + v1.match_algorithms = ['es'] + v1.save() + + # v1's sync prepares own's doc without vectors; v2 opts in and v2's sync runs before that write lands + with self.during_the_next_write(self.opt_in_and_sync(v2)): + self.sync(v1) + self.assert_no_vectors(self.own) # the race: v2 is semantic and own has no vectors + + self.assertEqual(self.sync(v2), {'docs': 1, 'filled': 1, 'stripped': 0}) # v2's recheck, minutes later + self.assert_vectors(self.own, ['Cholera', 'Asiatic cholera']) + self.assert_no_vectors(self.shared[0]) + + def test_an_opt_out_and_back_in_during_a_sync_is_repaired_by_the_recheck(self): + v1 = self.create_version('v1', LLM, [self.own]) + self.index(self.own) + Source.objects.filter(id=v1.id).update(match_algorithms=['es']) # v1 opts out + + # its sync prepares the strip; v1 opts back in and that sync runs before the strip lands + with self.during_the_next_write(self.opt_in_and_sync(v1)): + self.sync(v1) + self.assert_no_vectors(self.own) + + self.assertEqual(self.sync(v1)['filled'], 1) # the opt-in's recheck + self.assert_vectors(self.own, ['Cholera', 'Asiatic cholera']) + + def test_sync_sees_vectors_written_since_the_last_refresh(self): + v1 = self.create_version('v1', LLM, self.shared) + self.es().indices.put_settings(index=self.index_name(), settings={'index': {'refresh_interval': '-1'}}) + self.addCleanup( + self.es().indices.put_settings, index=self.index_name(), settings={'index': {'refresh_interval': None}}) + ConceptDocument().update(self.shared, refresh=False, parallel=False) # vectors written, not yet searchable + seq_nos = [self.seq_no(concept) for concept in self.shared] + self.encoder.reset() + + with override_settings(TEST_MODE=False): + summary = sync_concept_vectors(v1, parallel=False) + + self.assertEqual((summary['filled'], summary['stripped']), (0, 0)) + self.assertEqual([self.seq_no(concept) for concept in self.shared], seq_nos) + self.assertEqual(self.encoder.calls, []) + + def test_sync_does_nothing_in_test_mode(self): + v1 = self.create_version('v1', LLM, self.shared) + + self.assertIsNone(sync_concept_vectors(v1)) + + @patch('core.sources.models.index_source_mappings', Mock(__name__='index_source_mappings')) + @patch('core.sources.models.index_source_concepts', Mock(__name__='index_source_concepts')) + @patch('core.sources.models.Source.sync_concept_vectors_async') + def test_seeding_a_semantic_version_queues_its_sync(self, sync_async_mock): + v1 = self.create_version('v1', ['es'], []) + v2 = self.create_version('v2', ['es'], []) # v1 is no longer the latest release: nothing indexes it + Source.objects.filter(id=v1.id).update(match_algorithms=LLM) + + seed_children_to_new_version('source', v1.id, False) + + sync_async_mock.assert_called_once_with(v1.created_by) + self.assertEqual( + set(v1.concepts.values_list('id', flat=True)), {self.shared[0].id, self.shared[1].id, self.own.id}) + sync_async_mock.reset_mock() + seed_children_to_new_version('source', v2.id, False) + sync_async_mock.assert_not_called() # lexical + + @patch('core.sources.models.index_source_mappings', Mock(__name__='index_source_mappings')) + @patch('core.sources.models.index_source_concepts', Mock(__name__='index_source_concepts')) + @patch('core.sources.models.Source.sync_concept_vectors_async') + def test_seeding_reads_the_flag_after_seeding(self, sync_async_mock): + v1 = self.create_version('v1', ['es'], []) + + def opt_in_meanwhile(*_): # the opt-in's own sync, and its recheck, find no members yet + Source.objects.filter(id=v1.id).update(match_algorithms=LLM) + with patch('core.sources.models.Source.update_children_counts', side_effect=opt_in_meanwhile): + seed_children_to_new_version('source', v1.id, False) + + sync_async_mock.assert_called_once_with(v1.created_by) + + def test_a_seeded_semantic_release_gets_its_vectors_from_the_sync(self): + v1 = self.create_version('v1', ['es'], [*self.shared, self.own]) + self.index(*self.shared, self.own) # HEAD's docs as they were: no vectors + Source.objects.filter(id=v1.id).update(match_algorithms=LLM) # v1 was created vectorized + + self.assertEqual(self.sync(v1), {'docs': 3, 'filled': 3, 'stripped': 0}) # the sync its seeding queued + + self.assert_vectors(self.shared[0], ['Malaria', 'Paludism']) + self.assert_vectors(self.own, ['Cholera', 'Asiatic cholera']) + + +@override_settings(LM_MODEL_NAME=MODEL) +class ReindexKeepsSharedVectorsTest(VectorTestMixin, OCLTestCase): + """HEAD isn't semantic, and a release that shares HEAD's rows is: rebuilding those rows keeps their vectors.""" + def setUp(self): + super().setUp() + self.concept = self.create_concept('Malaria', 'Paludism') + self.plain = self.create_concept('Fever') + self.release = self.create_version('v1', LLM, [self.concept]) + self.index(self.concept, self.plain) + self.assert_vectors(self.concept, ['Malaria', 'Paludism']) + self.encoder.reset() + + def test_manual_reindex_keeps_the_vectors_of_rows_in_a_semantic_release(self): + with override_settings(TEST_MODE=False): + index_source_concepts(self.source.id) + + self.assert_vectors(self.concept, ['Malaria', 'Paludism']) + self.assert_no_vectors(self.plain) + self.assertEqual(self.encoder.texts, []) # reused + + def test_import_reindex_keeps_the_vectors_of_rows_in_a_semantic_release(self): + self.import_update_and_check(index=True) + + def test_small_import_reindex_keeps_the_vectors_of_rows_in_a_semantic_release(self): + self.import_update_and_check() # no `index`: imports of up to IMPORT_INDEX_MAX_LINES lines index anyway + + def test_edit_keeps_the_vectors_of_the_row_it_supersedes(self): + errors = Concept.create_new_version_for( + instance=self.concept.clone(), + data={'names': [{'locale': 'en', 'name': 'Malaria', 'locale_preferred': True}, + {'locale': 'en', 'name': 'Paludism'}, {'locale': 'en', 'name': 'Swamp fever'}]}, + user=self.concept.created_by + ) + self.assertEqual(errors, {}) + previous_row = Concept.objects.get(id=self.concept.id) + self.assertFalse(previous_row.is_latest_version) + + handle_save('concepts', 'Concept', previous_row.id) # what the edit queues for it (Concept.index) + + self.assert_vectors(previous_row, ['Malaria', 'Paludism']) # still in v1: not stripped + self.assertEqual(self.encoder.texts, []) + + def import_update_and_check(self, **importer_kwargs): + with patch('core.importers.models.batch_index_resources') as batch_index_mock, \ + patch('core.importers.models.index_concepts_mapped_codes') as mapped_codes_mock: + importer = BulkImportInline( + content=None, username='ocladmin', update_if_exists=True, input_list=[{ + 'type': 'Concept', 'id': self.concept.mnemonic, 'concept_class': 'Misc', 'datatype': 'None', + 'owner_type': 'Organization', 'owner': self.org.mnemonic, 'source': self.source.mnemonic, + 'names': [ + {'name': 'Malaria', 'locale': 'en', 'locale_preferred': True, 'name_type': 'Fully Specified'}, + {'name': 'Paludism', 'locale': 'en', 'name_type': 'Synonym'}, + {'name': 'Swamp fever', 'locale': 'en', 'name_type': 'Synonym'}, + ] + }], **importer_kwargs) + importer.run() + self.assertEqual(len(importer.updated), 1) + with override_settings(TEST_MODE=False): + for task_mock, task in [(batch_index_mock, batch_index_resources), + (mapped_codes_mock, index_concepts_mapped_codes)]: + for call_args in task_mock.apply_async.call_args_list: + task(*call_args[0][0]) + + previous_row = Concept.objects.get(id=self.concept.id) + self.assertFalse(previous_row.is_latest_version) + self.assert_vectors(previous_row, ['Malaria', 'Paludism']) # still in v1: not stripped + new_row = previous_row.versioned_object.get_latest_version() + self.assert_no_vectors(new_row) # in no semantic version + self.assertEqual(self.encoder.texts, []) # the previous row's vectors were reused + + def test_locale_change_rebuilds_rows_with_vectors_so_their_display_vector_follows(self): + ConceptNameFactory(concept=self.concept, name='Paludisme', locale='fr', locale_preferred=True) + self.index(Concept.objects.get(id=self.concept.id)) + self.encoder.reset() + self.source.default_locale = 'fr' + self.source.supported_locales = ['fr', 'en'] + self.source.save() + + with override_settings(TEST_MODE=False): + index_source_concepts(self.source.id, None, locales=['en', 'fr']) + + source = self.get_source(self.concept) + self.assertEqual(source['_embeddings']['text'], 'Paludisme') + self.assertEqual(source['name'], 'Paludisme') + self.assertEqual(self.encoder.texts, []) # every name already had a vector + self.assert_no_vectors(self.plain) + + +class VectorSyncQueuingTest(OCLTestCase): + def setUp(self): + super().setUp() + self.source = OrganizationSourceFactory(default_locale='en', supported_locales=['en']) + self.version = OrganizationSourceFactory( + mnemonic=self.source.mnemonic, organization=self.source.organization, version='v1', + match_algorithms=['es'], released=False) + + @patch('core.sources.models.Source.sync_concept_vectors_async') + @patch('core.sources.models.Source.index_concepts_async') + def test_flag_flip_queues_a_vector_sync_not_a_reindex(self, index_concepts_async_mock, sync_async_mock): + self.version.match_algorithms = LLM + errors = Source.persist_changes(self.version, self.version.created_by, None) + + self.assertEqual(errors, {}) + sync_async_mock.assert_called_once_with(self.version.updated_by) + index_concepts_async_mock.assert_not_called() + + @patch('core.sources.models.Source.sync_concept_vectors_async') + @patch('core.sources.models.Source.index_concepts_async') + def test_flag_flip_with_a_full_locale_reindex_queues_both(self, index_concepts_async_mock, sync_async_mock): + self.version.match_algorithms = LLM + self.version.supported_locales = None + self.version.default_locale = 'es' + errors = Source.persist_changes(self.version, self.version.created_by, None) + + self.assertEqual(errors, {}) + index_concepts_async_mock.assert_called_once_with(self.version.updated_by) # every doc + sync_async_mock.assert_called_once_with(self.version.updated_by) + + @patch('core.sources.models.Source.sync_concept_vectors_async') + @patch('core.sources.models.Source.index_resources_for_self_as_latest_released') + def test_releasing_and_vectorizing_in_one_change_does_both(self, index_released_mock, sync_async_mock): + self.version.match_algorithms = LLM + self.version.released = True + errors = Source.persist_changes(self.version, self.version.created_by, None) + + self.assertEqual(errors, {}) + index_released_mock.assert_called_once_with(only_update=True) + sync_async_mock.assert_called_once_with(self.version.updated_by) + + @patch('core.common.tasks.sync_concept_vectors') + @patch('core.sources.models.Source.sync_concept_vectors_async') + def test_sync_task_syncs_then_queues_one_delayed_recheck(self, sync_async_mock, sync_mock): + sync_source_concept_vectors(self.version.id) + + self.assertEqual(sync_mock.call_args[0][0].id, self.version.id) + sync_async_mock.assert_called_once_with(recheck=False, countdown=600) + + @patch('core.common.tasks.sync_concept_vectors') + @patch('core.sources.models.Source.sync_concept_vectors_async') + def test_the_recheck_queues_nothing_more(self, sync_async_mock, sync_mock): + sync_source_concept_vectors(self.version.id, False) + + sync_mock.assert_called_once() + sync_async_mock.assert_not_called() + + @patch('core.common.tasks.sync_concept_vectors', side_effect=BatchIndexingError('failed')) + @patch('core.sources.models.Source.sync_concept_vectors_async') + def test_a_failed_sync_still_queues_its_recheck(self, sync_async_mock, _): + with self.assertRaises(BatchIndexingError): + sync_source_concept_vectors(self.version.id) + + sync_async_mock.assert_called_once_with(recheck=False, countdown=600) + + @patch('celery.app.task.Task.apply_async') + def test_a_queued_sync_keeps_its_arguments_for_recovery(self, celery_apply_async_mock): + with override_settings(TEST_MODE=False): + self.version.sync_concept_vectors_async(self.version.created_by, recheck=False, countdown=600) + + args, kwargs = celery_apply_async_mock.call_args + self.assertEqual(args[0], (self.version.id, False)) + self.assertEqual(kwargs['countdown'], 600) + self.assertEqual(kwargs['queue'], 'indexing') + task = Task.objects.get(id=args[2]) + self.assertEqual(task.name, 'core.common.tasks.sync_source_concept_vectors') + self.assertEqual(list(task.args), [self.version.id, False]) # a rerun is still a sync + + @patch('core.sources.models.sync_source_concept_vectors', Mock(__name__='sync_source_concept_vectors')) + def test_in_test_mode_the_sync_runs_inline(self): + from core.sources import models as source_models + sync_task_mock = source_models.sync_source_concept_vectors + task = self.version.sync_concept_vectors_async() + + sync_task_mock.assert_called_once_with(self.version.id, True) + self.assertEqual(task.name, 'sync_source_concept_vectors') + + +class SourceConceptsIndexViewSyncVectorsTest(OCLAPITestCase): + """Staff can re-run a version's vector sync without changing its flag, e.g. after a sync task was lost.""" + def setUp(self): + super().setUp() + self.token = UserProfile.objects.filter(is_superuser=True).first().get_token() + self.source = OrganizationSourceFactory() + self.version = OrganizationSourceFactory( + mnemonic=self.source.mnemonic, organization=self.source.organization, version='v1', match_algorithms=LLM) + + def post(self, data, task_name): + with patch(f'core.sources.views.{task_name}') as task_mock: + task_mock.__name__ = task_name + response = self.client.post( + self.version.uri + 'concepts/indexes/', data, HTTP_AUTHORIZATION='Token ' + self.token, format='json') + self.assertEqual(response.status_code, 202, response.data) + return task_mock.apply_async.call_args[0][0] + + @patch('celery.app.task.Task.apply_async') + def test_sync_vectors_persists_its_arguments(self, celery_apply_async_mock): + with override_settings(TEST_MODE=False): + response = self.client.post( + self.version.uri + 'concepts/indexes/', {'sync_vectors': True}, + HTTP_AUTHORIZATION='Token ' + self.token, format='json') + + self.assertEqual(response.status_code, 202, response.data) + task = Task.objects.get(id=response.data['id']) + self.assertEqual(task.name, 'core.common.tasks.sync_source_concept_vectors') + self.assertEqual(list(task.args), [self.version.id, True]) # a rerun is still a sync + self.assertEqual(celery_apply_async_mock.call_args[0][0], (self.version.id, True)) + + def test_without_sync_vectors_queues_the_full_reindex_as_before(self): + self.assertEqual( + self.post({}, 'index_source_concepts'), (self.version.id, None, False, True, True, True)) + + +class VersionCreateInheritsVectorizationTest(OCLAPITestCase): + """Decision V1: a new version is vectorized when HEAD or the latest release is, unless the request says not.""" + def setUp(self): + super().setUp() + self.organization = Organization.objects.first() + self.token = UserProfile.objects.filter(is_superuser=True).first().get_token() + self.source = OrganizationSourceFactory(organization=self.organization) + + def create(self, token=None, **data): + with patch('core.sources.models.index_source_concepts', Mock(__name__='index_source_concepts')), \ + patch('core.sources.models.index_source_mappings', Mock(__name__='index_source_mappings')): + response = self.client.post( + f'/orgs/{self.organization.mnemonic}/sources/{self.source.mnemonic}/versions/', + {'id': 'v-new', 'description': 'new', 'released': True, **data}, + HTTP_AUTHORIZATION='Token ' + (token or self.token), format='json' + ) + self.assertEqual(response.status_code, 201, response.data) + return response.data['match_algorithms'], self.source.versions.get(version='v-new').match_algorithms + + def test_inherits_llm_from_semantic_head(self): + self.source.match_algorithms = LLM + self.source.save() + + api_value, stored = self.create() + + self.assertEqual(sorted(api_value), LLM) + self.assertEqual(sorted(stored), LLM) + + def test_inherits_llm_from_the_latest_semantic_release(self): + OrganizationSourceFactory( + mnemonic=self.source.mnemonic, organization=self.organization, version='v1', released=True, + match_algorithms=LLM) + + _, stored = self.create() + + self.assertEqual(sorted(stored), LLM) + + def test_does_not_inherit_from_an_older_release_than_the_latest(self): + OrganizationSourceFactory( + mnemonic=self.source.mnemonic, organization=self.organization, version='v1', released=True, + match_algorithms=LLM) + OrganizationSourceFactory( + mnemonic=self.source.mnemonic, organization=self.organization, version='v2', released=True, + match_algorithms=['es']) + + _, stored = self.create() + + self.assertEqual(stored, ['es']) + + def test_not_vectorized_when_neither_head_nor_the_latest_release_is(self): + _, stored = self.create() + + self.assertEqual(stored, ['es']) + + def test_owner_can_opt_a_new_release_out(self): + self.source.match_algorithms = LLM + self.source.save() + + _, stored = self.create(match_algorithms=['es']) + + self.assertEqual(stored, ['es']) + + def test_a_member_who_is_not_staff_inherits_it_too(self): + self.source.match_algorithms = LLM + self.source.save() + member = UserProfileFactory() + self.organization.members.add(member) + + _, stored = self.create(token=member.get_token()) + + self.assertEqual(sorted(stored), LLM) + + +class ConceptVectorMappingCommandTest(OCLTestCase): + """The deploy step that adds the vector provenance fields to an existing concepts index (mapped as before #247).""" + def setUp(self): + super().setUp() + self.es = connections.get_connection() + self.index = f'concepts-mapping-test-{uuid.uuid4().hex[:8]}' + vector_doc = {'type': 'nested', 'properties': {'vector': {'type': 'dense_vector'}, 'type': {'type': 'text'}}} + self.es.indices.create(index=self.index, mappings={'properties': { + '_embeddings': vector_doc, '_synonyms_embeddings': vector_doc, 'name': {'type': 'text'}}}) + self.addCleanup(self.es.indices.delete, index=self.index) + + def mapping(self): + return self.es.indices.get_mapping(index=self.index)[self.index]['mappings']['properties'] + + def call(self, *args): + out = StringIO() + call_command('concept_vector_mapping', '--index', self.index, *args, stdout=out) + return out.getvalue() + + def test_check_fails_while_the_fields_are_missing(self): + with self.assertRaisesRegex(CommandError, '_embeddings.text'): + self.call('--check') + self.assertNotIn('text', self.mapping()['_embeddings']['properties']) # --check changed nothing + + def test_adds_the_fields_and_leaves_the_rest(self): + output = self.call() + + mapping = self.mapping() + for field in ('_embeddings', '_synonyms_embeddings'): + self.assertEqual( + mapping[field]['properties']['text'], {'type': 'keyword', 'index': False, 'doc_values': False}) + self.assertEqual(mapping[field]['properties']['type'], {'type': 'text'}) + self.assertEqual(mapping['_embeddings_model'], {'type': 'keyword'}) + self.assertIn('ok', output) + self.assertIn('ok', self.call('--check')) + self.call() # again: nothing to change + + def test_refuses_a_field_already_mapped_differently(self): + self.es.indices.put_mapping(index=self.index, properties={'_embeddings_model': {'type': 'text'}}) + + with self.assertRaisesRegex(CommandError, '_embeddings_model'): + self.call() diff --git a/core/settings.py b/core/settings.py index 8daeeefa..a527d3ae 100644 --- a/core/settings.py +++ b/core/settings.py @@ -669,10 +669,15 @@ def get_set_from_env(name): ] ENCODER_MODEL_NAME = None ENCODER = None -LM_MODEL_NAME = None +# The model concept vectors are encoded with. Each concept doc records it, and a vector is only reused while it's +# unchanged (OpenConceptLab/ocl_online#247), so it's set even where the model isn't loaded. +LM_MODEL_NAME = 'all-MiniLM-L6-v2' +LM_ENCODE_BATCH_SIZE = int(os.environ.get('LM_ENCODE_BATCH_SIZE', 64)) +# A version's vector sync runs once more this long after it (OpenConceptLab/ocl_online#247), once any doc prepared from +# the flags as they were before has long been written +VECTOR_SYNC_RECHECK_SECONDS = int(os.environ.get('VECTOR_SYNC_RECHECK_SECONDS', 600)) LM = None if ENV not in ['ci', 'demo'] and not NO_LM: - LM_MODEL_NAME = 'all-MiniLM-L6-v2' LM = SentenceTransformer(LM_MODEL_NAME) if not NO_ENCODER: ENCODER_MODEL_NAME = "BAAI/bge-reranker-v2-m3" diff --git a/core/sources/models.py b/core/sources/models.py index d0a5438e..a7df65ce 100644 --- a/core/sources/models.py +++ b/core/sources/models.py @@ -15,7 +15,7 @@ from core.common.constants import HEAD from core.common.models import ConceptContainerModel from core.common.tasks import update_mappings_source, index_source_concepts, index_source_mappings, \ - resolve_url_registry_entries + sync_source_concept_vectors, resolve_url_registry_entries from core.common.utils import to_camel_case from core.common.validators import validate_non_negative from core.concepts.models import ConceptName, Concept @@ -544,10 +544,25 @@ def index_concepts_async(self, user, partial_doc=None, locales=None, exclude_loc except AlreadyQueued: pass - def get_concepts_reindex_filters(self, original): - if bool(self.has_semantic_match_algorithm) != bool(original.has_semantic_match_algorithm): - return {} + def sync_concept_vectors_async(self, user=None, recheck=True, countdown=None): + """ + Queues this version's vector sync (sync_source_concept_vectors) on the indexing queue, with its arguments + persisted, so that a rerun of the task is still a sync, and returns its Task. In TEST_MODE it runs inline. + """ + task = Task.new(queue='indexing', user=user or self.updated_by, name=sync_source_concept_vectors.__name__) + if get(settings, 'TEST_MODE', False): + sync_source_concept_vectors(self.id, recheck) + else: + sync_source_concept_vectors.apply_async( + (self.id, recheck), queue='indexing', persist_args=True, task_id=task.id, countdown=countdown) + return task + def get_concepts_reindex_filters(self, original): + """ + What a locale change to this repo version needs reindexed: None for nothing, {} for every concept doc, or the + narrowing to pass to index_concepts_async. A semantic flag change syncs the version's vectors instead + (persist_changes, sync_concept_vectors_async). + """ old_supported, new_supported = original.supported_locales, self.supported_locales default_changed = self.default_locale != original.default_locale nullness_changed = (old_supported is None) != (new_supported is None) @@ -569,6 +584,18 @@ def get_concepts_reindex_filters(self, original): return filters + def get_match_algorithms_for_new_version(self): + """ + The match algorithms a new version of this repo (HEAD) is created with unless the request sets them: HEAD's, + and vectorized when HEAD or the latest released version is (decision V1, OpenConceptLab/ocl_online#247). + """ + match_algorithms = list(self.match_algorithms or [self.TOKEN_MATCH_ALGORITHM]) # as clean_match_algorithms + if self.SEMANTIC_MATCH_ALGORITHM not in match_algorithms: + latest_released = self.get_latest_released_version() + if latest_released and latest_released.has_semantic_match_algorithm: + match_algorithms.append(self.SEMANTIC_MATCH_ALGORITHM) + return match_algorithms + def get_export_task(self): return Task.find(name__iendswith='export_source', args__contains=[self.id]) diff --git a/core/sources/tests/tests.py b/core/sources/tests/tests.py index c4238781..879c0f82 100644 --- a/core/sources/tests/tests.py +++ b/core/sources/tests/tests.py @@ -714,7 +714,11 @@ def filters(**changes): self.assertEqual(filters(default_locale='es'), {'locales': ['en', 'es']}) self.assertEqual( filters(default_locale='es', supported_locales=['de']), {'locales': ['de', 'en', 'es', 'fr']}) - self.assertEqual(filters(match_algorithms=['llm']), {}) + # a semantic flag change syncs the version's vectors instead (ocl_online#247): no reindex of its own + self.assertIsNone(filters(match_algorithms=['llm'])) + self.assertEqual( + filters(match_algorithms=['llm'], supported_locales=['fr', 'es']), + {'locales': ['es'], 'exclude_locale': 'en'}) def test_get_concepts_reindex_filters_null_supported_locales(self): source = OrganizationSourceFactory(default_locale='en', supported_locales=None) @@ -726,6 +730,7 @@ def test_get_concepts_reindex_filters_null_supported_locales(self): updated.default_locale = 'es' self.assertEqual(updated.get_concepts_reindex_filters(source), {}) + def test_source_version_create_positive(self): source = OrganizationSourceFactory() self.assertEqual(source.num_versions, 1) diff --git a/core/sources/views.py b/core/sources/views.py index aaec0842..30d86305 100644 --- a/core/sources/views.py +++ b/core/sources/views.py @@ -271,7 +271,9 @@ def create(self, request, *args, **kwargs): 'version': version, "meta": request.data.get('meta', head_object.meta), "properties": request.data.get('properties', head_object.properties), - "filters": request.data.get('filters', head_object.filters) + "filters": request.data.get('filters', head_object.filters), + "match_algorithms": request.data.get( + 'match_algorithms', head_object.get_match_algorithms_for_new_version()) } serializer = self.get_serializer(data=payload) if serializer.is_valid(): @@ -359,9 +361,20 @@ def get_task_args(self, instance): class SourceConceptsIndexView(SourceIndexBaseView): + """ + Reindexes the version's concepts. With `sync_vectors`, only syncs its vectors instead (sync_source_concept_vectors, + with its recheck): embeds the docs that need them and lack them, strips those no semantic version uses, and + leaves the rest. That re-runs a lost or failed sync without changing the version's flag (ocl_online#247). + """ def get_task_function(self): return index_source_concepts + def post(self, request, *args, **kwargs): + if (request.data or {}).get('sync_vectors', None) in get_truthy_values(): + task = self.get_object().sync_concept_vectors_async(request.user) + return Response(TaskBriefSerializer(task).data, status=status.HTTP_202_ACCEPTED) + return super().post(request, *args, **kwargs) + class SourceMappingsIndexView(SourceIndexBaseView): def get_task_function(self):