CoolFace
Modelpublic

Engram-protocol/engram

sourceHugging Faceapache-2.0updated 6mo agoView on Hugging Face
0likes
test_manifold_index.py98 linesDownload Raw Back to tests
1"""2ENGRAM Protocol — Manifold Index Tests3Tests for FAISS IndexFlatIP add/search/remove/persist (D2, D4).4"""5 6from __future__ import annotations7 8from pathlib import Path9 10import numpy as np11import pytest12import torch13 14from kvcos.core.manifold_index import IndexEntry, ManifoldIndex15 16 17def _entry(cid: str = "c1", model: str = "llama") -> IndexEntry:18    return IndexEntry(19        cache_id=cid, task_description="test",20        model_id=model, created_at="2026-01-01T00:00:00Z",21        context_len=256, l2_norm=1.0,22    )23 24 25class TestAddAndSearch:26    """Add vectors, search via MIPS."""27 28    def test_add_increments(self) -> None:29        idx = ManifoldIndex(dim=8)30        idx.add(torch.randn(8), _entry("a"))31        idx.add(torch.randn(8), _entry("b"))32        assert idx.n_entries == 233 34    def test_search_returns_correct_order(self) -> None:35        idx = ManifoldIndex(dim=4)36        v1 = torch.tensor([1.0, 0.0, 0.0, 0.0])37        v2 = torch.tensor([0.0, 1.0, 0.0, 0.0])38        idx.add(v1, _entry("close"))39        idx.add(v2, _entry("far"))40 41        query = torch.tensor([1.0, 0.0, 0.0, 0.0])42        results = idx.search(query, top_k=2)43        assert results[0]["cache_id"] == "close"44        assert results[0]["similarity"] > results[1]["similarity"]45 46    def test_search_empty_returns_empty(self) -> None:47        idx = ManifoldIndex(dim=4)48        results = idx.search(torch.randn(4), top_k=5)49        assert results == []50 51    def test_model_filter(self) -> None:52        idx = ManifoldIndex(dim=4)53        idx.add(torch.randn(4), _entry("a", model="llama"))54        idx.add(torch.randn(4), _entry("b", model="phi"))55        results = idx.search(torch.randn(4), top_k=10, model_id="phi")56        assert all(r["model_id"] == "phi" for r in results)57 58 59class TestRemoveAndRebuild:60    """Remove entries and rebuild index."""61 62    def test_remove_hides_from_search(self) -> None:63        idx = ManifoldIndex(dim=4)64        v = torch.tensor([1.0, 0.0, 0.0, 0.0])65        idx.add(v, _entry("target"))66        assert idx.remove("target")67        results = idx.search(v, top_k=1)68        assert len(results) == 069 70    def test_rebuild_compacts(self) -> None:71        idx = ManifoldIndex(dim=4)72        for i in range(5):73            idx.add(torch.randn(4), _entry(f"c{i}"))74        idx.remove("c1")75        idx.remove("c3")76        active = idx.rebuild()77        assert active == 378 79 80class TestPersistence:81    """Save/load round-trip (D2: serialize_index/deserialize_index)."""82 83    def test_save_load_round_trip(self, tmp_index_dir: Path) -> None:84        idx = ManifoldIndex(dim=4)85        v1 = torch.tensor([1.0, 0.0, 0.0, 0.0])86        idx.add(v1, _entry("persisted"))87        idx.save(tmp_index_dir / "test.faiss")88 89        idx2 = ManifoldIndex(dim=4, index_path=tmp_index_dir / "test.faiss")90        assert idx2.n_entries == 191        results = idx2.search(v1, top_k=1)92        assert results[0]["cache_id"] == "persisted"93 94    def test_dim_mismatch_raises(self) -> None:95        idx = ManifoldIndex(dim=4)96        with pytest.raises(ValueError, match="dim"):97            idx.add(torch.randn(8), _entry("wrong"))98