Spaces:
Running on Zero
Running on Zero
Download chat_app.py from Expanded-Repetition/Expanded_Repetition: direct link, hf CLI and curl.
- Browser
- Download file 49 kB
-
https://huggingface.co/spaces/Expanded-Repetition/Expanded_Repetition/resolve/main/chat_app.py
- Command line
-
hf download hf://spaces/Expanded-Repetition/Expanded_Repetition/chat_app.py
-
curl -L -o chat_app.py https://huggingface.co/spaces/Expanded-Repetition/Expanded_Repetition/resolve/main/chat_app.py
49 kB
| """ | |
| EIM Chat - a multimodal chat app on top of the EIM engine | |
| ============================================================ | |
| python chat_app.py start the app | |
| python chat_app.py --selftest check the core without a model, GPU or Gradio | |
| What it does | |
| * Chat with streaming answers, in the language of the user's message (Arabic or English). | |
| * Attach files in the same message: code, text, PDF, Word, Excel/CSV, notebooks, zip, images. | |
| Big files are not blindly truncated: a BM25 retriever picks the parts relevant to each question. | |
| * Images: shown in the chat; understood through an optional vision model (EIM_VISION_MODEL); without one the app | |
| says so honestly instead of inventing what is in the picture. OCR is used when pytesseract is installed. | |
| * Verified code: whenever the message contains tests (assert lines, a fenced test block, or an attached | |
| test_*.py), the EIM+ loop runs instead of plain chat: it can start from YOUR code, repairs it in the sandbox, | |
| checks the result on hidden tests, flags hard-coded answers, and gives you solution.py to download. | |
| * Mode switch: Auto / Chat / EIM. | |
| Environment | |
| EIM_MODEL code model used by the EIM loop (default from app.py) | |
| EIM_CHAT_MODEL model for conversation (default: the same model; a bigger instruct model chats much better, | |
| for example Qwen/Qwen2.5-7B-Instruct) | |
| EIM_VISION_MODEL optional vision-language model, for example Qwen/Qwen2-VL-2B-Instruct (not loaded unless set) | |
| EIM_CONTEXT_TOKENS prompt budget for chat history + file excerpts (default 6000, or 4000 on weak machines) | |
| EIM_PROFILE lean | balanced | full - force the hardware profile (default: detected from CPU / RAM / GPU) | |
| EIM_TIME_BUDGET seconds after which the EIM loop starts no new round and returns its best version (default: off) | |
| EIM_GPU_SECONDS ZeroGPU only: the `duration` of the @spaces.GPU call in app.py (default 60). The EIM loop stops | |
| starting new rounds at 65% of it and returns its best version instead of being cut off | |
| EIM_MEMORY_REPO optional private HF dataset ("user/name", needs HF_TOKEN) that keeps the learned memory across restarts | |
| EIM_REPAIR_LOG 0 disables the verified-repair log (eim_repairs.jsonl next to the memory file; EIM_REPAIR_LOG_PATH moves it) | |
| EIM_RESULT_CACHE 0 disables replaying identical, fully successful EIM runs from memory | |
| Security: uploaded files are only READ as text. Code is executed only by the EIM sandbox (see app.py for its limits; | |
| for public traffic run everything in a container without network access). | |
| """ | |
| from __future__ import annotations | |
| import ast | |
| import os | |
| import re | |
| import sys | |
| import tempfile | |
| import threading | |
| import time | |
| from dataclasses import dataclass, field | |
| from typing import Iterator, Protocol, Sequence | |
| from app import CFG, EUTV, Config, ExperienceMemory, Step, gpu, get_lm | |
| from eim_plus import (Attachment, EIMPlus, Report, SmartMemory, hardware_profile, is_test_file, read_attachment, | |
| select_context, tests_from_source) | |
| from correction_memory import CorrectionMemory | |
| def _accepts(function, name: str) -> bool: | |
| try: | |
| import inspect | |
| return name in inspect.signature(function).parameters | |
| except (TypeError, ValueError): | |
| return False | |
| _GEN_LOCK = threading.Lock() # one generation at a time on the shared model | |
| # ============================================================================ | |
| # 1. Language helpers | |
| # ============================================================================ | |
| STR = { | |
| "reading_image": ("🔍 Reading the image…", "🔍 جارٍ قراءة الصورة…"), | |
| "thinking": ("⏳ Thinking…", "⏳ جارٍ التفكير…"), | |
| "eim_start": ("⏳ Generating the first version and verifying it in the sandbox…", | |
| "⏳ جارٍ توليد النسخة الأولى والتحقق منها في البيئة المعزولة…"), | |
| "eim_needs_tests": ("EIM mode needs tests. Add lines such as `assert f(2) == 4`, or attach a `test_*.py` file.", | |
| "وضع EIM يحتاج اختبارات. أضف أسطراً مثل `assert f(2) == 4` أو أرفق ملف `test_*.py`."), | |
| "default_q": ("Summarise the attached files and tell me what is in them.", | |
| "لخّص الملفات المرفقة وأخبرني بما فيها."), | |
| "ctx_head": ("Content of the files the user attached (use it; say so if the answer is not in it):", | |
| "محتوى الملفات التي أرفقها المستخدم (استخدمه، وصرّح إن لم تكن الإجابة فيه):"), | |
| "ctx_trunc": ("(Large files: only the parts most relevant to the question are shown.)", | |
| "(الملفات كبيرة: عُرضت فقط الأجزاء الأقرب لسؤال المستخدم.)"), | |
| "no_vision": ("An image named {name} ({w}x{h}) was attached, but no vision model is configured, so its content is unknown. " | |
| "Do not guess what it shows.", | |
| "أُرفقت صورة باسم {name} ({w}x{h}) لكن لا يوجد نموذج رؤية مُفعّل، فمحتواها مجهول. لا تخمّن ما فيها."), | |
| "ocr": ("Text found in the image by OCR:", "نص مستخرج من الصورة بالـ OCR:"), | |
| "error": ("⚠️ Something went wrong: {err}", "⚠️ حدث خطأ: {err}"), | |
| "ok": ("✅ All visible tests pass", "✅ نجحت كل الاختبارات الظاهرة"), | |
| "partial": ("⚠️ Best version found: {p}/{t} visible tests", "⚠️ أفضل نسخة وُجدت: {p}/{t} من الاختبارات الظاهرة"), | |
| "reward": ("Reward", "المكافأة"), | |
| "visible": ("Visible tests", "الاختبارات الظاهرة"), | |
| "hidden": ("Hidden tests", "الاختبارات المخفية"), | |
| "iters": ("Iterations", "التكرارات"), | |
| "restarts": ("Restarts", "إعادات البدء"), | |
| "final_code": ("Final code", "الكود النهائي"), | |
| "file_note": ("Files you can download: solution.py", "ملف للتنزيل: solution.py"), | |
| "attached": ("[attached: {names}]", "[مرفق: {names}]"), | |
| "skipped": ("{n} test(s) using pytest fixtures/classes were skipped.", "تم تخطي {n} اختبار يعتمد على fixtures/classes الخاصة بـ pytest."), | |
| } | |
| def detect_lang(text: str) -> str: | |
| return "ar" if re.search(r"[\u0600-\u06FF]", text or "") else "en" | |
| def tr(key: str, lang: str, **kwargs) -> str: | |
| en, ar = STR[key] | |
| return (ar if lang == "ar" else en).format(**kwargs) | |
| SYS_CHAT = ( | |
| "You are EIM Assistant, a careful, capable assistant for programming and for working with documents. " | |
| "Always answer in the language of the user's latest message. " | |
| "When file excerpts are provided, ground your answer in them and mention the file name; if the answer is not in them, say so. " | |
| "Never invent what an image shows: rely only on the image description you are given. " | |
| "For code, give complete, runnable code in fenced blocks and explain briefly. " | |
| "Write general solutions that are correct for every valid input, not only for the examples shown: never embed " | |
| "expected values taken from tests, never special-case test inputs, and never build lookup tables of answers. " | |
| "Think about edge cases first (empty input, negatives, duplicates, boundaries). " | |
| "Keep the function names and signatures that the user's code or tests use, and when fixing code change only what " | |
| "is needed. Prefer the standard library and make no network calls at import time. " | |
| "Treat explicit user corrections and rejected approaches as constraints: do not repeat an approach the user said " | |
| "was wrong; use the correction in the next answer, and briefly state what changed. Before claiming code works, " | |
| "separate tests actually run from reasoning or unverified assumptions. " | |
| "If the user wants code you can be sure about, suggest they add `assert` tests: the app then runs the code against " | |
| "them in a sandbox and tries to repair it." | |
| ) | |
| # ============================================================================ | |
| # 2. Model backends | |
| # ============================================================================ | |
| class ChatBackend(Protocol): | |
| def stream(self, messages: list[dict], max_new_tokens: int, temperature: float) -> Iterator[str]: ... | |
| class VisionBackend(Protocol): | |
| def describe(self, image_path: str, question: str) -> str: ... | |
| class LockedCodeLM: | |
| """The EIM loop's model: lazy-loaded, and serialised with chat generation.""" | |
| def generate(self, prompts: Sequence[str], temperature: float, max_new_tokens: int) -> list[str]: | |
| with _GEN_LOCK: | |
| return get_lm().generate(prompts, temperature, max_new_tokens) | |
| def _remote_chat_prompt(messages: list[dict], max_chars: int, protect: int = 0) -> str: | |
| """Serialise role-tagged chat for APIs that currently accept a single user prompt.""" | |
| cleaned = [] | |
| for message in messages: | |
| role = str(message.get("role", "user")).upper()[:20] | |
| content = str(message.get("content", "")) | |
| cleaned.append((role, content)) | |
| if not cleaned: | |
| return "You are a helpful assistant." | |
| # Keep the system instruction and newest question; discard oldest history first. | |
| system = cleaned[0] if cleaned[0][0] == "SYSTEM" else ("SYSTEM", "You are a helpful assistant.") | |
| remaining = max(500, max_chars - len(system[1]) - 32) | |
| kept = [] | |
| for role, content in reversed(cleaned[1:] if cleaned[0][0] == "SYSTEM" else cleaned): | |
| if remaining <= 0: | |
| break | |
| if len(content) > remaining: | |
| if role == "USER" and protect > 0: | |
| tail = min(protect, len(content), remaining // 2) | |
| head = max(0, remaining - tail - 20) | |
| content = content[:head] + "\n[… earlier context trimmed …]\n" + content[-tail:] if tail else content[:remaining] | |
| else: | |
| content = content[-remaining:] | |
| kept.append((role, content)) | |
| remaining -= len(content) + len(role) + 8 | |
| kept.reverse() | |
| rows = [f"SYSTEM:\n{system[1]}"] | |
| rows.extend(f"{role}:\n{content}" for role, content in kept) | |
| rows.append("ASSISTANT:") | |
| return "\n\n".join(rows)[-max_chars - 64:] | |
| class HFChat: | |
| """Streaming chat on a Hugging Face causal LM. Reuses the EIM model when no separate chat model is configured.""" | |
| def __init__(self, name: str, max_context_tokens: int = 6000): | |
| self.name, self.max_context_tokens = name, max_context_tokens | |
| self._model = self._tok = None | |
| self._remote_client = None | |
| def _load(self) -> None: | |
| if self._model is not None or self._tok is not None or self._remote_client is not None: | |
| return | |
| if self.name == CFG.model_name: | |
| shared = get_lm() | |
| else: | |
| from app import CodeLM | |
| shared = CodeLM(self.name) | |
| if getattr(shared, "is_remote", False): | |
| self._remote_client = shared | |
| return | |
| self._model, self._tok = shared.model, shared.tokenizer | |
| def _length(self, messages: list[dict]) -> int: | |
| text = self._tok.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) | |
| return len(self._tok(text, add_special_tokens=False).input_ids) | |
| def _fit(self, messages: list[dict], protect: int = 0, budget: int | None = None) -> list[dict]: | |
| """Trim the prompt to `budget` tokens (default: max_context_tokens). | |
| protect = number of characters at the END of the newest message that must never be cut: that is the user's own | |
| question (file excerpts come before it). Excerpts are shrunk first; only if the question alone is still too long | |
| is it shortened, and then from the middle, so its beginning and its end both survive.""" | |
| limit = self.max_context_tokens if budget is None else budget | |
| msgs = list(messages) | |
| if len(msgs) > 2 and self._length(msgs) > limit: | |
| # Drop the oldest turns (keep system + newest). Same result as dropping one turn at a time, but found by | |
| # binary search: ~log2(n) tokenisations of the whole prompt instead of up to n - on a weak CPU that is | |
| # the difference between an instant and a multi-second wait before the first token. | |
| lo, hi = 1, len(msgs) - 2 # d = number of oldest turns dropped after the system turn | |
| while lo < hi: | |
| mid = (lo + hi) // 2 | |
| if self._length(msgs[:1] + msgs[1 + mid:]) <= limit: | |
| hi = mid | |
| else: | |
| lo = mid + 1 | |
| msgs = msgs[:1] + msgs[1 + lo:] | |
| for _ in range(8): # still too long: shrink the newest message | |
| if self._length(msgs) <= limit: | |
| break | |
| content = msgs[-1]["content"] | |
| keep = min(max(0, protect), len(content)) | |
| body, tail = content[: len(content) - keep], content[len(content) - keep:] | |
| if protect and body: | |
| content = body[: len(body) // 2] + tail # cut the excerpts, keep the question whole | |
| elif protect: | |
| quarter = max(1, len(tail) // 4) # nothing left to cut but the question itself | |
| content = tail[:quarter] + " ... " + tail[-quarter:] | |
| else: | |
| content = content[: len(content) // 2] | |
| msgs[-1] = {**msgs[-1], "content": content} | |
| return msgs | |
| def stream(self, messages: list[dict], max_new_tokens: int = 900, temperature: float = 0.4, | |
| protect: int = 0) -> Iterator[str]: | |
| self._load() | |
| if self._remote_client is not None: | |
| limit = max(2000, min(24000, int(os.environ.get("EIM_REMOTE_CONTEXT_CHARS", "14000")))) | |
| prompt = _remote_chat_prompt(messages, limit, protect) | |
| answer = self._remote_client.generate([prompt], temperature, max_new_tokens)[0] | |
| # Inference Providers return a completed response; chunk it for the same Gradio interface contract. | |
| for start in range(0, len(answer), 80): | |
| yield answer[start:start + 80] | |
| if not answer: | |
| yield "" | |
| return | |
| import torch | |
| from transformers import StoppingCriteria, StoppingCriteriaList, TextIteratorStreamer | |
| tok, model = self._tok, self._model | |
| window = getattr(getattr(model, "config", None), "max_position_embeddings", None) | |
| budget = self.max_context_tokens | |
| if isinstance(window, int) and window > 0: # prompt + answer must fit the model's own window | |
| budget = max(256, min(budget, window - int(max_new_tokens) - 64)) | |
| text = tok.apply_chat_template(self._fit(messages, protect, budget), tokenize=False, add_generation_prompt=True) | |
| inputs = tok(text, return_tensors="pt", add_special_tokens=False).to(model.device) | |
| class _Stop(StoppingCriteria): | |
| flag = False | |
| def __call__(self, input_ids, scores, **kwargs): | |
| return torch.full((input_ids.shape[0],), self.flag, dtype=torch.bool, device=input_ids.device) | |
| stop = _Stop() | |
| streamer = TextIteratorStreamer(tok, skip_prompt=True, skip_special_tokens=True, timeout=600) | |
| options = dict(**inputs, streamer=streamer, max_new_tokens=int(max_new_tokens), | |
| pad_token_id=tok.pad_token_id if tok.pad_token_id is not None else tok.eos_token_id, | |
| stopping_criteria=StoppingCriteriaList([stop])) | |
| if temperature > 0: | |
| options.update(do_sample=True, temperature=float(temperature), top_p=0.95) | |
| else: | |
| options.update(do_sample=False) | |
| errors: list[BaseException] = [] | |
| def work() -> None: | |
| try: | |
| with torch.inference_mode(): | |
| model.generate(**options) | |
| except BaseException as exc: # surface it instead of hanging the stream | |
| errors.append(exc) | |
| streamer.end() | |
| with _GEN_LOCK: | |
| thread = threading.Thread(target=work, daemon=True) | |
| thread.start() | |
| try: | |
| for piece in streamer: | |
| yield piece | |
| finally: # also runs when the user presses Stop | |
| # The flag stays True: this _Stop class belongs to this call only, and clearing it after a join timeout | |
| # would let a still-running generation continue in the background and overlap the next one. | |
| _Stop.flag = True | |
| thread.join(timeout=30) | |
| if errors: | |
| raise errors[0] | |
| class HFVision: | |
| """Optional vision-language model (Qwen2-VL style). NOT loaded unless EIM_VISION_MODEL is set.""" | |
| def __init__(self, name: str): | |
| self.name = name | |
| self._model = self._processor = None | |
| def _load(self) -> None: | |
| if self._model is not None: | |
| return | |
| import torch | |
| from transformers import AutoProcessor | |
| try: | |
| from transformers import AutoModelForImageTextToText as AutoVision | |
| except ImportError: | |
| from transformers import AutoModelForVision2Seq as AutoVision | |
| self._processor = AutoProcessor.from_pretrained(self.name) | |
| use_cuda = torch.cuda.is_available() or bool(os.environ.get("SPACES_ZERO_GPU")) | |
| self._model = AutoVision.from_pretrained(self.name, torch_dtype=torch.bfloat16 if use_cuda else torch.float32) | |
| if use_cuda: | |
| self._model.to("cuda") | |
| self._model.eval() | |
| def describe(self, image_path: str, question: str = "") -> str: | |
| import torch | |
| from PIL import Image | |
| self._load() | |
| image = Image.open(image_path).convert("RGB") | |
| image.thumbnail((1024, 1024)) | |
| prompt = (f"The user asks: {question}\nAnswer using the image, and also transcribe any visible text." | |
| if question.strip() else | |
| "Describe this image in detail. Transcribe any visible text exactly. Mention charts, tables and code if present.") | |
| messages = [{"role": "user", "content": [{"type": "image"}, {"type": "text", "text": prompt}]}] | |
| text = self._processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) | |
| inputs = self._processor(text=[text], images=[image], return_tensors="pt").to(self._model.device) | |
| with _GEN_LOCK, torch.inference_mode(): | |
| output = self._model.generate(**inputs, max_new_tokens=500, do_sample=False) | |
| return self._processor.batch_decode(output[:, inputs["input_ids"].shape[1]:], skip_special_tokens=True)[0].strip() | |
| def ocr_text(path: str) -> str: | |
| try: | |
| import pytesseract | |
| from PIL import Image | |
| with Image.open(path) as image: | |
| return pytesseract.image_to_string(image, lang="ara+eng").strip() | |
| except Exception: | |
| return "" | |
| # ============================================================================ | |
| # 3. Session, options, intent | |
| # ============================================================================ | |
| class Session: | |
| history: list[dict] = field(default_factory=list) # what the chat model sees (plain text turns) | |
| files: dict[str, Attachment] = field(default_factory=dict) | |
| working_code: str = "" # latest uploaded .py or EIM result | |
| workdir: str = "" | |
| result_path: str | None = None | |
| runs: int = 0 | |
| class Options: | |
| mode: str = "auto" # auto | chat | eim | |
| iterations: int = 4 | |
| candidates: int = 3 | |
| holdout: float = 0.25 | |
| temperature: float = 0.4 | |
| max_new_tokens: int = 900 | |
| context_chars: int = 6000 | |
| class Intent: | |
| kind: str # chat | eim | need_tests | |
| task: str = "" | |
| tests: list[str] = field(default_factory=list) | |
| initial_code: str = "" | |
| skipped: int = 0 | |
| _FENCE = re.compile(r"```[ \t]*([\w+-]*)[ \t]*\n(.*?)```", re.S) | |
| def split_message(text: str) -> tuple[str, list[str]]: | |
| """(prose without code blocks, python-ish fenced blocks)""" | |
| blocks = [m.group(2) for m in _FENCE.finditer(text) if m.group(1).lower() in ("", "python", "py")] | |
| return _FENCE.sub("", text).strip(), blocks | |
| def _looks_like_code(block: str) -> bool: | |
| return bool(re.search(r"^\s*(def |class |import |from \S+ import )", block, re.M)) | |
| def _defined_names(code: str) -> set[str]: | |
| try: | |
| return {n.name for n in ast.parse(code).body if isinstance(n, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef))} | |
| except (SyntaxError, ValueError): | |
| return set() | |
| def _used_names(tests: Sequence[str]) -> set[str]: | |
| used: set[str] = set() | |
| for test in tests: | |
| try: | |
| used |= {n.id for n in ast.walk(ast.parse(test)) if isinstance(n, ast.Name)} | |
| except (SyntaxError, ValueError): | |
| pass | |
| return used | |
| def detect_intent(text: str, new_atts: Sequence[Attachment], session: Session, mode: str) -> Intent: | |
| prose, blocks = split_message(text) | |
| tests: list[str] = [] | |
| skipped = 0 | |
| initial = "" | |
| for block in blocks: | |
| found, skip = tests_from_source(block) | |
| skipped += skip | |
| if found: | |
| tests += found | |
| elif _looks_like_code(block) and not initial: | |
| initial = block | |
| prose_lines = prose.splitlines() | |
| tests += [line.strip() for line in prose_lines if line.strip().startswith("assert ")] | |
| task = "\n".join(line for line in prose_lines if not line.strip().startswith("assert ")).strip() | |
| spec_parts = [] | |
| for att in new_atts: | |
| if att.kind == "code" and att.name.lower().endswith(".py"): | |
| found, skip = tests_from_source(att.text) | |
| skipped += skip | |
| if is_test_file(att.name, att.text, found): | |
| tests += found | |
| elif not initial: | |
| initial = att.text | |
| elif att.kind in ("text", "document") and att.text: | |
| spec_parts.append(f"Specification from {att.name}:\n{att.text[:3000]}") | |
| if mode == "chat": | |
| return Intent("chat") | |
| if not tests: | |
| return Intent("need_tests") if mode == "eim" else Intent("chat") | |
| if not initial and session.working_code and _defined_names(session.working_code) & _used_names(tests): | |
| initial = session.working_code # the tests talk about code we already have | |
| if not task: | |
| task = "Fix the code so that all tests pass." if initial else "Write code that passes all the tests." | |
| if spec_parts: | |
| task += "\n\n" + "\n\n".join(spec_parts) | |
| return Intent("eim", task, tests, initial, skipped) | |
| # ============================================================================ | |
| # 4. Assistant: the whole behaviour, independent of Gradio | |
| # ============================================================================ | |
| def render_progress(steps: Sequence[Step], lang: str, running: bool = True) -> str: | |
| rows = ["| # | Status | Strategy | Tests | Reward |", "|---|---|---|---|---|"] | |
| for s in steps: | |
| rows.append(f"| {s.iteration} | {s.label} | {s.strategy} | {s.result.passed}/{s.result.total} | {s.score.reward:.3f} |") | |
| accepted = [s for s in steps if s.accepted] | |
| code = f"\n\n```python\n{accepted[-1].code}\n```" if accepted else "" | |
| return ("\n".join(rows) + code + (f"\n\n{tr('thinking', lang)}" if running else "")) | |
| def render_final(report: Report, steps: Sequence[Step], lang: str, skipped: int) -> str: | |
| best = report.best | |
| ok = best.result.pass_rate == 1.0 | |
| head = tr("ok", lang) if ok else tr("partial", lang, p=best.result.passed, t=best.result.total) | |
| facts = [ | |
| f"**{tr('reward', lang)}:** {report.first.score.reward:.3f} → {best.score.reward:.3f}", | |
| f"**{tr('visible', lang)}:** {best.result.passed}/{best.result.total}", | |
| ] | |
| if report.hidden is not None: | |
| facts.append(f"**{tr('hidden', lang)}:** {report.hidden.passed}/{report.hidden.total}") | |
| facts.append(f"**{tr('iters', lang)}:** {report.iterations} · **{tr('restarts', lang)}:** {report.restarts}") | |
| notes = [f"- ⚠️ {w}" for w in report.warnings] | |
| if skipped: | |
| notes.append(f"- ℹ️ {tr('skipped', lang, n=skipped)}") | |
| parts = [f"### {head}", " · ".join(facts)] | |
| if notes: | |
| parts.append("\n".join(notes)) | |
| parts += [f"#### {tr('final_code', lang)}", f"```python\n{best.code}\n```", f"_{tr('file_note', lang)}_", | |
| "<details><summary>log</summary>\n\n" + render_progress(steps, lang, running=False) + "\n\n</details>"] | |
| return "\n\n".join(parts) | |
| class Assistant: | |
| def __init__(self, chat: ChatBackend, vision: VisionBackend | None = None, code_lm=None, | |
| verifier: EUTV | None = None, memory: ExperienceMemory | None = None, cfg: Config = CFG): | |
| self.chat, self.vision, self.cfg = chat, vision, cfg | |
| self.code_lm = code_lm or LockedCodeLM() | |
| self.verifier = verifier or EUTV(cfg) | |
| self.memory = memory if memory is not None else SmartMemory(path=cfg.memory_path) | |
| self.corrections = CorrectionMemory() | |
| # -------------------------------------------------------------- public | |
| def handle(self, session: Session, text: str, paths: Sequence[str], opts: Options) -> Iterator[str]: | |
| """Yields the FULL assistant message so far, each time it grows.""" | |
| text = (text or "").strip() | |
| lang = detect_lang(text) if text else "en" | |
| new_atts = self._ingest(session, paths) | |
| if not text and not new_atts: | |
| return | |
| try: | |
| for att in new_atts: | |
| if att.kind == "image": | |
| yield tr("reading_image", lang) | |
| att.text = self._describe(att, text, lang) | |
| if not text: | |
| text = tr("default_q", lang) | |
| if not paths and not text: | |
| return | |
| intent = detect_intent(text, new_atts, session, opts.mode) | |
| if intent.kind == "need_tests": | |
| yield tr("eim_needs_tests", lang) | |
| elif intent.kind == "eim": | |
| yield from self._run_eim(session, text, intent, opts, lang, new_atts) | |
| else: | |
| yield from self._run_chat(session, text, opts, lang, new_atts) | |
| except Exception as exc: | |
| yield tr("error", lang, err=f"{type(exc).__name__}: {exc}") | |
| # -------------------------------------------------------------- internals | |
| def _ingest(self, session: Session, paths: Sequence[str]) -> list[Attachment]: | |
| out = [] | |
| for path in paths: | |
| att = read_attachment(path) | |
| base, ext, n = att.name, os.path.splitext(att.name)[1], 1 | |
| while att.name in session.files: # same name uploaded twice: keep both | |
| n += 1 | |
| att.name = f"{os.path.splitext(base)[0]} ({n}){ext}" | |
| session.files[att.name] = att | |
| if att.kind == "code" and att.name.lower().endswith(".py"): | |
| found, _ = tests_from_source(att.text) | |
| if not is_test_file(att.name, att.text, found): | |
| session.working_code = att.text | |
| out.append(att) | |
| return out | |
| def _describe(self, att: Attachment, question: str, lang: str) -> str: | |
| width, height = att.image_size or (0, 0) | |
| if self.vision is not None: | |
| try: | |
| return f"[image {att.name}] " + self.vision.describe(att.path, question) | |
| except Exception as exc: | |
| return f"[image {att.name}: the vision model failed ({type(exc).__name__}: {exc})]" | |
| text = tr("no_vision", lang, name=att.name, w=width, h=height) | |
| ocr = ocr_text(att.path) | |
| return text + (f"\n{tr('ocr', lang)}\n{ocr}" if ocr else "") | |
| def _run_chat(self, session: Session, text: str, opts: Options, lang: str, new_atts) -> Iterator[str]: | |
| # Persist explicit rejections/corrections before building this turn's prompt. | |
| self.corrections.add(text) | |
| context, truncated = select_context(list(session.files.values()), text, opts.context_chars) | |
| parts = [] | |
| if context: | |
| parts.append(f"{tr('ctx_head', lang)}\n{context}") | |
| if truncated: | |
| parts.append(tr("ctx_trunc", lang)) | |
| parts.append(text) | |
| correction_context = self.corrections.prompt(text) | |
| system_text = SYS_CHAT + ("\n\n" + correction_context if correction_context else "") | |
| messages = [{"role": "system", "content": system_text}] + session.history[-40:] + \ | |
| [{"role": "user", "content": "\n\n".join(parts)}] | |
| yield tr("thinking", lang) | |
| answer = "" | |
| last = 0.0 | |
| extra = {"protect": len(text)} if _accepts(self.chat.stream, "protect") else {} # custom chat models may not | |
| for piece in self.chat.stream(messages, opts.max_new_tokens, opts.temperature, **extra): | |
| answer += piece | |
| now = time.monotonic() | |
| if now - last > 0.05: # throttle UI updates | |
| last = now | |
| yield answer | |
| yield answer | |
| names = ", ".join(a.name for a in new_atts) | |
| shown = text + (("\n" + tr("attached", lang, names=names)) if names else "") | |
| session.history += [{"role": "user", "content": shown}, {"role": "assistant", "content": answer}] | |
| def _run_eim(self, session: Session, text: str, intent: Intent, opts: Options, lang: str, new_atts) -> Iterator[str]: | |
| yield tr("eim_start", lang) | |
| self.corrections.add(text) | |
| correction_context = self.corrections.prompt(text) | |
| task = intent.task + ("\n\n" + correction_context if correction_context else "") | |
| engine = EIMPlus(self.code_lm, self.verifier, self.cfg, self.memory) | |
| steps: list[Step] = [] | |
| report: Report | None = None | |
| for event in engine.run(task, intent.tests, intent.initial_code or None, | |
| int(opts.iterations), int(opts.candidates), float(opts.holdout)): | |
| if isinstance(event, Report): | |
| report = event | |
| break | |
| steps.append(event) | |
| yield render_progress(steps, lang, running=True) | |
| assert report is not None | |
| final = render_final(report, steps, lang, intent.skipped) | |
| session.runs += 1 | |
| if not session.workdir: | |
| session.workdir = tempfile.mkdtemp(prefix="eimchat_") | |
| path = os.path.join(session.workdir, f"solution_{session.runs}.py") | |
| with open(path, "w", encoding="utf-8") as handle: | |
| handle.write(report.best.code + "\n") | |
| session.result_path = path | |
| session.working_code = report.best.code | |
| yield final | |
| names = ", ".join(a.name for a in new_atts) | |
| shown = text + (("\n" + tr("attached", lang, names=names)) if names else "") | |
| summary = f"{tr('final_code', lang)} ({report.best.result.passed}/{report.best.result.total}):\n```python\n{report.best.code}\n```" | |
| session.history += [{"role": "user", "content": shown}, {"role": "assistant", "content": summary}] | |
| # ============================================================================ | |
| # 5. Gradio UI | |
| # ============================================================================ | |
| CSS = """ | |
| #chat .message, #chat .prose, #chat .message-content { unicode-bidi: plaintext; text-align: start; } | |
| #chat pre, #chat code { direction: ltr; unicode-bidi: isolate; text-align: left; } | |
| footer { display: none !important; } | |
| """ | |
| ACCEPTED_FILES = [".py", ".txt", ".md", ".json", ".csv", ".tsv", ".pdf", ".docx", ".xlsx", ".ipynb", ".zip", | |
| ".js", ".ts", ".java", ".c", ".cpp", ".go", ".rs", ".html", ".css", ".sql", ".yaml", ".yml", | |
| ".png", ".jpg", ".jpeg", ".webp", ".gif", ".bmp"] | |
| _ASSISTANT: Assistant | None = None | |
| _ASSISTANT_LOCK = threading.Lock() | |
| def get_assistant() -> Assistant: | |
| global _ASSISTANT | |
| with _ASSISTANT_LOCK: | |
| if _ASSISTANT is None: | |
| # prompt prefill dominates latency on a CPU: weak machines get a smaller default window (env still wins) | |
| default_context = "4000" if hardware_profile().name == "lean" else "6000" | |
| chat = HFChat(os.environ.get("EIM_CHAT_MODEL", CFG.model_name), | |
| int(os.environ.get("EIM_CONTEXT_TOKENS", default_context))) | |
| vision_name = os.environ.get("EIM_VISION_MODEL", "") | |
| _ASSISTANT = Assistant(chat, HFVision(vision_name) if vision_name else None) | |
| return _ASSISTANT | |
| def _paths(files) -> list[str]: | |
| out = [] | |
| for item in files or []: | |
| path = item.get("path") if isinstance(item, dict) else getattr(item, "path", item) | |
| if isinstance(path, str) and path: | |
| out.append(path) | |
| return out | |
| def _files_markdown(session: Session | None) -> str: | |
| if not session or not session.files: | |
| return "_no files yet / لا توجد ملفات بعد_" | |
| lines = [] | |
| for att in session.files.values(): | |
| extra = f" — {att.note}" if att.note else "" | |
| size = f"{att.size / 1024:.0f} KB" | |
| lines.append(f"- 📎 **{att.name}** · {att.kind} · {size}{extra}") | |
| return "\n".join(lines) | |
| def _chatbot(gr): | |
| kwargs = dict(elem_id="chat", height=560, show_copy_button=True, render_markdown=True, | |
| latex_delimiters=[{"left": "$$", "right": "$$", "display": True}]) | |
| for attempt in (dict(type="messages", **kwargs), kwargs, dict(elem_id="chat", height=560)): | |
| try: | |
| return gr.Chatbot(**attempt) | |
| except TypeError: | |
| continue | |
| return gr.Chatbot() | |
| def respond(message, history, session, mode, iterations, candidates, holdout, temperature, max_tokens): | |
| import gradio as gr | |
| session = session or Session() | |
| history = list(history or []) | |
| text = ((message or {}).get("text") or "").strip() | |
| paths = _paths((message or {}).get("files")) | |
| if not text and not paths: | |
| yield history, session, gr.MultimodalTextbox(value=None, interactive=True), _files_markdown(session), session.result_path | |
| return | |
| for path in paths: # show the files in the chat bubble | |
| history.append({"role": "user", "content": {"path": path}}) | |
| if text: | |
| history.append({"role": "user", "content": text}) | |
| history.append({"role": "assistant", "content": ""}) | |
| box = gr.MultimodalTextbox(value=None, interactive=False) | |
| opts = Options(mode=str(mode).lower(), iterations=int(iterations), candidates=int(candidates), | |
| holdout=float(holdout), temperature=float(temperature), max_new_tokens=int(max_tokens)) | |
| yield history, session, box, _files_markdown(session), session.result_path | |
| for snapshot in get_assistant().handle(session, text, paths, opts): | |
| history[-1] = {"role": "assistant", "content": snapshot} | |
| yield history, session, box, _files_markdown(session), session.result_path | |
| yield history, session, box, _files_markdown(session), session.result_path | |
| def build_ui(): | |
| import gradio as gr | |
| # Gradio 6 moved theme/css from Blocks(...) to launch(...). Keep both API generations supported. | |
| try: | |
| major = int(str(getattr(gr, "__version__", "5")).split(".", 1)[0]) | |
| except (TypeError, ValueError): | |
| major = 5 | |
| if major >= 6: | |
| blocks = gr.Blocks(title="EIM Chat") | |
| launch_kwargs = {"theme": gr.themes.Soft(), "css": CSS} | |
| else: | |
| blocks = gr.Blocks(title="EIM Chat", theme=gr.themes.Soft(), css=CSS) | |
| launch_kwargs = {} | |
| setattr(blocks, "_eim_launch_kwargs", launch_kwargs) | |
| with blocks as demo: | |
| gr.Markdown("## EIM Chat\nمحادثة + ملفات + صور + كود يُتحقَّق منه ويُصلَح تلقائياً · " | |
| "Chat + files + images + code that is verified and repaired automatically") | |
| session = gr.State(None) | |
| with gr.Row(): | |
| with gr.Column(scale=4): | |
| chat = _chatbot(gr) | |
| box = gr.MultimodalTextbox( | |
| file_count="multiple", file_types=ACCEPTED_FILES, show_label=False, interactive=True, | |
| placeholder="اكتب رسالتك أو أرفق ملفات وصوراً… · Type, or attach files and images… " | |
| "(add `assert f(2) == 4` lines to verify code)") | |
| with gr.Column(scale=1, min_width=260): | |
| mode = gr.Radio(["Auto", "Chat", "EIM"], value="Auto", label="Mode") | |
| with gr.Accordion("EIM", open=False): | |
| iterations = gr.Slider(1, 8, value=4, step=1, label="Max iterations") | |
| candidates = gr.Slider(1, 6, value=3, step=1, label="Candidates per iteration") | |
| holdout = gr.Slider(0, 0.5, value=0.25, step=0.05, label="Hidden tests share (overfitting check)") | |
| with gr.Accordion("Chat", open=False): | |
| temperature = gr.Slider(0, 1.2, value=0.4, step=0.05, label="Temperature") | |
| max_tokens = gr.Slider(128, 2048, value=900, step=64, label="Max new tokens") | |
| files_view = gr.Markdown(_files_markdown(None)) | |
| result = gr.File(label="solution.py", interactive=False) | |
| clear = gr.Button("New chat / محادثة جديدة") | |
| inputs = [box, chat, session, mode, iterations, candidates, holdout, temperature, max_tokens] | |
| outputs = [chat, session, box, files_view, result] | |
| box.submit(respond, inputs, outputs).then(lambda: gr.MultimodalTextbox(interactive=True), None, [box]) | |
| clear.click(lambda: ([], None, gr.MultimodalTextbox(value=None, interactive=True), _files_markdown(None), None), | |
| None, outputs) | |
| return demo | |
| # ============================================================================ | |
| # 6. Self-test (no model, no GPU, no Gradio): python chat_app.py --selftest | |
| # ============================================================================ | |
| class _FakeChat: | |
| def __init__(self, answer: str = "Fake answer."): | |
| self.answer, self.seen = answer, [] | |
| def stream(self, messages, max_new_tokens=900, temperature=0.4): | |
| self.seen.append(messages) | |
| for i in range(0, len(self.answer), 7): | |
| yield self.answer[i:i + 7] | |
| class _FakeVision: | |
| def describe(self, image_path, question=""): | |
| return "A red square on a white background with the text HELLO." | |
| def run_selftest() -> int: | |
| import shutil | |
| from app import _ScriptedLM | |
| checks = 0 | |
| def check(name: str, condition: bool) -> None: | |
| nonlocal checks | |
| checks += 1 | |
| print(("PASS " if condition else "FAIL ") + name) | |
| if not condition: | |
| raise SystemExit(1) | |
| tmp = tempfile.mkdtemp(prefix="eimchat_test_") | |
| fast = Config(per_test_seconds=1, wall_seconds=8.0, cpu_seconds=6, memory_path="") | |
| def write(name: str, text: str) -> str: | |
| path = os.path.join(tmp, name) | |
| with open(path, "w", encoding="utf-8") as handle: | |
| handle.write(text) | |
| return path | |
| try: | |
| # ---- intent detection | |
| s = Session() | |
| msg = "Fix this:\n```python\ndef add(a, b):\n return a - b\n```\nassert add(1, 2) == 3\nassert add(2, 2) == 4" | |
| intent = detect_intent(msg, [], s, "auto") | |
| check("fenced code + assert lines -> EIM from the user's code", | |
| intent.kind == "eim" and len(intent.tests) == 2 and "return a - b" in intent.initial_code) | |
| check("plain question -> chat", detect_intent("what is a decorator?", [], s, "auto").kind == "chat") | |
| check("mode=chat wins over tests", detect_intent(msg, [], s, "chat").kind == "chat") | |
| check("mode=eim without tests asks for tests", detect_intent("hello", [], s, "eim").kind == "need_tests") | |
| sol = read_attachment(write("solution.py", "def add(a, b):\n return a - b\n")) | |
| tst = read_attachment(write("test_solution.py", "from solution import add\n\ndef test_a():\n assert add(1, 2) == 3\n")) | |
| intent = detect_intent("make it pass", [sol, tst], Session(), "auto") | |
| check("attached solution.py + test_solution.py -> EIM", | |
| intent.kind == "eim" and intent.initial_code.startswith("def add") and len(intent.tests) == 1) | |
| s2 = Session(working_code="def add(a, b):\n return a - b") | |
| check("tests about already-known code reuse it", | |
| detect_intent("assert add(1, 2) == 3", [], s2, "auto").initial_code.startswith("def add")) | |
| check("unrelated tests do NOT reuse old code", | |
| detect_intent("assert mul(1, 2) == 2", [], s2, "auto").initial_code == "") | |
| check("language detection", detect_lang("مرحبا") == "ar" and detect_lang("hello") == "en") | |
| # Remote API chat must not try to access .model/.tokenizer or load Transformers weights. | |
| class _FakeRemote: | |
| is_remote = True | |
| def generate(self, prompts, temperature, max_new_tokens): | |
| self.prompts = list(prompts) | |
| return ["remote response"] | |
| old_get_lm = globals()["get_lm"] | |
| fake_remote = _FakeRemote() | |
| try: | |
| globals()["get_lm"] = lambda: fake_remote | |
| remote = HFChat(CFG.model_name) | |
| remote_text = "".join(remote.stream([{"role": "system", "content": "Keep safe"}, | |
| {"role": "user", "content": "Hello"}], 64, 0.0)) | |
| check("remote chat backend works without local model weights", | |
| remote_text == "remote response" and "SYSTEM:" in fake_remote.prompts[0] | |
| and "Hello" in fake_remote.prompts[0]) | |
| finally: | |
| globals()["get_lm"] = old_get_lm | |
| # ---- chat with files (retrieval across turns) | |
| filler = "\n".join(f"row {i}: nothing relevant here at all" for i in range(700)) | |
| needle = "The warranty period for the X200 device is 26 months." | |
| path = write("manual.txt", filler + "\n" + needle + "\n" + filler) | |
| chat = _FakeChat("The warranty is 26 months.") | |
| bot = Assistant(chat, None, _ScriptedLM(["x = 1"]), EUTV(fast), SmartMemory(path=""), fast) | |
| s = Session() | |
| out = list(bot.handle(s, "What is the warranty period for the X200?", [path], Options())) | |
| check("chat streams and ends with the full answer", out[-1] == "The warranty is 26 months.") | |
| check("the relevant chunk of a big file reached the model", needle in chat.seen[0][-1]["content"]) | |
| check("history keeps the plain turn, not the file dump", | |
| "row 5:" not in s.history[0]["content"] and "[attached: manual.txt]" in s.history[0]["content"]) | |
| list(bot.handle(s, "And what about the warranty again?", [], Options())) | |
| check("follow-up retrieves from earlier files", needle in chat.seen[1][-1]["content"] and len(chat.seen[1]) == 4) | |
| out = list(bot.handle(Session(), "ما هي فترة الضمان؟", [write("m2.txt", "فترة الضمان ستة وعشرون شهرا")], Options())) | |
| check("Arabic message gets the Arabic status line", out[0] == STR["thinking"][1]) | |
| # ---- images | |
| from PIL import Image | |
| Image.new("RGB", (20, 10), "red").save(os.path.join(tmp, "pic.png")) | |
| chat = _FakeChat("It is a red square.") | |
| bot_v = Assistant(chat, _FakeVision(), _ScriptedLM(["x = 1"]), EUTV(fast), SmartMemory(path=""), fast) | |
| out = list(bot_v.handle(Session(), "What is in the picture?", [os.path.join(tmp, "pic.png")], Options())) | |
| check("vision description reaches the chat model", "red square on a white background" in chat.seen[0][-1]["content"]) | |
| chat = _FakeChat("I cannot see it.") | |
| bot_nv = Assistant(chat, None, _ScriptedLM(["x = 1"]), EUTV(fast), SmartMemory(path=""), fast) | |
| list(bot_nv.handle(Session(), "What is in the picture?", [os.path.join(tmp, "pic.png")], Options())) | |
| check("without a vision model the prompt says the image is unknown", | |
| "no vision model is configured" in chat.seen[0][-1]["content"]) | |
| out = list(bot_nv.handle(Session(), "", [os.path.join(tmp, "pic.png")], Options())) | |
| check("files without text still get an answer", out and out[-1] == "I cannot see it.") | |
| # ---- history trimming: binary search == dropping one turn at a time | |
| class _WordTok: | |
| def apply_chat_template(self, messages, tokenize=False, add_generation_prompt=True): | |
| return " ".join(f"<{m['role']}> {m['content']}" for m in messages) | |
| def __call__(self, text, add_special_tokens=False): | |
| return type("Enc", (), {"input_ids": text.split()})() | |
| def linear_fit(hf, messages): | |
| msgs = list(messages) | |
| while hf._length(msgs) > hf.max_context_tokens and len(msgs) > 2: | |
| del msgs[1] | |
| return msgs | |
| same = True | |
| for budget in (8, 40, 120, 400, 5000): | |
| hf = HFChat("x", budget) | |
| hf._tok = _WordTok() | |
| convo = [{"role": "system", "content": "sys prompt"}] + [ | |
| {"role": "user" if i % 2 == 0 else "assistant", "content": "word " * (3 + i * 2)} for i in range(14)] | |
| expected = linear_fit(hf, convo) | |
| for _ in range(6): # the original's shrink step, applied to both | |
| if hf._length(expected) <= budget: | |
| break | |
| expected[-1] = {**expected[-1], "content": expected[-1]["content"][: len(expected[-1]["content"]) // 2]} | |
| same &= hf._fit(convo) == expected | |
| check("history trimming gives the same result as the one-by-one version", same) | |
| hf = HFChat("x", 60) | |
| hf._tok = _WordTok() | |
| question = "QUESTION why does the loop never end please answer" | |
| long_ctx = " ".join(f"ctx{i}" for i in range(400)) | |
| convo = [{"role": "system", "content": "sys"}, {"role": "user", "content": long_ctx + "\n\n" + question}] | |
| cut = hf._fit(convo, protect=len(question)) | |
| check("trimming cuts file excerpts and keeps the user's question whole", | |
| cut[-1]["content"].endswith(question) and hf._length(cut) <= 60) | |
| old_way = hf._fit(convo) | |
| check("without protect the old behaviour is unchanged (end is cut)", not old_way[-1]["content"].endswith(question)) | |
| huge_q = " ".join(f"q{i}" for i in range(300)) | |
| cut = hf._fit([{"role": "system", "content": "sys"}, {"role": "user", "content": huge_q}], protect=len(huge_q)) | |
| check("an oversized question keeps its beginning AND its end", | |
| cut[-1]["content"].startswith("q0") and cut[-1]["content"].endswith("q299") and hf._length(cut) <= 60) | |
| check("a smaller explicit budget is honoured", hf._length(hf._fit(convo, protect=len(question), budget=30)) <= 30) | |
| # ---- unsupported / broken files never crash the turn | |
| out = list(bot.handle(Session(), "read this", [write("x.bin", "\x00\x00\x00")], Options())) | |
| check("a binary file does not crash the turn", out and not out[-1].startswith("⚠️ Something went wrong")) | |
| # ---- EIM flow | |
| fixed = "def add(a, b):\n return a + b" | |
| bot_e = Assistant(_FakeChat(), None, _ScriptedLM([fixed]), EUTV(fast), SmartMemory(path=""), fast) | |
| s = Session() | |
| out = list(bot_e.handle(s, msg, [], Options(iterations=3, candidates=2))) | |
| check("EIM flow streams progress then a final report", | |
| "| # | Status" in out[1] and out[-1].startswith("### ✅") and "return a + b" in out[-1]) | |
| check("solution.py is produced and holds the repaired code", | |
| s.result_path and "return a + b" in open(s.result_path, encoding="utf-8").read()) | |
| check("working code and history are updated", s.working_code.strip() == fixed and "Final code" in s.history[-1]["content"]) | |
| out = list(bot_e.handle(Session(), "اكتب دالة الجمع\nassert add(1, 2) == 3\nassert add(2, 2) == 4", [], Options())) | |
| check("EIM report speaks Arabic when asked in Arabic", "الاختبارات الظاهرة" in out[-1]) | |
| out = list(bot_e.handle(Session(), "hello", [], Options(mode="eim"))) | |
| check("EIM mode without tests explains what to add", "assert" in out[-1]) | |
| finally: | |
| shutil.rmtree(tmp, ignore_errors=True) | |
| print(f"\nAll {checks} checks passed.") | |
| return 0 | |
| if __name__ == "__main__" and "--selftest" in sys.argv: | |
| raise SystemExit(run_selftest()) | |
| if __name__ == "__main__": | |
| demo = build_ui() | |
| demo.queue().launch(allowed_paths=[tempfile.gettempdir()], **getattr(demo, "_eim_launch_kwargs", {})) | |