diff --git a/antigravity_engine.py b/antigravity_engine.py index 1ea3374..1f4c9f3 100644 --- a/antigravity_engine.py +++ b/antigravity_engine.py @@ -1,6 +1,9 @@ import copy import hashlib +import os +import tempfile import time +from pathlib import Path import numpy as np from qdrant_client import QdrantClient @@ -20,6 +23,40 @@ from embedding_backend import create_embedding_backend from vector_store import create_vector_store + +def _pre_cycle_adapter_bytes(adapter, adapter_path) -> bytes: + """Bytes of the adapter file from before this cycle's trained save. + + When the file is missing, the current in-memory weights are that + pre-cycle state and are saved once so a later overwrite can be undone. + """ + path = Path(adapter_path) + if not path.exists(): + adapter.save(path) + return path.read_bytes() + + +def _write_adapter_bytes(adapter_path, payload: bytes) -> None: + """Replace adapter_path with payload using a same-directory temp file.""" + destination = Path(adapter_path) + descriptor, temporary_name = tempfile.mkstemp( + prefix="." + destination.name + ".", + suffix=".tmp", + dir=destination.parent, + ) + temporary = Path(temporary_name) + try: + with os.fdopen(descriptor, "wb") as handle: + handle.write(payload) + handle.flush() + os.fsync(handle.fileno()) + os.replace(temporary, destination) + except Exception: + if temporary.exists(): + temporary.unlink() + raise + + class AntigravityEngine: def __init__(self, qdrant_location=":memory:", chelation_p=ChelationConfig.DEFAULT_CHELATION_P, model_name='ollama:nomic-embed-text', use_centering=False, use_quantization=False, training_mode: str = "baseline", teacher_model_name: Optional[str] = None, teacher_models=None, teacher_weight: float = 0.5, store_full_text_payload: Optional[bool] = None): """ @@ -1903,6 +1940,13 @@ def _finish_annealing_cycle(): # --- TRAINING LOOP (wrapped in SafeTrainingContext for F-043) --- self.logger.log_training_start(num_samples=len(training_inputs), learning_rate=learning_rate, epochs=epochs, threshold=threshold) + failed_updates = 0 + total_updates = 0 + # create_checkpoint copies the file only when it already exists. + # Saving here when it does not records the pre-cycle weights. Saving + # when it does would replace that file with a later in-memory state. + if not Path(self.adapter_path).exists(): + self.adapter.save(self.adapter_path) with SafeTrainingContext( self.checkpoint_manager, self.adapter_path, @@ -2064,9 +2108,10 @@ def _finish_annealing_cycle(): new_vectors_np = self.adapter(input_tensor).numpy() # Batch Update Logic using shared helper (F-031: pass cached payload_map) - total_updates, failed_updates = sync_vectors_to_qdrant( + total_updates, failed_updates, compensated = sync_vectors_to_qdrant( self.qdrant, self.collection_name, ordered_ids, - new_vectors_np, chunk_size, self.logger, payload_map + new_vectors_np, chunk_size, self.logger, payload_map, + original_vectors_np=input_tensor.detach().cpu().numpy(), ) # Mark success only if no failed vector updates (F-043) @@ -2077,16 +2122,25 @@ def _finish_annealing_cycle(): f"Training cycle completed successfully. Updated {total_updates} vectors.", vectors_updated=total_updates ) + self.chelation_log.clear() else: + failure_message = ( + f"Training completed but {failed_updates} vector updates failed." + ) + if compensated: + failure_message += ( + " Pre-cycle vectors were written back for the chunks that had already been stored." + ) self.logger.log_error( "sedimentation_partial_failure", - f"Training completed but {failed_updates} vector updates failed. Rolling back.", + failure_message, vectors_updated=total_updates, vectors_failed=failed_updates ) + if failed_updates: + self.adapter.load(self.adapter_path) self.logger.log_training_complete(final_loss=final_loss, vectors_updated=total_updates, vectors_failed=failed_updates) - self.chelation_log.clear() _finish_annealing_cycle() def run_offline_distillation(self, batch_size: int = 100, learning_rate: float = None, epochs: int = None): @@ -2249,6 +2303,9 @@ def run_offline_distillation(self, batch_size: int = 100, learning_rate: float = self.adapter.train() final_loss = 0.0 + # Snapshot before either trained save below. Reading the file after + # those saves would restore the weights this cycle just wrote. + prior_adapter = _pre_cycle_adapter_bytes(self.adapter, self.adapter_path) if getattr(self, '_sedimentation_optimizer_type', 'adam') == "eggroll_es": from evolution_strategies_optimizer import ( @@ -2375,50 +2432,18 @@ def run_offline_distillation(self, batch_size: int = 100, learning_rate: float = self.logger.log_event("offline_distillation_update", "Updating corpus vectors in Qdrant") with torch.no_grad(): + original_vectors_np = input_tensor.detach().cpu().numpy() new_vectors_np = self.adapter(input_tensor).numpy() chunk_size = ChelationConfig.CHUNK_SIZE - total_updates = 0 - failed_updates = 0 - - for i in range(0, len(ordered_ids), chunk_size): - chunk_ids = ordered_ids[i:i+chunk_size] - chunk_vectors = new_vectors_np[i:i+chunk_size] - - try: - existing_points = self.qdrant.retrieve( - collection_name=self.collection_name, - ids=chunk_ids, - with_vectors=False - ) - payload_map = {p.id: p.payload for p in existing_points} - - batch_points = [] - for j, doc_id in enumerate(chunk_ids): - vec = chunk_vectors[j].tolist() - pay = payload_map.get(doc_id, {}) - - batch_points.append(PointStruct( - id=doc_id, - vector=vec, - payload=pay - )) - - if batch_points: - self.qdrant.upsert( - collection_name=self.collection_name, - points=batch_points - ) - total_updates += len(batch_points) - except Exception as e: - self.logger.log_error( - "offline_distillation_update_failed", - f"Failed to update batch {i//chunk_size}", - exception=e, - batch_num=i//chunk_size - ) - failed_updates += len(chunk_ids) - + total_updates, failed_updates, _compensated = sync_vectors_to_qdrant( + self.qdrant, self.collection_name, ordered_ids, + new_vectors_np, chunk_size, self.logger, None, + original_vectors_np=original_vectors_np, + ) + if failed_updates: + _write_adapter_bytes(self.adapter_path, prior_adapter) + self.adapter.load(self.adapter_path) self.logger.log_training_complete( final_loss=final_loss, vectors_updated=total_updates, diff --git a/sedimentation.py b/sedimentation.py index 8028967..6a460c8 100644 --- a/sedimentation.py +++ b/sedimentation.py @@ -12,6 +12,7 @@ import numpy as np import torch import torch.optim as optim +from pathlib import Path from typing import List, Tuple from chelation_logger import get_logger @@ -126,6 +127,14 @@ def run_hierarchical_sedimentation(self, threshold=3, learning_rate=0.001, epoch # --- Phase 1 & 2: Training wrapped in SafeTrainingContext (F-043) --- print(f"Phase 1: Per-cluster training ({len(clusters)} clusters, {cluster_epochs} epochs)...") + failed_updates = 0 + total_updates = 0 + final_loss = 0.0 + # create_checkpoint copies the file only when it already exists. + # Saving here when it does not records the pre-cycle weights. Saving + # when it does would replace that file with a later in-memory state. + if not Path(self.engine.adapter_path).exists(): + self.engine.adapter.save(self.engine.adapter_path) with SafeTrainingContext( self.checkpoint_manager, self.engine.adapter_path, @@ -187,9 +196,10 @@ def run_hierarchical_sedimentation(self, threshold=3, learning_rate=0.001, epoch new_vectors_np = self.engine.adapter(input_tensor).numpy() # Use shared helper for Qdrant sync (F-031: pass cached payload_map) - total_updates, failed_updates = sync_vectors_to_qdrant( + total_updates, failed_updates, compensated = sync_vectors_to_qdrant( self.engine.qdrant, self.engine.collection_name, ordered_ids, - new_vectors_np, chunk_size, self.logger, payload_map + new_vectors_np, chunk_size, self.logger, payload_map, + original_vectors_np=input_tensor.detach().cpu().numpy(), ) # Mark success only if no failed vector updates (F-043) @@ -200,21 +210,29 @@ def run_hierarchical_sedimentation(self, threshold=3, learning_rate=0.001, epoch f"Hierarchical training completed successfully. Updated {total_updates} vectors.", vectors_updated=total_updates ) + self.engine.chelation_log.clear() else: + failure_message = ( + f"Training completed but {failed_updates} vector updates failed." + ) + if compensated: + failure_message += ( + " Pre-cycle vectors were written back for the chunks that had already been stored." + ) self.logger.log_error( "hierarchical_sedimentation_partial_failure", - f"Training completed but {failed_updates} vector updates failed. Rolling back.", + failure_message, vectors_updated=total_updates, vectors_failed=failed_updates ) + if failed_updates: + self.engine.adapter.load(self.engine.adapter_path) self.logger.log_training_complete( final_loss=final_loss, vectors_updated=total_updates, vectors_failed=failed_updates, ) - - self.engine.chelation_log.clear() print(f"Hierarchical sedimentation complete. Updated {total_updates} vectors, {failed_updates} failed.") print("--- HIERARCHICAL SLEEP CYCLE COMPLETE ---") diff --git a/sedimentation_trainer.py b/sedimentation_trainer.py index 81c0bfa..39542ab 100644 --- a/sedimentation_trainer.py +++ b/sedimentation_trainer.py @@ -68,7 +68,8 @@ def compute_homeostatic_target(current_vec: np.ndarray, noise_vectors: List[np.n def sync_vectors_to_qdrant(qdrant: Any, collection_name: str, ordered_ids: List, new_vectors_np: np.ndarray, chunk_size: int, logger: Any, - payload_map: dict = None) -> tuple: + payload_map: dict = None, + original_vectors_np: np.ndarray = None) -> tuple: """ Synchronize updated vectors to Qdrant in batches, preserving existing payloads. @@ -86,11 +87,59 @@ def sync_vectors_to_qdrant(qdrant: Any, collection_name: str, ordered_ids: List, skips qdrant.retrieve for payload lookup (F-031 optimization) Returns: - Tuple of (total_updates, failed_updates) counts + Tuple of (total_updates, failed_updates, compensated). + compensated is True only when a chunk failed, at least one id from + this call was already written, and the compensating upsert returned. + It is False when nothing failed, when nothing was written, when the + pre-cycle vectors were not provided, or when the compensation + retrieve or upsert raised. """ total_updates = 0 failed_updates = 0 - + written_ids = [] + id_to_index = {doc_id: idx for idx, doc_id in enumerate(ordered_ids)} + + def _restore_written() -> bool: + if not written_ids: + return True + if original_vectors_np is None: + logger.log_error( + "corpus_not_restored", + "Corpus was not restored because the pre-cycle vectors were not provided.", + ) + return False + try: + # Payload retrieve stays in this try. An exception here must not + # escape, or the caller would skip the pre-cycle adapter reload. + source_payloads = payload_map + if source_payloads is None: + existing = qdrant.retrieve( + collection_name=collection_name, + ids=written_ids, + with_vectors=False, + ) + source_payloads = {point.id: point.payload for point in existing} + points = [] + for doc_id in written_ids: + idx = id_to_index[doc_id] + pay = source_payloads.get(doc_id, {}) or {} + points.append(PointStruct( + id=doc_id, + vector=original_vectors_np[idx].tolist(), + payload=pay, + )) + if not points: + return False + qdrant.upsert(collection_name=collection_name, points=points) + return True + except Exception as restore_error: + logger.log_error( + "corpus_not_restored", + "Corpus was not restored.", + exception=restore_error, + ) + return False + for i in range(0, len(ordered_ids), chunk_size): chunk_ids = ordered_ids[i:i + chunk_size] chunk_vectors = new_vectors_np[i:i + chunk_size] @@ -124,15 +173,14 @@ def sync_vectors_to_qdrant(qdrant: Any, collection_name: str, ordered_ids: List, points=batch_points ) total_updates += len(batch_points) - except ValueError as e: - logger.log_error("database_update", - f"Invalid vector data in batch {i//chunk_size}", - exception=e, batch_num=i//chunk_size) - failed_updates += len(chunk_ids) + written_ids.extend(chunk_ids) except Exception as e: - logger.log_error("database_update", - f"Update batch {i//chunk_size} failed", + logger.log_error("database_update", + f"Update batch {i//chunk_size} failed", exception=e, batch_num=i//chunk_size) - failed_updates += len(chunk_ids) - - return total_updates, failed_updates + failed_updates += len(ordered_ids) - len(written_ids) + restored = _restore_written() + compensated = bool(written_ids) and restored + return total_updates, failed_updates, compensated + + return total_updates, failed_updates, False diff --git a/test_antigravity_engine.py b/test_antigravity_engine.py index 8087b32..32b15ce 100644 --- a/test_antigravity_engine.py +++ b/test_antigravity_engine.py @@ -344,7 +344,7 @@ def test_sedimentation_uses_safe_training_context(self): mock_stc_cls.return_value = mock_stc # Mock sync to return success (no failures) - mock_sync.return_value = (5, 0) + mock_sync.return_value = (5, 0, False) engine = self._make_engine() train_param = torch.nn.Parameter(torch.tensor(1.0)) diff --git a/test_recursive_decomposer.py b/test_recursive_decomposer.py index 48f1b45..4e0021a 100644 --- a/test_recursive_decomposer.py +++ b/test_recursive_decomposer.py @@ -805,7 +805,7 @@ def test_hierarchical_sedimentation_uses_safe_training_context(self): mock_stc_cls.return_value = mock_stc # Mock sync to return success (no failures) - mock_sync.return_value = (10, 0) + mock_sync.return_value = (10, 0, False) # Create mock engine with required attributes mock_engine = MagicMock() diff --git a/test_sedimentation_corpus_rollback.py b/test_sedimentation_corpus_rollback.py new file mode 100644 index 0000000..50d0f96 --- /dev/null +++ b/test_sedimentation_corpus_rollback.py @@ -0,0 +1,505 @@ +"""Partial Qdrant upserts write the pre-cycle vectors back.""" + +import os +import tempfile +import unittest +from pathlib import Path + +import numpy as np + +from sedimentation_trainer import sync_vectors_to_qdrant + +try: + import torch +except ImportError: + torch = None + + +class _Point: + def __init__(self, point_id, vector, payload): + self.id = point_id + self.vector = vector + self.payload = payload + + +class _Client: + def __init__(self): + self.points = {} + self.upserts = 0 + + def upsert(self, collection_name, points): + self.upserts += 1 + if any(point.id == "b" for point in points): + raise RuntimeError("second chunk failed") + for point in points: + self.points[point.id] = (list(point.vector), dict(point.payload or {})) + + def retrieve(self, collection_name, ids, with_vectors=False): + return [] + + +class _ScriptedClient: + """Succeeds a fixed number of upserts, then raises.""" + + def __init__(self, succeed_upserts, fail_retrieve_after_upserts=None): + self.succeed_upserts = succeed_upserts + self.fail_retrieve_after_upserts = fail_retrieve_after_upserts + self.upserts = 0 + self.points = {} + + def upsert(self, collection_name, points): + self.upserts += 1 + if self.upserts > self.succeed_upserts: + raise RuntimeError("upsert failed") + for point in points: + self.points[point.id] = (list(point.vector), dict(point.payload or {})) + + def retrieve(self, collection_name, ids, with_vectors=False): + if ( + self.fail_retrieve_after_upserts is not None + and self.upserts >= self.fail_retrieve_after_upserts + ): + raise RuntimeError("injected retrieve failure during compensation") + found = [] + for doc_id in ids: + if doc_id not in self.points: + continue + vector, payload = self.points[doc_id] + found.append(_Point(doc_id, vector, payload)) + return found + + +class _Logger: + def __init__(self): + self.errors = [] + + def log_error(self, kind, message, **kwargs): + self.errors.append((kind, message)) + + def log_event(self, *args, **kwargs): + return None + + def log_training_start(self, *args, **kwargs): + return None + + def log_training_epoch(self, *args, **kwargs): + return None + + def log_training_complete(self, *args, **kwargs): + return None + + def log_checkpoint(self, *args, **kwargs): + return None + + +class TestCorpusRollback(unittest.TestCase): + def test_second_chunk_failure_restores_the_first_id(self): + client = _Client() + logger = _Logger() + originals = np.array([[0.25], [0.5]], dtype=np.float32) + adapted = np.array([[9.0], [9.0]], dtype=np.float32) + total, failed, compensated = sync_vectors_to_qdrant( + client, + "docs", + ["a", "b"], + adapted, + 1, + logger, + {"a": {"text": "keep"}, "b": {"text": "later"}}, + original_vectors_np=originals, + ) + self.assertEqual(failed, 1) + self.assertTrue(compensated) + self.assertEqual(client.points["a"][0], [0.25]) + self.assertEqual(client.points["a"][1], {"text": "keep"}) + self.assertNotIn("b", client.points) + self.assertFalse(any("Rolling back" in message for _, message in logger.errors)) + self.assertGreater(total, 0) + + def test_compensating_upsert_leaves_the_adapted_vector(self): + client = _ScriptedClient(succeed_upserts=1) + logger = _Logger() + originals = np.array([[0.25], [0.5]], dtype=np.float32) + adapted = np.array([[9.0], [9.0]], dtype=np.float32) + total, failed, compensated = sync_vectors_to_qdrant( + client, + "docs", + ["a", "b"], + adapted, + 1, + logger, + {"a": {"text": "keep"}, "b": {"text": "later"}}, + original_vectors_np=originals, + ) + self.assertEqual(failed, 1) + self.assertGreater(total, 0) + self.assertFalse(compensated) + self.assertEqual(client.points["a"][0], [9.0]) + self.assertNotIn("b", client.points) + messages = [message for _, message in logger.errors] + self.assertIn("Corpus was not restored.", messages) + self.assertFalse(any("written back" in message for message in messages)) + self.assertFalse(any("Rolling back" in message for message in messages)) + + def test_compensation_retrieve_failure_does_not_escape(self): + client = _ScriptedClient(succeed_upserts=1, fail_retrieve_after_upserts=2) + logger = _Logger() + originals = np.array([[0.25], [0.5]], dtype=np.float32) + adapted = np.array([[9.0], [9.0]], dtype=np.float32) + total, failed, compensated = sync_vectors_to_qdrant( + client, + "docs", + ["a", "b"], + adapted, + 1, + logger, + None, + original_vectors_np=originals, + ) + self.assertEqual(failed, 1) + self.assertGreater(total, 0) + self.assertFalse(compensated) + self.assertEqual(client.points["a"][0], [9.0]) + messages = [message for _, message in logger.errors] + self.assertIn("Corpus was not restored.", messages) + self.assertFalse(any("written back" in message for message in messages)) + + +class _StoredPoint: + def __init__(self, point_id, vector, payload): + self.id = point_id + self.vector = [float(value) for value in vector] + self.payload = dict(payload or {}) + + +class _Corpus: + def __init__(self, rows, adapter_path, fail_on_id="b", fail_compensation=False, + fail_compensation_retrieve=False): + self.points = { + row_id: _StoredPoint(row_id, vector, payload) + for row_id, vector, payload in rows + } + self.adapter_path = adapter_path + self.fail_on_id = fail_on_id + self.fail_compensation = fail_compensation + self.fail_compensation_retrieve = fail_compensation_retrieve + self.second_failed = False + self.trained_adapter_bytes = None + self.first_written = {} + self.upsert_ids = [] + + def scroll(self, collection_name, limit, with_vectors, with_payload, offset): + return list(self.points.values()), None + + def retrieve(self, collection_name, ids, with_vectors=False): + if self.fail_compensation_retrieve and self.second_failed and not with_vectors: + raise RuntimeError("injected retrieve failure during compensation") + return [self.points[doc_id] for doc_id in ids] + + def upsert(self, collection_name, points): + if self.trained_adapter_bytes is None: + self.trained_adapter_bytes = Path(self.adapter_path).read_bytes() + ids = [point.id for point in points] + self.upsert_ids.append(ids) + if self.fail_on_id is not None and any(point.id == self.fail_on_id for point in points): + self.second_failed = True + raise RuntimeError("second chunk failed") + if self.fail_compensation and self.second_failed: + raise RuntimeError("compensation failed") + for point in points: + vector = [float(value) for value in point.vector] + if point.id not in self.first_written: + self.first_written[point.id] = vector + self.points[point.id] = _StoredPoint(point.id, vector, point.payload) + + +def _clone_state(adapter): + return {key: value.detach().cpu().clone() for key, value in adapter.state_dict().items()} + + +def _assert_adapter_state(test, adapter, expected): + current = adapter.state_dict() + test.assertEqual(set(current), set(expected)) + for key, value in expected.items(): + test.assertTrue(torch.equal(current[key].detach().cpu(), value), key) + + +def _assert_file_state(test, path, expected): + loaded = torch.load(path, weights_only=True) + test.assertEqual(set(loaded), set(expected)) + for key, value in expected.items(): + test.assertTrue(torch.equal(loaded[key].detach().cpu(), value), key) + + +class _Teacher: + _projection_enabled = False + _projection = None + + def check_dimension_compatibility(self, student_dim): + return student_dim == 4 + + def generate_distillation_targets(self, texts, current_embeddings, teacher_weight=1.0): + current = np.asarray(current_embeddings, dtype=np.float32) + target = np.zeros_like(current) + target[:, -1] = 1.0 + return target + + +@unittest.skipUnless(torch is not None, "adapter reload tests need torch") +class TestProductionAdapterRollback(unittest.TestCase): + DIM = 4 + + def setUp(self): + from config import ChelationConfig + self._chunk_size = ChelationConfig.CHUNK_SIZE + ChelationConfig.CHUNK_SIZE = 1 + self._tmp = tempfile.TemporaryDirectory() + self.tmp = Path(self._tmp.name) + + def tearDown(self): + from config import ChelationConfig + ChelationConfig.CHUNK_SIZE = self._chunk_size + self._tmp.cleanup() + + def _basis(self, index): + values = [0.0] * self.DIM + values[index] = 1.0 + return values + + def _rows(self): + payload = {"a": {"text": "alpha"}, "b": {"text": "beta"}} + return [ + (doc_id, self._basis(0), payload[doc_id]) + for doc_id in ("a", "b") + ] + + def _new_adapter(self): + from chelation_adapter import ChelationAdapter + torch.manual_seed(0) + return ChelationAdapter(self.DIM) + + def _run_sedimentation(self, client, adapter, path, logger): + from types import SimpleNamespace + from antigravity_engine import AntigravityEngine + from checkpoint_manager import CheckpointManager + noise = np.array(self._basis(1), dtype=np.float32) + chelation_log = { + "a": [noise.copy()], + "b": [noise.copy()], + } + engine = SimpleNamespace( + training_mode="baseline", + teacher_helper=None, + chelation_log=chelation_log, + logger=logger, + qdrant=client, + collection_name="docs", + adapter=adapter, + adapter_path=path, + checkpoint_manager=CheckpointManager(self.tmp / "checkpoints"), + _stability_tracker=None, + _sedimentation_optimizer_type="adam", + _sedimentation_loss_type="mse", + _sedimentation_loss_kwargs={}, + _convergence_enabled=False, + _kalman_lr_enabled=False, + _weight_scheduler=None, + _annealing_controller=None, + ) + AntigravityEngine.run_sedimentation_cycle( + engine, + threshold=1, + learning_rate=1.0, + epochs=5, + noise_injection=0.0, + ) + return chelation_log + + def _assert_failure_keeps_pre_cycle_weights(self, adapter, path, pre_state, trained_bytes): + self.assertNotEqual(path.read_bytes(), trained_bytes) + _assert_file_state(self, path, pre_state) + _assert_adapter_state(self, adapter, pre_state) + + def test_failed_sedimentation_restores_prefix_and_keeps_log(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path) + chelation_log = self._run_sedimentation(client, adapter, path, logger) + messages = [message for _, message in logger.errors] + self.assertIn("a", chelation_log) + self.assertIn("b", chelation_log) + self.assertTrue(any("written back" in message for message in messages)) + self.assertFalse(any("Rolling back" in message for message in messages)) + self.assertIn("sedimentation_partial_failure", [kind for kind, _ in logger.errors]) + np.testing.assert_allclose(client.points["a"].vector, self._basis(0), atol=1e-5) + np.testing.assert_allclose(client.points["b"].vector, self._basis(0), atol=1e-5) + self._assert_failure_keeps_pre_cycle_weights( + adapter, path, pre_state, client.trained_adapter_bytes + ) + + def test_compensation_failure_does_not_claim_writeback(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path, fail_compensation=True) + chelation_log = self._run_sedimentation(client, adapter, path, logger) + messages = [message for _, message in logger.errors] + kinds = [kind for kind, _ in logger.errors] + self.assertIn("a", chelation_log) + self.assertIn("corpus_not_restored", kinds) + self.assertIn("sedimentation_partial_failure", kinds) + self.assertIn("Corpus was not restored.", messages) + self.assertFalse(any("written back" in message for message in messages)) + self.assertFalse(any("Rolling back" in message for message in messages)) + self.assertEqual(client.points["a"].vector, client.first_written["a"]) + self.assertFalse(np.allclose(client.points["a"].vector, self._basis(0), atol=1e-5)) + self._assert_failure_keeps_pre_cycle_weights( + adapter, path, pre_state, client.trained_adapter_bytes + ) + + def test_successful_sedimentation_clears_chelation_log(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path, fail_on_id=None) + chelation_log = self._run_sedimentation(client, adapter, path, logger) + self.assertEqual(chelation_log, {}) + self.assertFalse(logger.errors) + self.assertEqual(client.upsert_ids, [["a"], ["b"]]) + self.assertEqual(path.read_bytes(), client.trained_adapter_bytes) + loaded = torch.load(path, weights_only=True) + self.assertFalse(torch.equal( + loaded["correction_net.0.weight"].detach().cpu(), + pre_state["correction_net.0.weight"], + )) + + def test_existing_adapter_file_is_not_overwritten_before_checkpoint(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + adapter.save(path) + pre_bytes = path.read_bytes() + with torch.no_grad(): + adapter.correction_net[0].weight.add_(1.0) + logger = _Logger() + client = _Corpus(self._rows(), path, fail_compensation=True) + self._run_sedimentation(client, adapter, path, logger) + self.assertEqual(path.read_bytes(), pre_bytes) + _assert_adapter_state(self, adapter, torch.load(path, weights_only=True)) + self.assertFalse(any("written back" in message for _, message in logger.errors)) + + def test_hierarchical_without_adapter_file_reloads_pre_step_weights(self): + from checkpoint_manager import CheckpointManager + from sedimentation import HierarchicalSedimentationEngine + path = self.tmp / "hierarchical.pt" + adapter = self._new_adapter() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path) + noise = np.array(self._basis(1), dtype=np.float32) + engine = type("Engine", (), {})() + engine.chelation_log = {"a": [noise.copy()], "b": [noise.copy()]} + engine.qdrant = client + engine.collection_name = "docs" + engine.adapter = adapter + engine.adapter_path = path + cwd = os.getcwd() + os.chdir(self.tmp) + try: + hierarchical = HierarchicalSedimentationEngine(engine) + hierarchical.logger = logger + hierarchical.checkpoint_manager = CheckpointManager(self.tmp / "hier-checkpoints") + hierarchical.run_hierarchical_sedimentation( + threshold=1, + learning_rate=1.0, + epochs=4, + ) + finally: + os.chdir(cwd) + self.assertIn("a", engine.chelation_log) + self.assertIn("b", engine.chelation_log) + kinds = [kind for kind, _ in logger.errors] + self.assertIn("hierarchical_sedimentation_partial_failure", kinds) + self.assertFalse(any("written back" in message for _, message in logger.errors)) + self.assertFalse(any("Rolling back" in message for _, message in logger.errors)) + self._assert_failure_keeps_pre_cycle_weights( + adapter, path, pre_state, client.trained_adapter_bytes + ) + + def _run_offline(self, client, adapter, path, logger): + from types import SimpleNamespace + from antigravity_engine import AntigravityEngine + engine = SimpleNamespace( + teacher_helper=_Teacher(), + logger=logger, + training_mode="offline", + vector_size=self.DIM, + qdrant=client, + collection_name="docs", + adapter=adapter, + adapter_path=path, + _sedimentation_optimizer_type="adam", + _sedimentation_loss_type="mse", + _sedimentation_loss_kwargs={}, + _convergence_enabled=False, + _kalman_lr_enabled=False, + _weight_scheduler=None, + ) + AntigravityEngine.run_offline_distillation( + engine, + batch_size=10, + learning_rate=1.0, + epochs=5, + ) + + def test_offline_failed_second_upsert_reloads_existing_file(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + adapter.save(path) + pre_bytes = path.read_bytes() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path) + self._run_offline(client, adapter, path, logger) + self.assertNotEqual(client.trained_adapter_bytes, pre_bytes) + self.assertEqual(path.read_bytes(), pre_bytes) + _assert_file_state(self, path, pre_state) + _assert_adapter_state(self, adapter, pre_state) + np.testing.assert_allclose(client.points["a"].vector, self._basis(0), atol=1e-5) + + def test_offline_missing_file_reloads_weights_saved_before_training(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path) + self.assertFalse(path.exists()) + self._run_offline(client, adapter, path, logger) + self.assertTrue(path.exists()) + self._assert_failure_keeps_pre_cycle_weights( + adapter, path, pre_state, client.trained_adapter_bytes + ) + + def test_offline_compensation_retrieve_failure_reloads_adapter(self): + path = self.tmp / "adapter.pt" + adapter = self._new_adapter() + adapter.save(path) + pre_bytes = path.read_bytes() + pre_state = _clone_state(adapter) + logger = _Logger() + client = _Corpus(self._rows(), path, fail_compensation_retrieve=True) + self._run_offline(client, adapter, path, logger) + messages = [message for _, message in logger.errors] + self.assertIn("Corpus was not restored.", messages) + self.assertFalse(any("written back" in message for message in messages)) + self.assertEqual(path.read_bytes(), pre_bytes) + _assert_adapter_state(self, adapter, pre_state) + self.assertEqual(client.points["a"].vector, client.first_written["a"]) + self.assertFalse(np.allclose(client.points["a"].vector, self._basis(0), atol=1e-5)) + + +if __name__ == "__main__": + unittest.main() diff --git a/test_sedimentation_trainer.py b/test_sedimentation_trainer.py index c2fde62..2864b7f 100644 --- a/test_sedimentation_trainer.py +++ b/test_sedimentation_trainer.py @@ -187,13 +187,14 @@ def test_successful_sync(self): ordered_ids = ["doc1"] new_vectors = np.array([[0.1, 0.2, 0.3]]) - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=10, logger=mock_logger ) self.assertEqual(total, 1) self.assertEqual(failed, 0) + self.assertFalse(compensated) mock_qdrant.retrieve.assert_called_once() mock_qdrant.upsert.assert_called_once() @@ -221,7 +222,7 @@ def mock_retrieve(collection_name, ids, with_vectors): mock_qdrant.upsert.return_value = None chunk_size = 10 - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=chunk_size, logger=mock_logger ) @@ -231,6 +232,7 @@ def mock_retrieve(collection_name, ids, with_vectors): self.assertEqual(mock_qdrant.upsert.call_count, 3) self.assertEqual(total, n_docs) self.assertEqual(failed, 0) + self.assertFalse(compensated) def test_payload_preservation(self): """Test that existing payloads are preserved.""" @@ -272,17 +274,18 @@ def test_value_error_handling(self): ordered_ids = ["doc1", "doc2"] new_vectors = np.array([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]) - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=10, logger=mock_logger ) self.assertEqual(total, 0) self.assertEqual(failed, 2) + self.assertFalse(compensated) mock_logger.log_error.assert_called_once() call_args = mock_logger.log_error.call_args self.assertEqual(call_args[0][0], "database_update") - self.assertIn("Invalid vector data", call_args[0][1]) + self.assertEqual(call_args[0][1], "Update batch 0 failed") def test_generic_exception_handling(self): """Test that generic exceptions are caught and logged.""" @@ -299,13 +302,14 @@ def test_generic_exception_handling(self): ordered_ids = ["doc1"] new_vectors = np.array([[0.1, 0.2, 0.3]]) - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=10, logger=mock_logger ) self.assertEqual(total, 0) self.assertEqual(failed, 1) + self.assertFalse(compensated) mock_logger.log_error.assert_called_once() call_args = mock_logger.log_error.call_args self.assertEqual(call_args[0][0], "database_update") @@ -336,14 +340,17 @@ def side_effect_retrieve(*args, **kwargs): ordered_ids = ["doc1", "doc2"] new_vectors = np.array([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]) - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=1, logger=mock_logger ) self.assertEqual(total, 1) # First chunk succeeded self.assertEqual(failed, 1) # Second chunk failed - self.assertEqual(mock_logger.log_error.call_count, 1) + self.assertFalse(compensated) + kinds = [call[0][0] for call in mock_logger.log_error.call_args_list] + self.assertIn("database_update", kinds) + self.assertIn("corpus_not_restored", kinds) def test_empty_batch(self): """Test handling of empty batches.""" @@ -353,13 +360,14 @@ def test_empty_batch(self): ordered_ids = [] new_vectors = np.array([]).reshape(0, 384) - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=10, logger=mock_logger ) self.assertEqual(total, 0) self.assertEqual(failed, 0) + self.assertFalse(compensated) mock_qdrant.retrieve.assert_not_called() mock_qdrant.upsert.assert_not_called() @@ -382,7 +390,7 @@ def capture_upsert(collection_name, points): upserted_points.extend(points) mock_qdrant.upsert.side_effect = capture_upsert - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=10, logger=mock_logger, payload_map=payload_map @@ -395,6 +403,7 @@ def capture_upsert(collection_name, points): self.assertEqual(mock_qdrant.upsert.call_count, 1) self.assertEqual(total, 2) self.assertEqual(failed, 0) + self.assertFalse(compensated) # Verify payloads from provided map were used self.assertEqual(len(upserted_points), 2) @@ -423,7 +432,7 @@ def mock_retrieve(collection_name, ids, with_vectors): new_vectors = np.array([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]) # Call WITHOUT payload_map - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=10, logger=mock_logger ) @@ -432,6 +441,7 @@ def mock_retrieve(collection_name, ids, with_vectors): mock_qdrant.retrieve.assert_called_once() self.assertEqual(total, 2) self.assertEqual(failed, 0) + self.assertFalse(compensated) def test_payload_map_chunking(self): """Test F-031: payload_map works correctly with multiple chunks.""" @@ -447,7 +457,7 @@ def test_payload_map_chunking(self): mock_qdrant.upsert.return_value = None chunk_size = 10 - total, failed = sync_vectors_to_qdrant( + total, failed, compensated = sync_vectors_to_qdrant( mock_qdrant, "test_collection", ordered_ids, new_vectors, chunk_size=chunk_size, logger=mock_logger, payload_map=payload_map @@ -460,6 +470,7 @@ def test_payload_map_chunking(self): self.assertEqual(mock_qdrant.upsert.call_count, 3) self.assertEqual(total, n_docs) self.assertEqual(failed, 0) + self.assertFalse(compensated) if __name__ == "__main__":