Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
115 changes: 70 additions & 45 deletions antigravity_engine.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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):
"""
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand All @@ -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):
Expand Down Expand Up @@ -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 (
Expand Down Expand Up @@ -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,
Expand Down
28 changes: 23 additions & 5 deletions sedimentation.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand All @@ -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 ---")

Expand Down
74 changes: 61 additions & 13 deletions sedimentation_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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]
Expand Down Expand Up @@ -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
2 changes: 1 addition & 1 deletion test_antigravity_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
2 changes: 1 addition & 1 deletion test_recursive_decomposer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
Loading
Loading