Cesium2 / src /regex_features.py
MORPH-AI
feat: dynamic MoE expansion, multi-head CoT, plugin architecture, improved MoD
82f262a
Raw
History Blame Contribute Delete
5.03 kB
"""
RegexFeatureExtractor - token feature gating.
Extends build_code_features (4-dim) to a 7-dim per-token feature vector plus
a running SyntaxState (bracket stack, quote open/closed) used as a cheap
pre-model sanity gate.
Dimensions per token:
[0] is_code_like (indent / brackets / operators / newlines)
[1] indent_depth (normalized leading whitespace)
[2] bracket_balance (1.0 open, 0.5 neutral, 0.0 close)
[3] has_newline
[4] keyword_hit (NEW)
[5] quote_state (running open/closed string literal) (NEW)
[6] numeric_literal (NEW)
"""
import re
import torch
from dataclasses import dataclass, field
from typing import List, Optional, Tuple
CODE_CHARS = set("{}()[];=<>!&|+-*/%'\"`#@.,:")
KEYWORD_RE = re.compile(
r"\b(def|class|import|return|if|else|elif|for|while|try|except|finally|"
r"lambda|pass|with|as|yield|from|async|await)\b"
)
NUMERIC_RE = re.compile(r"\b\d+(\.\d+)?\b")
OPEN_BRACKETS = {"{": 1, "[": 1, "(": 1}
CLOSE_BRACKETS = {"}": 1, "]": 1, ")": 1}
@dataclass
class SyntaxState:
bracket_stack: List[str] = field(default_factory=list)
quote_open: Optional[str] = None # '"' or "'" while a string literal is open
depth: int = 0
def is_balanced(self) -> bool:
return not self.bracket_stack and self.quote_open is None
def to_dict(self) -> dict:
return {
"balanced": self.is_balanced(),
"bracket_depth": len(self.bracket_stack),
"quote_open": self.quote_open,
}
class RegexFeatureExtractor:
def __init__(self, num_features: int = 7):
self.num_features = num_features
def extract(self, tokenizer, input_ids: torch.Tensor) -> Tuple[torch.Tensor, List[SyntaxState]]:
"""Per-token features (B, T, F) + one SyntaxState per row."""
feats: List[List[List[float]]] = []
states: List[SyntaxState] = []
for row in input_ids.tolist():
tokens = tokenizer.convert_ids_to_tokens(row)
state = SyntaxState()
row_feats = []
for tok in tokens:
is_code = any(c in CODE_CHARS for c in tok)
indent = 0.0
stripped = tok.lstrip()
if stripped and tok != stripped:
indent = min((len(tok) - len(stripped)) / 8.0, 1.0)
is_code = True
bal = 0.5
for c in tok:
if c in OPEN_BRACKETS:
bal = 1.0
state.bracket_stack.append(c)
elif c in CLOSE_BRACKETS:
bal = 0.0
if state.bracket_stack:
state.bracket_stack.pop()
# running quote state
for c in tok:
if c in ('"', "'"):
if state.quote_open is None:
state.quote_open = c
elif state.quote_open == c:
state.quote_open = None
quote = 1.0 if state.quote_open is not None else 0.0
newline = 1.0 if "\n" in tok else 0.0
kw = 1.0 if KEYWORD_RE.search(tok) else 0.0
num = 1.0 if NUMERIC_RE.search(tok) else 0.0
row_feats.append([1.0 if is_code else 0.0, indent, bal, newline, kw, quote, num])
state.depth = len(state.bracket_stack)
row_feats = row_feats[: input_ids.shape[1]]
feats.append(row_feats)
states.append(state)
max_len = max(len(r) for r in feats)
padded = [
r + [[0.0, 0.0, 0.5, 0.0, 0.0, 0.0, 0.0]] * (max_len - len(r))
for r in feats
]
t = torch.tensor(padded, dtype=torch.float32)
if t.shape[-1] > self.num_features:
t = t[..., : self.num_features]
return t, states
def extract_text(self, text: str) -> SyntaxState:
"""Run a string-only pass for the pre-model gate (no tokenizer)."""
state = SyntaxState()
for c in text:
if c in OPEN_BRACKETS:
state.bracket_stack.append(c)
elif c in CLOSE_BRACKETS and state.bracket_stack:
state.bracket_stack.pop()
elif c in ('"', "'"):
if state.quote_open is None:
state.quote_open = c
elif state.quote_open == c:
state.quote_open = None
state.depth = len(state.bracket_stack)
return state
def gate(self, syntax_state: SyntaxState) -> str:
"""Returns 'pass' | 'warn' | 'block'."""
if syntax_state.is_balanced():
return "pass"
return "block" if syntax_state.depth > 4 else "warn"
# drop-in replacement for architecture.build_code_features with 7-dim output
def build_code_features_v2(tokenizer, input_ids: torch.Tensor) -> torch.Tensor:
return RegexFeatureExtractor(num_features=7).extract(tokenizer, input_ids)[0]