Spaces:
Running on Zero
Running on Zero
Download correction_memory.py from Expanded-Repetition/Expanded_Repetition: direct link, hf CLI and curl.
- Browser
- Download file 7.31 kB
-
https://huggingface.co/spaces/Expanded-Repetition/Expanded_Repetition/resolve/main/correction_memory.py
- Command line
-
hf download hf://spaces/Expanded-Repetition/Expanded_Repetition/correction_memory.py
-
curl -L -o correction_memory.py https://huggingface.co/spaces/Expanded-Repetition/Expanded_Repetition/resolve/main/correction_memory.py
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() | |
| 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 | |