ThinkSplitter v3: narration gate + paragraph classifier + degenerate-output auto-recovery; num_predict 1100
This commit is contained in:
240
draco_core.py
240
draco_core.py
@@ -142,81 +142,217 @@ def build_prompt(question, k=6):
|
||||
# ---------------------------------------------------------------- ollama
|
||||
import requests as http
|
||||
|
||||
OPEN, CLOSE = "<think>", "</think>"
|
||||
|
||||
def strip_cot(text):
|
||||
text = re.sub(r"(?s)<think>.*?</think>", "", text)
|
||||
if "</think>" in text:
|
||||
text = text.split("</think>")[-1]
|
||||
if "<think>" in text: # unbalanced opener: keep what follows
|
||||
text = text.split("<think>")[-1]
|
||||
return text.strip()
|
||||
"""Non-streaming: return (reasoning, answer). Handles tagged AND untagged CoT."""
|
||||
if OPEN in text:
|
||||
pre, rest = text.split(OPEN, 1)
|
||||
if CLOSE in rest:
|
||||
think, answer = rest.split(CLOSE, 1)
|
||||
return (pre + think).strip(), answer.strip()
|
||||
return "", rest.strip() # unbalanced opener
|
||||
if CLOSE in text: # missing opener: preamble before close is reasoning
|
||||
think, answer = text.split(CLOSE, 1)
|
||||
return think.strip(), answer.strip()
|
||||
return "", text.strip() # no tags: answer as-is
|
||||
|
||||
def ask_ollama(prompt):
|
||||
r = http.post(f"{CFG['ollama_url']}/api/generate",
|
||||
json={"model": CFG["model"], "prompt": prompt, "stream": False,
|
||||
"options": {"temperature": 0.4, "num_predict": 700,
|
||||
"num_ctx": CFG.get("num_ctx", 8192)}},
|
||||
timeout=180)
|
||||
r.raise_for_status()
|
||||
return strip_cot(r.json().get("response", ""))
|
||||
|
||||
class CoTFilter:
|
||||
"""Streaming filter that drops <think>...</think> blocks token-by-token."""
|
||||
OPEN, CLOSE = "<think>", "</think>"
|
||||
class ThinkSplitter:
|
||||
"""Streaming: yield (channel, piece) where channel is 'think' or 'answer'.
|
||||
Handles: <think>..</think>, missing opener (</think> closes), or fully untagged
|
||||
output — where meta-narration ("The user asks... Let me check the excerpts...")
|
||||
is detected by a gated prefix + per-paragraph classifier and routed to 'think',
|
||||
while real content paragraphs start the 'answer'."""
|
||||
PRE, THINK, NARR, ANSWER = range(4)
|
||||
import re as _re
|
||||
_BOLD = r"\*\*[^*\n]{1,60}\*\*"
|
||||
NARR_START = _re.compile(
|
||||
r"^\s*(?:the user|user asks|user want|let me|i need|i'll|i will|i should|"
|
||||
r"okay\b|ok,|alright|hmm\b|to answer|so i|so,|we need|looking at|"
|
||||
r"first,|step 1[:.]|" + _BOLD + r")", _re.IGNORECASE)
|
||||
NARR_CONT = _re.compile(
|
||||
r"^\s*(?:" + _BOLD +
|
||||
r"|the (?:user|excerpts?|passages?|books?|sources?|quotes?)|none of"
|
||||
r"|excerpt \[|excerpt \d|passage \d|quote \d|source \d"
|
||||
r"|there (?:is|are) no|no (?:book|passage|excerpt|relevant)"
|
||||
r"|based on (?:the|these|my)|so,?(?: i| the| this| none)"
|
||||
r"|i (?:should|will|would|need|can't|cannot|don't|do not|think|see|notice)"
|
||||
r"|to answer|as (?:the|these|per)|these (?:excerpts?|passages?|books?)"
|
||||
r"|okay\b|alright|hmm|let me|however,? (?:the|none|there)|therefore,? (?:the|i)"
|
||||
r"|in (?:short|conclusion|summary)|overall|finally,? (?:i|the))", _re.IGNORECASE)
|
||||
GATE_CHARS = 140
|
||||
GATE_MAX = 400 # hard cap: decide even without a paragraph boundary
|
||||
|
||||
def __init__(self):
|
||||
self.buf, self.in_think = "", False
|
||||
self.state = self.PRE
|
||||
self.buf = "" # scanner buffer (PRE/THINK) OR paragraph buffer (NARR)
|
||||
self.ans_buf = "" # text held by the undecided gate
|
||||
self.gated = False # True once answer-vs-narration has been decided
|
||||
|
||||
def _emit(self, channel, s):
|
||||
return [(channel, s)] if s else []
|
||||
|
||||
def _take(self, n=None):
|
||||
if n is None:
|
||||
piece, self.buf = self.buf, ""
|
||||
else:
|
||||
piece, self.buf = self.buf[:n], self.buf[n:]
|
||||
return piece
|
||||
|
||||
def _gate_feed(self, s):
|
||||
"""PRE-state text: hold until the start-of-answer gate can judge.
|
||||
Decides at the first paragraph boundary after GATE_CHARS (or at GATE_MAX)."""
|
||||
if self.gated:
|
||||
return self._emit("answer", s)
|
||||
self.ans_buf += s
|
||||
boundary = self.ans_buf.find("\n\n", self.GATE_CHARS)
|
||||
if boundary == -1 and len(self.ans_buf) < self.GATE_MAX:
|
||||
return []
|
||||
cut = boundary + 2 if boundary != -1 else len(self.ans_buf)
|
||||
held, self.ans_buf = self.ans_buf[:cut], self.ans_buf[cut:]
|
||||
self.gated = True
|
||||
tail, self.buf = self.buf, "" # tail is NEWER than ans_buf; stream order = held + rest + tail
|
||||
if self.NARR_START.match(held):
|
||||
self.state = self.NARR
|
||||
out = self._emit("think", held)
|
||||
out += self._narr_feed(self.ans_buf + tail)
|
||||
self.ans_buf = ""
|
||||
return out
|
||||
self.state = self.ANSWER
|
||||
return self._emit("answer", held + self.ans_buf + tail)
|
||||
|
||||
def _narr_feed(self, s):
|
||||
"""NARR state: classify complete paragraphs; first content paragraph = answer."""
|
||||
self.buf += s
|
||||
out = []
|
||||
while "\n\n" in self.buf:
|
||||
para, rest = self.buf.split("\n\n", 1)
|
||||
if para.strip() and self.NARR_CONT.match(para):
|
||||
out += self._emit("think", para + "\n\n")
|
||||
self.buf = rest
|
||||
else:
|
||||
self.state = self.ANSWER
|
||||
self.gated = True
|
||||
out += self._emit("answer", self.buf)
|
||||
self.buf = ""
|
||||
return out
|
||||
return out
|
||||
|
||||
def feed(self, s):
|
||||
self.buf += s
|
||||
out = ""
|
||||
out = []
|
||||
while True:
|
||||
if self.in_think:
|
||||
i = self.buf.find(self.CLOSE)
|
||||
if i == -1:
|
||||
keep = min(len(self.buf), len(self.CLOSE) - 1)
|
||||
self.buf = self.buf[-keep:] if keep else ""
|
||||
if self.state == self.ANSWER:
|
||||
out += self._emit("answer", self._take())
|
||||
return out
|
||||
if self.state == self.NARR:
|
||||
out += self._narr_feed(self._take())
|
||||
if self.state == self.NARR:
|
||||
return out
|
||||
self.in_think = False
|
||||
self.buf = self.buf[i + len(self.CLOSE):]
|
||||
else:
|
||||
i_open = self.buf.find(self.OPEN)
|
||||
i_close = self.buf.find(self.CLOSE)
|
||||
# close-tag with no open-tag seen: drop the CoT preamble wholesale
|
||||
if i_close != -1 and (i_open == -1 or i_close < i_open):
|
||||
self.buf = self.buf[i_close + len(self.CLOSE):]
|
||||
self.in_think = False
|
||||
continue
|
||||
i_open, i_close = self.buf.find(OPEN), self.buf.find(CLOSE)
|
||||
if self.state == self.PRE:
|
||||
if i_open != -1 and (i_close == -1 or i_open < i_close):
|
||||
out += self._gate_feed(self._take(i_open))
|
||||
self.gated = True # tagged CoT follows; gate no longer needed
|
||||
self.buf = self.buf[len(OPEN):]
|
||||
self.state = self.THINK
|
||||
continue
|
||||
if i_open == -1:
|
||||
keep = min(len(self.buf), max(len(self.OPEN), len(self.CLOSE)) - 1)
|
||||
cut = len(self.buf) - keep
|
||||
out += self.buf[:cut]
|
||||
self.buf = self.buf[cut:]
|
||||
return out
|
||||
out += self.buf[:i_open]
|
||||
self.in_think = True
|
||||
self.buf = self.buf[i_open + len(self.OPEN):]
|
||||
if i_close != -1: # missing opener — preamble is reasoning
|
||||
held, self.ans_buf = self.ans_buf, ""
|
||||
out += self._emit("think", held + self._take(i_close))
|
||||
self.buf = self.buf[len(CLOSE):]
|
||||
self.state = self.ANSWER
|
||||
continue
|
||||
elif self.state == self.THINK:
|
||||
if i_close != -1:
|
||||
held, self.ans_buf = self.ans_buf, "" # rare held PRE text merges
|
||||
out += self._emit("think", held + self._take(i_close))
|
||||
self.buf = self.buf[len(CLOSE):]
|
||||
self.state = self.ANSWER
|
||||
continue
|
||||
# no marker found: emit all but a tag-length tail
|
||||
keep = min(len(self.buf), max(len(OPEN), len(CLOSE)) - 1)
|
||||
cut = len(self.buf) - keep
|
||||
if cut > 0:
|
||||
piece = self._take(cut)
|
||||
if self.state == self.THINK:
|
||||
out += self._emit("think", piece)
|
||||
else: # PRE: route through the gate
|
||||
out += self._gate_feed(piece)
|
||||
return out
|
||||
|
||||
def flush(self):
|
||||
out, self.buf = ("" if self.in_think else self.buf), ""
|
||||
out = []
|
||||
if self.state == self.PRE and (self.ans_buf or self.buf):
|
||||
rest = self.ans_buf + self.buf # stream order: held text, then tail
|
||||
self.ans_buf = self.buf = ""
|
||||
if self.NARR_START.match(rest):
|
||||
self.state = self.NARR
|
||||
out += self._narr_feed(rest) # classify paragraphs; content → answer
|
||||
else:
|
||||
out += self._emit("answer", rest)
|
||||
return out
|
||||
if self.state == self.NARR:
|
||||
residual, self.buf = self.buf, ""
|
||||
if residual.strip():
|
||||
out += self._emit("think" if self.NARR_CONT.match(residual) else "answer", residual)
|
||||
return out
|
||||
if self.buf:
|
||||
piece = self._take()
|
||||
if self.state == self.THINK:
|
||||
out += self._emit("think", piece + self.ans_buf)
|
||||
self.ans_buf = ""
|
||||
else:
|
||||
out += self._gate_feed(piece)
|
||||
if self.ans_buf: # gate decided held text still pending
|
||||
held, self.ans_buf = self.ans_buf, ""
|
||||
out += self._emit("think" if self.NARR_START.match(held) else "answer", held)
|
||||
return out
|
||||
|
||||
def stream_ollama(prompt):
|
||||
"""Yield (channel, piece) tuples: channel 'think' or 'answer'."""
|
||||
r = http.post(f"{CFG['ollama_url']}/api/generate",
|
||||
json={"model": CFG["model"], "prompt": prompt, "stream": True,
|
||||
"options": {"temperature": 0.4, "num_predict": 700,
|
||||
"options": {"temperature": 0.4, "num_predict": CFG.get("num_predict", 1100),
|
||||
"num_ctx": CFG.get("num_ctx", 8192)}},
|
||||
timeout=(5, None), stream=True)
|
||||
flt = CoTFilter()
|
||||
sp = ThinkSplitter()
|
||||
for line in r.iter_lines():
|
||||
if line:
|
||||
tok = json.loads(line).get("response", "")
|
||||
if tok:
|
||||
piece = flt.feed(tok)
|
||||
if piece:
|
||||
yield piece
|
||||
tail = flt.flush()
|
||||
if tail:
|
||||
yield tail
|
||||
for channel, piece in sp.feed(tok):
|
||||
yield channel, piece
|
||||
for channel, piece in sp.flush():
|
||||
yield channel, piece
|
||||
|
||||
def ask_ollama(prompt):
|
||||
"""Non-streaming RAG answer (reasoning stripped, tagged or not)."""
|
||||
r = http.post(f"{CFG['ollama_url']}/api/generate",
|
||||
json={"model": CFG["model"], "prompt": prompt, "stream": False,
|
||||
"options": {"temperature": 0.4, "num_predict": CFG.get("num_predict", 1100),
|
||||
"num_ctx": CFG.get("num_ctx", 8192)}},
|
||||
timeout=180)
|
||||
r.raise_for_status()
|
||||
_, answer = strip_cot(r.json().get("response", ""))
|
||||
if degenerate(answer):
|
||||
unload_model()
|
||||
raise RuntimeError("model returned degenerate output — model unloaded, retry")
|
||||
return answer
|
||||
|
||||
def degenerate(s):
|
||||
"""vornith VRAM/session corruption signature: run of '?' chars."""
|
||||
s = s.strip()
|
||||
return len(s) >= 40 and s.count("?") / len(s) > 0.4
|
||||
|
||||
def unload_model():
|
||||
"""Drop the model from VRAM so the next request reloads clean."""
|
||||
try:
|
||||
http.post(f"{CFG['ollama_url']}/api/generate",
|
||||
json={"model": CFG["model"], "keep_alive": 0}, timeout=10)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# ---------------------------------------------------------------- auth + metering
|
||||
def auth_user(req):
|
||||
|
||||
Reference in New Issue
Block a user