Franek-Le's picture
Do something idk.
4a3e249
Raw
History Blame Contribute Delete
4.35 kB
import sqlite3
import json
import time
import numpy as np
from typing import List, Dict, Any, Optional, Tuple
from sentence_transformers import SentenceTransformer
class EmbeddingMemory:
"""
Long-term memory system for AI VTuber with embeddings.
Core features:
- Every message is saved
- Every message is embedded
- Semantic search (meaning-based recall)
- SQLite persistence
"""
def __init__(
self,
db_path: str = "vtuber_memory.db",
embedding_model: str = "all-MiniLM-L6-v2"
):
self.db_path = db_path
self.conn = sqlite3.connect(
self.db_path,
check_same_thread=False
)
self.model = SentenceTransformer(embedding_model)
self._create_tables()
def _create_tables(self):
cursor = self.conn.cursor()
cursor.execute("""
CREATE TABLE IF NOT EXISTS messages (
id INTEGER PRIMARY KEY AUTOINCREMENT,
timestamp REAL,
session_id TEXT,
role TEXT,
content TEXT,
embedding BLOB,
metadata TEXT
)
""")
self.conn.commit()
def _embed(self, text: str) -> np.ndarray:
return self.model.encode(text)
def _serialize_embedding(self, emb: np.ndarray) -> bytes:
return emb.astype(np.float32).tobytes()
def _deserialize_embedding(self, blob: bytes) -> np.ndarray:
return np.frombuffer(blob, dtype=np.float32)
def add_message(
self,
session_id: str,
role: str,
content: str,
metadata: Optional[Dict[str, Any]] = None
):
emb = self._embed(content)
cursor = self.conn.cursor()
cursor.execute("""
INSERT INTO messages (timestamp, session_id, role, content, embedding, metadata)
VALUES (?, ?, ?, ?, ?, ?)
""", (
time.time(),
session_id,
role,
content,
self._serialize_embedding(emb),
json.dumps(metadata or {})
))
self.conn.commit()
def search(
self,
query: str,
session_id: Optional[str] = None,
top_k: int = 5
) -> List[Dict]:
query_emb = self._embed(query)
cursor = self.conn.cursor()
if session_id:
cursor.execute("""
SELECT timestamp, role, content, embedding
FROM messages
WHERE session_id = ?
""", (session_id,))
else:
cursor.execute("""
SELECT timestamp, role, content, embedding
FROM messages
""")
rows = cursor.fetchall()
scored: List[Tuple[float, Dict]] = []
for ts, role, content, emb_blob in rows:
emb = self._deserialize_embedding(emb_blob)
# cosine similarity
score = self._cosine_similarity(query_emb, emb)
scored.append((score, {
"timestamp": ts,
"role": role,
"content": content
}))
scored.sort(key=lambda x: x[0], reverse=True)
return [item for _, item in scored[:top_k]]
def _cosine_similarity(self, a: np.ndarray, b: np.ndarray) -> float:
a = a / (np.linalg.norm(a) + 1e-8)
b = b / (np.linalg.norm(b) + 1e-8)
return float(np.dot(a, b))
def build_context(self, session_id: str, query: str, k: int = 5) -> str:
"""
Builds memory context using semantic recall.
"""
relevant = self.search_similar(query, session_id=session_id, top_k=k)
formatted = []
for m in relevant:
formatted.append(f"{m['role'].upper()}: {m['content']}")
return "\n".join(formatted)
def get_recent(self, session_id: str, limit: int = 20) -> List[Dict]:
cursor = self.conn.cursor()
cursor.execute("""
SELECT timestamp, role, content
FROM messages
WHERE session_id = ?
ORDER BY id DESC
LIMIT ?
""", (session_id, limit))
rows = cursor.fetchall()
return [
{"timestamp": r[0], "role": r[1], "content": r[2]}
for r in reversed(rows)
]
def close(self):
self.conn.close()