Expanded_Repetition / correction_memory.py
Expanded-Repetition's picture
Upload 12 files
1e214ed verified
Raw History Blame Contribute Delete
7.31 kB
"""Persistent memory for explicit user corrections, with optional private HF Dataset sync.
Local-only by default (stdlib). Set EIM_CORRECTIONS_REPO or EIM_MEMORY_REPO to a
private Hugging Face Dataset repo and HF_TOKEN to a token with write access to keep
corrections across ephemeral Space restarts. Remote sync failures never break chat.
"""
from __future__ import annotations
import json
import os
import re
import shutil
import threading
import time
from pathlib import Path
_EN_MARKERS = (
"that's wrong", "that is wrong", "wrong answer", "not what i asked", "i said",
"i told you", "you repeated", "don't repeat", "do not repeat", "you made the same",
"not correct", "you forgot", "i already said", "stop doing", "instead of",
)
_AR_MARKERS = (
"غلط", "مو هذا", "مو هيج", "مو هيچ", "قلتلك", "كلتلك", "نفس الخطأ", "نفس الاخطاء",
"لا تكرر", "لا تعيد", "كررت", "نسيت", "مو اللي طلبته", "ما طلبت", "مو صحيح", "خطأ",
)
_STOP = set("the a an and or to of for in on with this that it is are was were be do did you your i me my we please fix make write code answer about from into using use".split())
def _tokens(text: str) -> set[str]:
words = re.findall(r"[a-zA-Z0-9_]+|[\u0600-\u06FF]+", (text or "").lower())
return {w for w in words if len(w) > 1 and w not in _STOP}
class CorrectionMemory:
MAX_RECORDS = 300
PUSH_DELAY = 2.0
def __init__(self, path: str | None = None, repo: str | None = None):
self.path = os.path.abspath(path or os.environ.get("EIM_CORRECTIONS_PATH", "eim_corrections.jsonl"))
self.repo = (repo if repo is not None else (
os.environ.get("EIM_CORRECTIONS_REPO") or os.environ.get("EIM_MEMORY_REPO", "")
)).strip()
self.token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGINGFACEHUB_API_TOKEN") or None
self._lock = threading.RLock()
self._timer: threading.Timer | None = None
self.sync_status = "disabled" if not self.repo else "configured"
Path(self.path).parent.mkdir(parents=True, exist_ok=True)
if self.repo:
self._pull_remote()
@staticmethod
def is_correction(text: str) -> bool:
low = (text or "").lower()
return any(x in low for x in _EN_MARKERS + _AR_MARKERS)
def add(self, text: str) -> bool:
text = (text or "").strip()
if not text or not self.is_correction(text):
return False
stored_text = text[:2000]
norm = " ".join(stored_text.split()).casefold()
with self._lock:
existing = self._load()
if any(" ".join(r.get("text", "").split()).casefold() == norm for r in existing):
return False
rows = (existing + [{"ts": int(time.time()), "text": stored_text}])[-self.MAX_RECORDS:]
temporary = f"{self.path}.{os.getpid()}.{threading.get_ident()}.tmp"
try:
with open(temporary, "w", encoding="utf-8") as f:
for row in rows:
f.write(json.dumps(row, ensure_ascii=False) + "\n")
os.replace(temporary, self.path)
finally:
try:
os.unlink(temporary)
except OSError:
pass
self._schedule_push()
return True
def _load(self) -> list[dict]:
rows = []
try:
with open(self.path, encoding="utf-8") as f:
for line in f:
try:
row = json.loads(line)
if isinstance(row, dict) and isinstance(row.get("text"), str):
rows.append(row)
except (ValueError, TypeError):
continue
except OSError:
pass
return rows[-self.MAX_RECORDS:]
def relevant(self, query: str, limit: int = 4) -> list[str]:
rows = self._load()
if not rows:
return []
q = _tokens(query)
scored = []
for i, row in enumerate(rows):
words = _tokens(row.get("text", ""))
overlap = len(q & words) / max(1, len(q | words))
# Recency is a tie-breaker, not a replacement for relevance.
score = overlap + 0.015 * (i / max(1, len(rows) - 1))
if overlap > 0 or i >= len(rows) - 3:
scored.append((score, i, row.get("text", "")))
scored.sort(reverse=True)
return [text for _, _, text in scored[:max(1, limit)] if text]
def prompt(self, query: str, limit: int = 4) -> str:
lessons = self.relevant(query, limit)
if not lessons:
return ""
bullets = "\n".join(f"- {item}" for item in lessons)
return ("Persistent user corrections from earlier turns. Treat these as constraints; do not repeat "
"rejected approaches. If a correction conflicts with the current explicit request, follow the current request.\n"
+ bullets)
def _pull_remote(self) -> None:
"""Pull corrections from a configured dataset repo; tolerate missing repo/file/offline mode."""
try:
from huggingface_hub import hf_hub_download
local = hf_hub_download(
repo_id=self.repo,
repo_type="dataset",
filename=os.path.basename(self.path),
token=self.token,
local_dir=os.path.dirname(self.path),
)
if os.path.abspath(local) != self.path and os.path.isfile(local):
shutil.copyfile(local, self.path)
self.sync_status = "pulled"
except Exception as exc: # A missing file/new repo is normal on first launch.
self.sync_status = f"pull-unavailable:{type(exc).__name__}"
def _schedule_push(self) -> None:
if not self.repo:
return
with self._lock:
if self._timer is not None:
self._timer.cancel()
self._timer = threading.Timer(self.PUSH_DELAY, self.sync_now)
self._timer.daemon = True
self._timer.start()
def sync_now(self) -> bool:
"""Push the JSONL file to the configured dataset. Returns True only after upload succeeds."""
if not self.repo:
self.sync_status = "disabled"
return False
if not self.token:
self.sync_status = "push-unavailable:missing-token"
return False
try:
from huggingface_hub import HfApi
with self._lock:
api = HfApi(token=self.token)
api.create_repo(repo_id=self.repo, repo_type="dataset", private=True, exist_ok=True)
api.upload_file(
path_or_fileobj=self.path,
path_in_repo=os.path.basename(self.path),
repo_id=self.repo,
repo_type="dataset",
commit_message="Update EIM user-correction memory",
)
self.sync_status = "pushed"
return True
except Exception as exc:
self.sync_status = f"push-unavailable:{type(exc).__name__}"
return False